ThejasRao/openenv-content-moderation
0
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 