ashucode/metaxhuggingfacehackathon
0
1 2from __future__ import annotations3 4import uuid5from typing import Any, Dict, List, Optional6 7from fastapi import FastAPI, HTTPException8from pydantic import BaseModel, Field9 10from supermart_env import (11 SupermarketEnv,12 PRODUCT_CATALOGUE,13 ALL_PRODUCTS,14 ACTION_NAMES,15 N_ACTIONS,16 LEVEL_CONFIG,17)18 19app = FastAPI(20 title="SupermarketNav OpenEnv",21 description="Client-Server RL environment for LLM agent evaluation.",22 version="1.0.0",23)24 25 26_sessions: Dict[str, SupermarketEnv] = {}27 28 29 30 31 32class ResetRequest(BaseModel):33 level: str = Field("easy", description="easy | medium | hard")34 products: Optional[List[str]] = Field(35 None,36 description=(37 "Optional list of exact product names to use. "38 "Leave null to let the env pick randomly."39 ),40 )41 seed: Optional[int] = Field(None, description="RNG seed for reproducibility.")42 43 44class StepRequest(BaseModel):45 action: int = Field(46 ...,47 ge=0,48 le=N_ACTIONS - 1,49 description=f"Integer action 0-{N_ACTIONS - 1}. Names: {ACTION_NAMES}",50 )51 52 53class StateResponse(BaseModel):54 session_id: str55 task_status: str56 phase: int57 phase_label: str58 agent_pos: List[int]59 targets: List[str]60 inventory: List[str]61 closest_target: Optional[str]62 steps_taken: int63 steps_remaining: int64 total_reward: float65 normalised_score: float66 action_mask: List[bool]67 level: str68 closest_rule: bool69 event: str70 71 72def _env_state(session_id: str, env: SupermarketEnv, event: str = "") -> Dict[str, Any]:73 info = env._build_info(event, env._total_reward)74 return {75 "session_id": session_id,76 "task_status": info["task_status"],77 "phase": info["phase"],78 "phase_label": info["phase_label"],79 "agent_pos": info["agent_pos"],80 "targets": info["targets"],81 "inventory": info["inventory"],82 "closest_target": info["closest_target"],83 "steps_taken": info["steps_taken"],84 "steps_remaining": info["steps_remaining"],85 "total_reward": info["total_reward"],86 "normalised_score": info["normalised_score"],87 "action_mask": info["action_mask"],88 "level": info["level"],89 "closest_rule": info["closest_rule"],90 "event": event,91 }92 93 94def _get_env(session_id: str) -> SupermarketEnv:95 env = _sessions.get(session_id)96 if env is None:97 raise HTTPException(status_code=404, detail=f"Session '{session_id}' not found. Call /reset first.")98 return env99 100 101 102 103@app.get("/healthz", tags=["meta"])104def healthz():105 return {"status": "ok", "active_sessions": len(_sessions)}106 107 108@app.get("/catalogue", tags=["meta"])109def catalogue():110 111 return {"catalogue": PRODUCT_CATALOGUE, "all_products": list(ALL_PRODUCTS.keys())}112 113 114@app.post("/reset", tags=["env"])115def reset(body: Optional[ResetRequest] = None):116 # If the automated grader sends an empty request, use defaults117 if body is None:118 body = ResetRequest()119 120 if body.level not in LEVEL_CONFIG:121 raise HTTPException(122 status_code=422,123 detail=f"level must be one of {list(LEVEL_CONFIG)}",124 )125 126 session_id = str(uuid.uuid4())127 env = SupermarketEnv(level=body.level, products=body.products)128 129 try:130 _, info = env.reset(seed=body.seed)131 except ValueError as exc:132 raise HTTPException(status_code=422, detail=str(exc))133 134 _sessions[session_id] = env135 136 return {137 **_env_state(session_id, env, "Episode started"),138 "action_names": ACTION_NAMES,139 "level_config": {140 k: v for k, v in LEVEL_CONFIG[body.level].items()141 if k not in ("q_lr", "q_gamma", "eps_start", "eps_end", "eps_decay", "train_episodes")142 },143 }144 145 146@app.post("/step/{session_id}", tags=["env"])147def step(session_id: str, body: StepRequest):148 149 env = _get_env(session_id)150 151 if env.task_status != "In-Progress":152 raise HTTPException(153 status_code=409,154 detail=f"Episode already ended with status '{env.task_status}'. Call /reset to start a new one.",155 )156 157 obs, raw_reward, terminated, truncated, info = env.step(body.action)158 159 response = {160 **_env_state(session_id, env, info["event"]),161 "step_reward": info["step_reward"],162 "terminated": terminated,163 "truncated": truncated,164 "done": terminated or truncated,165 }166 167 168 if terminated or truncated:169 _sessions.pop(session_id, None)170 171 return response172 173 174@app.get("/state/{session_id}", tags=["env"])175def state(session_id: str):176 177 env = _get_env(session_id)178 return _env_state(session_id, env, "state query")