David Baba
backend
0437edf
Raw
History Blame Contribute Delete
5.21 kB
"""Creator Vision API.
Phase 1: POST /analyze (BackgroundTasks + in-memory job store) and
GET /status/{job_id}. The analysis is currently stubbed (see analysis.py);
Phase 2 wires real SAM 3. See DesignDoc.md for the result schema.
"""
import json
import os
import shutil
import uuid
from pathlib import Path
from dotenv import load_dotenv
from fastapi import BackgroundTasks, FastAPI, File, Form, HTTPException, UploadFile
from fastapi.middleware.cors import CORSMiddleware
from .analysis import (
BuildOpts,
ProductInput,
owlv2_needs_images,
requires_name,
requires_reference,
run_analysis,
)
from .jobs import store
from .schemas import DetectionMode, JobResponse, StatusResponse
# Load backend/.env (and fall back to repo-root .env).
load_dotenv()
UPLOAD_DIR = Path(__file__).resolve().parent.parent / "uploads"
UPLOAD_DIR.mkdir(exist_ok=True)
app = FastAPI(title="Creator Vision API", version="0.1.0")
# Allowed browser origins. Defaults to the local Next.js dev server; in prod set
# ALLOWED_ORIGINS to your deployed frontend URL(s), comma-separated.
_origins = os.getenv("ALLOWED_ORIGINS", "http://localhost:3000")
app.add_middleware(
CORSMiddleware,
allow_origins=[o.strip() for o in _origins.split(",") if o.strip()],
allow_methods=["*"],
allow_headers=["*"],
)
@app.get("/health")
def health() -> dict:
"""Liveness + config sanity check (does the SAM3 key exist?)."""
return {
"status": "ok",
"fal_key_configured": bool(os.getenv("FAL_KEY")),
}
def _save_upload(upload: UploadFile, dest_dir: Path) -> str:
"""Persist an UploadFile under dest_dir with a collision-proof name."""
suffix = Path(upload.filename or "").suffix
dest = dest_dir / f"{uuid.uuid4().hex}{suffix}"
with dest.open("wb") as f:
shutil.copyfileobj(upload.file, f)
return str(dest)
@app.post("/analyze", response_model=JobResponse)
def analyze(
background_tasks: BackgroundTasks,
video: UploadFile = File(...),
# `products` is JSON: [{"name": str, "image_count": int}, ...]. `exemplars` is a
# flat file list in product order, sliced back per product by `image_count`.
products: str = Form(...),
exemplars: list[UploadFile] = File(default=[]),
caption: str = Form(default=""),
mention_keywords: str = Form(default=""),
mode: DetectionMode = Form(default=DetectionMode.sam3_text),
split_on_cut: bool = Form(default=False),
dino_variant: str = Form(default="v2"),
owl_ref_type: str = Form(default="text"), # owlv2: text | image | both
owl_dino: str = Form(default="none"), # owlv2 DINO-on-top: none | v2 | v3
) -> JobResponse:
"""Accept the uploads, kick off background analysis, return a job id."""
if not video.filename:
raise HTTPException(status_code=400, detail="A video file is required.")
try:
product_meta = json.loads(products)
assert isinstance(product_meta, list) and product_meta
except (json.JSONDecodeError, AssertionError):
raise HTTPException(status_code=400, detail="`products` must be a non-empty JSON list.")
job = store.create()
job_dir = UPLOAD_DIR / job.id
job_dir.mkdir(parents=True, exist_ok=True)
video_path = _save_upload(video, job_dir)
exemplar_files = [ex for ex in exemplars if ex.filename]
# Slice the flat exemplar list back into per-product groups by image_count.
product_inputs: list[ProductInput] = []
cursor = 0
for p in product_meta:
name = str(p.get("name", "")).strip()
count = int(p.get("image_count", 0))
group = exemplar_files[cursor : cursor + count]
cursor += count
paths = [_save_upload(ex, job_dir) for ex in group]
product_inputs.append(ProductInput(name=name, exemplar_paths=paths))
# Validation is driven by the mode registry (single source of truth).
opts = BuildOpts(dino_variant=dino_variant, owl_ref_type=owl_ref_type, owl_dino=owl_dino)
if requires_name(mode) and any(not p.name for p in product_inputs):
raise HTTPException(status_code=400, detail="Every product needs a name in this mode.")
if (requires_reference(mode) or owlv2_needs_images(mode, opts)) and any(
not p.exemplar_paths for p in product_inputs
):
raise HTTPException(
status_code=400,
detail=f"{mode.value} mode requires a reference image for every product.",
)
variant = dino_variant if dino_variant in ("v2", "v3") else "v2"
owl_dino_v = owl_dino if owl_dino in ("none", "v2", "v3") else "none"
owl_ref = owl_ref_type if owl_ref_type in ("text", "image", "both") else "text"
background_tasks.add_task(
run_analysis, job.id, video_path, product_inputs, caption,
mention_keywords, mode, split_on_cut, variant, owl_ref, owl_dino_v,
)
return JobResponse(job_id=job.id)
@app.get("/status/{job_id}", response_model=StatusResponse)
def status(job_id: str) -> StatusResponse:
job = store.get(job_id)
if job is None:
raise HTTPException(status_code=404, detail="Unknown job_id.")
return StatusResponse(
job_id=job.id, status=job.status, result=job.result, error=job.error
)