CoolFace
Apppublic

ThejasRao/openenv-content-moderation

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
app.py305 linesDownload Raw Back to api
1"""2OpenENV Moderation Environment — FastAPI application.3 4Standard OpenEnv endpoints:5  WS   /ws         — persistent WebSocket session (primary client interface)6  GET  /health     — liveness check7  POST /reset      — start a new episode8  POST /step       — take an action9  GET  /state      — current observation / state10  GET  /docs       — OpenAPI documentation (auto-generated)11 12Custom endpoints:13  GET  /tasks      — available tasks14  GET  /grader     — final episode score15  GET  /baseline   — run rule-based baseline agent and return its score16  POST /agent/run  — run selected LLM agent on a full episode17"""18from __future__ import annotations19 20import json21import logging22 23from dotenv import load_dotenv24load_dotenv()  # loads .env from project root before anything else25 26from fastapi import FastAPI, HTTPException, Body, WebSocket, WebSocketDisconnect27from fastapi.middleware.cors import CORSMiddleware28 29from openenv.core.env_server.types import (30    HealthResponse,31    HealthStatus,32    ResetRequest as OEResetRequest,33    ResetResponse,34    StepRequest,35    StepResponse,36    WSObservationResponse,37    WSStateResponse,38    WSErrorResponse,39    WSErrorCode,40)41 42from data.tasks import TASKS43from env.grader import Grader44from env.state_manager import StateManager45from models.schemas import (46    Action,47    BaselineResult,48    EpisodeScore,49    ResetRequest,50    TaskConfig,51)52 53logging.basicConfig(level=logging.INFO)54logger = logging.getLogger(__name__)55 56app = FastAPI(57    title="OpenENV — Content Moderation Environment",58    description=(59        "A multi-step RL environment for AI content moderation agents. "60        "Agents receive partial observations and must investigate context, "61        "classify violations, and make final moderation decisions."62    ),63    version="1.0.0",64)65 66app.add_middleware(67    CORSMiddleware,68    allow_origins=["*"],   # Open for HF Spaces + local dev69    allow_methods=["*"],70    allow_headers=["*"],71)72 73# Single shared state manager (single-threaded MVP)74_state_manager = StateManager()75_grader = Grader()76 77 78# ---------------------------------------------------------------------------79# Endpoints80# ---------------------------------------------------------------------------81 82@app.get("/health", response_model=HealthResponse)83def health() -> HealthResponse:84    return HealthResponse(status=HealthStatus.HEALTHY)85 86 87@app.get("/tasks")88def list_tasks() -> dict[str, TaskConfig]:89    return TASKS90 91 92@app.post("/reset", response_model=ResetResponse)93def reset(request: OEResetRequest | None = Body(default=None)) -> ResetResponse:94    # task_id passed as extra field; fall back to episode_id or default95    extra = (request.model_extra or {}) if request else {}96    task_id = extra.get("task_id") or (request.episode_id if request else None) or "easy_harassment"97    seed = (request.seed if request else None) or 4298 99    if task_id not in TASKS:100        raise HTTPException(101            status_code=400,102            detail=f"Unknown task_id '{task_id}'. Available: {list(TASKS.keys())}",103        )104 105    task = TASKS[task_id]106    task = task.model_copy(update={"seed": seed})107 108    obs = _state_manager.reset(task)109    return ResetResponse(observation=obs.model_dump(), reward=None, done=obs.done)110 111 112@app.post("/step", response_model=StepResponse)113def step(request: StepRequest) -> StepResponse:114    if not _state_manager.has_active_episode():115        raise HTTPException(status_code=400, detail="No active episode. Call /reset first.")116 117    try:118        action = Action(**request.action)119    except Exception as exc:120        raise HTTPException(status_code=422, detail=str(exc))121 122    try:123        result = _state_manager.step(action)124    except ValueError as exc:125        raise HTTPException(status_code=400, detail=str(exc))126 127    logger.info(128        "Step %d: action=%s reward=%.3f done=%s",129        result.observation.step,130        action.action_type.value,131        result.reward,132        result.done,133    )134    return StepResponse(135        observation=result.observation.model_dump(),136        reward=result.reward,137        done=result.done,138    )139 140 141@app.get("/state")142def get_state() -> dict:143    if not _state_manager.has_active_episode():144        raise HTTPException(status_code=400, detail="No active episode. Call /reset first.")145    return _state_manager.get_state().model_dump()146 147 148@app.websocket("/ws")149async def websocket_endpoint(websocket: WebSocket) -> None:150    await websocket.accept()151    try:152        while True:153            try:154                raw = await websocket.receive_text()155                data = json.loads(raw)156            except json.JSONDecodeError:157                await websocket.send_text(158                    WSErrorResponse(data={"message": "Invalid JSON", "code": WSErrorCode.INVALID_JSON}).model_dump_json()159                )160                continue161 162            msg_type = data.get("type")163 164            if msg_type == "reset":165                reset_data = data.get("data", {})166                task_id = reset_data.get("task_id") or reset_data.get("episode_id") or "easy_harassment"167                seed = reset_data.get("seed") or 42168 169                if task_id not in TASKS:170                    await websocket.send_text(171                        WSErrorResponse(data={"message": f"Unknown task_id '{task_id}'", "code": WSErrorCode.VALIDATION_ERROR}).model_dump_json()172                    )173                    continue174 175                task = TASKS[task_id].model_copy(update={"seed": seed})176                obs = _state_manager.reset(task)177                await websocket.send_text(178                    WSObservationResponse(data={"observation": obs.model_dump(), "reward": None, "done": obs.done}).model_dump_json()179                )180 181            elif msg_type == "step":182                if not _state_manager.has_active_episode():183                    await websocket.send_text(184                        WSErrorResponse(data={"message": "No active episode. Send reset first.", "code": WSErrorCode.SESSION_ERROR}).model_dump_json()185                    )186                    continue187 188                action_data = data.get("data", {})189                try:190                    action = Action(**action_data)191                except Exception as exc:192                    await websocket.send_text(193                        WSErrorResponse(data={"message": str(exc), "code": WSErrorCode.VALIDATION_ERROR}).model_dump_json()194                    )195                    continue196 197                try:198                    result = _state_manager.step(action)199                except ValueError as exc:200                    await websocket.send_text(201                        WSErrorResponse(data={"message": str(exc), "code": WSErrorCode.EXECUTION_ERROR}).model_dump_json()202                    )203                    continue204 205                await websocket.send_text(206                    WSObservationResponse(data={"observation": result.observation.model_dump(), "reward": result.reward, "done": result.done}).model_dump_json()207                )208 209            elif msg_type == "state":210                if not _state_manager.has_active_episode():211                    await websocket.send_text(212                        WSErrorResponse(data={"message": "No active episode.", "code": WSErrorCode.SESSION_ERROR}).model_dump_json()213                    )214                    continue215                obs = _state_manager.get_state()216                await websocket.send_text(217                    WSStateResponse(data=obs.model_dump()).model_dump_json()218                )219 220            elif msg_type == "close":221                break222 223            else:224                await websocket.send_text(225                    WSErrorResponse(data={"message": f"Unknown message type: {msg_type!r}", "code": WSErrorCode.UNKNOWN_TYPE}).model_dump_json()226                )227 228    except WebSocketDisconnect:229        pass230 231 232@app.get("/grader", response_model=EpisodeScore)233def grade() -> EpisodeScore:234    if not _state_manager.has_active_episode():235        raise HTTPException(status_code=400, detail="No active episode. Call /reset first.")236 237    episode = _state_manager.get_episode_state()238    if not episode.observation.done:239        raise HTTPException(240            status_code=400,241            detail="Episode is not finished yet. Complete the episode before grading.",242        )243 244    score = _grader.score(episode)245    logger.info("Graded episode: total=%.4f", score.total)246    return score247 248 249@app.get("/baseline", response_model=BaselineResult)250def baseline(task_id: str = "easy_harassment", seed: int | None = None) -> BaselineResult:251    """Run the built-in rule-based baseline agent and return its score."""252    from baseline.agent import BaselineAgent253 254    if task_id not in TASKS:255        raise HTTPException(256            status_code=400,257            detail=f"Unknown task_id '{task_id}'. Available: {list(TASKS.keys())}",258        )259 260    task = TASKS[task_id]261    if seed is not None:262        task = task.model_copy(update={"seed": seed})263 264    agent = BaselineAgent(state_manager=_state_manager, grader=_grader)265    result = agent.run(task)266    return result267 268 269@app.post("/agent/run", response_model=BaselineResult)270def agent_run(request: ResetRequest) -> BaselineResult:271    """272    Run the selected LLM agent (OpenAI or Gemini) on a full episode and return the graded result.273 274    Requires OPENAI_API_KEY, or GOOGLE_API_KEY/GEMINI_API_KEY depending on LLM_PROVIDER.275    """276    import os277    from agent.openai_agent import OpenAIAgent278    from agent.gemini_agent import GeminiAgent279    280    provider = os.getenv("LLM_PROVIDER", "openai").lower()281 282    if request.task_id not in TASKS:283        raise HTTPException(284            status_code=400,285            detail=f"Unknown task_id '{request.task_id}'. Available: {list(TASKS.keys())}",286        )287 288    task = TASKS[request.task_id]289    if request.seed is not None:290        task = task.model_copy(update={"seed": request.seed})291 292    try:293        if provider == "gemini":294            agent = GeminiAgent(state_manager=_state_manager, grader=_grader)295        else:296            agent = OpenAIAgent(state_manager=_state_manager, grader=_grader)297    except EnvironmentError as exc:298        raise HTTPException(status_code=500, detail=str(exc))299 300    result = agent.run(task)301    logger.info(302        "%s agent finished: task=%s total=%.4f", provider.capitalize(), task.task_id, result.score.total303    )304    return result305