Jayant2304/commitment-os
0
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 