from __future__ import annotations from unittest import IsolatedAsyncioTestCase from unittest.mock import AsyncMock, patch from uuid import UUID from pydantic import ValidationError from .config import Settings from .services import continuous_improvement_agentic as agentic from .services import continuous_improvement_producer as producer from .services import continuous_improvement_trigger as trigger from .services.guardrail import mask_pii DRIFT_REPORT_ID = UUID("82000000-0000-0000-0000-000000000001") def _drift_row(**overrides: object) -> dict[str, object]: payload: dict[str, object] = { "drift_report_id": DRIFT_REPORT_ID, "content_hash": "a" * 64, "matched_count": 12, "status": "drift_flagged", "baseline_accuracy": 0.92, "candidate_accuracy": 0.67, "accuracy_delta": -0.25, "alerts": ["overall_accuracy_regression"], "subgroup_metrics": [], "data_classification": "synthetic_educational", "clinical_claim_allowed": False, } payload.update(overrides) return payload class AcquireContext: def __init__(self, conn: AsyncMock) -> None: self.conn = conn async def __aenter__(self) -> AsyncMock: return self.conn async def __aexit__(self, *_: object) -> None: return None class ContinuousImprovementDriftTriggerTest(IsolatedAsyncioTestCase): def test_drift_ledger_trigger_is_independently_disabled_by_default(self) -> None: self.assertFalse( Settings.model_fields[ "continuous_improvement_drift_trigger_enabled" ].default ) async def test_loader_reads_only_metadata_from_research_drift_ledger(self) -> None: conn = AsyncMock() conn.fetch.return_value = [_drift_row()] signals, rejected = await trigger.load_triggerable_drift_signals(conn, limit=3) self.assertEqual(rejected, 0) self.assertEqual(len(signals), 1) sql = str(conn.fetch.await_args.args[0]).lower() self.assertIn("from app.supervision_drift_report", sql) self.assertIn("status = 'drift_flagged'", sql) self.assertIn("data_classification = 'synthetic_educational'", sql) self.assertIn("clinical_claim_allowed = false", sql) self.assertIn("app.ci_agentic_job", sql) self.assertIn("app.supervision_drift_subgroup_metric", sql) self.assertIn("report.matched_count >= $2", sql) self.assertEqual( conn.fetch.await_args.args[2], trigger.supervision_research.MIN_DRIFT_MATCHES, ) self.assertNotIn("for update", sql) for forbidden in ( "app.sessions", "app.turns", "learner_id", "session_id", "cohort_id", "disagreement_case_refs", "evidence_pointer_ids", "raw_transcript", "text_masked", ): self.assertNotIn(forbidden, sql) def test_signal_builds_deterministic_metadata_only_incident_and_job(self) -> None: signal = trigger.DriftBenchmarkSignal.model_validate(_drift_row()) first = trigger.build_drift_adversarial_trigger(signal) second = trigger.build_drift_adversarial_trigger(signal) self.assertEqual(first, second) self.assertFalse(first.incident.pii_included) self.assertEqual(first.job_spec.trigger_kind, "scheduled_incident") self.assertEqual( first.job_spec.job_key, f"oas-g8-job-g6-drift-{DRIFT_REPORT_ID.hex}", ) self.assertEqual(first.job_spec.content_kind, "benchmark") self.assertEqual(first.job_spec.data_classification, producer.DATA_CLASSIFICATION) self.assertEqual(len(first.job_spec.source_packs), 1) self.assertEqual( first.job_spec.source_packs[0], agentic.source_pack_from_operational_incident(first.incident), ) serialized = first.job_spec.source_packs[0].content.lower() self.assertEqual(mask_pii(serialized).entities, []) for forbidden in ( "learner_id", "session_id", "cohort_id", "raw_transcript", "clinical_claim_allowed", "diagnosis", ): self.assertNotIn(forbidden, serialized) self.assertRegex( first.incident.evidence_refs[0], r"^db://app/supervision-drift-report/[a-p]{32}$", ) self.assertNotRegex(first.incident.evidence_refs[0], r"[0-9]") def test_non_triggerable_or_unsafe_ledger_rows_fail_closed(self) -> None: rejected_rows = ( _drift_row(status="stable", alerts=[]), _drift_row(data_classification="production_transcript"), _drift_row(clinical_claim_allowed=True), _drift_row(alerts=["unknown_runtime_alert"]), _drift_row(matched_count=0), _drift_row(accuracy_delta=0.25), _drift_row( candidate_accuracy=0.91, accuracy_delta=-0.01, ), _drift_row( candidate_accuracy=0.92, accuracy_delta=0.0, alerts=["synthetic_subgroup_regression:synthetic-a"], ), ) for row in rejected_rows: with self.subTest(row=row): with self.assertRaises(ValidationError): trigger.DriftBenchmarkSignal.model_validate(row) def test_subgroup_alert_requires_the_canonical_g6_metric_threshold(self) -> None: signal = trigger.DriftBenchmarkSignal.model_validate( _drift_row( candidate_accuracy=0.92, accuracy_delta=0.0, alerts=["synthetic_subgroup_regression:synthetic-a"], subgroup_metrics=[ { "subgroup": "synthetic-a", "matched_count": 2, "baseline_accuracy": 1.0, "candidate_accuracy": 0.5, "accuracy_delta": -0.5, } ], ) ) self.assertEqual( signal.alerts, ("synthetic_subgroup_regression:synthetic-a",), ) with self.assertRaises(ValidationError): trigger.DriftBenchmarkSignal.model_validate( _drift_row( candidate_accuracy=0.92, accuracy_delta=0.0, alerts=["synthetic_subgroup_regression:synthetic-a"], subgroup_metrics=[ { "subgroup": "synthetic-a", "matched_count": 2, "baseline_accuracy": 1.0, "candidate_accuracy": 0.9, "accuracy_delta": -0.1, } ], ) ) async def test_invalid_rows_are_counted_without_crossing_write_boundary(self) -> None: conn = AsyncMock() conn.fetch.return_value = [ _drift_row(status="stable", alerts=[]), _drift_row(drift_report_id=UUID(int=2)), ] signals, rejected = await trigger.load_triggerable_drift_signals(conn, limit=5) self.assertEqual(rejected, 1) self.assertEqual([item.drift_report_id for item in signals], [UUID(int=2)]) conn.execute.assert_not_awaited() async def test_trigger_persists_incident_then_enqueues_job_in_same_rls_scope(self) -> None: conn = AsyncMock() conn.fetch.return_value = [_drift_row()] submit = AsyncMock(return_value={"idempotent_replay": False}) enqueue = AsyncMock(return_value=UUID(int=9)) with ( patch.object(trigger.db, "acquire", return_value=AcquireContext(conn)) as acquire, patch.object(trigger.continuous_improvement_store, "submit_incident_dag", submit), patch.object(producer, "enqueue_agentic_job", enqueue), ): result = await trigger.enqueue_drift_adversarial_jobs_once(limit=4) acquire.assert_called_once_with(ai_view="research", ai_context=True) self.assertEqual(result.scanned, 1) self.assertEqual(result.incidents_created, 1) self.assertEqual(result.jobs_enqueued, 1) self.assertEqual(result.invalid_signals, 0) submit.assert_awaited_once() enqueue.assert_awaited_once() self.assertIs(submit.await_args.kwargs["conn"], conn) self.assertIs(enqueue.await_args.args[0], conn) spec = enqueue.await_args.args[1] self.assertEqual(spec.trigger_kind, "scheduled_incident") self.assertNotIn("catalog", str(submit.await_args).lower()) self.assertNotIn("approval", str(submit.await_args).lower()) async def test_incident_replay_still_repairs_missing_job_idempotently(self) -> None: conn = AsyncMock() conn.fetch.return_value = [_drift_row()] submit = AsyncMock(return_value={"idempotent_replay": True}) enqueue = AsyncMock(return_value=UUID(int=9)) with ( patch.object(trigger.db, "acquire", return_value=AcquireContext(conn)), patch.object(trigger.continuous_improvement_store, "submit_incident_dag", submit), patch.object(producer, "enqueue_agentic_job", enqueue), ): result = await trigger.enqueue_drift_adversarial_jobs_once(limit=1) self.assertEqual(result.incident_replays, 1) self.assertEqual(result.incidents_created, 0) self.assertEqual(result.jobs_enqueued, 1) async def test_existing_opt_in_scheduler_feeds_trigger_into_lease_retry_queue(self) -> None: trigger_result = trigger.DriftTriggerCycleResult( scanned=1, invalid_signals=0, incidents_created=1, incident_replays=0, jobs_enqueued=1, ) with ( patch.object( trigger, "enqueue_drift_adversarial_jobs_once", AsyncMock(return_value=trigger_result), ) as run_trigger, patch.object(producer, "ensure_repo_approved_job", AsyncMock(return_value=UUID(int=1))), patch.object(producer, "claim_next_agentic_job", AsyncMock(return_value=None)), patch.object( producer.settings, "continuous_improvement_drift_trigger_enabled", True, ), ): result = await producer.produce_queued_agentic_jobs_once() run_trigger.assert_awaited_once() self.assertEqual(result["drift_jobs_enqueued"], 1) self.assertEqual(result["drift_trigger_failed"], 0) self.assertEqual(result["claimed"], 0) async def test_producer_cycle_does_not_read_drift_ledger_without_trigger_opt_in( self, ) -> None: run_trigger = AsyncMock() with ( patch.object( trigger, "enqueue_drift_adversarial_jobs_once", run_trigger, ), patch.object(producer, "ensure_repo_approved_job", AsyncMock(return_value=UUID(int=1))), patch.object(producer, "claim_next_agentic_job", AsyncMock(return_value=None)), patch.object( producer.settings, "continuous_improvement_drift_trigger_enabled", False, ), ): result = await producer.produce_queued_agentic_jobs_once() run_trigger.assert_not_awaited() self.assertEqual(result["drift_signals_scanned"], 0) self.assertEqual(result["drift_jobs_enqueued"], 0) self.assertEqual(result["drift_trigger_failed"], 0) async def test_trigger_failure_is_reported_without_claiming_it_succeeded(self) -> None: with ( patch.object( trigger, "enqueue_drift_adversarial_jobs_once", AsyncMock(side_effect=RuntimeError("drift ledger unavailable")), ) as run_trigger, patch.object( producer, "ensure_repo_approved_job", AsyncMock(return_value=UUID(int=1)), ), patch.object( producer, "claim_next_agentic_job", AsyncMock(return_value=None), ), patch.object( producer.settings, "continuous_improvement_drift_trigger_enabled", True, ), ): result = await producer.produce_queued_agentic_jobs_once() run_trigger.assert_awaited_once() self.assertEqual(result["drift_signals_scanned"], 0) self.assertEqual(result["drift_jobs_enqueued"], 0) self.assertEqual(result["drift_trigger_failed"], 1) if __name__ == "__main__": import unittest unittest.main()