CoolFace
Apppublic

ashucode/metaxhuggingfacehackathon

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
app .py178 linesDownload Raw Back to root
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")