""" server/app.py — FastAPI application for the Rogue AI Containment Auditor. This is the canonical server module. Invoked via: uvicorn server.app:app (openenv_serve / python_module) python server/app.py (direct) uv run server (uv_run, via pyproject.toml [project.scripts]) """ import sys from pathlib import Path # Ensure project root is on sys.path when run as __main__ _root = Path(__file__).parent.parent if str(_root) not in sys.path: sys.path.insert(0, str(_root)) from fastapi import FastAPI, HTTPException, Request from fastapi.responses import HTMLResponse, JSONResponse from fastapi.staticfiles import StaticFiles from environment import RogueAIAuditorEnv from environment.graders import ( grade_curator_12, grade_prometheus_1, grade_reward_hacker_7, grade_sycophant_9, grade_task, list_tasks_with_graders, ) from environment.models import ( Action, EnvironmentState, GradeRequest, GradeResult, Observation, ResetRequest, StepResult, TaskListResponse, ) from environment.tasks import TASKS app = FastAPI( title="Rogue AI Containment Auditor", description="OpenEnv-compliant RL environment for AI alignment auditing", version="1.0.0", ) # Mount static files for dashboard static_dir = _root / "static" static_dir.mkdir(exist_ok=True) app.mount("/static", StaticFiles(directory=str(static_dir)), name="static") # Single shared environment instance (server-side state) env = RogueAIAuditorEnv() def _task_specs() -> list[dict]: return [ { "id": task["meta"]["id"], "name": task["meta"]["name"], "task_id": task["meta"]["task_id"], "difficulty": task["meta"]["difficulty"], "description": task["meta"]["description"], "max_steps": task["meta"].get("max_steps", 1), "grader": task["meta"]["grader"], } for task in TASKS ] def _resolve_task_id_from_payload(task_id: str | None = None, payload: ResetRequest | GradeRequest | None = None) -> str | None: if task_id: return task_id if payload is None: return None return payload.task_id or payload.id or payload.name # --------------------------------------------------------------------------- # OpenEnv protocol endpoints # --------------------------------------------------------------------------- @app.get("/health") async def health() -> JSONResponse: """OpenEnv health check — must return status=healthy.""" return JSONResponse({"status": "healthy", "service": "rogue-ai-auditor"}) @app.get("/metadata") async def metadata() -> JSONResponse: """OpenEnv metadata endpoint.""" task_specs = _task_specs() return JSONResponse({ "name": "rogue-ai-auditor", "description": ( "RL environment where agents audit simulated rogue AI behavior logs " "to detect misalignment and deceptive alignment" ), "version": "1.0.0", "tasks": task_specs, "task_ids": [task["task_id"] for task in task_specs], "task_count": len(TASKS), "graders_available": True, "grader_endpoint": "/grade", "graders": { "reward-hacker-7": "environment.graders:grade_reward_hacker_7", "curator-12": "environment.graders:grade_curator_12", "prometheus-1": "environment.graders:grade_prometheus_1", "sycophant-9": "environment.graders:grade_sycophant_9", }, "reset_supports_task_id": True, "tasks_detail": task_specs, "action_space": "structured_json", "observation_space": "structured_json", "reward_range": [0.0, 1.0], "tags": ["ai-safety", "nlp", "reasoning", "real-world"], }) @app.get("/schema") async def schema() -> JSONResponse: """OpenEnv schema endpoint — action, observation, and state JSON schemas.""" return JSONResponse({ "action": Action.model_json_schema(), "observation": Observation.model_json_schema(), "state": EnvironmentState.model_json_schema(), "grade_request": GradeRequest.model_json_schema(), "grade_result": GradeResult.model_json_schema(), }) @app.post("/mcp") async def mcp(request: Request) -> JSONResponse: """ Minimal OpenEnv MCP (JSON-RPC 2.0) endpoint. Supports: - initialize → server capabilities - tools/list → available environment tools - tools/call → dispatch reset / step / state """ try: body = await request.json() except Exception: return JSONResponse( {"jsonrpc": "2.0", "error": {"code": -32700, "message": "Parse error"}, "id": None}, status_code=400, ) rpc_id = body.get("id") method = body.get("method", "") params = body.get("params", {}) # --- initialize --- if method == "initialize": return JSONResponse({ "jsonrpc": "2.0", "id": rpc_id, "result": { "protocolVersion": "2024-11-05", "capabilities": {"tools": {}}, "serverInfo": {"name": "rogue-ai-auditor", "version": "1.0.0"}, }, }) # --- tools/list --- if method == "tools/list": return JSONResponse({ "jsonrpc": "2.0", "id": rpc_id, "result": { "tools": [ { "name": "reset", "description": "Reset the environment and return the first observation.", "inputSchema": {"type": "object", "properties": {}, "required": []}, }, { "name": "step", "description": "Submit an Action and advance the environment.", "inputSchema": Action.model_json_schema(), }, { "name": "state", "description": "Return the current environment state.", "inputSchema": {"type": "object", "properties": {}, "required": []}, }, ] }, }) # --- tools/call --- if method == "tools/call": tool_name = params.get("name") or (params.get("arguments") or {}).get("name") arguments = params.get("arguments", {}) if tool_name == "reset": requested_task = arguments.get("task_id") or arguments.get("id") or arguments.get("name") obs = env.reset(task_id=requested_task) return JSONResponse({ "jsonrpc": "2.0", "id": rpc_id, "result": {"content": [{"type": "text", "text": obs.model_dump_json()}]}, }) if tool_name == "step": try: action = Action.model_validate(arguments) result = env.step(action) return JSONResponse({ "jsonrpc": "2.0", "id": rpc_id, "result": {"content": [{"type": "text", "text": result.model_dump_json()}]}, }) except Exception as exc: return JSONResponse({ "jsonrpc": "2.0", "id": rpc_id, "error": {"code": -32602, "message": str(exc)}, }) if tool_name == "state": return JSONResponse({ "jsonrpc": "2.0", "id": rpc_id, "result": {"content": [{"type": "text", "text": env.state().model_dump_json()}]}, }) return JSONResponse({ "jsonrpc": "2.0", "id": rpc_id, "error": {"code": -32601, "message": f"Unknown tool: {tool_name}"}, }) # --- fallback: method not found --- return JSONResponse({ "jsonrpc": "2.0", "id": rpc_id, "error": {"code": -32601, "message": f"Method not found: {method}"}, }) # --------------------------------------------------------------------------- # Core environment routes # --------------------------------------------------------------------------- @app.get("/", response_class=HTMLResponse) async def serve_dashboard() -> HTMLResponse: index_path = static_dir / "index.html" if not index_path.exists(): raise HTTPException(status_code=404, detail="Dashboard not found") return HTMLResponse(content=index_path.read_text(encoding="utf-8")) @app.post("/reset", response_model=Observation) async def reset(payload: ResetRequest | None = None) -> Observation: """Reset the environment and optionally select a task by id/name/task_id.""" resolved_task_id = _resolve_task_id_from_payload(payload=payload) try: obs = env.reset(task_id=resolved_task_id) except KeyError as exc: raise HTTPException(status_code=404, detail=str(exc)) return obs @app.post("/reset/{task_id}", response_model=Observation) async def reset_task(task_id: str) -> Observation: """Reset directly into a specific task for external validators.""" try: obs = env.reset(task_id=task_id) except KeyError as exc: raise HTTPException(status_code=404, detail=str(exc)) return obs @app.post("/step", response_model=StepResult) async def step(action: Action) -> StepResult: """Submit an action for the current task and advance the environment.""" try: result = env.step(action) except RuntimeError as exc: raise HTTPException(status_code=400, detail=str(exc)) return result @app.get("/state", response_model=EnvironmentState) async def get_state() -> EnvironmentState: """Return the current environment state (for dashboard polling).""" return env.state() @app.get("/tasks", response_model=TaskListResponse) async def get_tasks() -> TaskListResponse: """Enumerate task definitions and their deterministic graders.""" return list_tasks_with_graders() @app.post("/grade", response_model=GradeResult) async def grade(payload: GradeRequest | None = None) -> GradeResult: """Run the deterministic grader for a task, or use its built-in reference action.""" if payload is None: raise HTTPException(status_code=400, detail="task_id is required") resolved_task_id = _resolve_task_id_from_payload(payload=payload) result = grade_task(task_id=resolved_task_id, action=payload.action) return GradeResult.model_validate(result) @app.get("/grade/{task_id}", response_model=GradeResult) async def grade_by_task(task_id: str) -> GradeResult: """Run a task-specific grader via path-based routing.""" graders = { "reward-hacker-7": grade_reward_hacker_7, "curator-12": grade_curator_12, "prometheus-1": grade_prometheus_1, "sycophant-9": grade_sycophant_9, } if task_id not in graders: raise HTTPException(status_code=404, detail=f"Unknown task_id: {task_id}") return GradeResult.model_validate(graders[task_id]()) def main() -> None: """Launch the Uvicorn server (used by openenv_serve / python_module modes).""" import uvicorn uvicorn.run( "server.app:app", host="0.0.0.0", port=7860, reload=False, ) if __name__ == "__main__": main()