"""ELECTION-01 integrity gates (contract §5). No model fitting. MODEL_CHANGE=NO. Executable negative-control primitives used by the harness. These are synthetic/stub checks that prove interface separation and rejection rules — not empirical election forecasts. """ from __future__ import annotations from dataclasses import dataclass from datetime import date, datetime, timezone from typing import Any, Iterable, Sequence def _as_date(value: date | datetime | str) -> date: if isinstance(value, datetime): return value.date() if isinstance(value, date): return value s = str(value).strip() if "T" in s: return datetime.fromisoformat(s.replace("Z", "+00:00")).date() return date.fromisoformat(s[:10]) @dataclass(frozen=True) class RecordDecision: accepted: bool reason: str record_id: str def reject_future_release( record_id: str, release_date: date | datetime | str, cutoff: date | datetime | str, ) -> RecordDecision: """§5.1 Future-release records are rejected when release_date > cutoff.""" rel = _as_date(release_date) cut = _as_date(cutoff) if rel > cut: return RecordDecision(False, "future_release", record_id) return RecordDecision(True, "ok", record_id) def reject_unavailable_revision( record_id: str, available_asof: date | datetime | str | None, cutoff: date | datetime | str, *, known_unavailable: bool = False, ) -> RecordDecision: """§5.2 Unavailable revisions are rejected. Rejects when available_asof is missing, marked unavailable, or after cutoff. """ if known_unavailable or available_asof is None: return RecordDecision(False, "unavailable_revision", record_id) avail = _as_date(available_asof) cut = _as_date(cutoff) if avail > cut: return RecordDecision(False, "unavailable_revision", record_id) return RecordDecision(True, "ok", record_id) def assert_no_train_test_cross( train_elections: Sequence[str | int], test_elections: Sequence[str | int], ) -> dict[str, Any]: """§5.4 No election crosses training/test boundaries.""" train = {str(x) for x in train_elections} test = {str(x) for x in test_elections} overlap = sorted(train & test) return { "ok": len(overlap) == 0, "overlap": overlap, "train": sorted(train), "test": sorted(test), } @dataclass(frozen=True) class PollSample: poll_id: str sample_id: str field_start: date | datetime | str field_end: date | datetime | str house: str = "" def _overlap_days(a0: date, a1: date, b0: date, b1: date) -> int: start = max(a0, b0) end = min(a1, b1) if end < start: return 0 return (end - start).days + 1 def collapse_overlapping_samples( polls: Sequence[PollSample], *, min_overlap_days: int = 1, ) -> dict[str, Any]: """§5.5 Overlapping samples are not counted as independent polls. Rule: each distinct sample_id is at most one independent unit (same sample never inflates N). Additionally, distinct sample_ids whose fieldwork windows overlap by >= min_overlap_days are merged (overlapping samples). Returns independent units (one representative poll_id per group). """ # First collapse by sample_id (hard rule) by_sample: dict[str, list[PollSample]] = {} for p in polls: by_sample.setdefault(p.sample_id, []).append(p) sample_reps: list[tuple[str, PollSample, list[PollSample]]] = [] for sample_id, members in sorted(by_sample.items()): rep = sorted(members, key=lambda m: m.poll_id)[0] sample_reps.append((sample_id, rep, list(members))) # Then union-find across sample_ids on fieldwork overlap of their spans n = len(sample_reps) parent = list(range(n)) def find(i: int) -> int: while parent[i] != i: parent[i] = parent[parent[i]] i = parent[i] return i def union(i: int, j: int) -> None: ri, rj = find(i), find(j) if ri != rj: parent[rj] = ri spans = [] for sample_id, rep, members in sample_reps: starts = [_as_date(m.field_start) for m in members] ends = [_as_date(m.field_end) for m in members] spans.append((min(starts), max(ends))) for i in range(n): for j in range(i + 1, n): a0, a1 = spans[i] b0, b1 = spans[j] if _overlap_days(a0, a1, b0, b1) >= min_overlap_days: union(i, j) clusters: dict[int, list[int]] = {} for i in range(n): clusters.setdefault(find(i), []).append(i) independent: list[str] = [] groups: list[dict[str, Any]] = [] for root, idxs in sorted(clusters.items()): all_members: list[PollSample] = [] sample_ids = [] for i in idxs: sid, rep, members = sample_reps[i] sample_ids.append(sid) all_members.extend(members) rep = sorted(all_members, key=lambda m: m.poll_id)[0] independent.append(rep.poll_id) groups.append( { "sample_ids": sorted(sample_ids), "representative_poll_id": rep.poll_id, "member_poll_ids": sorted(m.poll_id for m in all_members), "collapsed_count": len(all_members), } ) raw_count = len(polls) indep_count = len(independent) return { "ok": indep_count < raw_count or raw_count == 0, "raw_poll_count": raw_count, "independent_count": indep_count, "independent_poll_ids": independent, "groups": groups, "no_false_independence": indep_count <= raw_count, } def identical_eligible_test_cycles( cycles_a: Sequence[str | int], cycles_b: Sequence[str | int], ) -> dict[str, Any]: """§5.6 Comparisons use identical eligible test cycles.""" a = [str(x) for x in cycles_a] b = [str(x) for x in cycles_b] set_a, set_b = set(a), set(b) return { "ok": a == b, # identical ordered eligible list "equal_as_sets": set_a == set_b, "cycles_a": a, "cycles_b": b, "only_in_a": sorted(set_a - set_b), "only_in_b": sorted(set_b - set_a), } def utc_now_iso() -> str: return datetime.now(timezone.utc).replace(microsecond=0).isoformat().replace("+00:00", "Z")