| 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) |
|
|