2026-05-12 11:16:56 +02:00
|
|
|
|
"""
|
|
|
|
|
|
API Server — Detect + Classify with Polling
|
|
|
|
|
|
=============================================
|
|
|
|
|
|
POST /jobs — submit an image + class filter, get back a job_id
|
|
|
|
|
|
GET /jobs/{job_id} — poll for status / progress / results
|
|
|
|
|
|
|
|
|
|
|
|
Polling pattern
|
|
|
|
|
|
---------------
|
|
|
|
|
|
1. Client POSTs image (base64) + interested_classes list
|
|
|
|
|
|
2. Server returns { "job_id": "<uuid>" } immediately
|
|
|
|
|
|
3. Client polls GET /jobs/{job_id} every 3 s
|
|
|
|
|
|
4. Server returns:
|
|
|
|
|
|
{
|
|
|
|
|
|
"status": "queued" | "processing" | "done" | "failed",
|
|
|
|
|
|
"progress": 0‥100,
|
|
|
|
|
|
"results": {
|
|
|
|
|
|
"annotated_image": "<base64 jpg>",
|
|
|
|
|
|
"counts": { "cola_can": 3, "pepsi_can": 1, … },
|
|
|
|
|
|
"detections": [
|
|
|
|
|
|
{ "class": "cola_can", "confidence": 0.91,
|
|
|
|
|
|
"bbox": [x1, y1, x2, y2] },
|
|
|
|
|
|
…
|
|
|
|
|
|
]
|
|
|
|
|
|
} // null while processing
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
Run
|
|
|
|
|
|
---
|
|
|
|
|
|
pip install fastapi uvicorn python-multipart
|
|
|
|
|
|
|
|
|
|
|
|
uvicorn api_server:app --host 0.0.0.0 --port 8000 --workers 1
|
|
|
|
|
|
|
|
|
|
|
|
Environment variables (all optional)
|
|
|
|
|
|
-------------------------------------
|
|
|
|
|
|
DETECTOR_WEIGHTS path to YOLOv8n .pt (default: detector/best.pt)
|
|
|
|
|
|
CLASSIFIER_WEIGHTS path to classifier .pt (default: runs/classify/best.pt)
|
|
|
|
|
|
DET_CONF float (default: 0.25)
|
|
|
|
|
|
CLS_CONF float (default: 0.50)
|
|
|
|
|
|
DET_IOU float (default: 0.45)
|
|
|
|
|
|
PADDING int px (default: 10)
|
|
|
|
|
|
MAX_JOBS max jobs kept in memory (default: 200)
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
import base64
|
|
|
|
|
|
import io
|
|
|
|
|
|
import json
|
|
|
|
|
|
import os
|
|
|
|
|
|
import uuid
|
|
|
|
|
|
from collections import OrderedDict
|
|
|
|
|
|
from concurrent.futures import ThreadPoolExecutor
|
|
|
|
|
|
from enum import Enum
|
|
|
|
|
|
from typing import Optional
|
|
|
|
|
|
|
|
|
|
|
|
import cv2
|
|
|
|
|
|
import numpy as np
|
|
|
|
|
|
import torch
|
|
|
|
|
|
import torch.nn.functional as F
|
|
|
|
|
|
from fastapi import FastAPI, HTTPException, BackgroundTasks
|
|
|
|
|
|
from fastapi.middleware.cors import CORSMiddleware
|
|
|
|
|
|
from pydantic import BaseModel, Field
|
|
|
|
|
|
from torchvision import models, transforms
|
|
|
|
|
|
from ultralytics import YOLO
|
|
|
|
|
|
|
|
|
|
|
|
# ─────────────────────────── config ──────────────────────────────────────────
|
|
|
|
|
|
|
|
|
|
|
|
DETECTOR_WEIGHTS = os.getenv("DETECTOR_WEIGHTS", "detector/best.pt")
|
2026-05-14 15:48:37 +02:00
|
|
|
|
CLASSIFIER_WEIGHTS = os.getenv("CLASSIFIER_WEIGHTS", "runs/classify/best.pt")
|
2026-05-12 11:16:56 +02:00
|
|
|
|
DET_CONF = float(os.getenv("DET_CONF", "0.25"))
|
|
|
|
|
|
CLS_CONF = float(os.getenv("CLS_CONF", "0.75"))
|
|
|
|
|
|
DET_IOU = float(os.getenv("DET_IOU", "0.45"))
|
|
|
|
|
|
PADDING = int(os.getenv("PADDING", "10"))
|
|
|
|
|
|
MAX_JOBS = int(os.getenv("MAX_JOBS", "200"))
|
|
|
|
|
|
|
|
|
|
|
|
# ─────────────────────────── device ──────────────────────────────────────────
|
|
|
|
|
|
|
|
|
|
|
|
def get_device() -> torch.device:
|
|
|
|
|
|
if torch.cuda.is_available():
|
|
|
|
|
|
return torch.device("cuda")
|
|
|
|
|
|
if torch.backends.mps.is_available():
|
|
|
|
|
|
return torch.device("mps")
|
|
|
|
|
|
return torch.device("cpu")
|
|
|
|
|
|
|
|
|
|
|
|
DEVICE = get_device()
|
|
|
|
|
|
|
|
|
|
|
|
# ─────────────────────────── OOM helpers (same as inference script) ──────────
|
|
|
|
|
|
|
|
|
|
|
|
def _is_oom(exc: BaseException) -> bool:
|
|
|
|
|
|
if isinstance(exc, MemoryError):
|
|
|
|
|
|
return True
|
|
|
|
|
|
if torch.cuda.is_available() and isinstance(exc, torch.cuda.OutOfMemoryError):
|
|
|
|
|
|
return True
|
|
|
|
|
|
if isinstance(exc, RuntimeError):
|
|
|
|
|
|
msg = str(exc).lower()
|
|
|
|
|
|
return "out of memory" in msg or ("memory" in msg and "alloc" in msg)
|
|
|
|
|
|
return False
|
|
|
|
|
|
|
|
|
|
|
|
def _free_device_cache(device: torch.device) -> None:
|
|
|
|
|
|
if device.type == "cuda":
|
|
|
|
|
|
torch.cuda.empty_cache()
|
|
|
|
|
|
elif device.type == "mps" and hasattr(torch.mps, "empty_cache"):
|
|
|
|
|
|
torch.mps.empty_cache()
|
|
|
|
|
|
|
|
|
|
|
|
def _sync_device(device: torch.device) -> None:
|
|
|
|
|
|
if device.type == "cuda":
|
|
|
|
|
|
torch.cuda.synchronize()
|
|
|
|
|
|
elif device.type == "mps" and hasattr(torch.mps, "synchronize"):
|
|
|
|
|
|
torch.mps.synchronize()
|
|
|
|
|
|
|
|
|
|
|
|
# ─────────────────────────── classifier ──────────────────────────────────────
|
|
|
|
|
|
|
|
|
|
|
|
class ProductClassifier:
|
|
|
|
|
|
_BUILDERS = {
|
|
|
|
|
|
"efficientnet_b0": models.efficientnet_b0,
|
|
|
|
|
|
"efficientnet_b2": models.efficientnet_b2,
|
|
|
|
|
|
"mobilenet_v3_small": models.mobilenet_v3_small,
|
|
|
|
|
|
"resnet50": models.resnet50,
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
def __init__(self, weights_path: str, device: torch.device, img_size: int = 224):
|
|
|
|
|
|
ckpt = torch.load(weights_path, map_location=device)
|
|
|
|
|
|
self.class_names: list[str] = ckpt["class_names"]
|
|
|
|
|
|
num_classes = len(self.class_names)
|
|
|
|
|
|
model_name = ckpt.get("args", {}).get("model", "efficientnet_b0")
|
|
|
|
|
|
|
|
|
|
|
|
base = self._BUILDERS[model_name](weights=None)
|
|
|
|
|
|
# Attach correct head
|
|
|
|
|
|
if "efficientnet" in model_name:
|
|
|
|
|
|
base.classifier[1] = torch.nn.Linear(base.classifier[1].in_features, num_classes)
|
|
|
|
|
|
elif "mobilenet" in model_name:
|
|
|
|
|
|
base.classifier[3] = torch.nn.Linear(base.classifier[3].in_features, num_classes)
|
|
|
|
|
|
else:
|
|
|
|
|
|
base.fc = torch.nn.Linear(base.fc.in_features, num_classes)
|
|
|
|
|
|
|
|
|
|
|
|
base.load_state_dict(ckpt["model_state"])
|
|
|
|
|
|
self.model = base.to(device).eval()
|
|
|
|
|
|
self.device = device
|
|
|
|
|
|
|
|
|
|
|
|
mean = [0.485, 0.456, 0.406]
|
|
|
|
|
|
std = [0.229, 0.224, 0.225]
|
|
|
|
|
|
self.transform = transforms.Compose([
|
|
|
|
|
|
transforms.ToPILImage(),
|
|
|
|
|
|
transforms.Resize(int(img_size * 1.15)),
|
|
|
|
|
|
transforms.CenterCrop(img_size),
|
|
|
|
|
|
transforms.ToTensor(),
|
|
|
|
|
|
transforms.Normalize(mean, std),
|
|
|
|
|
|
])
|
|
|
|
|
|
self._batch_size: int = self._calibrate(img_size)
|
|
|
|
|
|
|
|
|
|
|
|
def _calibrate(self, img_size: int) -> int:
|
|
|
|
|
|
if self.device.type == "cpu":
|
|
|
|
|
|
return 64
|
|
|
|
|
|
dummy = torch.zeros(1, 3, img_size, img_size, device=self.device)
|
|
|
|
|
|
safe, probe = 1, 2
|
|
|
|
|
|
with torch.no_grad():
|
|
|
|
|
|
while probe <= 512:
|
|
|
|
|
|
try:
|
|
|
|
|
|
self.model(dummy.expand(probe, -1, -1, -1))
|
|
|
|
|
|
_sync_device(self.device)
|
|
|
|
|
|
safe, probe = probe, probe * 2
|
|
|
|
|
|
except Exception as exc:
|
|
|
|
|
|
if _is_oom(exc):
|
|
|
|
|
|
break
|
|
|
|
|
|
raise
|
|
|
|
|
|
finally:
|
|
|
|
|
|
_free_device_cache(self.device)
|
|
|
|
|
|
return max(1, int(safe * 0.8))
|
|
|
|
|
|
|
|
|
|
|
|
@torch.no_grad()
|
|
|
|
|
|
def predict_batch(self, crops_bgr: list[np.ndarray]) -> list[tuple[str, float]]:
|
|
|
|
|
|
if not crops_bgr:
|
|
|
|
|
|
return []
|
|
|
|
|
|
tensors = [self.transform(cv2.cvtColor(c, cv2.COLOR_BGR2RGB)) for c in crops_bgr]
|
|
|
|
|
|
results: list[tuple[str, float]] = []
|
|
|
|
|
|
chunk_size = self._batch_size
|
|
|
|
|
|
i = 0
|
|
|
|
|
|
while i < len(tensors):
|
|
|
|
|
|
batch = torch.stack(tensors[i : i + chunk_size]).to(self.device)
|
|
|
|
|
|
try:
|
|
|
|
|
|
probs = F.softmax(self.model(batch), dim=1)
|
|
|
|
|
|
confs, idxs = probs.max(dim=1)
|
|
|
|
|
|
for idx, conf in zip(idxs, confs):
|
|
|
|
|
|
results.append((self.class_names[idx.item()], conf.item()))
|
|
|
|
|
|
i += chunk_size
|
|
|
|
|
|
except Exception as exc:
|
|
|
|
|
|
if not _is_oom(exc):
|
|
|
|
|
|
raise
|
|
|
|
|
|
_free_device_cache(self.device)
|
|
|
|
|
|
chunk_size = self._batch_size = max(1, chunk_size // 2)
|
|
|
|
|
|
finally:
|
|
|
|
|
|
del batch
|
|
|
|
|
|
return results
|
|
|
|
|
|
|
|
|
|
|
|
# ─────────────────────────── model singletons ────────────────────────────────
|
|
|
|
|
|
|
|
|
|
|
|
print(f"[startup] Device: {DEVICE}")
|
|
|
|
|
|
print(f"[startup] Loading detector : {DETECTOR_WEIGHTS}")
|
|
|
|
|
|
DETECTOR = YOLO(DETECTOR_WEIGHTS)
|
|
|
|
|
|
print(f"[startup] Loading classifier: {CLASSIFIER_WEIGHTS}")
|
|
|
|
|
|
CLASSIFIER = ProductClassifier(CLASSIFIER_WEIGHTS, DEVICE)
|
|
|
|
|
|
print(f"[startup] Ready — classes: {CLASSIFIER.class_names}")
|
|
|
|
|
|
|
|
|
|
|
|
# Thread pool: one worker keeps GPU access serialised; increase if CPU-only
|
|
|
|
|
|
EXECUTOR = ThreadPoolExecutor(max_workers=1)
|
|
|
|
|
|
|
|
|
|
|
|
# ─────────────────────────── job store ───────────────────────────────────────
|
|
|
|
|
|
|
|
|
|
|
|
class Status(str, Enum):
|
|
|
|
|
|
queued = "queued"
|
|
|
|
|
|
processing = "processing"
|
|
|
|
|
|
done = "done"
|
|
|
|
|
|
failed = "failed"
|
|
|
|
|
|
|
|
|
|
|
|
class Job:
|
|
|
|
|
|
def __init__(self, job_id: str, image_bytes: bytes,
|
|
|
|
|
|
interested_classes: list[str], cls_conf: float):
|
|
|
|
|
|
self.job_id = job_id
|
|
|
|
|
|
self.image_bytes = image_bytes
|
|
|
|
|
|
self.interested_classes = [c.lower() for c in interested_classes]
|
|
|
|
|
|
self.cls_conf = cls_conf
|
|
|
|
|
|
self.status = Status.queued
|
|
|
|
|
|
self.progress = 0
|
|
|
|
|
|
self.results = None
|
|
|
|
|
|
self.error = None
|
|
|
|
|
|
|
|
|
|
|
|
# OrderedDict so we can evict oldest jobs when MAX_JOBS is reached
|
|
|
|
|
|
JOB_STORE: OrderedDict[str, Job] = OrderedDict()
|
|
|
|
|
|
|
|
|
|
|
|
def store_job(job: Job) -> None:
|
|
|
|
|
|
if len(JOB_STORE) >= MAX_JOBS:
|
|
|
|
|
|
JOB_STORE.popitem(last=False) # evict oldest
|
|
|
|
|
|
JOB_STORE[job.job_id] = job
|
|
|
|
|
|
|
|
|
|
|
|
def get_job(job_id: str) -> Job:
|
|
|
|
|
|
job = JOB_STORE.get(job_id)
|
|
|
|
|
|
if job is None:
|
|
|
|
|
|
raise HTTPException(status_code=404, detail=f"Job '{job_id}' not found")
|
|
|
|
|
|
return job
|
|
|
|
|
|
|
|
|
|
|
|
# ─────────────────────────── drawing ─────────────────────────────────────────
|
|
|
|
|
|
|
|
|
|
|
|
_COLOUR_CACHE: dict[str, tuple] = {}
|
|
|
|
|
|
|
|
|
|
|
|
def class_colour(name: str) -> tuple:
|
|
|
|
|
|
if name not in _COLOUR_CACHE:
|
|
|
|
|
|
h = hash(name) % 180 # OpenCV hue: 0–179
|
|
|
|
|
|
col = cv2.cvtColor(
|
|
|
|
|
|
np.array([[[h, 220, 200]]], dtype=np.uint8), cv2.COLOR_HSV2BGR
|
|
|
|
|
|
)[0][0]
|
|
|
|
|
|
_COLOUR_CACHE[name] = tuple(int(c) for c in col)
|
|
|
|
|
|
return _COLOUR_CACHE[name]
|
|
|
|
|
|
|
|
|
|
|
|
def draw_box(frame: np.ndarray, x1: int, y1: int, x2: int, y2: int,
|
|
|
|
|
|
label: str, conf: float) -> None:
|
|
|
|
|
|
colour = class_colour(label)
|
|
|
|
|
|
cv2.rectangle(frame, (x1, y1), (x2, y2), colour, 2)
|
|
|
|
|
|
text = f"{label} {conf:.2f}"
|
|
|
|
|
|
(tw, th), bl = cv2.getTextSize(text, cv2.FONT_HERSHEY_SIMPLEX, 0.55, 1)
|
|
|
|
|
|
cv2.rectangle(frame, (x1, y1 - th - bl - 4), (x1 + tw + 4, y1), colour, -1)
|
|
|
|
|
|
cv2.putText(frame, text, (x1 + 2, y1 - bl - 2),
|
|
|
|
|
|
cv2.FONT_HERSHEY_SIMPLEX, 0.55, (255, 255, 255), 1, cv2.LINE_AA)
|
|
|
|
|
|
|
|
|
|
|
|
# ─────────────────────────── pipeline ────────────────────────────────────────
|
|
|
|
|
|
|
|
|
|
|
|
def _encode_image(frame: np.ndarray) -> str:
|
|
|
|
|
|
"""Encode a BGR numpy array as a base64 JPEG string."""
|
|
|
|
|
|
_, buf = cv2.imencode(".jpg", frame, [cv2.IMWRITE_JPEG_QUALITY, 90])
|
|
|
|
|
|
return base64.b64encode(buf).decode("utf-8")
|
|
|
|
|
|
|
|
|
|
|
|
def run_pipeline(job: Job) -> None:
|
|
|
|
|
|
"""Blocking function — runs in the thread-pool executor."""
|
|
|
|
|
|
try:
|
|
|
|
|
|
job.status = Status.processing
|
|
|
|
|
|
job.progress = 5
|
|
|
|
|
|
|
|
|
|
|
|
# ── Decode image ─────────────────────────────────────────────────────
|
|
|
|
|
|
arr = np.frombuffer(job.image_bytes, np.uint8)
|
|
|
|
|
|
frame = cv2.imdecode(arr, cv2.IMREAD_COLOR)
|
|
|
|
|
|
if frame is None:
|
|
|
|
|
|
raise ValueError("Could not decode image — unsupported format or corrupt data")
|
|
|
|
|
|
|
|
|
|
|
|
job.progress = 10
|
|
|
|
|
|
|
|
|
|
|
|
# ── Detection ────────────────────────────────────────────────────────
|
|
|
|
|
|
h, w = frame.shape[:2]
|
|
|
|
|
|
det_results = DETECTOR(
|
|
|
|
|
|
frame, conf=DET_CONF, iou=DET_IOU, verbose=False
|
|
|
|
|
|
)[0]
|
|
|
|
|
|
|
|
|
|
|
|
job.progress = 30
|
|
|
|
|
|
|
|
|
|
|
|
# ── Crop each detection ───────────────────────────────────────────────
|
|
|
|
|
|
crops, boxes = [], []
|
|
|
|
|
|
for box in det_results.boxes:
|
|
|
|
|
|
x1, y1, x2, y2 = map(int, box.xyxy[0].tolist())
|
|
|
|
|
|
x1 = max(0, x1 - PADDING); y1 = max(0, y1 - PADDING)
|
|
|
|
|
|
x2 = min(w, x2 + PADDING); y2 = min(h, y2 + PADDING)
|
|
|
|
|
|
crop = frame[y1:y2, x1:x2]
|
|
|
|
|
|
if crop.size == 0:
|
|
|
|
|
|
continue
|
|
|
|
|
|
crops.append(crop)
|
|
|
|
|
|
boxes.append((x1, y1, x2, y2))
|
|
|
|
|
|
|
|
|
|
|
|
job.progress = 50
|
|
|
|
|
|
|
|
|
|
|
|
# ── Classify ─────────────────────────────────────────────────────────
|
|
|
|
|
|
predictions = CLASSIFIER.predict_batch(crops)
|
|
|
|
|
|
|
|
|
|
|
|
job.progress = 75
|
|
|
|
|
|
|
|
|
|
|
|
# ── Filter by interested_classes + cls_conf ───────────────────────────
|
|
|
|
|
|
annotated = frame.copy()
|
|
|
|
|
|
counts: dict[str, int] = {}
|
|
|
|
|
|
detections: list[dict] = []
|
|
|
|
|
|
|
|
|
|
|
|
# If the user sends an empty list → show everything above cls_conf
|
|
|
|
|
|
filter_active = len(job.interested_classes) > 0
|
|
|
|
|
|
|
|
|
|
|
|
for (x1, y1, x2, y2), (cls_name, conf) in zip(boxes, predictions):
|
|
|
|
|
|
if conf < job.cls_conf:
|
|
|
|
|
|
continue
|
|
|
|
|
|
if filter_active and cls_name.lower() not in job.interested_classes:
|
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
|
|
draw_box(annotated, x1, y1, x2, y2, cls_name, conf)
|
|
|
|
|
|
counts[cls_name] = counts.get(cls_name, 0) + 1
|
|
|
|
|
|
detections.append({
|
|
|
|
|
|
"class": cls_name,
|
|
|
|
|
|
"confidence": round(conf, 4),
|
|
|
|
|
|
"bbox": [x1, y1, x2, y2],
|
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
|
|
job.progress = 90
|
|
|
|
|
|
|
|
|
|
|
|
# ── Encode output image ───────────────────────────────────────────────
|
|
|
|
|
|
annotated_b64 = _encode_image(annotated)
|
|
|
|
|
|
|
|
|
|
|
|
job.results = {
|
|
|
|
|
|
"annotated_image": annotated_b64,
|
|
|
|
|
|
"counts": counts,
|
|
|
|
|
|
"detections": detections,
|
|
|
|
|
|
}
|
|
|
|
|
|
job.status = Status.done
|
|
|
|
|
|
job.progress = 100
|
|
|
|
|
|
|
|
|
|
|
|
except Exception as exc:
|
|
|
|
|
|
job.status = Status.failed
|
|
|
|
|
|
job.error = str(exc)
|
|
|
|
|
|
raise
|
|
|
|
|
|
|
|
|
|
|
|
# ─────────────────────────── schemas ─────────────────────────────────────────
|
|
|
|
|
|
|
|
|
|
|
|
class SubmitRequest(BaseModel):
|
|
|
|
|
|
# Base64-encoded image (JPEG / PNG / BMP …)
|
|
|
|
|
|
image: str = Field(..., description="Base64-encoded image bytes")
|
|
|
|
|
|
# Classes to keep in results; empty list = keep all classes
|
|
|
|
|
|
interested_classes: list[str] = Field(
|
|
|
|
|
|
default=[],
|
|
|
|
|
|
description="Product class names to include. Leave empty to show all.",
|
|
|
|
|
|
examples=[["cola_can", "pepsi_can"]],
|
|
|
|
|
|
)
|
|
|
|
|
|
# Per-request confidence override (optional)
|
|
|
|
|
|
cls_conf: Optional[float] = Field(
|
|
|
|
|
|
default=None,
|
|
|
|
|
|
ge=0.0, le=1.0,
|
|
|
|
|
|
description="Min classifier confidence (overrides server default)",
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
class SubmitResponse(BaseModel):
|
|
|
|
|
|
job_id: str
|
|
|
|
|
|
|
|
|
|
|
|
class PollResponse(BaseModel):
|
|
|
|
|
|
status: Status
|
|
|
|
|
|
progress: int
|
|
|
|
|
|
results: Optional[dict] = None
|
|
|
|
|
|
error: Optional[str] = None
|
|
|
|
|
|
|
|
|
|
|
|
# ─────────────────────────── app ─────────────────────────────────────────────
|
|
|
|
|
|
|
|
|
|
|
|
app = FastAPI(
|
|
|
|
|
|
title="Product Detection API",
|
|
|
|
|
|
description="Detect and classify shelf products using YOLOv8 + EfficientNet",
|
|
|
|
|
|
version="1.0.0",
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
app.add_middleware(
|
|
|
|
|
|
CORSMiddleware,
|
|
|
|
|
|
allow_origins=["*"],
|
|
|
|
|
|
allow_methods=["*"],
|
|
|
|
|
|
allow_headers=["*"],
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
# ── POST /jobs ────────────────────────────────────────────────────────────────
|
|
|
|
|
|
|
|
|
|
|
|
@app.post("/jobs", response_model=SubmitResponse, status_code=202)
|
|
|
|
|
|
async def submit_job(body: SubmitRequest, background_tasks: BackgroundTasks):
|
|
|
|
|
|
"""
|
|
|
|
|
|
Submit an image for detection + classification.
|
|
|
|
|
|
|
|
|
|
|
|
Returns a job_id to poll with GET /jobs/{job_id}.
|
|
|
|
|
|
"""
|
|
|
|
|
|
# Decode base64 → raw bytes
|
|
|
|
|
|
try:
|
|
|
|
|
|
image_bytes = base64.b64decode(body.image)
|
|
|
|
|
|
except Exception:
|
|
|
|
|
|
raise HTTPException(status_code=422, detail="Invalid base64 in 'image' field")
|
|
|
|
|
|
|
|
|
|
|
|
job_id = str(uuid.uuid4())
|
|
|
|
|
|
job = Job(
|
|
|
|
|
|
job_id = job_id,
|
|
|
|
|
|
image_bytes = image_bytes,
|
|
|
|
|
|
interested_classes = body.interested_classes,
|
|
|
|
|
|
cls_conf = body.cls_conf if body.cls_conf is not None else CLS_CONF,
|
|
|
|
|
|
)
|
|
|
|
|
|
store_job(job)
|
|
|
|
|
|
|
|
|
|
|
|
# Submit to thread pool — does not block the event loop
|
|
|
|
|
|
EXECUTOR.submit(run_pipeline, job)
|
|
|
|
|
|
|
|
|
|
|
|
return SubmitResponse(job_id=job_id)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ── GET /jobs/{job_id} ────────────────────────────────────────────────────────
|
|
|
|
|
|
|
|
|
|
|
|
@app.get("/jobs/{job_id}", response_model=PollResponse)
|
|
|
|
|
|
async def poll_job(job_id: str):
|
|
|
|
|
|
"""
|
|
|
|
|
|
Poll the status of a submitted job.
|
|
|
|
|
|
|
|
|
|
|
|
Call every ~3 seconds until status is 'done' or 'failed'.
|
|
|
|
|
|
"""
|
|
|
|
|
|
job = get_job(job_id)
|
|
|
|
|
|
return PollResponse(
|
|
|
|
|
|
status = job.status,
|
|
|
|
|
|
progress = job.progress,
|
|
|
|
|
|
results = job.results if job.status == Status.done else None,
|
|
|
|
|
|
error = job.error,
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ── GET /health ───────────────────────────────────────────────────────────────
|
|
|
|
|
|
|
|
|
|
|
|
@app.get("/health")
|
|
|
|
|
|
async def health():
|
|
|
|
|
|
return {
|
|
|
|
|
|
"status": "ok",
|
|
|
|
|
|
"device": str(DEVICE),
|
|
|
|
|
|
"classes": CLASSIFIER.class_names,
|
|
|
|
|
|
"active_jobs": sum(1 for j in JOB_STORE.values()
|
|
|
|
|
|
if j.status in (Status.queued, Status.processing)),
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ── GET /classes ──────────────────────────────────────────────────────────────
|
|
|
|
|
|
|
|
|
|
|
|
@app.get("/classes")
|
|
|
|
|
|
async def list_classes():
|
|
|
|
|
|
"""Return all class names the classifier knows about."""
|
|
|
|
|
|
return {"classes": CLASSIFIER.class_names}
|