import uuid from fastapi import FastAPI, HTTPException import uvicorn from environment.actions import ContextCorruptionAction, EpisodeObservation from environment.env import ContextCorruptionEnv app = FastAPI(title="ContextCorruption-Env") # session_id -> env instance _sessions: dict[str, ContextCorruptionEnv] = {} _MAX_SESSIONS = 64 @app.post("/reset", response_model=EpisodeObservation) def reset(session_id: str | None = None): if session_id is None: if len(_sessions) >= _MAX_SESSIONS: raise HTTPException(status_code=503, detail="Max concurrent sessions reached") session_id = str(uuid.uuid4()) if session_id not in _sessions: _sessions[session_id] = ContextCorruptionEnv() obs = _sessions[session_id].reset() return obs @app.post("/step/{session_id}", response_model=EpisodeObservation) def step(session_id: str, action: ContextCorruptionAction): if session_id not in _sessions: raise HTTPException(status_code=404, detail="Session not found") obs = _sessions[session_id].step(action) if obs.episode_done: del _sessions[session_id] return obs @app.get("/state/{session_id}") def state(session_id: str): if session_id not in _sessions: raise HTTPException(status_code=404, detail="Session not found") return _sessions[session_id].state() @app.delete("/session/{session_id}") def close_session(session_id: str): _sessions.pop(session_id, None) return {"status": "closed"} @app.get("/health") def health(): return {"status": "ok", "active_sessions": len(_sessions)} if __name__ == "__main__": uvicorn.run("environment.server:app", host="0.0.0.0", port=8000, reload=False)