RealBhupesh
Fix /reset to accept empty POST body
60e6d0d
Raw
History Blame Contribute Delete
2.94 kB
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)