Team Ai
Apppublic

Jayant2304/commitment-os

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
app.py126 linesDownload Raw Back to server
1"""FastAPI composition root — wires environment, MCP, and custom endpoints."""2 3from __future__ import annotations4 5import os6from threading import Lock7 8from openenv.core.env_server import create_fastapi_app9from fastapi import Query10from pydantic import BaseModel11 12from constants import PROJECT_DESCRIPTION, VERSION13from models import CommitmentAction, CommitmentObservation, CommitmentState14from server.environment import CommitmentEnvironment15from server.mcp import router as mcp_router16from server.tasks import get_scenario_ids_grouped17 18_DEFAULT_SESSION_ID = "default"19_env_store: dict[str, CommitmentEnvironment] = {20    _DEFAULT_SESSION_ID: CommitmentEnvironment(),21}22_env_store_lock = Lock()23 24 25def _get_env(session_id: str) -> CommitmentEnvironment:26    """Return a per-session environment instance.27 28    This avoids cross-user state bleed from a single shared mutable environment.29    Clients can pass ``episode_id`` query param to isolate sessions.30    """31    with _env_store_lock:32        env = _env_store.get(session_id)33        if env is None:34            env = CommitmentEnvironment()35            _env_store[session_id] = env36        return env37 38 39class StepPayload(BaseModel):40    action: CommitmentAction41 42app = create_fastapi_app(43    env=lambda: _get_env(_DEFAULT_SESSION_ID),44    action_cls=CommitmentAction,45    observation_cls=CommitmentObservation,46)47 48app.title = "CommitmentOS"49app.description = PROJECT_DESCRIPTION50app.version = VERSION51 52app.routes[:] = [53    r for r in app.routes54    if not (hasattr(r, "path") and r.path in ("/state", "/mcp", "/reset", "/step"))55]56 57 58@app.post("/reset")59def reset_episode(60    task_id: str | None = Query(default=None),61    difficulty: str | None = Query(default=None),62    seed: int | None = Query(default=None),63    episode_id: str | None = Query(default=None),64) -> dict[str, object]:65    """Reset endpoint with explicit query-param support.66 67    The default OpenEnv route did not reliably propagate ``task_id`` from68    query params in this deployment setup, which made scenario selection69    non-deterministic for demos/evaluations.70    """71    session_id = episode_id or _DEFAULT_SESSION_ID72    env = _get_env(session_id)73    obs = env.reset(74        seed=seed,75        episode_id=session_id,76        task_id=task_id,77        difficulty=difficulty,78    )79    return {80        "observation": obs.model_dump(),81        "reward": float(obs.reward),82        "done": bool(obs.done),83        "episode_id": session_id,84    }85 86 87@app.post("/step")88def step_episode(89    payload: StepPayload,90    episode_id: str | None = Query(default=None),91) -> dict[str, object]:92    session_id = episode_id or _DEFAULT_SESSION_ID93    env = _get_env(session_id)94    obs = env.step(payload.action)95    return {96        "observation": obs.model_dump(),97        "reward": float(obs.reward),98        "done": bool(obs.done),99        "episode_id": session_id,100    }101 102 103@app.get("/state", response_model=CommitmentState)104def get_state(episode_id: str | None = Query(default=None)) -> CommitmentState:105    session_id = episode_id or _DEFAULT_SESSION_ID106    return _get_env(session_id).state107 108 109@app.get("/tasks")110def list_tasks() -> dict[str, list[str]]:111    return get_scenario_ids_grouped()112 113 114app.include_router(mcp_router)115 116 117def main() -> None:118    import uvicorn119 120    port = int(os.environ.get("PORT", 7860))121    uvicorn.run(app, host="0.0.0.0", port=port)122 123 124if __name__ == "__main__":125    main()126