sankar-raul/ICD-10-code-predictor-env
0
1"""Typed models for the Medical Coding Assistant environment."""2 3from typing import Literal4 5from openenv.core.env_server.types import Action, Observation, State6from pydantic import BaseModel, Field7 8 9Difficulty = Literal["easy", "medium", "hard"]10 11 12class RewardBreakdown(BaseModel):13 """Typed reward components for deterministic per-step feedback."""14 15 score_delta: float = Field(16 default=0.0,17 description="Positive reward from improvement in deterministic grader score.",18 )19 hint_penalty: float = Field(20 default=0.0,21 description="Penalty applied when requesting hints.",22 )23 invalid_code_penalty: float = Field(24 default=0.0,25 description="Penalty for codes outside the task's allowed set.",26 )27 loop_penalty: float = Field(28 default=0.0,29 description="Penalty for repeating identical actions.",30 )31 timeout_penalty: float = Field(32 default=0.0,33 description="Penalty when the episode ends by step-budget exhaustion.",34 )35 total: float = Field(36 default=0.0,37 description="Final per-step reward after all components.",38 )39 40 41class MedicalCodingAction(Action):42 """Action submitted by the agent on each environment step."""43 44 primary_code: str = Field(45 default="",46 description="Proposed primary ICD-10 code. Leave blank to keep the current draft.",47 )48 secondary_codes: list[str] = Field(49 default_factory=list,50 description="Proposed secondary ICD-10 codes for supporting findings or status codes.",51 )52 needs_review: bool = Field(53 default=False,54 description="Whether the chart should be escalated for human review.",55 )56 request_hint: bool = Field(57 default=False,58 description="Ask the environment for an additional deterministic hint.",59 )60 finalize: bool = Field(61 default=False,62 description="End the current task and score the current draft.",63 )64 65 66class MedicalCodingObservation(Observation):67 """Observation returned after reset and each environment step."""68 69 task_id: str = Field(..., description="Current task identifier.")70 difficulty: Difficulty = Field(..., description="Task difficulty level.")71 objective: str = Field(..., description="Task objective for the agent.")72 encounter_text: str = Field(..., description="Clinical note excerpt to code.")73 allowed_codes: list[str] = Field(74 default_factory=list,75 description="Closed code set allowed for this task.",76 )77 revealed_hints: list[str] = Field(78 default_factory=list,79 description="Hints revealed so far by the environment.",80 )81 current_primary_code: str = Field(82 default="",83 description="The agent's current draft primary code.",84 )85 current_secondary_codes: list[str] = Field(86 default_factory=list,87 description="The agent's current draft secondary codes.",88 )89 current_needs_review: bool = Field(90 default=False,91 description="The agent's current draft review flag.",92 )93 attempts_remaining: int = Field(94 default=0,95 description="How many steps remain in the current task.",96 )97 progress_score: float = Field(98 default=0.0,99 ge=0.0,100 le=1.0,101 description="Best deterministic grader score achieved so far.",102 )103 grader_feedback: list[str] = Field(104 default_factory=list,105 description="Deterministic feedback strings explaining the current draft.",106 )107 reward_breakdown: RewardBreakdown = Field(108 default_factory=RewardBreakdown,109 description="Typed per-step reward details.",110 )111 112 113class MedicalCodingState(State):114 """Internal state exposed through the OpenEnv state endpoint."""115 116 current_task_id: str = Field(default="", description="Active task identifier.")117 difficulty: Difficulty = Field(default="easy", description="Task difficulty.")118 current_primary_code: str = Field(default="", description="Draft primary code.")119 current_secondary_codes: list[str] = Field(120 default_factory=list,121 description="Draft secondary codes.",122 )123 current_needs_review: bool = Field(124 default=False,125 description="Current review flag draft.",126 )127 best_score: float = Field(128 default=0.0,129 ge=0.0,130 le=1.0,131 description="Best grader score seen in this episode.",132 )133 hints_used: int = Field(default=0, ge=0, description="Number of hints revealed.")134 repeated_actions: int = Field(135 default=0,136 ge=0,137 description="How many times the agent repeated the same draft.",138 )139 completed: bool = Field(default=False, description="Whether the task has finished.")140 