shashanks/medical_coding
0
1"""2Tests for environment reset/step lifecycle and observation schema.3"""4 5import sys6import os7 8# Ensure project root is importable9sys.path.insert(0, os.path.join(os.path.dirname(__file__), ".."))10 11from server.environment import MedicalCodingEnvironment12from models import MedicalCodingAction, MedicalCodingObservation13 14TASK_IDS = [15 "easy_demographic",16 "medium_ncci_conflict",17 "medium_excludes1",18 "hard_specificity_untraceable",19 "expert_multi_error",20]21 22 23def test_reset_returns_valid_observation():24 """Reset should return a valid MedicalCodingObservation for every task."""25 env = MedicalCodingEnvironment()26 for task_id in TASK_IDS:27 obs = env.reset(task_id=task_id)28 assert isinstance(obs, MedicalCodingObservation)29 assert obs.task_id == task_id30 assert obs.clinical_note != ""31 assert len(obs.proposed_codes) > 032 assert obs.done is False33 assert obs.grader_score is None34 35 36def test_easy_demographic_patient_is_male():37 """Easy task patient should be male (maternity code mismatch)."""38 env = MedicalCodingEnvironment()39 obs = env.reset(task_id="easy_demographic")40 assert obs.patient_demographics["sex"] == "male"41 assert obs.patient_demographics["age"] == 3442 43 44def test_query_guideline_returns_tool_result():45 """Querying a valid proposed code should return guideline text."""46 env = MedicalCodingEnvironment()47 obs = env.reset(task_id="easy_demographic")48 action = MedicalCodingAction(action_type="query_guideline", code="O80")49 obs = env.step(action)50 assert obs.tool_result != ""51 assert "O80" in obs.codes_queried52 assert obs.reward > 0 # first valid query = +0.1053 54 55def test_query_hallucinated_code_penalized():56 """Querying a code NOT in proposed set should get -0.50 penalty."""57 env = MedicalCodingEnvironment()58 env.reset(task_id="easy_demographic")59 action = MedicalCodingAction(action_type="query_guideline", code="X99.99")60 obs = env.step(action)61 assert obs.reward == -0.5062 63 64def test_submit_audit_ends_episode():65 """Submitting audit should end the episode and set grader_score."""66 env = MedicalCodingEnvironment()67 env.reset(task_id="easy_demographic")68 action = MedicalCodingAction(action_type="submit_audit")69 obs = env.step(action)70 assert obs.done is True71 assert obs.grader_score is not None72 assert 0.0 <= obs.grader_score <= 1.073 74 75def test_full_correct_easy_task():76 """Complete the easy task correctly: query O80 -> flag -> submit."""77 env = MedicalCodingEnvironment()78 env.reset(task_id="easy_demographic")79 80 # Query guideline for O8081 obs = env.step(MedicalCodingAction(action_type="query_guideline", code="O80"))82 assert obs.reward > 083 84 # Flag the error85 obs = env.step(MedicalCodingAction(86 action_type="flag_error",87 code="O80",88 error_type="demographic_mismatch",89 justification="O80 is a female-only maternity code. Patient is male.",90 ))91 assert obs.reward > 092 93 # Submit94 obs = env.step(MedicalCodingAction(action_type="submit_audit"))95 assert obs.done is True96 assert obs.grader_score is not None97 assert obs.grader_score >= 0.5 # should pass threshold98 99 100def test_episode_metrics_populated_on_done():101 """Episode metrics should be populated when done=True."""102 env = MedicalCodingEnvironment()103 env.reset(task_id="easy_demographic")104 obs = env.step(MedicalCodingAction(action_type="submit_audit"))105 assert obs.done is True106 assert obs.episode_metrics is not None107 assert "trajectory_length" in obs.episode_metrics108 assert "flag_precision" in obs.episode_metrics109 assert "flag_recall" in obs.episode_metrics110 111 112def test_step_after_done_raises_or_noop():113 """Stepping after episode is done should not crash."""114 env = MedicalCodingEnvironment()115 env.reset(task_id="easy_demographic")116 env.step(MedicalCodingAction(action_type="submit_audit"))117 # A second step should either raise or return a done observation118 try:119 obs = env.step(MedicalCodingAction(action_type="query_guideline", code="O80"))120 # If it doesn't raise, it should still be done121 assert obs.done is True122 except Exception:123 pass # Raising is also acceptable124 125 126def test_all_tasks_have_proposed_codes():127 """Every task should have at least 2 proposed codes."""128 env = MedicalCodingEnvironment()129 for task_id in TASK_IDS:130 obs = env.reset(task_id=task_id)131 assert len(obs.proposed_codes) >= 2, f"{task_id} has too few codes"132 