3v324v23's picture
fix: add uv.lock, server entry point and openenv-core dependency for openenv validate
ca25a0a
Raw
History Blame Contribute Delete
5.32 kB
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from fastapi.openapi.docs import get_swagger_ui_html
from pydantic import BaseModel, model_validator
from typing import Optional
from environment import CustomerSupportEnv, STEP_ORDER
from graders.base_grader import BaseGrader, HardTaskGrader
from tasks.easy_task import EASY_TASK
from tasks.medium_task import MEDIUM_TASK
from tasks.hard_task import HARD_TASK
app = FastAPI(title="Customer Support AI Environment")
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_methods=["*"],
allow_headers=["*"],
)
_env: Optional[CustomerSupportEnv] = None
TASK_MAP = {
"easy": EASY_TASK,
"medium": MEDIUM_TASK,
"hard": HARD_TASK,
}
class ResetRequest(BaseModel):
task: str = "easy"
class StepRequest(BaseModel):
action: Optional[str] = None # accepts {"action": "..."}
response: Optional[str] = None # accepts {"response": "..."}
@model_validator(mode="after")
def resolve_action(self):
self.action = self.action or self.response
if not self.action:
raise ValueError("Provide either 'action' or 'response' field with the agent reply.")
return self
def _build_observation():
ep = _env.episode
task = _env.task
step_index = _env._step_index
current_step_name = (
STEP_ORDER[step_index].value if step_index < len(STEP_ORDER) else "done"
)
return {
"task_id": task.task_id,
"difficulty": task.difficulty.value,
"customer_emotion": task.customer_emotion,
"customer_message": task.customer_message,
"scenario_context": task.scenario_context,
"current_step": current_step_name,
"step_number": step_index + 1,
"total_steps": len(STEP_ORDER),
"episode_status": ep.status.value,
"total_reward": round(ep.total_reward, 3),
"wrong_step_count": ep.wrong_step_count,
"steps_completed": [
{
"step": s.step.value,
"correct": s.correct,
"detected_action": s.detected_action,
"reward": round(s.reward, 3),
"penalty": round(s.penalty, 3),
"penalty_reasons": s.penalty_reasons,
}
for s in ep.steps
],
}
@app.get("/")
def health_check():
return {"status": "ok", "message": "Customer Support AI Environment"}
@app.post("/reset")
def reset_env(body: ResetRequest = None):
global _env
task_key = (body.task if body else "easy").lower()
if task_key not in TASK_MAP:
raise HTTPException(
status_code=400,
detail=f"Unknown task '{task_key}'. Choose from: {list(TASK_MAP.keys())}"
)
task = TASK_MAP[task_key]
grader = HardTaskGrader() if task_key == "hard" else BaseGrader()
_env = CustomerSupportEnv(task=task, grader=grader)
_env.reset()
return {"observation": _build_observation()}
@app.post("/step")
def step_env(body: StepRequest):
global _env
if _env is None:
raise HTTPException(status_code=400, detail="Environment not initialized. Call /reset first.")
if _env.episode.status.value != "running":
raise HTTPException(status_code=400, detail=f"Episode already ended: {_env.episode.status.value}")
result, done = _env.step(body.action)
return {
"observation": _build_observation(),
"reward": round(result.reward, 3),
"done": done,
"info": {
"step": result.step.value,
"correct": result.correct,
"detected_action": result.detected_action,
"expected_action": result.expected_action,
"penalty": round(result.penalty, 3),
"penalty_reasons": result.penalty_reasons,
"fail_triggered": result.fail_triggered,
"fail_reason": result.fail_reason,
},
}
@app.get("/state")
def get_state():
if _env is None:
raise HTTPException(status_code=400, detail="Environment not initialized. Call /reset first.")
return _env.summary()
@app.get("/observation_space")
def observation_space():
return {
"type": "Dict",
"fields": {
"task_id": "str",
"difficulty": "str (easy|medium|hard)",
"customer_emotion": "str",
"customer_message": "str",
"scenario_context": "str",
"current_step": "str (empathy|collect_info|investigate|resolution|done)",
"step_number": "int",
"total_steps": "int",
"episode_status": "str (running|success|fail)",
"total_reward": "float",
"wrong_step_count": "int",
"steps_completed": "List[Dict]",
}
}
@app.get("/action_space")
def action_space():
return {
"type": "Text",
"description": "Agent free-text reply to customer",
"min_words": 6,
}
@app.get("/docs", include_in_schema=False)
def custom_docs():
return get_swagger_ui_html(openapi_url="/openapi.json", title="API Docs")
def main():
"""Entry point for 'server' script defined in pyproject.toml."""
import uvicorn
uvicorn.run("api:app", host="0.0.0.0", port=7860)
if __name__ == "__main__":
main()