import os from fastapi import FastAPI, HTTPException from pydantic import BaseModel, Field from src.predictor import ArticleTopicPredictor MODEL_DIR = os.getenv("MODEL_DIR", "artifacts/article_topic_model") app = FastAPI(title="Article Topic Classifier API", version="1.0.0") _predictor: ArticleTopicPredictor | None = None class PredictRequest(BaseModel): title: str = Field(default="", max_length=600) abstract: str = Field(default="", max_length=8000) top95_threshold: float = Field(default=0.95, ge=0.5, le=0.999) @app.on_event("startup") def startup_event() -> None: global _predictor _predictor = ArticleTopicPredictor(model_dir=MODEL_DIR) @app.get("/health") def health() -> dict: return {"status": "ok", "model_dir": MODEL_DIR} @app.post("/predict") def predict(request: PredictRequest) -> dict: if not request.title.strip() and not request.abstract.strip(): raise HTTPException(status_code=400, detail="Provide at least title or abstract.") if _predictor is None: raise HTTPException(status_code=503, detail="Model is not loaded yet.") try: return _predictor.predict( title=request.title, abstract=request.abstract, top95_threshold=request.top95_threshold, ) except Exception as exc: # noqa: BLE001 raise HTTPException(status_code=500, detail=str(exc)) from exc