import io import os import time import logging from contextlib import asynccontextmanager from fastapi import FastAPI, File, UploadFile, HTTPException, Security, Depends from fastapi.security.api_key import APIKeyHeader from fastapi.middleware.cors import CORSMiddleware from PIL import Image from schemas import PredictResponse, Detection, BoundingBox, HealthResponse from model_loader import load_model, get_model, get_class_names, get_model_version logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) # ── API key auth ────────────────────────────────────────────────────────── API_KEY_NAME = "X-API-Key" api_key_header = APIKeyHeader(name=API_KEY_NAME, auto_error=False) API_SECRET_KEY = os.environ.get("API_SECRET_KEY", "") def verify_api_key(api_key: str = Security(api_key_header)): if not API_SECRET_KEY: raise RuntimeError( "API_SECRET_KEY is not configured. " "Add it under Space Settings → Repository secrets." ) if api_key != API_SECRET_KEY: raise HTTPException(status_code=403, detail="Invalid or missing API key.") return api_key # ── Load model at startup ───────────────────────────────────────────────── @asynccontextmanager async def lifespan(app: FastAPI): logger.info("Startup: loading model...") load_model() logger.info("Model ready. Accepting requests.") yield logger.info("Shutdown.") # ── App ─────────────────────────────────────────────────────────────────── app = FastAPI( title = "Fabric Defect Detection API", description = "YOLOv8 ONNX — detects defects in fabric images.", version = "1.0.0", lifespan = lifespan, ) app.add_middleware( CORSMiddleware, allow_origins = ["*"], # lock down to your React domain in production allow_methods = ["POST", "GET"], allow_headers = ["*"], ) # ── Config from env (optional overrides) ───────────────────────────────── CONF_THRESH = float(os.environ.get("CONF_THRESH", "0.35")) IOU_THRESH = float(os.environ.get("IOU_THRESH", "0.45")) MAX_SIZE_MB = 10 ALLOWED_MIME = {"image/jpeg", "image/png", "image/webp", "image/bmp"} def _severity(conf: float) -> str: if conf >= 0.75: return "high" if conf >= 0.50: return "medium" return "low" # ── Routes ──────────────────────────────────────────────────────────────── @app.get("/", tags=["System"]) def root(): return { "api" : "Fabric Defect Detection API", "version": "1.0.0", "docs" : "/docs", } @app.get("/health", response_model=HealthResponse, tags=["System"]) def health(): try: m = get_model() loaded = m is not None classes = list(get_class_names().values()) except Exception: loaded = False classes = [] return HealthResponse( status = "ok" if loaded else "model_not_loaded", model_loaded = loaded, model_version = get_model_version(), classes = classes, ) @app.post("/predict", response_model=PredictResponse, tags=["Inference"]) async def predict( file: UploadFile = File(..., description="Fabric image — JPEG/PNG/WebP, max 10 MB"), _key: str = Depends(verify_api_key), ): # Validate MIME type if file.content_type not in ALLOWED_MIME: raise HTTPException( status_code=415, detail=f"Unsupported file type '{file.content_type}'. Use JPEG, PNG, or WebP." ) # Read and size-check raw = await file.read() if len(raw) > MAX_SIZE_MB * 1024 * 1024: raise HTTPException( status_code=413, detail=f"File too large. Maximum is {MAX_SIZE_MB} MB." ) # Decode image try: image = Image.open(io.BytesIO(raw)).convert("RGB") except Exception: raise HTTPException(status_code=422, detail="Could not decode image file.") # Run inference model = get_model() class_names = get_class_names() t0 = time.perf_counter() results = model(image, conf=CONF_THRESH, iou=IOU_THRESH, verbose=False)[0] elapsed = round((time.perf_counter() - t0) * 1000, 2) # Parse boxes detections = [] if results.boxes is not None: for box in results.boxes: cls_id = int(box.cls) conf = round(float(box.conf), 4) x1,y1,x2,y2 = map(int, box.xyxy[0].tolist()) defect_type = class_names.get(cls_id, f"class_{cls_id}") if defect_type == "no_defect": continue detections.append(Detection( defect_type = defect_type, confidence = conf, bbox = BoundingBox(x1=x1, y1=y1, x2=x2, y2=y2), severity = _severity(conf), )) return PredictResponse( pass_fail = "FAIL" if detections else "PASS", total_defects = len(detections), defects = detections, processing_time_ms = elapsed, model_version = get_model_version(), )