from __future__ import annotations import random import unittest from datetime import date, datetime from scripts.rectification.candidate_contrast import ( MIN_DISCRIMINATOR_DOMAINS, MIN_DISCRIMINATOR_EVENTS, PROBE_PHASE_CANDIDATE_DISCRIMINATOR, SIGNATURE_LAYERS, discriminator_gate_open, distinguish_contract_errors, feature_signature, select_signature_representatives, training_scoreable_stats, ) from scripts.rectification.event_probes import ( ANSWER_PRIORS, REFRESH_MAX_PROBES, REFRESH_MAX_PROBES_PER_DOMAIN, candidate_contrast_opportunities, discriminating_event_probes, event_clarification_probes, evidence_collection_probes, _dominant_existence_prior, _partition_ranked_probes, _probe_caps, ) from scripts.rectification.refinement_packet import window_scan PLANETS = { "Sun": 12.0, "Moon": 100.0, "Mars": 40.0, "Mercury": 20.0, "Jupiter": 80.0, "Venus": 50.0, "Saturn": 200.0, "Rahu": 310.0, "Ketu": 130.0, } def _varga(asc: int, planet_sign: int) -> dict: return { "Ascendant": {"sign_idx": asc}, **{name: {"sign_idx": planet_sign} for name in PLANETS}, } def _context( time: str, *, d4_asc: int, d9_asc: int = 1, d10_asc: int = 1, d12_asc: int = 1, d24_asc: int = 1, sun_house: int = 10, sun_varga_sign: int = 9, moon: float = 100.0, ) -> dict: hour, minute = (int(part) for part in time.split(":")) planets = {**PLANETS, "Moon": moon} natal_planets = { name: {"house": sun_house if name != "Moon" else 4, "lon": lon} for name, lon in planets.items() } return { "candidate_at": datetime(1997, 8, 8, hour, minute), "chart": {"ascendant": {"lon": 10.0, "sign": "Aries"}, "planets": natal_planets}, "planet_longitudes": {name: lon for name, lon in planets.items()}, "ascendant_index": 0, "varga_charts": { "D4": _varga(d4_asc, sun_varga_sign), "D9": _varga(d9_asc, 1), "D10": _varga(d10_asc, 1), "D5": _varga(1, 1), "D24": _varga(d24_asc, 1), "D12": _varga(d12_asc, 1), "D7": _varga(1, 1), "D3": _varga(1, 1), }, "arudha_padas": {}, "feature": { "time": time, "ascendant_sign_index": 0, "varga_ascendants": { "D4": d4_asc, "D9": d9_asc, "D10": d10_asc, "D5": 1, "D24": d24_asc, "D12": d12_asc, }, }, } def _gate_events() -> list[dict]: return [ {"id": "e1", "domain": "education", "event_kind": "education_start", "date": "2014-09-01", "precision": "month"}, {"id": "e2", "domain": "education", "event_kind": "education_completion", "date": "2017-06-01", "precision": "month"}, {"id": "e3", "domain": "career", "event_kind": "career_entry", "date": "2018-07-01", "precision": "month"}, {"id": "e4", "domain": "relationship", "event_kind": "relationship_start", "date": "2021-08-01", "precision": "month"}, ] def _request(**extra: object) -> dict: return { "birth_date": "1997-08-08", "events": _gate_events(), **extra, } class DiscriminatorContractTest(unittest.TestCase): def test_ci_forbids_invalid_distinguish_payloads(self) -> None: self.assertEqual( distinguish_contract_errors({ "role": "distinguish", "information_gain": 0, "candidate_ids": ["05:00", "05:20"], "expected_outcomes": [ {"answer_class": "yes", "supports": ["05:00"], "conflicts": ["05:20"]}, {"answer_class": "no", "supports": ["05:20"], "conflicts": ["05:00"]}, ], }), ["distinguish_non_positive_information_gain"], ) self.assertEqual( distinguish_contract_errors({ "role": "distinguish", "information_gain": 0.4, "candidate_ids": [], "expected_outcomes": [ {"answer_class": "yes", "supports": [], "conflicts": []}, {"answer_class": "no", "supports": [], "conflicts": []}, ], }), ["distinguish_empty_candidate_ids"], ) self.assertEqual( distinguish_contract_errors({ "role": "distinguish", "information_gain": 0.4, "candidate_ids": ["05:00", "05:20"], "expected_outcomes": [], }), ["distinguish_empty_expected_outcomes"], ) def test_quality_clarify_stays_and_anchored_distinguish_is_gated(self) -> None: built = { "static_contexts": [ _context("05:13", d4_asc=1, d9_asc=1), _context("05:40", d4_asc=2, d9_asc=4), ] } events = _gate_events() + [{ "id": "exam", "domain": "education", "event_kind": "education_milestone", "summary": "入学考试", "date": "2015-06-01", "precision": "year", }] request = _request(events=events) probes = discriminating_event_probes( request, built, scan=window_scan(built), candidate_times=["05:13", "05:40"], representative_time="05:13", today=date(2026, 8, 22), ) self.assertFalse(any(item.get("role") == "distinguish" and distinguish_contract_errors(item) for item in probes)) quality = [item for item in probes if item.get("source") == "known_event_quality"] for item in quality: self.assertEqual(item.get("role"), "distinguish") self.assertTrue(item.get("target_evidence_id")) self.assertEqual(item.get("choice_kind"), "event_quality") clarification = event_clarification_probes(request) self.assertTrue(any(item.get("source") == "known_event_quality" for item in clarification)) self.assertTrue(all(item.get("phase") == "event_clarification" for item in clarification)) self.assertFalse(any(item.get("role") == "distinguish" for item in clarification)) def test_gate_blocks_discriminator_until_three_events_two_domains(self) -> None: built = { "static_contexts": [ _context("05:13", d4_asc=1), _context("05:40", d4_asc=2), ] } too_few = discriminating_event_probes( {"birth_date": "1997-08-08", "events": _gate_events()[:2]}, built, scan=window_scan(built), candidate_times=["05:13", "05:40"], representative_time="05:13", today=date(2026, 8, 22), ) self.assertEqual(too_few, []) collection = evidence_collection_probes({"birth_date": "1997-08-08", "events": _gate_events()[:2]}) self.assertTrue(collection) self.assertTrue(all(item.get("phase") == "evidence_collection" for item in collection)) self.assertGreaterEqual(MIN_DISCRIMINATOR_EVENTS, 3) self.assertGreaterEqual(MIN_DISCRIMINATOR_DOMAINS, 2) def test_training_gate_opens_at_three_events_when_holdout_is_not_reserved(self) -> None: two = _gate_events()[:2] three = _gate_events()[:3] four = _gate_events() # 原值: 3 件留 holdout 后门关;4 件才开 # 新值: 3 件全训练、门开;4 件才留 1 件 holdout # 原因: BUG-647 self.assertFalse(discriminator_gate_open(two)) self.assertTrue(discriminator_gate_open(three)) two_count, _, _ = training_scoreable_stats(two) three_count, three_domains, _ = training_scoreable_stats(three) four_count, four_domains, _ = training_scoreable_stats(four) self.assertLess(two_count, MIN_DISCRIMINATOR_EVENTS) self.assertGreaterEqual(three_count, MIN_DISCRIMINATOR_EVENTS) self.assertGreaterEqual(four_count, MIN_DISCRIMINATOR_EVENTS) self.assertGreaterEqual(three_domains, MIN_DISCRIMINATOR_DOMAINS) self.assertGreaterEqual(four_domains, MIN_DISCRIMINATOR_DOMAINS) from scripts.rectification.case_holdout import holdout_event_ids, holdout_reservation_status self.assertEqual(holdout_reservation_status(three), "not_reserved_min_events") self.assertEqual(len(holdout_event_ids(three)), 0) self.assertEqual(holdout_reservation_status(four), "reserved") self.assertEqual(len(holdout_event_ids(four)), 1) built = { "static_contexts": [ _context("05:13", d4_asc=0, sun_house=4, sun_varga_sign=3), _context("05:40", d4_asc=1, sun_house=10, sun_varga_sign=9), ] } three_probes = discriminating_event_probes( {"birth_date": "1997-08-08", "events": three}, built, scan=window_scan(built), candidate_times=["05:13", "05:40"], representative_time="05:13", today=date(2026, 8, 22), ) four_probes = discriminating_event_probes( {"birth_date": "1997-08-08", "events": four}, built, scan=window_scan(built), candidate_times=["05:13", "05:40"], representative_time="05:13", today=date(2026, 8, 22), ) self.assertTrue(three_probes) self.assertTrue(four_probes) self.assertGreater(three_domains, 0) def test_signature_clusters_are_not_three_adjacent_minutes(self) -> None: rows = [ {"time": "05:13", "score": 20}, {"time": "05:14", "score": 19}, {"time": "05:15", "score": 18}, {"time": "05:40", "score": 12}, ] contexts = [ _context("05:13", d4_asc=1, d9_asc=1), _context("05:14", d4_asc=1, d9_asc=1), _context("05:15", d4_asc=1, d9_asc=1), _context("05:40", d4_asc=2, d9_asc=4), ] public = select_signature_representatives(rows, contexts) times = [row["time"] for row in public] self.assertIn("05:40", times) self.assertLessEqual(sum(1 for time in times if time in {"05:13", "05:14", "05:15"}), 1) self.assertNotEqual(feature_signature(contexts[0]), feature_signature(contexts[3])) self.assertEqual(SIGNATURE_LAYERS[:6], ("d1", "d9", "d10", "d24", "d4", "d12")) self.assertIn("md", SIGNATURE_LAYERS) def test_staging_quick_gate_runs_this_contract(self) -> None: from pathlib import Path text = Path("scripts/run_quality_gate.py").read_text(encoding="utf-8") self.assertIn('"tests/test_candidate_discriminator_contract.py"', text) def test_randomized_hidden_mutated_splits_keep_mapping_and_gain(self) -> None: rng = random.Random(20260826) built = { "static_contexts": [ _context("04:50", d4_asc=0, d9_asc=1, d10_asc=2, moon=99.0), _context("05:20", d4_asc=3, d9_asc=6, d10_asc=8, moon=101.5), ] } for _ in range(12): events = list(_gate_events()) rng.shuffle(events) for event in events: event = dict(event) event["summary"] = rng.choice(["记不清细节", "家里提过", "档案上有"]) request = _request(events=events) probes = discriminating_event_probes( request, built, scan=window_scan(built), candidate_times=["04:50", "05:20"], representative_time="04:50", today=date(2026, 8, 22), ) self.assertFalse(any( item.get("source") == "known_event_quality" and not item.get("target_evidence_id") for item in probes )) for probe in probes: self.assertEqual(distinguish_contract_errors(probe), []) self.assertGreater(float(probe["information_gain"]), 0) self.assertGreaterEqual(len(probe["candidate_ids"]), 2) self.assertGreaterEqual(len(probe["expected_outcomes"]), 2) self.assertTrue(probe["candidate_set_version"]) self.assertTrue(probe["candidate_split_hash"]) self.assertNotEqual(probe["candidate_split_hash"], f"{probe['domain']}:{probe['year']}") opportunities = candidate_contrast_opportunities( request, built, scan=window_scan(built), candidate_times=["04:50", "05:20"], representative_time="04:50", today=date(2026, 8, 22), ) for opportunity in opportunities: self.assertGreater(float(opportunity["information_gain"]), 0) self.assertGreaterEqual(len(opportunity["candidate_groups"]), 2) self.assertGreaterEqual(len(opportunity["expected_outcomes"]), 2) self.assertTrue(opportunity["domain"]) self.assertTrue(opportunity["source_features"]) def test_collection_reserves_holdout_out_of_scoring_and_probes(self) -> None: from scripts.rectification.case_holdout import holdout_domain_years, holdout_event_ids from scripts.rectification.scoring_service import score_from_matrix events = _gate_events() holdout = holdout_event_ids(events) self.assertEqual(len(holdout), 1) holdout_id = next(iter(holdout)) built = { "candidate_times": ["05:00", "05:20"], "matrix": { "e1": { "05:00": {"points": 10, "rule_ids": []}, "05:20": {"points": 1, "rule_ids": []}, }, "e2": { "05:00": {"points": 10, "rule_ids": []}, "05:20": {"points": 1, "rule_ids": []}, }, "e3": { "05:00": {"points": 100, "rule_ids": []}, "05:20": {"points": 0, "rule_ids": []}, }, "e4": { "05:00": {"points": 4, "rule_ids": []}, "05:20": {"points": 1, "rule_ids": []}, }, }, "missing_layers": [], } request = { "birth_date": "1997-08-08", "start_time": "04:50", "end_time": "05:30", "lat": 31.2, "lon": 121.5, "tz": 8.0, "events": [ { "id": event["id"], "domain": event["domain"], "event_kind": event["event_kind"], "date_start": event["date"], "date_end": event["date"], "precision": event["precision"], "summary": "dated", } for event in events ], } rows = score_from_matrix(request, built) by_time = {row["time"]: row for row in rows} training_ids = {event["id"] for event in events} - holdout expected = sum( built["matrix"][event_id]["05:00"]["points"] for event_id in training_ids if event_id in built["matrix"] ) self.assertEqual(by_time["05:00"]["score"], expected) self.assertFalse(any(item["event_id"] == holdout_id for item in by_time["05:00"]["evidence"])) self.assertIn(holdout_id, {event["id"] for event in events}) probes = discriminating_event_probes( _request(events=events), { "static_contexts": [ _context("05:13", d4_asc=1, d9_asc=1), _context("05:40", d4_asc=2, d9_asc=4), ] }, scan=window_scan({ "static_contexts": [ _context("05:13", d4_asc=1, d9_asc=1), _context("05:40", d4_asc=2, d9_asc=4), ] }), candidate_times=["05:13", "05:40"], representative_time="05:13", today=date(2026, 8, 22), ) blocked = holdout_domain_years(events) self.assertTrue(blocked) for probe in probes: self.assertNotIn(f"{probe['domain']}:{probe['year']}", blocked) def test_refresh_caps_rise_only_for_five_or_fewer_remaining(self) -> None: self.assertEqual(_probe_caps(refresh=False, remaining_count=5), (8, 3)) self.assertEqual(_probe_caps(refresh=True, remaining_count=6), (8, 3)) self.assertEqual( _probe_caps(refresh=True, remaining_count=5), (REFRESH_MAX_PROBES, REFRESH_MAX_PROBES_PER_DOMAIN), ) def test_monthly_family_dasha_boundary_is_ranked_not_dropped(self) -> None: priors = dict(ANSWER_PRIORS[("family", "existence")]) probe = { "role": "distinguish", "phase": PROBE_PHASE_CANDIDATE_DISCRIMINATOR, "source": "dasha_boundary", "domain": "family", "year": 2018, "month": 5, "choice_kind": "existence", "information_gain": 0.4, "semantic_key": "family.2018.05", "candidate_ids": ["04:50", "05:06"], "expected_outcomes": [ {"answer_class": "yes", "supports": ["04:50"], "conflicts": ["05:06"]}, {"answer_class": "no", "supports": ["05:06"], "conflicts": ["04:50"]}, ], } self.assertFalse(_dominant_existence_prior(probe, priors)) public, dropped = _partition_ranked_probes([probe]) self.assertEqual([item["semantic_key"] for item in public], ["family.2018.05"]) self.assertFalse(any(item.get("reason") == "dominant_answer_prior" for item in dropped)) yearless = {**probe, "month": None, "source": "age_band", "year": 0, "semantic_key": "family.age"} self.assertTrue(_dominant_existence_prior(yearless, priors)) def test_remaining_five_cluster_refresh_keeps_family_monthly_boundary(self) -> None: asked = [ "career.2020", "career.2023", "education.2016", "education.2020", "relationship.2023", "relocation.2015", ] built = { "static_contexts": [ _context("04:48", d4_asc=1, d9_asc=1, d10_asc=1, d12_asc=1, d24_asc=1), _context("04:53", d4_asc=1, d9_asc=2, d10_asc=1, d12_asc=1, d24_asc=1), _context("04:59", d4_asc=1, d9_asc=2, d10_asc=1, d12_asc=1, d24_asc=1), _context("05:06", d4_asc=2, d9_asc=2, d10_asc=3, d12_asc=4, d24_asc=2), _context("05:07", d4_asc=2, d9_asc=2, d10_asc=3, d12_asc=4, d24_asc=3), ] } times = ["04:48", "04:53", "04:59", "05:06", "05:07"] request = _request(asked_probe_keys=asked, refresh_probes=True) probes = discriminating_event_probes( request, built, scan=window_scan(built), candidate_times=times, representative_time="04:53", today=date(2026, 9, 11), ) self.assertIsInstance(probes, list) family_probe = { "role": "distinguish", "phase": PROBE_PHASE_CANDIDATE_DISCRIMINATOR, "source": "dasha_boundary", "domain": "family", "year": 2018, "month": 5, "choice_kind": "existence", "information_gain": 0.42, "semantic_key": "family.2018.05.dasha_boundary", "candidate_ids": times, "expected_outcomes": [ {"answer_class": "yes", "supports": ["04:48", "04:53"], "conflicts": ["05:06", "05:07"]}, {"answer_class": "no", "supports": ["05:06", "05:07"], "conflicts": ["04:48", "04:53"]}, ], } public, dropped = _partition_ranked_probes( [family_probe, *probes], max_probes=REFRESH_MAX_PROBES, ) family = [ item for item in public if item.get("domain") == "family" and item.get("source") == "dasha_boundary" and isinstance(item.get("month"), int) ] self.assertTrue(family, [item.get("semantic_key") for item in public]) self.assertTrue(all(1 <= int(item["month"]) <= 12 for item in family)) self.assertFalse(any( item.get("semantic_key") == "family.2018.05.dasha_boundary" and item.get("reason") == "dominant_answer_prior" for item in dropped )) if __name__ == "__main__": unittest.main()