diff --git a/src/political_event_tracking_research/feed_primitives.py b/src/political_event_tracking_research/feed_primitives.py new file mode 100644 index 0000000..5a31a61 --- /dev/null +++ b/src/political_event_tracking_research/feed_primitives.py @@ -0,0 +1,334 @@ +"""Validated producer rows and canonical per-feed status contract.""" +from __future__ import annotations + +import hashlib +import json +import re +from collections.abc import Iterable, Mapping +from dataclasses import dataclass +from typing import Any + +STATUS_VERSION = "pert.feed_primitives.v1" +MAX_ROWS_PER_FEED = 10_000 +MAX_SAFE_JSON_INTEGER = 2**53 - 1 +_ROW_KEYS = frozenset({"item_id", "published_at", "source_type", "source_url", "author", "text"}) +_FEED_KEYS = frozenset({"feed_id", "feed_url", "kind", "state", "rows", "error_code"}) +_FEED_WIRE_KEYS = frozenset( + {"feed_id", "feed_url", "kind", "state", "accepted_row_count", "rejected_row_count", "row_digest", "error_code"} +) +_STATUS_KEYS = frozenset( + { + "status_version", + "configured_feed_count", + "feed_count", + "successful_feed_count", + "failed_feed_count", + "quarantined_feed_count", + "accepted_row_count", + "rejected_row_count", + "publication_complete", + "eligible_for_live_publication", + "aggregate_row_digest", + "feeds", + } +) +_STATES = frozenset({"accepted", "failed", "quarantined"}) +_KINDS = frozenset({"rss", "atom", "unknown"}) +_SAFE_ERROR = re.compile(r"^[a-z][a-z0-9_]*$") + + +class PrimitiveStatusError(ValueError): + def __init__(self, code: str) -> None: + self.code = code + super().__init__(code) + + +def _fail(code: str) -> PrimitiveStatusError: + return PrimitiveStatusError(code) + + +def _canonical_bytes(value: object) -> bytes: + try: + return json.dumps( + value, ensure_ascii=False, sort_keys=True, separators=(",", ":"), allow_nan=False + ).encode("utf-8") + except (TypeError, ValueError, UnicodeError, OverflowError, RecursionError): + raise _fail("status_serialization_invalid") from None + + +def _digest(value: object) -> str: + return hashlib.sha256(_canonical_bytes(value)).hexdigest() + + +def _string(value: object, code: str, *, allow_empty: bool = True) -> str: + if type(value) is not str or (not allow_empty and not value): + raise _fail(code) + return value + + +def _safe_int(value: object, code: str) -> int: + if type(value) is not int or value < 0 or value > MAX_SAFE_JSON_INTEGER: + raise _fail(code) + return value + + +@dataclass(frozen=True, slots=True) +class PrimitiveRow: + item_id: str + published_at: str + source_type: str + source_url: str + author: str + text: str + + @classmethod + def from_mapping(cls, value: object) -> "PrimitiveRow": + if isinstance(value, cls): + value = value.to_mapping() + if not isinstance(value, Mapping) or set(value) != _ROW_KEYS: + raise _fail("row_shape_invalid") + try: + snapshot = dict(value) + except (TypeError, ValueError, RuntimeError): + raise _fail("row_shape_invalid") from None + return cls( + item_id=_string(snapshot["item_id"], "row_invalid", allow_empty=False), + published_at=_string(snapshot["published_at"], "row_invalid", allow_empty=False), + source_type=_string(snapshot["source_type"], "row_invalid", allow_empty=False), + source_url=_string(snapshot["source_url"], "row_invalid", allow_empty=False), + author=_string(snapshot["author"], "row_invalid"), + text=_string(snapshot["text"], "row_invalid"), + ) + + def to_mapping(self) -> dict[str, str]: + return { + "item_id": self.item_id, + "published_at": self.published_at, + "source_type": self.source_type, + "source_url": self.source_url, + "author": self.author, + "text": self.text, + } + + +def _snapshot_rows(value: object) -> tuple[PrimitiveRow, ...]: + if isinstance(value, (str, bytes, Mapping)) or not isinstance(value, Iterable): + raise _fail("rows_shape_invalid") + rows: list[PrimitiveRow] = [] + try: + for item in value: + if len(rows) >= MAX_ROWS_PER_FEED: + raise _fail("rows_limit_exceeded") + rows.append(PrimitiveRow.from_mapping(item)) + except PrimitiveStatusError: + raise + except (TypeError, ValueError, RuntimeError, RecursionError): + raise _fail("rows_invalid") from None + return tuple(rows) + + +def _snapshot_feed(value: object) -> dict[str, Any]: + if not isinstance(value, Mapping) or set(value) != _FEED_KEYS: + raise _fail("feed_shape_invalid") + try: + snapshot = dict(value) + except (TypeError, ValueError, RuntimeError): + raise _fail("feed_shape_invalid") from None + feed_id = _string(snapshot["feed_id"], "feed_invalid", allow_empty=False) + feed_url = _string(snapshot["feed_url"], "feed_invalid", allow_empty=False) + kind = _string(snapshot["kind"], "feed_invalid", allow_empty=False) + state = _string(snapshot["state"], "feed_invalid", allow_empty=False) + error_code = snapshot["error_code"] + if kind not in _KINDS or state not in _STATES or ( + error_code is not None and (type(error_code) is not str or not _SAFE_ERROR.fullmatch(error_code)) + ): + raise _fail("feed_invalid") + rows = _snapshot_rows(snapshot["rows"]) + if state == "accepted" and (not rows or error_code is not None): + raise _fail("feed_state_invalid") + if state == "failed" and (rows or not error_code): + raise _fail("feed_state_invalid") + if state == "quarantined" and not error_code: + raise _fail("feed_state_invalid") + return { + "feed_id": feed_id, + "feed_url": feed_url, + "kind": kind, + "state": state, + "rows": rows, + "error_code": error_code, + } + + +def _feed_wire(feed: Mapping[str, Any]) -> dict[str, object]: + rows = [row.to_mapping() for row in feed["rows"]] + return { + "feed_id": feed["feed_id"], + "feed_url": feed["feed_url"], + "kind": feed["kind"], + "state": feed["state"], + "accepted_row_count": len(rows), + "rejected_row_count": 0, + "row_digest": _digest(rows), + "error_code": feed["error_code"], + } + + +def build_status(feed_records: Iterable[Mapping[str, object]]) -> dict[str, object]: + if isinstance(feed_records, (str, bytes, Mapping)): + raise _fail("feed_records_invalid") + try: + records = [_snapshot_feed(item) for item in feed_records] + except PrimitiveStatusError: + raise + except (TypeError, ValueError, RuntimeError, RecursionError): + raise _fail("feed_records_invalid") from None + if not records: + raise _fail("configured_feed_empty") + if len({item["feed_id"] for item in records}) != len(records): + raise _fail("feed_duplicate") + if len(records) > MAX_SAFE_JSON_INTEGER: + raise _fail("feed_count_overflow") + records.sort(key=lambda item: (item["feed_id"], item["feed_url"])) + feeds = [_feed_wire(item) for item in records] + accepted = sum(item["state"] == "accepted" for item in records) + failed = sum(item["state"] == "failed" for item in records) + quarantined = sum(item["state"] == "quarantined" for item in records) + rows = [row.to_mapping() for item in records for row in item["rows"]] + complete = accepted == len(records) and failed == 0 and quarantined == 0 and bool(rows) + return { + "status_version": STATUS_VERSION, + "configured_feed_count": len(records), + "feed_count": len(records), + "successful_feed_count": accepted, + "failed_feed_count": failed, + "quarantined_feed_count": quarantined, + "accepted_row_count": len(rows), + "rejected_row_count": 0, + "publication_complete": complete, + "eligible_for_live_publication": complete, + "aggregate_row_digest": _digest(rows), + "feeds": feeds, + } + + +def _validate_wire(value: object) -> dict[str, object]: + if not isinstance(value, Mapping) or set(value) != _STATUS_KEYS: + raise _fail("status_shape_invalid") + try: + snapshot = dict(value) + except (TypeError, ValueError, RuntimeError): + raise _fail("status_shape_invalid") from None + if snapshot["status_version"] != STATUS_VERSION or type(snapshot["feeds"]) is not list: + raise _fail("status_shape_invalid") + integer_keys = ( + "configured_feed_count", + "feed_count", + "successful_feed_count", + "failed_feed_count", + "quarantined_feed_count", + "accepted_row_count", + "rejected_row_count", + ) + if any(_safe_int(snapshot[key], "status_counter_invalid") != snapshot[key] for key in integer_keys): + raise _fail("status_counter_invalid") + if any(type(snapshot[key]) is not bool for key in ("publication_complete", "eligible_for_live_publication")): + raise _fail("status_counter_invalid") + if snapshot["publication_complete"] != snapshot["eligible_for_live_publication"] or not _is_digest( + snapshot["aggregate_row_digest"] + ): + raise _fail("status_integrity_invalid") + if snapshot["configured_feed_count"] != snapshot["feed_count"] or snapshot["feed_count"] != len(snapshot["feeds"]): + raise _fail("status_counter_mismatch") + expected_counts = { + "successful_feed_count": 0, + "failed_feed_count": 0, + "quarantined_feed_count": 0, + "accepted_row_count": 0, + "rejected_row_count": 0, + } + feed_ids: set[str] = set() + feed_order: list[tuple[str, str]] = [] + for feed in snapshot["feeds"]: + if not isinstance(feed, Mapping) or set(feed) != _FEED_WIRE_KEYS: + raise _fail("feed_wire_shape_invalid") + if ( + any(type(feed[key]) is not str or not feed[key] for key in ("feed_id", "feed_url", "kind", "state")) + or feed["kind"] not in _KINDS + or feed["state"] not in _STATES + ): + raise _fail("feed_wire_shape_invalid") + if feed["feed_id"] in feed_ids: + raise _fail("feed_duplicate") + feed_ids.add(feed["feed_id"]) + feed_order.append((feed["feed_id"], feed["feed_url"])) + if ( + _safe_int(feed["accepted_row_count"], "feed_counter_invalid") != feed["accepted_row_count"] + or _safe_int(feed["rejected_row_count"], "feed_counter_invalid") != feed["rejected_row_count"] + ): + raise _fail("feed_counter_invalid") + if not _is_digest(feed["row_digest"]): + raise _fail("feed_digest_invalid") + state = feed["state"] + error = feed["error_code"] + if state not in _STATES or (error is not None and (type(error) is not str or not _SAFE_ERROR.fullmatch(error))): + raise _fail("feed_state_invalid") + if state == "accepted" and ( + feed["accepted_row_count"] <= 0 or feed["rejected_row_count"] != 0 or error is not None + ): + raise _fail("feed_state_invalid") + if state == "failed" and (feed["accepted_row_count"] != 0 or feed["rejected_row_count"] != 0 or not error): + raise _fail("feed_state_invalid") + if state == "quarantined" and not error: + raise _fail("feed_state_invalid") + expected_counts[f"{state if state != 'accepted' else 'successful'}_feed_count"] += 1 + expected_counts["accepted_row_count"] += feed["accepted_row_count"] + expected_counts["rejected_row_count"] += feed["rejected_row_count"] + if feed_order != sorted(feed_order): + raise _fail("feed_order_invalid") + if any(snapshot[key] != value for key, value in expected_counts.items()): + raise _fail("status_counter_mismatch") + complete = ( + snapshot["successful_feed_count"] == snapshot["feed_count"] + and snapshot["feed_count"] > 0 + and snapshot["accepted_row_count"] > 0 + ) + if snapshot["publication_complete"] != complete: + raise _fail("status_counter_mismatch") + return snapshot + + +def _is_digest(value: object) -> bool: + return type(value) is str and len(value) == 64 and all(char in "0123456789abcdef" for char in value) + + +def serialize_status(value: Mapping[str, object]) -> bytes: + return _canonical_bytes(_validate_wire(value)) + + +def parse_status_bytes(wire: bytes) -> dict[str, object]: + if type(wire) is not bytes: + raise _fail("status_wire_invalid") + def pairs(items: list[tuple[str, object]]) -> dict[str, object]: + result: dict[str, object] = {} + for key, item in items: + if key in result: + raise _fail("status_duplicate_key") + result[key] = item + return result + try: + value = json.loads(wire.decode("utf-8"), object_pairs_hook=pairs) + except (UnicodeError, json.JSONDecodeError, TypeError, ValueError, RecursionError): + raise _fail("status_wire_invalid") from None + parsed = _validate_wire(value) + if serialize_status(parsed) != wire: + raise _fail("status_noncanonical") + return parsed + + +def status_for_rows(wire: bytes, feed_records: Iterable[Mapping[str, object]]) -> dict[str, object]: + parse_status_bytes(wire) + expected = build_status(feed_records) + if serialize_status(expected) != wire: + raise _fail("status_integrity_mismatch") + return expected diff --git a/src/political_event_tracking_research/rss_source_fetch.py b/src/political_event_tracking_research/rss_source_fetch.py index 0884243..c45b03f 100644 --- a/src/political_event_tracking_research/rss_source_fetch.py +++ b/src/political_event_tracking_research/rss_source_fetch.py @@ -5,7 +5,6 @@ import email.utils import hashlib import html -import json import re import urllib.request from collections.abc import Callable @@ -16,6 +15,7 @@ from defusedxml.common import DefusedXmlException from .csv_utils import read_csv_rows, write_csv_rows +from .feed_primitives import PrimitiveStatusError, PrimitiveRow, build_status, serialize_status USER_AGENT = ( @@ -33,24 +33,6 @@ class FeedConfig: author: str -@dataclass(frozen=True) -class FeedFetchStatus: - feed_id: str - feed_url: str - ok: bool - item_count: int - error: str = "" - - def to_json(self) -> dict[str, object]: - return { - "feed_id": self.feed_id, - "feed_url": self.feed_url, - "ok": self.ok, - "item_count": self.item_count, - "error": self.error, - } - - class FeedXmlError(ValueError): """Sanitized producer-boundary XML failure.""" @@ -131,7 +113,9 @@ def stable_item_id(feed_id: str, link: str, title: str) -> str: return f"{feed_id}-{digest}" -def parse_feed_items(feed_bytes: bytes, feed: FeedConfig, *, max_items: int = 25) -> list[dict[str, str]]: +def parse_feed_items( + feed_bytes: bytes, feed: FeedConfig, *, max_items: int = 25, include_kind: bool = False +) -> list[dict[str, str]] | tuple[str, list[dict[str, str]]]: if type(feed_bytes) is not bytes: raise FeedXmlError("feed_xml_invalid") if len(feed_bytes) > MAX_XML_BYTES: @@ -160,7 +144,7 @@ def parse_feed_items(feed_bytes: bytes, feed: FeedConfig, *, max_items: int = 25 "text": text, } ) - return rows + return ("rss", rows) if include_kind else rows atom_entries = root.findall("{http://www.w3.org/2005/Atom}entry") for entry in atom_entries[:max_items]: @@ -182,25 +166,14 @@ def parse_feed_items(feed_bytes: bytes, feed: FeedConfig, *, max_items: int = 25 "text": text, } ) - return rows - + return ("atom", rows) if include_kind else rows -def utc_now_iso() -> str: - return dt.datetime.now(dt.UTC).replace(microsecond=0).isoformat().replace("+00:00", "Z") - -def write_fetch_status(path: str | Path, statuses: list[FeedFetchStatus], *, item_count: int) -> None: - payload = { - "generated_at": utc_now_iso(), - "feed_count": len(statuses), - "successful_feed_count": sum(1 for item in statuses if item.ok), - "failed_feed_count": sum(1 for item in statuses if not item.ok), - "item_count": item_count, - "feeds": [item.to_json() for item in statuses], - } +def write_fetch_status(path: str | Path, feed_records: list[dict[str, object]]) -> None: + payload = serialize_status(build_status(feed_records)) output_path = Path(path) output_path.parent.mkdir(parents=True, exist_ok=True) - output_path.write_text(json.dumps(payload, ensure_ascii=False, indent=2, sort_keys=True) + "\n", encoding="utf-8") + output_path.write_bytes(payload) def fetch_rss_sources( @@ -213,43 +186,54 @@ def fetch_rss_sources( fetcher: Callable[[str], bytes] = fetch_url, ) -> list[dict[str, str]]: rows: list[dict[str, str]] = [] - statuses: list[FeedFetchStatus] = [] + feed_records: list[dict[str, object]] = [] for feed in load_feed_config(feeds_path): try: - feed_rows = parse_feed_items(fetcher(feed.feed_url), feed, max_items=max_items_per_feed) + kind, parsed_rows = parse_feed_items( + fetcher(feed.feed_url), feed, max_items=max_items_per_feed, include_kind=True + ) + feed_rows = [_row_mapping(PrimitiveRow.from_mapping(row)) for row in parsed_rows] + feed_records.append( + { + "feed_id": feed.feed_id, + "feed_url": feed.feed_url, + "kind": kind, + "state": "accepted" if feed_rows else "quarantined", + "rows": feed_rows, + "error_code": None if feed_rows else "zero_entries", + } + ) except Exception as exc: - statuses.append( - FeedFetchStatus( - feed_id=feed.feed_id, - feed_url=feed.feed_url, - ok=False, - item_count=0, - error=f"{type(exc).__name__}: {exc}", - ) + error_code = exc.code if isinstance(exc, (FeedXmlError, PrimitiveStatusError)) else "fetch_failed" + feed_records.append( + { + "feed_id": feed.feed_id, + "feed_url": feed.feed_url, + "kind": "unknown", + "state": "failed", + "rows": [], + "error_code": error_code, + } ) if not continue_on_feed_error: raise continue - rows.extend(feed_rows) - statuses.append( - FeedFetchStatus( - feed_id=feed.feed_id, - feed_url=feed.feed_url, - ok=True, - item_count=len(feed_rows), - ) - ) - if statuses and not any(item.ok for item in statuses): + if feed_records and not any(item["state"] == "accepted" for item in feed_records): if status_output: - write_fetch_status(status_output, statuses, item_count=0) + write_fetch_status(status_output, feed_records) raise RuntimeError("all configured RSS/Atom feeds failed") + rows = [row for record in feed_records for row in record["rows"]] rows.sort(key=lambda row: (row["published_at"], row["item_id"])) write_csv_rows(output_path, ["item_id", "published_at", "source_type", "source_url", "author", "text"], rows) if status_output: - write_fetch_status(status_output, statuses, item_count=len(rows)) + write_fetch_status(status_output, feed_records) return rows +def _row_mapping(row: PrimitiveRow) -> dict[str, str]: + return row.to_mapping() + + def build_arg_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(description="Fetch RSS/Atom feeds into source_items CSV schema.") parser.add_argument("--feeds", required=True, help="Feed config CSV.") diff --git a/tests/test_feed_primitives.py b/tests/test_feed_primitives.py new file mode 100644 index 0000000..e77b47f --- /dev/null +++ b/tests/test_feed_primitives.py @@ -0,0 +1,147 @@ +from __future__ import annotations + +import json + +import pytest + +from political_event_tracking_research.feed_primitives import ( + MAX_SAFE_JSON_INTEGER, + PrimitiveStatusError, + build_status, + parse_status_bytes, + serialize_status, + status_for_rows, +) + + +ROW = { + "item_id": "feed-a-1", + "published_at": "2026-05-01T12:30:00Z", + "source_type": "official_remarks", + "source_url": "https://example.test/item/1", + "author": "Example", + "text": "Policy mention", +} + + +def feed( + feed_id: str, + *, + state: str = "accepted", + rows: list[dict[str, str]] | None = None, + error_code: str | None = None, +) -> dict[str, object]: + return { + "feed_id": feed_id, + "feed_url": f"https://example.test/{feed_id}", + "kind": "rss", + "state": state, + "rows": rows if rows is not None else [ROW], + "error_code": error_code, + } + + +def test_status_binds_rows_and_is_canonical() -> None: + status = build_status([feed("b"), feed("a")]) + wire = serialize_status(status) + assert parse_status_bytes(wire) == status + assert status_for_rows(wire, [feed("a"), feed("b")]) == status + assert status["accepted_row_count"] == 2 + assert status["publication_complete"] is True + + +def test_feed_input_order_is_not_semantic() -> None: + left = build_status([feed("a"), feed("b")]) + right = build_status([feed("b"), feed("a")]) + assert left == right + assert serialize_status(left) == serialize_status(right) + + +def test_row_order_is_producer_semantic() -> None: + other = {**ROW, "item_id": "feed-a-2"} + assert build_status([feed("a", rows=[ROW, other])]) != build_status([feed("a", rows=[other, ROW])]) + + +@pytest.mark.parametrize( + "record", + [ + feed("a", state="failed", rows=[], error_code=None), + feed("a", state="failed", error_code="network"), + feed("a", state="quarantined", rows=[], error_code=None), + feed("a", state="accepted", rows=[], error_code=None), + feed("a", state="accepted", error_code="unexpected"), + feed("a", state="stale", error_code="old"), + ], +) +def test_state_contract_is_closed(record: dict[str, object]) -> None: + with pytest.raises(PrimitiveStatusError): + build_status([record]) + + +def test_failed_and_quarantined_are_not_eligible() -> None: + status = build_status([feed("good"), feed("bad", state="failed", rows=[], error_code="network")]) + assert status["publication_complete"] is False + assert status["eligible_for_live_publication"] is False + assert status["failed_feed_count"] == 1 + + empty = build_status([feed("empty", state="quarantined", rows=[], error_code="zero_entries")]) + assert empty["accepted_row_count"] == 0 + assert empty["publication_complete"] is False + + +def test_status_digest_mismatch_is_rejected() -> None: + status = build_status([feed("a")]) + wire = json.loads(serialize_status(status)) + wire["feeds"][0]["row_digest"] = "0" * 64 + with pytest.raises(PrimitiveStatusError, match="status_integrity_mismatch"): + status_for_rows(serialize_status(wire), [feed("a")]) + + +def test_malformed_mapping_generator_and_stateful_rows_are_sanitized() -> None: + with pytest.raises(PrimitiveStatusError, match="feed_shape_invalid"): + build_status([{"feed_id": "a"}]) + + def records(): + yield feed("a") + + assert build_status(records())["feed_count"] == 1 + + class DivergingRows: + def __iter__(self): + yield ROW + yield {**ROW, "item_id": "changed"} + + record = feed("a") + record["rows"] = DivergingRows() + with pytest.raises(PrimitiveStatusError, match="status_integrity_mismatch"): + status_for_rows(serialize_status(build_status([feed("a")])), [record]) + + +def test_unknown_keys_duplicate_wire_and_unsafe_types_fail_closed() -> None: + status = build_status([feed("a")]) + payload = json.loads(serialize_status(status)) + payload["unknown"] = True + with pytest.raises(PrimitiveStatusError): + serialize_status(payload) + + payload = json.loads(serialize_status(status)) + payload["feed_count"] = True + with pytest.raises(PrimitiveStatusError): + serialize_status(payload) + + duplicate = serialize_status(status).replace(b'"status_version":', b'"status_version":') + with pytest.raises(PrimitiveStatusError): + parse_status_bytes(duplicate[:-1] + b',"status_version":"other"}') + + +def test_safe_integer_bound_and_feed_order_are_strict() -> None: + status = build_status([feed("a"), feed("b")]) + payload = json.loads(serialize_status(status)) + payload["feed_count"] = MAX_SAFE_JSON_INTEGER + 1 + with pytest.raises(PrimitiveStatusError, match="status_counter_invalid"): + serialize_status(payload) + + payload = json.loads(serialize_status(status)) + payload["feeds"] = list(reversed(payload["feeds"])) + with pytest.raises(PrimitiveStatusError, match="feed_order_invalid"): + serialize_status(payload) diff --git a/tests/test_rss_source_fetch.py b/tests/test_rss_source_fetch.py index d1f324e..7fee812 100644 --- a/tests/test_rss_source_fetch.py +++ b/tests/test_rss_source_fetch.py @@ -7,6 +7,7 @@ import pytest from political_event_tracking_research import rss_source_fetch +from political_event_tracking_research.feed_primitives import parse_status_bytes from political_event_tracking_research.rss_source_fetch import FeedConfig, fetch_rss_sources, parse_feed_items @@ -169,10 +170,11 @@ def fake_fetch(url: str) -> bytes: assert len(rows) == 1 payload = json.loads(status.read_text(encoding="utf-8")) + parse_status_bytes(status.read_bytes()) assert payload["successful_feed_count"] == 1 assert payload["failed_feed_count"] == 1 - assert payload["feeds"][1]["feed_id"] == "bad" - assert "RuntimeError" in payload["feeds"][1]["error"] + failed = next(item for item in payload["feeds"] if item["feed_id"] == "bad") + assert failed["error_code"] == "fetch_failed" def test_fetch_rss_sources_fails_when_all_feeds_fail(tmp_path: Path) -> None: