Ayushmaan03's picture
Add root redirect to docs to fix 404 error
9dd8eee
Raw
History Blame Contribute Delete
5.06 kB
"""
Log Analysis OpenEnv - FastAPI Server
Sessions are isolated: each /reset call creates a new session and returns a
session_id. Callers pass ?session_id=<id> to /step and /state.
If no session_id is supplied, the server falls back to a default global session
so that the plain `openenv validate` flow (which doesn't pass session IDs) still
works correctly.
"""
import os
import sys
# ---------------------------------------------------------------------------
# Path fix: make the project root importable when running via `uvicorn server.app:app`
# from the project root directory. This is the standard approach for projects
# that aren't installed as editable packages via `pip install -e .`.
# ---------------------------------------------------------------------------
_project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if _project_root not in sys.path:
sys.path.insert(0, _project_root)
from uuid import uuid4
from typing import Dict, Optional
from fastapi import FastAPI, Query
from fastapi.responses import RedirectResponse
from pydantic import BaseModel
from models import Action, Observation, State, PublicState
from environment import LogAnalysisEnv
app = FastAPI(title="Log Analysis OpenEnv")
@app.get("/")
def read_root():
"""Redirect to the API documentation."""
return RedirectResponse(url="/docs")
# ---------------------------------------------------------------------------
# Session store — maps session_id -> LogAnalysisEnv instance.
# A default "global" session is kept for tool compatibility (openenv validate).
# ---------------------------------------------------------------------------
_DEFAULT_SESSION = "default"
_sessions: Dict[str, LogAnalysisEnv] = {_DEFAULT_SESSION: LogAnalysisEnv()}
def _get_or_create_session(session_id: Optional[str]) -> tuple[str, LogAnalysisEnv]:
"""Return (session_id, env) for the given id, creating if absent."""
sid = session_id or _DEFAULT_SESSION
if sid not in _sessions:
_sessions[sid] = LogAnalysisEnv()
return sid, _sessions[sid]
# ---------------------------------------------------------------------------
# Response model
# ---------------------------------------------------------------------------
class StepResult(BaseModel):
session_id: str
observation: Observation
reward: float
done: bool
info: str = ""
# ---------------------------------------------------------------------------
# Endpoints
# ---------------------------------------------------------------------------
@app.post("/reset", response_model=StepResult)
def reset_endpoint(
session_id: Optional[str] = Query(default=None),
seed: Optional[int] = Query(default=None),
) -> StepResult:
"""
Start a new episode.
- Omit session_id to use a fresh auto-generated session.
- Pass a UUID session_id to run an isolated episode.
- Pass seed to make the scenario deterministic (for reproducible baselines).
Returns the initial observation plus the session_id to use for subsequent calls.
"""
sid = session_id or str(uuid4())
env = LogAnalysisEnv()
_sessions[sid] = env
obs = env.reset(seed=seed)
return StepResult(session_id=sid, observation=obs, reward=0.0, done=False, info=obs.info)
@app.post("/step", response_model=StepResult)
def step_endpoint(
action: Action,
session_id: Optional[str] = Query(default=None),
) -> StepResult:
"""Advance the episode by one action."""
sid, env = _get_or_create_session(session_id)
obs, reward, done = env.step(action)
return StepResult(session_id=sid, observation=obs, reward=reward, done=done, info=obs.info)
@app.get("/state", response_model=PublicState)
def state_endpoint(session_id: Optional[str] = Query(default=None)) -> PublicState:
"""
Return the current public state of the episode.
ground_truth is intentionally excluded to prevent trivial agent exploits.
"""
sid, env = _get_or_create_session(session_id)
if env._state is None:
env.reset()
s = env.state()
return PublicState(
episode_id=s.episode_id,
step_count=s.step_count,
scenario_name=s.scenario_name,
total_reward=s.total_reward,
is_done=s.is_done,
)
@app.get("/health")
def health_endpoint():
return {"status": "healthy"}
@app.get("/debug/ground_truth")
def debug_ground_truth_endpoint(session_id: Optional[str] = Query(default=None)):
"""
Returns ground truth for post-episode grading in inference.py.
This is intentionally separate from /state to prevent in-episode leakage.
Only call this AFTER the episode is done.
"""
sid, env = _get_or_create_session(session_id)
if env._state is None:
return {"error": "No episode running for this session."}
return env._current_scenario["ground_truth"]
def main():
"""Entry point when the package script `server` is invoked."""
import uvicorn
uvicorn.run("server.app:app", host="0.0.0.0", port=8000, reload=False)
if __name__ == "__main__":
main()