File size: 4,521 Bytes
0cb452d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
"""
Module: client.py
Purpose: OpenEnv client for connecting to the TriageFlow environment.
Part of: Medical Triage Assistant — OpenEnv Round 1
Author: Team Squirrel

Overview:
    Implements the TriageFlowEnv client that inherits from OpenEnv's EnvClient.
    Handles conversion between typed Python objects and the WebSocket wire format.
    Users import this client to interact with a remote or local TriageFlow server.

Dependencies:
    - openenv.core.env_client: EnvClient base class
    - openenv.core.client_types: StepResult
    - models: TriageAction, TriageObservation, TriageState

Usage:
    from client import TriageFlowEnv
    from models import TriageAction

    async with TriageFlowEnv(base_url="http://localhost:8000") as env:
        result = await env.reset(task_name="basic-triage")
        result = await env.step(TriageAction(action_type="assign_priority", ...))
"""

from openenv.core.env_client import EnvClient
from openenv.core.client_types import StepResult
from models import TriageAction, TriageObservation, TriageState


class TriageFlowEnv(EnvClient[TriageAction, TriageObservation, TriageState]):
    """
    Client for the TriageFlow medical triage environment.

    Provides a type-safe interface for interacting with the TriageFlow
    server via WebSocket. Handles serialization of actions and
    deserialization of observations and state.

    Notes:
        Use .sync() for synchronous access in scripts and notebooks.
    """

    def _step_payload(self, action: TriageAction) -> dict:
        """
        Convert a TriageAction into JSON payload for the server.

        Args:
            action (TriageAction): The typed action object.

        Returns:
            dict: JSON-serializable payload.
        """
        payload = {
            "action_type": action.action_type.value if hasattr(action.action_type, 'value') else str(action.action_type),
            "patient_id": action.patient_id,
        }
        if action.priority_level is not None:
            payload["priority_level"] = action.priority_level.value if hasattr(action.priority_level, 'value') else str(action.priority_level)
        if action.info_field is not None:
            payload["info_field"] = action.info_field.value if hasattr(action.info_field, 'value') else str(action.info_field)
        if action.escalation_reason is not None:
            payload["escalation_reason"] = action.escalation_reason
        return payload

    def _parse_result(self, payload: dict) -> StepResult:
        """
        Parse the server's JSON response into a typed StepResult.

        Args:
            payload (dict): Raw JSON response from the server.

        Returns:
            StepResult: Typed result containing observation, reward, and done flag.
        """
        obs_data = payload.get("observation", payload)
        return StepResult(
            observation=TriageObservation(
                done=payload.get("done", False),
                reward=payload.get("reward"),
                current_patient=obs_data.get("current_patient"),
                queue_length=obs_data.get("queue_length", 0),
                queue_position=obs_data.get("queue_position", 0),
                missing_fields=obs_data.get("missing_fields", []),
                previous_action_feedback=obs_data.get("previous_action_feedback"),
                step_number=obs_data.get("step_number", 0),
                task_name=obs_data.get("task_name", ""),
            ),
            reward=payload.get("reward"),
            done=payload.get("done", False),
        )

    def _parse_state(self, payload: dict) -> TriageState:
        """
        Parse the server's state response into a typed TriageState.

        Args:
            payload (dict): Raw JSON state from the server.

        Returns:
            TriageState: Full internal state object.
        """
        return TriageState(
            episode_id=payload.get("episode_id"),
            step_count=payload.get("step_count", 0),
            task_name=payload.get("task_name", ""),
            patients=payload.get("patients", []),
            assignments=payload.get("assignments", {}),
            escalations=payload.get("escalations", {}),
            info_requests=payload.get("info_requests", []),
            action_history=payload.get("action_history", []),
            current_index=payload.get("current_index", 0),
            max_steps=payload.get("max_steps", 20),
            queue_cleared=payload.get("queue_cleared", False),
        )