diff --git a/triage_env/server/app.py b/triage_env/server/app.py index aa7a2f9..f018042 100644 --- a/triage_env/server/app.py +++ b/triage_env/server/app.py @@ -39,6 +39,14 @@ from fastapi.openapi.docs import get_redoc_html, get_swagger_ui_html from fastapi.openapi.utils import get_openapi from fastapi.responses import JSONResponse, RedirectResponse +try: + import gradio as gr +except Exception: # pragma: no cover + gr = None + +GRADIO_PATH = "/gradio" +_GRADIO_MOUNT_ERROR: str | None = None + try: from ..models import TriageAction, TriageObservation from .triage_env_environment import TriageEnvironment @@ -56,22 +64,60 @@ app = create_app( max_concurrent_envs=1, # increase this number to allow more concurrent WebSocket sessions ) +if gr is not None: + try: + try: + from .gradio_ui import build_gradio_ui + except ModuleNotFoundError: + from server.gradio_ui import build_gradio_ui + app = gr.mount_gradio_app(app, build_gradio_ui(), path=GRADIO_PATH) + except Exception as exc: + # Keep API serving even if Gradio UI mount fails, but retain error for diagnostics. + _GRADIO_MOUNT_ERROR = f"{type(exc).__name__}: {exc}" + def _has_route(path: str) -> bool: return any(getattr(route, "path", None) == path for route in app.routes) +def _has_route_variant(path: str) -> bool: + return _has_route(path) or _has_route(f"{path}/") + + @app.get("/", include_in_schema=False) def root(): - # On Spaces, users land on "/" first. Redirect to docs when available. + # On Spaces, users land on "/" first. Prefer Gradio UI, then docs. + if _has_route_variant(GRADIO_PATH): + return RedirectResponse(url=GRADIO_PATH) if _has_route("/docs"): return RedirectResponse(url="/docs") - return { + payload = { "message": "MedicalTriage API is running", "health": "/health", "openapi": "/openapi.json", + "ui": GRADIO_PATH, "docs": "/docs", } + if _GRADIO_MOUNT_ERROR is not None: + payload["gradio_mount_error"] = _GRADIO_MOUNT_ERROR + return payload + + +if not _has_route("/ui"): + + @app.get("/ui", include_in_schema=False) + def ui_redirect(): + if _has_route_variant(GRADIO_PATH): + return RedirectResponse(url=GRADIO_PATH) + if _has_route("/docs"): + return RedirectResponse(url="/docs") + return JSONResponse( + { + "message": "Gradio UI is unavailable", + "gradio_mount_error": _GRADIO_MOUNT_ERROR, + }, + status_code=503, + ) if not _has_route("/openapi.json"): diff --git a/triage_env/server/gradio_ui.py b/triage_env/server/gradio_ui.py new file mode 100644 index 0000000..d5dfd04 --- /dev/null +++ b/triage_env/server/gradio_ui.py @@ -0,0 +1,115 @@ +from __future__ import annotations + +import threading +import traceback +from typing import Any + +import gradio as gr + +try: + from ..models import TriageAction + from .triage_env_environment import TriageEnvironment +except ModuleNotFoundError: + from models import TriageAction + from server.triage_env_environment import TriageEnvironment + + +_ENV_LOCK = threading.Lock() +_ENV: TriageEnvironment | None = None + + +def _ensure_env(task: str | None = None) -> TriageEnvironment: + global _ENV + if _ENV is None: + initial_task = task or "task2" + _ENV = TriageEnvironment(task=initial_task) + _ENV.reset(task=initial_task) + return _ENV + + +def _obs_payload(env: TriageEnvironment, observation: Any) -> dict[str, Any]: + return { + "observation": observation.model_dump(mode="json"), + "state": env.state.model_dump(mode="json"), + } + + +def _state_payload(env: TriageEnvironment) -> dict[str, Any]: + return { + "task": env.task_name, + "step_count": env.state.step_count, + "state": env.state.model_dump(mode="json"), + } + + +def build_gradio_ui() -> gr.Blocks: + with gr.Blocks(title="Medical Triage") as demo: + gr.Markdown("## Medical Triage\nSimple controls for reset, step, and getState.") + + with gr.Row(): + task = gr.Dropdown( + choices=["task1", "task2", "task3"], + value="task2", + label="Task", + ) + action_type = gr.Dropdown( + choices=["treat", "allocate_ventilator", "wait"], + value="wait", + label="Action Type", + ) + patient_id = gr.Number(value=-1, precision=0, label="Patient ID") + + with gr.Row(): + reset_btn = gr.Button("Reset", variant="primary") + step_btn = gr.Button("Step") + get_state_btn = gr.Button("Get State") + + output = gr.JSON(label="Response") + + def on_reset(selected_task: str): + try: + with _ENV_LOCK: + env = _ensure_env(selected_task) + obs = env.reset(task=selected_task) + return _obs_payload(env, obs) + except Exception as exc: + return {"error": str(exc), "traceback": traceback.format_exc()} + + def on_step( + selected_task: str, + selected_action: str, + selected_patient_id: float, + ): + try: + with _ENV_LOCK: + env = _ensure_env(selected_task) + if env.task_name != selected_task: + env.reset(task=selected_task) + + pid = -1 if selected_action == "wait" else int(selected_patient_id or 0) + action = TriageAction(action_type=selected_action, patient_id=pid) + obs = env.step(action) + return _obs_payload(env, obs) + except Exception as exc: + return {"error": str(exc), "traceback": traceback.format_exc()} + + def on_get_state(): + try: + with _ENV_LOCK: + env = _ensure_env("task2") + return _state_payload(env) + except Exception as exc: + return {"error": str(exc), "traceback": traceback.format_exc()} + + reset_btn.click(on_reset, inputs=[task], outputs=[output], queue=False) + step_btn.click( + on_step, + inputs=[task, action_type, patient_id], + outputs=[output], + queue=False, + ) + get_state_btn.click(on_get_state, inputs=[], outputs=[output], queue=False) + + demo.load(on_get_state, inputs=[], outputs=[output], queue=False) + + return demo diff --git a/triage_env/server/triage_env_environment.py b/triage_env/server/triage_env_environment.py index 429f208..7e7dcd6 100644 --- a/triage_env/server/triage_env_environment.py +++ b/triage_env/server/triage_env_environment.py @@ -1,3 +1,4 @@ +import math from uuid import uuid4 from openenv.core.env_server.interfaces import Environment @@ -12,6 +13,7 @@ except ImportError: class TriageEnvironment(Environment): SUPPORTS_CONCURRENT_SESSIONS: bool = True + STEP_REWARD_SCALE: float = 25.0 def __init__( self, @@ -253,6 +255,8 @@ class TriageEnvironment(Environment): reward += terminal_reward reward_breakdown["episode_terminal_reward"] = terminal_reward + raw_reward = reward + reward = self._normalize_step_reward(raw_reward) self._state.total_reward += reward components = { @@ -274,7 +278,11 @@ class TriageEnvironment(Environment): reward=reward, reward_detail=TriageReward(value=float(reward), components=components, penalties=penalties), message=message, - metadata=self._build_metadata(reward_breakdown=reward_breakdown), + metadata=self._build_metadata( + reward_breakdown=reward_breakdown, + raw_reward=raw_reward, + normalized_reward=reward, + ), ) @property @@ -385,6 +393,16 @@ class TriageEnvironment(Environment): return terminal_reward, success_achieved + def _normalize_step_reward(self, raw_reward: float) -> float: + """Normalize raw reward to (0,1) and round to two decimals.""" + if not math.isfinite(raw_reward): + raw_reward = 0.0 + + scaled = 0.5 + (0.5 * math.tanh(raw_reward / self.STEP_REWARD_SCALE)) + rounded_cents = int(math.floor((scaled * 100.0) + 0.5)) + bounded_cents = min(99, max(1, rounded_cents)) + return bounded_cents / 100.0 + def _is_done(self) -> bool: if self._state.step_count >= self._state.max_steps: return True @@ -398,12 +416,17 @@ class TriageEnvironment(Environment): return False - def _build_metadata(self, reward_breakdown: dict) -> dict: + def _build_metadata( + self, + reward_breakdown: dict, + raw_reward: float | None = None, + normalized_reward: float | None = None, + ) -> dict: alive_patients = [p for p in self._state.patients if p.alive] critical_patients = [p for p in self._state.patients if p.severity == "critical"] surviving_critical = [p for p in critical_patients if p.alive] - return { + metadata = { "episode_id": self._state.episode_id, "task": self.task_name, "total_reward": self._state.total_reward, @@ -418,4 +441,11 @@ class TriageEnvironment(Environment): "ventilators_available": self.task_config.ventilators_available, }, "reward_breakdown": reward_breakdown, - } \ No newline at end of file + } + + if raw_reward is not None: + metadata["raw_reward"] = float(raw_reward) + if normalized_reward is not None: + metadata["normalized_reward"] = float(normalized_reward) + + return metadata \ No newline at end of file diff --git a/triage_env/tests/test_environment.py b/triage_env/tests/test_environment.py index c2a6a0e..eaba830 100644 --- a/triage_env/tests/test_environment.py +++ b/triage_env/tests/test_environment.py @@ -5,6 +5,11 @@ from triage_env.tasks import TASK_CONFIGS from triage_env.models import TriageAction +def _assert_step_reward_contract(value: float) -> None: + assert 0.0 < value < 1.0 + assert value == pytest.approx(round(value, 2), abs=1e-12) + + @pytest.fixture def env(): environment = TriageEnvironment(max_steps=20) @@ -44,6 +49,8 @@ def test_treat_more_urgent_patient_better_than_wait(): env2.reset() wait_obs = env2.step(TriageAction(action_type="wait")) + _assert_step_reward_contract(treat_obs.reward) + _assert_step_reward_contract(wait_obs.reward) assert treat_obs.reward > wait_obs.reward @@ -60,13 +67,16 @@ def test_treat_critical_better_than_moderate(): TriageAction(action_type="treat", patient_id=2) ).reward + _assert_step_reward_contract(critical_reward) + _assert_step_reward_contract(moderate_reward) assert critical_reward > moderate_reward def test_invalid_treatment_gets_penalty(env): obs = env.step(TriageAction(action_type="treat", patient_id=999)) - assert obs.reward < 0 + _assert_step_reward_contract(obs.reward) + assert obs.reward < 0.5 assert obs.message == "Invalid treatment action" @@ -83,13 +93,16 @@ def test_allocate_ventilator_to_critical_is_better_than_moderate(): TriageAction(action_type="allocate_ventilator", patient_id=2) ) + _assert_step_reward_contract(critical_obs.reward) + _assert_step_reward_contract(moderate_obs.reward) assert critical_obs.reward > moderate_obs.reward def test_wait_penalty_when_urgent_patient_exists(env): obs = env.step(TriageAction(action_type="wait")) - assert obs.reward < 0 + _assert_step_reward_contract(obs.reward) + assert obs.reward < 0.5 assert obs.message == "Waited one step" @@ -105,6 +118,8 @@ def test_ignoring_critical_is_bad(): wait_env.reset() wait_reward = wait_env.step(TriageAction(action_type="wait")).reward + _assert_step_reward_contract(moderate_reward) + _assert_step_reward_contract(wait_reward) assert moderate_reward > wait_reward @@ -127,15 +142,30 @@ def test_patient_can_die_if_ignored(): def test_patient_death_penalty(): - env = TriageEnvironment(max_steps=10) - env.reset() + wait_env = TriageEnvironment(max_steps=10) + wait_env.reset() + + wait_total = 0.0 + for _ in range(10): + obs = wait_env.step(TriageAction(action_type="wait")) + _assert_step_reward_contract(obs.reward) + wait_total += obs.reward + + treat_env = TriageEnvironment(max_steps=10) + treat_env.reset() - total_reward = 0.0 + treat_total = 0.0 for _ in range(10): - obs = env.step(TriageAction(action_type="wait")) - total_reward += obs.reward + alive = [p for p in treat_env.state.patients if p.alive] + if not alive: + break + + target = min(alive, key=lambda p: p.health) + obs = treat_env.step(TriageAction(action_type="treat", patient_id=target.id)) + _assert_step_reward_contract(obs.reward) + treat_total += obs.reward - assert total_reward < -20 + assert wait_total < treat_total def test_environment_runs_multiple_steps_without_crashing(): @@ -154,10 +184,26 @@ def test_environment_runs_multiple_steps_without_crashing(): assert obs is not None assert obs.step_count >= 1 assert isinstance(obs.reward, float) + _assert_step_reward_contract(obs.reward) def test_reward_breakdown_present(env): obs = env.step(TriageAction(action_type="treat", patient_id=0)) assert "reward_breakdown" in obs.metadata - assert isinstance(obs.metadata["reward_breakdown"], dict) \ No newline at end of file + assert isinstance(obs.metadata["reward_breakdown"], dict) + + +@pytest.mark.parametrize("task_name", ["task1", "task2", "task3"]) +def test_step_reward_contract_across_all_tasks(task_name): + env = TriageEnvironment(task=task_name) + env.reset(task=task_name) + + actions = [ + TriageAction(action_type="wait", patient_id=-1), + TriageAction(action_type="treat", patient_id=0), + TriageAction(action_type="allocate_ventilator", patient_id=0), + ] + for action in actions: + obs = env.step(action) + _assert_step_reward_contract(obs.reward) \ No newline at end of file diff --git a/triage_env/tests/test_step.py b/triage_env/tests/test_step.py index 1433214..3cc480c 100644 --- a/triage_env/tests/test_step.py +++ b/triage_env/tests/test_step.py @@ -32,12 +32,18 @@ from triage_env.server.triage_env_environment import TriageEnvironment from triage_env.models import TriageAction +def _assert_step_reward_contract(value: float) -> None: + assert 0.0 < value < 1.0 + assert value == round(value, 2) + + def test_step_increments_step_count(): env = TriageEnvironment() env.reset() obs = env.step(TriageAction(action_type="wait", patient_id=-1)) assert obs.step_count == 1 + _assert_step_reward_contract(obs.reward) def test_treat_action_returns_observation(): @@ -46,4 +52,5 @@ def test_treat_action_returns_observation(): obs = env.step(TriageAction(action_type="treat", patient_id=0)) assert isinstance(obs.reward, float) + _assert_step_reward_contract(obs.reward) assert hasattr(obs, "patients") \ No newline at end of file diff --git a/triage_env/tests/test_task_reward_scaling.py b/triage_env/tests/test_task_reward_scaling.py index 4bc5a65..823aaf7 100644 --- a/triage_env/tests/test_task_reward_scaling.py +++ b/triage_env/tests/test_task_reward_scaling.py @@ -2,6 +2,11 @@ from triage_env.models import TriageAction from triage_env.server.triage_env_environment import TriageEnvironment +def _assert_step_reward_contract(value: float) -> None: + assert 0.0 < value < 1.0 + assert value == round(value, 2) + + def test_wait_penalty_harsher_on_task3_than_task1(): env_easy = TriageEnvironment(task="task1") env_easy.reset(task="task1") @@ -11,4 +16,6 @@ def test_wait_penalty_harsher_on_task3_than_task1(): env_hard.reset(task="task3") reward_hard = env_hard.step(TriageAction(action_type="wait", patient_id=-1)).reward + _assert_step_reward_contract(reward_easy) + _assert_step_reward_contract(reward_hard) assert reward_hard < reward_easy