Siddh12334's picture
feat: implement actions, reward, env, server
6801c6a
Raw
History Blame
1.71 kB
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)