heatherlead's picture
Add moderation backend with LFS tracking for large files
26c8f44
Raw History Blame Contribute Delete
6.31 kB
"""
app.py
======
FastAPI application for comment moderation.
Endpoints:
POST /moderate — Binary moderation (flagged/safe)
POST /moderate/detail — Detailed per-category scores
GET /health — Health check
Usage:
uvicorn api.app:app --host 0.0.0.0 --port 8000
uvicorn api.app:app --host 0.0.0.0 --port 8000 --reload # dev mode
"""
import logging
import sys
from contextlib import asynccontextmanager
from pathlib import Path
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
# Add project root to path for imports
PROJECT_ROOT = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(PROJECT_ROOT))
from api.middleware import TimingMiddleware
from api.schemas import (
ErrorResponse,
HealthResponse,
ModerateDetailResponse,
ModerateRequest,
ModerateResponse,
)
# ---------------------------------------------------------------------------
# Setup
# ---------------------------------------------------------------------------
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s | %(levelname)s | %(name)s | %(message)s",
)
logger = logging.getLogger(__name__)
# Global moderator instance (loaded at startup)
moderator = None
# ---------------------------------------------------------------------------
# Application lifespan (startup/shutdown)
# ---------------------------------------------------------------------------
@asynccontextmanager
async def lifespan(app: FastAPI):
"""Load the moderation model on startup."""
global moderator
logger.info("🚀 Starting Comment Moderation API...")
try:
from inference.predictor import CommentModerator
moderator = CommentModerator()
logger.info("✅ Model loaded successfully!")
except Exception as e:
logger.error(f"❌ Failed to load model: {e}")
logger.error("The API will start but moderation endpoints will return errors.")
logger.error("Make sure you've trained the model and exported to ONNX:")
logger.error(" python training/train.py")
logger.error(" python inference/export_onnx.py")
yield
# Cleanup
logger.info("👋 Shutting down...")
moderator = None
# ---------------------------------------------------------------------------
# FastAPI app
# ---------------------------------------------------------------------------
app = FastAPI(
title="Comment Moderation API",
description=(
"AI-powered comment moderation system supporting English and Hinglish. "
"Uses fine-tuned XLM-RoBERTa with ONNX inference for <20ms latency."
),
version="1.0.0",
lifespan=lifespan,
docs_url="/docs",
redoc_url="/redoc",
)
# --- Middleware ---
app.add_middleware(TimingMiddleware)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"], # Restrict in production
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# ---------------------------------------------------------------------------
# Helper
# ---------------------------------------------------------------------------
def _check_moderator():
"""Raise 503 if model is not loaded."""
if moderator is None:
raise HTTPException(
status_code=503,
detail="Model not loaded. Train and export the model first.",
)
# ---------------------------------------------------------------------------
# Endpoints
# ---------------------------------------------------------------------------
@app.get(
"/health",
response_model=HealthResponse,
tags=["System"],
summary="Health check",
)
async def health():
"""Check if the service is running and the model is loaded."""
return HealthResponse(
status="healthy" if moderator is not None else "degraded",
model="xlm-roberta-base-moderation",
version="1.0.0",
)
@app.post(
"/moderate",
response_model=ModerateResponse,
tags=["Moderation"],
summary="Moderate a comment (binary)",
responses={
200: {
"description": "Moderation result",
"content": {
"application/json": {
"examples": {
"flagged": {
"summary": "Harmful comment",
"value": {"flagged": True},
},
"safe": {
"summary": "Safe comment",
"value": {"flagged": False},
},
}
}
},
},
503: {"model": ErrorResponse, "description": "Model not loaded"},
},
)
async def moderate(request: ModerateRequest):
"""
Classify a comment as harmful or safe.
Returns `{"flagged": true}` if the comment is harmful,
`{"flagged": false}` otherwise.
"""
_check_moderator()
result = moderator.moderate(request.comment)
return ModerateResponse(**result)
@app.post(
"/moderate/detail",
response_model=ModerateDetailResponse,
tags=["Moderation"],
summary="Moderate a comment (detailed)",
responses={
200: {"description": "Detailed moderation result"},
503: {"model": ErrorResponse, "description": "Model not loaded"},
},
)
async def moderate_detail(request: ModerateRequest):
"""
Classify a comment with detailed per-category probability scores.
Returns flagged status, confidence, top harmful category, and
individual category probabilities.
"""
_check_moderator()
result = moderator.predict(request.comment)
return ModerateDetailResponse(
flagged=result["flagged"],
confidence=result["confidence"],
top_category=result["top_category"],
categories=result["categories"],
threshold=moderator.threshold,
latency_ms=result["latency_ms"],
)
# ---------------------------------------------------------------------------
# Run directly
# ---------------------------------------------------------------------------
if __name__ == "__main__":
import uvicorn
uvicorn.run(
"api.app:app",
host="0.0.0.0",
port=8000,
reload=True,
log_level="info",
)