""" FastAPI app with all model endpoints. """ import os import uuid import cv2 import numpy as np from pathlib import Path from contextlib import asynccontextmanager from fastapi import FastAPI, UploadFile, File, Form, HTTPException from fastapi.responses import Response, JSONResponse from fastapi.middleware.cors import CORSMiddleware from fastapi.staticfiles import StaticFiles from app.config import PipelineConfig, ModelPaths, Quality, ContentType, Difficulty, LineArtStyle from app.pipeline import ColorByNumberPipeline pipeline: ColorByNumberPipeline | None = None @asynccontextmanager async def lifespan(app: FastAPI): global pipeline paths = ModelPaths() # Create directories if they don't exist for d in ["outputs", "outputs/debug", "models"]: Path(d).mkdir(exist_ok=True) # Note: We don't crash if models are missing, the pipeline gracefully degrades # but we print their status. pipeline = ColorByNumberPipeline( model_paths=paths, generation_api_key=os.getenv("GENERATION_API_KEY", ""), generation_provider=os.getenv("GENERATION_PROVIDER", "replicate"), ) print("✓ Pro Pipeline loaded") print(f" Model Status: {pipeline.registry.status()}") yield app = FastAPI(title="Color-by-Number Pro API", version="3.0.0", lifespan=lifespan) app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) # Mount outputs for static access app.mount("/static", StaticFiles(directory="outputs"), name="static") @app.get("/health") async def health(): return {"status": "ok", "model_loaded": pipeline is not None} @app.get("/models/status") async def model_status(): if pipeline is None: raise HTTPException(503, "Not initialized") return pipeline.registry.status() @app.post("/generate") async def generate( image: UploadFile = File(...), difficulty: str = Form(default="auto"), quality: str = Form(default="balanced"), content_type: str = Form(default="auto"), line_style: str = Form(default="auto"), k_colors: int = Form(default=0, ge=0, le=30), remove_background: bool = Form(default=False), use_sam: bool = Form(default=True), use_depth: bool = Form(default=True), save_debug: bool = Form(default=False), max_dimension: int = Form(default=1024, ge=256, le=2048), response_format: str = Form(default="json"), ): if pipeline is None: raise HTTPException(503, "Not initialized") ct = image.content_type or "" if not ct.startswith("image/"): raise HTTPException(400, f"Expected image, got {ct}") raw = await image.read() arr = np.frombuffer(raw, dtype=np.uint8) img = cv2.imdecode(arr, cv2.IMREAD_COLOR) if img is None: raise HTTPException(400, "Could not decode image") job_id = uuid.uuid4().hex[:12] config = PipelineConfig( content_type=ContentType(content_type), difficulty=Difficulty(difficulty), quality=Quality(quality), line_style=LineArtStyle(line_style), k_colors=k_colors if k_colors > 0 else None, remove_background=remove_background, use_sam=use_sam, use_depth=use_depth, save_debug=save_debug, debug_dir=f"outputs/debug/{job_id}", max_dimension=max_dimension, ) result = pipeline.process( image=img, config=config, output_path=f"outputs/{job_id}.svg", ) response = { "job_id": job_id, "svg_url": f"/static/{job_id}.svg", "outline_url": f"/static/{job_id}_outline.svg", "palette": result["palette"], "num_regions": result["num_regions"], "content_type": result["content_type"], "models_used": result["models_used"], "k_colors_used": result["k_colors_used"], "density_used": result["density_used"], } if "suggested_order" in result: response["suggested_order"] = result["suggested_order"] if "difficulty_layers" in result: response["difficulty_layers"] = result["difficulty_layers"] if "regions" in result: response["regions"] = result["regions"] if save_debug: response["debug_dir"] = f"/static/debug/{job_id}/" if response_format == "svg": return Response(content=result["svg_string"], media_type="image/svg+xml") elif response_format == "outline": return Response(content=result["svg_outline_string"], media_type="image/svg+xml") response["svg_string"] = result["svg_string"] return JSONResponse(response) @app.post("/generate-from-prompt") async def generate_from_prompt( prompt: str = Form(..., description="What to generate"), style: str = Form(default="coloring_book"), difficulty: str = Form(default="medium"), quality: str = Form(default="balanced"), response_format: str = Form(default="json"), ): if pipeline is None: raise HTTPException(503, "Not initialized") if not pipeline.generator.available: raise HTTPException( 501, "Generation API not configured. Set GENERATION_API_KEY env var.", ) job_id = uuid.uuid4().hex[:12] config = PipelineConfig( difficulty=Difficulty(difficulty), quality=Quality(quality), content_type=ContentType.ILLUSTRATION, ) result = await pipeline.generate_from_prompt( prompt=prompt, style=style, config=config, output_path=f"outputs/{job_id}.svg", ) if result is None: raise HTTPException(500, "Image generation failed") response = { "job_id": job_id, "prompt": prompt, "style": style, "svg_url": f"/static/{job_id}.svg", "palette": result["palette"], "num_regions": result["num_regions"], "models_used": result["models_used"], } if response_format == "svg": return Response(content=result["svg_string"], media_type="image/svg+xml") response["svg_string"] = result["svg_string"] return JSONResponse(response)