from __future__ import annotations from pathlib import Path from typing import Any from fastapi import Body, FastAPI, HTTPException from fastapi.responses import PlainTextResponse from pydantic import BaseModel, Field from submission_common import add_project_to_path add_project_to_path() from env.email_triage_env import EmailTriageEnv from env.task_config import TASK_CONFIGS class ResetRequest(BaseModel): task_id: str = "email_resolution" seed: int = 42 session_id: str = "default" scenario_id: str | None = None class StepRequest(BaseModel): session_id: str = "default" action: dict[str, Any] = Field(default_factory=dict) app = FastAPI( title="EmailTriageEnv", description="Deterministic customer-support email benchmark for OpenEnv.", version="1.0.0", ) _sessions: dict[str, EmailTriageEnv] = {} _ROOT = Path(__file__).resolve().parent @app.get("/") def root() -> dict[str, Any]: return { "env": "EmailTriageEnv", "version": "1.0.0", "tasks": ["email_classification", "email_triage", "email_resolution"], "endpoints": ["/health", "/reset", "/step", "/state", "/tasks", "/openenv.yaml"], } @app.get("/health") def health() -> dict[str, str]: return {"status": "ok"} @app.get("/tasks") def tasks() -> list[dict[str, Any]]: return [ { "id": cfg.id, "description": cfg.description, "difficulty": cfg.difficulty, "max_steps": cfg.max_steps, "allowed_actions": [action.value for action in cfg.allowed_actions], } for cfg in TASK_CONFIGS.values() ] @app.post("/reset") def reset(request: ResetRequest | None = Body(default=None)) -> dict[str, Any]: request = request or ResetRequest() env = EmailTriageEnv(task_id=request.task_id, seed=request.seed) _sessions[request.session_id] = env observation = env.reset(scenario_id=request.scenario_id) return observation.model_dump(mode="json") @app.post("/step") def step(request: StepRequest) -> dict[str, Any]: env = _sessions.get(request.session_id) if env is None: raise HTTPException(status_code=404, detail="Session not found. Call /reset first.") observation, reward, done, info = env.step(request.action) return { "observation": observation.model_dump(mode="json"), "reward": reward, "done": done, "info": info, } @app.get("/state") def state(session_id: str = "default") -> dict[str, Any]: env = _sessions.get(session_id) if env is None: raise HTTPException(status_code=404, detail="Session not found. Call /reset first.") return env.state() @app.get("/openenv.yaml", response_class=PlainTextResponse) def get_openenv_yaml() -> str: return (_ROOT / "openenv.yaml").read_text(encoding="utf-8") if __name__ == "__main__": import uvicorn uvicorn.run(app, host="0.0.0.0", port=7860)