aruntandra's picture
Upload main.py
668c77a verified
Raw
History Blame Contribute Delete
7.21 kB
import sys
import os
import io
import traceback
from pathlib import Path
from typing import Any
import numpy as np
import cv2
from PIL import Image
from fastapi import FastAPI, Request, File, UploadFile, HTTPException
from fastapi.staticfiles import StaticFiles
from fastapi.templating import Jinja2Templates
from fastapi.middleware.cors import CORSMiddleware
# Add parent path so imports work (kept from your code)
sys.path.append(str(Path(__file__).resolve().parent.parent))
from app.config import settings
from app import __version__
from app.Hackathon_setup import face_recognition, exp_recognition
# -----------------------------------------------------------------------------
# App setup
# -----------------------------------------------------------------------------
app = FastAPI(
title=settings.PROJECT_NAME,
openapi_url=f"{settings.API_V1_STR}/openapi.json",
)
# Static and templates
STATIC_DIR = "app/static"
os.makedirs(STATIC_DIR, exist_ok=True)
app.mount("/static", StaticFiles(directory=STATIC_DIR), name="static")
templates = Jinja2Templates(directory="app/templates")
# -----------------------------------------------------------------------------
# Helpers
# -----------------------------------------------------------------------------
async def save_and_load_rgb(upload: UploadFile):
"""
Read upload ONCE -> save to app/static -> return (relative_name_for_html, rgb_np)
rgb_np: HxWx3 uint8 RGB numpy array
"""
if upload is None:
raise ValueError("No file uploaded")
if not upload.content_type or "image" not in upload.content_type:
raise ValueError(f"Invalid content_type: {upload.content_type}")
data = await upload.read()
if not data:
raise ValueError("Uploaded file is empty")
# Safe filename (avoid path traversal)
filename = os.path.basename(upload.filename) if upload.filename else "uploaded.png"
save_path = os.path.join(STATIC_DIR, filename)
# Save bytes for template display
with open(save_path, "wb") as f:
f.write(data)
# Decode image in memory for inference (ALWAYS convert to RGB)
pil_img = Image.open(io.BytesIO(data)).convert("RGB")
rgb_np = np.asarray(pil_img, dtype=np.uint8) # HxWx3 RGB
# Return the relative path expected by templates
rel_for_html = "../static/" + filename
return rel_for_html, rgb_np
def to_bgr(rgb_np: np.ndarray) -> np.ndarray:
"""Optional: convert RGB numpy to BGR for OpenCV-based pipelines."""
return cv2.cvtColor(rgb_np, cv2.COLOR_RGB2BGR)
# -----------------------------------------------------------------------------
# Home
# -----------------------------------------------------------------------------
@app.get("/")
async def root(request: Request):
return templates.TemplateResponse("index.html", {"request": request})
# -----------------------------------------------------------------------------
# Face Similarity
# -----------------------------------------------------------------------------
@app.get("/similarity/")
async def similarity_root(request: Request):
return templates.TemplateResponse("similarity.html", {"request": request})
@app.post("/predict_similarity/")
async def predict_similarity(request: Request, file1: UploadFile = File(...), file2: UploadFile = File(...)):
try:
simi_path1, img1_rgb = await save_and_load_rgb(file1)
simi_path2, img2_rgb = await save_and_load_rgb(file2)
# If your downstream expects BGR, uncomment these lines:
# img1 = to_bgr(img1_rgb)
# img2 = to_bgr(img2_rgb)
# If your face_recognition module handles RGB->BGR internally, keep as-is:
img1 = img1_rgb
img2 = img2_rgb
result = face_recognition.get_similarity(img1, img2)
return templates.TemplateResponse(
"predict_similarity.html",
{
"request": request,
"result": float(np.round(result, 3)),
"simi_filename1": simi_path1,
"simi_filename2": simi_path2,
},
)
except Exception as e:
print("ERROR in /predict_similarity/:", repr(e))
traceback.print_exc()
raise HTTPException(status_code=500, detail=str(e))
# -----------------------------------------------------------------------------
# Face Recognition
# -----------------------------------------------------------------------------
@app.get("/face_recognition/")
async def face_recognition_root(request: Request):
return templates.TemplateResponse("face_recognition.html", {"request": request})
@app.post("/predict_face_recognition/")
async def predict_face_recognition(request: Request, file3: UploadFile = File(...)):
try:
face_path, img_rgb = await save_and_load_rgb(file3)
# Optional BGR conversion if needed:
# img = to_bgr(img_rgb)
img = img_rgb
result = face_recognition.get_face_class(img)
return templates.TemplateResponse(
"predict_face_recognition.html",
{
"request": request,
"result": result,
"face_rec_filename": face_path,
},
)
except Exception as e:
print("ERROR in /predict_face_recognition/:", repr(e))
traceback.print_exc()
raise HTTPException(status_code=500, detail=str(e))
# -----------------------------------------------------------------------------
# Expression Recognition
# -----------------------------------------------------------------------------
@app.get("/expr_recognition/")
async def expr_recognition_root(request: Request):
return templates.TemplateResponse("expr_recognition.html", {"request": request})
@app.post("/predict_expr_recognition/")
async def predict_expr_recognition(request: Request, file4: UploadFile = File(...)):
try:
expr_path, img_rgb = await save_and_load_rgb(file4)
# Optional BGR conversion if needed:
# img = to_bgr(img_rgb)
img = img_rgb
result = exp_recognition.get_expression(img)
return templates.TemplateResponse(
"predict_expr_recognition.html",
{
"request": request,
"result": result,
"expr_rec_filename": expr_path,
},
)
except Exception as e:
print("ERROR in /predict_expr_recognition/:", repr(e))
traceback.print_exc()
raise HTTPException(status_code=500, detail=str(e))
# -----------------------------------------------------------------------------
# CORS
# -----------------------------------------------------------------------------
if settings.BACKEND_CORS_ORIGINS:
app.add_middleware(
CORSMiddleware,
allow_origins=[str(origin) for origin in settings.BACKEND_CORS_ORIGINS],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# -----------------------------------------------------------------------------
# Run locally
# -----------------------------------------------------------------------------
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8001)