CoolFace
Apppublic

shashanks/medical_coding

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
test_env_cycle.py132 linesDownload Raw Back to tests
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