from __future__ import annotations import unittest from pathlib import Path from uuid import uuid4 from pydantic import ValidationError from .contracts.rupture_repair import ( RUPTURE_TYPES, RepairAttemptObservation, RuptureBenchmarkPack, RuptureDetectionSignal, RuptureEpisodeAssessment, RuptureEpisodeInput, ) from .services.rupture_repair import ( assess_rupture_episode, evaluate_rupture_benchmark, load_rupture_benchmark, render_rupture_benchmark_report, ) BENCHMARK_PATH = ( Path(__file__).resolve().parent / "data" / "rupture_repair_benchmark_g3.v1.json" ) def _all_keys(value: object) -> set[str]: if isinstance(value, dict): return set(value) | set().union(*(_all_keys(item) for item in value.values())) if isinstance(value, (list, tuple)): return set().union(*(_all_keys(item) for item in value)) if value else set() return set() class RuptureRepairContractTests(unittest.TestCase): @classmethod def setUpClass(cls) -> None: cls.pack = load_rupture_benchmark(BENCHMARK_PATH) def test_benchmark_is_version_fixed_synthetic_and_covers_every_type(self) -> None: self.assertEqual(self.pack.version, "1.0.0") self.assertEqual(self.pack.data_classification, "synthetic_educational") self.assertFalse(self.pack.clinical_claim_allowed) self.assertEqual( { case.expected.rupture_type for case in self.pack.cases if case.expected.detected }, set(RUPTURE_TYPES), ) def test_detected_signal_requires_type_confidence_and_evidence(self) -> None: with self.assertRaisesRegex( ValidationError, "requires type, confidence, and evidence" ): RuptureDetectionSignal( signal_id="invalid", loop="deep", status="detected", rupture_type="withdrawal", confidence=0.8, uncertainty=0.2, observed_at_turn=1, source_kind="simulated_state", perspective="client_simulation", ) def test_not_detected_signal_cannot_carry_a_type_or_confidence(self) -> None: with self.assertRaisesRegex(ValidationError, "must remain type/scoreless"): RuptureDetectionSignal( signal_id="invalid-negative", loop="deep", status="not_detected", rupture_type="withdrawal", confidence=0.2, uncertainty=0.8, observed_at_turn=1, source_kind="simulated_state", perspective="client_simulation", ) def test_model_detection_requires_model_run_provenance(self) -> None: base = { "signal_id": "model-signal", "loop": "deep", "status": "detected", "rupture_type": "withdrawal", "confidence": 0.8, "uncertainty": 0.2, "observed_at_turn": 1, "source_kind": "model_inferred", "perspective": "independent_observer", "evidence_refs": [ {"ref_id": "turn-1", "turn_index": 1, "speaker": "client"} ], } with self.assertRaisesRegex(ValidationError, "requires model_run_id"): RuptureDetectionSignal.model_validate(base) base["model_run_id"] = str(uuid4()) self.assertIsNotNone(RuptureDetectionSignal.model_validate(base).model_run_id) def test_detection_rejects_source_perspective_layer_mixing(self) -> None: with self.assertRaisesRegex(ValidationError, "mixes source and perspective"): RuptureDetectionSignal( signal_id="layer-mix", loop="deep", status="detected", rupture_type="withdrawal", confidence=0.8, uncertainty=0.2, observed_at_turn=1, source_kind="learner_reported", perspective="independent_observer", evidence_refs=( {"ref_id": "turn-1", "turn_index": 1, "speaker": "client"}, ), ) def test_client_response_evidence_must_follow_attempt(self) -> None: with self.assertRaisesRegex(ValidationError, "must follow the attempt"): RepairAttemptObservation( attempt_id="bad-order", turn_index=2, behaviors=("curiosity",), client_response="mixed", evidence_refs=( {"ref_id": "learner-2", "turn_index": 2, "speaker": "learner"}, ), response_evidence_refs=( {"ref_id": "client-2", "turn_index": 2, "speaker": "client"}, ), uncertainty=0.3, ) def test_repair_attempt_requires_prior_recognition(self) -> None: payload = self.pack.cases[0].episode.model_dump(mode="json") payload["recognized_at_turn"] = None payload["recognition_evidence_refs"] = [] with self.assertRaisesRegex( ValidationError, "repair attempts require rupture recognition" ): RuptureEpisodeInput.model_validate(payload) def test_fast_warning_must_reference_detected_fast_signal(self) -> None: payload = self.pack.cases[0].episode.model_dump(mode="json") payload["fast_warning"]["signal_id"] = "b001-deep" with self.assertRaisesRegex(ValidationError, "detected fast-loop signal"): RuptureEpisodeInput.model_validate(payload) def test_assessment_contract_rejects_compensating_total_score(self) -> None: payload = assess_rupture_episode(self.pack.cases[0].episode).model_dump(mode="json") payload["total_score"] = 1.0 with self.assertRaises(ValidationError): RuptureEpisodeAssessment.model_validate(payload) def test_benchmark_requires_same_template_to_have_contextual_outcomes(self) -> None: payload = self.pack.model_dump(mode="json") for case in payload["cases"]: for attempt in case["episode"]["repair_attempts"]: if attempt.get("utterance_template_id") == "magic-repair-v1": attempt["utterance_template_id"] = case["case_id"] with self.assertRaisesRegex( ValidationError, "memorized template must have different contextual outcomes" ): RuptureBenchmarkPack.model_validate(payload) class RuptureRepairStateMachineTests(unittest.TestCase): @classmethod def setUpClass(cls) -> None: cls.pack = load_rupture_benchmark(BENCHMARK_PATH) def _result(self, case_index: int): return assess_rupture_episode(self.pack.cases[case_index].episode) def test_fast_warning_is_superseded_when_follow_up_repairs_the_rupture(self) -> None: result = self._result(0) self.assertEqual(result.final_status, "resolved") self.assertEqual(result.reconciliation.disposition, "superseded_resolved") self.assertEqual( [entry.event_name for entry in result.ledger], [ "rupture.detected", "rupture.recognized", "repair.attempted", "repair.resolved", "rupture.reconciled", ], ) self.assertEqual( [entry.sequence_no for entry in result.ledger], list(range(1, len(result.ledger) + 1)), ) self.assertEqual(result.ledger[-1].reconciles_event_id, "b001-warning") def test_unrecognized_rupture_ends_missed_with_explicit_counterevidence(self) -> None: result = self._result(2) self.assertEqual(result.final_status, "missed") self.assertEqual(result.ledger[-1].event_name, "rupture.missed") self.assertIn("no_recognition_evidence", result.ledger[-1].counterevidence) def test_incomplete_behavior_and_mixed_response_stay_partial(self) -> None: result = self._result(1) attempt = result.repair_attempts[0] self.assertEqual(result.final_status, "partial") self.assertEqual(attempt.outcome, "partial") self.assertEqual(attempt.missing_behaviors, ("follow_up_check",)) self.assertIn( "required_repair_behavior_incomplete", attempt.counterevidence ) def test_formulaic_compliance_is_never_resolved(self) -> None: result = self._result(10) control = self._result(3) self.assertEqual(result.final_status, "missed") self.assertEqual(result.repair_attempts[0].client_response, "compliance_only") self.assertIn( "formulaic_language_without_observed_repair_impact", result.counterevidence, ) self.assertEqual(control.final_status, "resolved") self.assertEqual( self.pack.cases[10].episode.repair_attempts[0].utterance_template_id, self.pack.cases[3].episode.repair_attempts[0].utterance_template_id, ) def test_deep_review_can_dismiss_a_fast_false_positive(self) -> None: result = self._result(9) self.assertFalse(result.detected) self.assertEqual(result.final_status, "not_applicable") self.assertEqual(result.reconciliation.disposition, "dismissed") self.assertEqual(result.ledger, ()) self.assertIsNone(result.confidence) def test_all_sensor_errors_fail_closed_as_insufficient_evidence(self) -> None: episode = RuptureEpisodeInput( episode_id="oas-g3-episode-all-sensors-error", detection_signals=( RuptureDetectionSignal( signal_id="deep-error", loop="deep", status="error", uncertainty=1.0, observed_at_turn=2, source_kind="simulated_state", perspective="client_simulation", error_code="structured_output_invalid", ), ), ) result = assess_rupture_episode(episode) self.assertEqual(result.assessment_status, "error") self.assertFalse(result.detected) self.assertEqual(result.final_status, "insufficient_evidence") self.assertEqual(result.uncertainty, 1.0) self.assertEqual(result.ledger, ()) self.assertIn("all_detection_signals_failed", result.counterevidence[0]) def test_safety_is_returned_but_never_changes_repair_state(self) -> None: episode = self.pack.cases[6].episode with_safety = assess_rupture_episode(episode) without_safety = assess_rupture_episode( episode.model_copy(update={"safety_signals": ()}) ) self.assertEqual(with_safety.final_status, without_safety.final_status) self.assertEqual(with_safety.ledger, without_safety.ledger) self.assertEqual(len(with_safety.safety_signals), 1) self.assertEqual(without_safety.safety_signals, ()) def test_payload_has_no_total_or_score_key(self) -> None: keys = _all_keys(self._result(0).model_dump(mode="json")) self.assertFalse(any("total" in key for key in keys)) self.assertFalse(any("score" in key for key in keys)) self.assertIn("counterevidence", keys) self.assertIn("uncertainty", keys) self.assertIn("evidence_refs", keys) class RuptureRepairBenchmarkTests(unittest.TestCase): @classmethod def setUpClass(cls) -> None: cls.pack = load_rupture_benchmark(BENCHMARK_PATH) cls.report = evaluate_rupture_benchmark(cls.pack) def test_benchmark_meets_type_and_repair_gates(self) -> None: self.assertEqual(self.report["case_count"], 11) self.assertGreaterEqual(self.report["rupture_type_macro_f1"], 0.85) self.assertEqual(self.report["rupture_type_macro_f1"], 1.0) self.assertEqual(self.report["repair_status_accuracy"], 1.0) self.assertEqual(self.report["critical_miss_count"], 0) self.assertTrue( all(value == 1.0 for value in self.report["rupture_type_f1"].values()) ) def test_benchmark_has_zero_judge_gaming_or_memorized_phrase_regression(self) -> None: self.assertEqual(self.report["judge_gaming_regressions"], 0) self.assertEqual(self.report["memorized_phrase_false_resolutions"], 0) self.assertEqual( self.report["detection_confusion"], { "true_positive": 10, "false_positive": 0, "false_negative": 0, "true_negative": 1, }, ) def test_report_exposes_reconciliation_uncertainty_and_evidence(self) -> None: first = self.report["rows"][0] self.assertEqual(first["reconciliation"], "superseded_resolved") self.assertIsInstance(first["uncertainty"], float) self.assertTrue(first["evidence_refs"]) self.assertIn("counterevidence", first) def test_report_keeps_synthetic_scope_and_safety_separate(self) -> None: rendered = render_rupture_benchmark_report(self.report) self.assertIn('"data_classification": "synthetic_educational"', rendered) self.assertIn('"clinical_claim_allowed": false', rendered) safety_row = next( row for row in self.report["rows"] if row["case_id"] == "oas-g3-bench-007" ) self.assertEqual(safety_row["safety_signal_count"], 1) self.assertEqual(safety_row["actual_status"], "missed") if __name__ == "__main__": unittest.main()