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..557d208 --- /dev/null +++ b/src/political_event_tracking_research/feed_primitives.py @@ -0,0 +1,300 @@ +from __future__ import annotations + +import hashlib +import json +import re +from collections.abc import Iterable, Mapping +from dataclasses import dataclass + + +STATUS_VERSION = "pert.feed_primitives.v1" +MAX_SAFE_JSON_INTEGER = 2**53 - 1 +MAX_ROWS_PER_FEED = 10_000 +_DIGEST = re.compile(r"^[0-9a-f]{64}$") +_ERROR = re.compile(r"^[a-z][a-z0-9_]{0,63}$") +_ROW_KEYS = ("item_id", "published_at", "source_type", "source_url", "author", "text") +_RECORD_KEYS = frozenset({"feed_id", "feed_url", "kind", "state", "rows", "error_code"}) +_WIRE_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", + } +) +_FEED_KEYS = frozenset( + { + "feed_id", + "feed_url", + "kind", + "state", + "accepted_row_count", + "rejected_row_count", + "row_digest", + "error_code", + } +) + + +class PrimitiveStatusError(ValueError): + """Sanitized primitive/status contract error.""" + + def __init__(self, code: str): + super().__init__(code) + self.code = code + + +def _error(code: str) -> None: + raise PrimitiveStatusError(code) + + +def _mapping(value: object, keys: frozenset[str], code: str) -> dict[str, object]: + if not isinstance(value, Mapping): + _error(code) + try: + result = dict(value) + except (AttributeError, KeyError, OverflowError, RuntimeError, TypeError, UnicodeError, ValueError): + _error(code) + if set(result) != keys or any(type(key) is not str for key in result): + _error(code) + return result + + +def _text(value: object, code: str, *, allow_empty: bool = False) -> str: + if type(value) is not str or (not allow_empty and not value) or any(ord(char) < 0x20 for char in value): + _error(code) + return value + + +def _counter(value: object, code: str) -> int: + if type(value) is not int or value < 0 or value > MAX_SAFE_JSON_INTEGER: + _error(code) + return value + + +@dataclass(frozen=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: + data = _mapping(value, frozenset(_ROW_KEYS), "row_invalid") + values = [_text(data[key], "row_invalid", allow_empty=key == "author") for key in _ROW_KEYS] + return cls(*values) + + def to_mapping(self) -> dict[str, str]: + return {key: getattr(self, key) for key in _ROW_KEYS} + + +def _row_key(row: PrimitiveRow) -> tuple[str, str]: + return row.published_at, row.item_id + + +def _canonical(value: object, code: str) -> bytes: + try: + return json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":"), allow_nan=False).encode( + "utf-8" + ) + except (RecursionError, TypeError, UnicodeError, ValueError): + _error(code) + + +def _rows_digest(rows: list[PrimitiveRow]) -> str: + ordered = sorted(rows, key=_row_key) + return hashlib.sha256(_canonical([row.to_mapping() for row in ordered], "row_digest_invalid")).hexdigest() + + +def _record(value: object) -> tuple[dict[str, object], list[PrimitiveRow]]: + data = _mapping(value, _RECORD_KEYS, "feed_invalid") + feed_id = _text(data["feed_id"], "feed_invalid") + feed_url = _text(data["feed_url"], "feed_invalid") + kind = data["kind"] + if type(kind) is not str or kind not in {"rss2", "atom", "unknown"}: + _error("feed_kind_invalid") + state = data["state"] + if type(state) is not str or state not in {"accepted", "failed", "quarantined"}: + _error("feed_state_invalid") + rows_value = data["rows"] + if not isinstance(rows_value, (list, tuple)) or len(rows_value) > MAX_ROWS_PER_FEED: + _error("feed_rows_invalid") + rows = [PrimitiveRow.from_mapping(row) for row in rows_value] + error_code = data["error_code"] + if error_code is not None and (type(error_code) is not str or not _ERROR.fullmatch(error_code)): + _error("feed_error_invalid") + if state == "accepted" and (kind not in {"rss2", "atom"} or not rows or error_code is not None): + _error("feed_state_invalid") + if state == "quarantined" and (kind not in {"rss2", "atom"} or rows or error_code is None): + _error("feed_state_invalid") + if state == "failed" and (rows or error_code is None): + _error("feed_state_invalid") + return {"feed_id": feed_id, "feed_url": feed_url, "kind": kind, "state": state, "error_code": error_code}, rows + + +def _wire_record(data: dict[str, object], rows: list[PrimitiveRow]) -> dict[str, object]: + accepted = data["state"] == "accepted" + return { + "feed_id": data["feed_id"], + "feed_url": data["feed_url"], + "kind": data["kind"], + "state": data["state"], + "accepted_row_count": len(rows) if accepted else 0, + "rejected_row_count": 0, + "row_digest": _rows_digest(rows) if accepted else hashlib.sha256(b"[]").hexdigest(), + "error_code": data["error_code"], + } + + +def build_status(feed_records: Iterable[Mapping[str, object]]) -> dict[str, object]: + try: + records = list(feed_records) + except (AttributeError, RuntimeError, TypeError, ValueError): + _error("status_invalid") + parsed = [_record(record) for record in records] + if len({data["feed_id"] for data, _ in parsed}) != len(parsed): + _error("feed_duplicate") + parsed.sort(key=lambda pair: (pair[0]["feed_id"], pair[0]["feed_url"])) + feeds = [_wire_record(data, rows) for data, rows in parsed] + accepted_rows = sorted( + [row for data, rows in parsed if data["state"] == "accepted" for row in rows], key=_row_key + ) + failed = sum(data["state"] == "failed" for data, _ in parsed) + quarantined = sum(data["state"] == "quarantined" for data, _ in parsed) + complete = bool(parsed) and failed == 0 and quarantined == 0 + return { + "status_version": STATUS_VERSION, + "configured_feed_count": len(parsed), + "feed_count": len(parsed), + "successful_feed_count": len(parsed) - failed - quarantined, + "failed_feed_count": failed, + "quarantined_feed_count": quarantined, + "accepted_row_count": len(accepted_rows), + "rejected_row_count": 0, + "publication_complete": complete, + "eligible_for_live_publication": complete, + "aggregate_row_digest": _rows_digest(accepted_rows), + "feeds": feeds, + } + + +def _validate_wire(value: object) -> dict[str, object]: + data = _mapping(value, _WIRE_KEYS, "status_invalid") + if data["status_version"] != STATUS_VERSION: + _error("status_version_invalid") + counter_keys = _WIRE_KEYS - { + "status_version", + "publication_complete", + "eligible_for_live_publication", + "aggregate_row_digest", + "feeds", + } + for key in counter_keys: + _counter(data[key], "status_counter_invalid") + for key in ("publication_complete", "eligible_for_live_publication"): + if type(data[key]) is not bool: + _error("status_flag_invalid") + aggregate = _text(data["aggregate_row_digest"], "status_digest_invalid") + if not _DIGEST.fullmatch(aggregate): + _error("status_digest_invalid") + feeds = data["feeds"] + if not isinstance(feeds, list): + _error("feed_invalid") + previous: tuple[str, str] | None = None + feed_ids: set[str] = set() + for value in feeds: + item = _mapping(value, _FEED_KEYS, "feed_invalid") + feed_id = _text(item["feed_id"], "feed_invalid") + feed_url = _text(item["feed_url"], "feed_invalid") + key = (feed_id, feed_url) + if feed_id in feed_ids: + _error("feed_duplicate") + if previous is not None and key <= previous: + _error("feed_order_invalid") + feed_ids.add(feed_id) + previous = key + kind = item["kind"] + state = item["state"] + if type(kind) is not str or kind not in {"rss2", "atom", "unknown"}: + _error("feed_kind_invalid") + if type(state) is not str or state not in {"accepted", "failed", "quarantined"}: + _error("feed_state_invalid") + if state in {"accepted", "quarantined"} and kind not in {"rss2", "atom"}: + _error("feed_kind_invalid") + if state == "failed" and item["accepted_row_count"] != 0: + _error("feed_state_invalid") + accepted_count = _counter(item["accepted_row_count"], "feed_counter_invalid") + _counter(item["rejected_row_count"], "feed_counter_invalid") + if state == "accepted" and accepted_count == 0: + _error("feed_state_invalid") + if state != "accepted" and accepted_count != 0: + _error("feed_state_invalid") + digest = _text(item["row_digest"], "status_digest_invalid") + if not _DIGEST.fullmatch(digest): + _error("status_digest_invalid") + error_code = item["error_code"] + if state == "accepted" and error_code is not None: + _error("feed_state_invalid") + if error_code is None and state != "accepted": + _error("feed_error_invalid") + if error_code is not None and (type(error_code) is not str or not _ERROR.fullmatch(error_code)): + _error("feed_error_invalid") + if data["configured_feed_count"] != len(feeds) or data["feed_count"] != len(feeds): + _error("status_counter_invalid") + accepted = sum(item["state"] == "accepted" for item in feeds) + failed = sum(item["state"] == "failed" for item in feeds) + quarantined = sum(item["state"] == "quarantined" for item in feeds) + rows = sum(item["accepted_row_count"] for item in feeds) + if data["successful_feed_count"] != accepted or data["failed_feed_count"] != failed: + _error("status_counter_invalid") + if data["quarantined_feed_count"] != quarantined or data["accepted_row_count"] != rows: + _error("status_counter_invalid") + complete = bool(feeds) and failed == 0 and quarantined == 0 + if data["publication_complete"] != complete or data["eligible_for_live_publication"] != complete: + _error("status_flag_invalid") + return data + + +def serialize_status(status: Mapping[str, object]) -> bytes: + return _canonical(_validate_wire(status), "status_invalid") + + +def _reject_duplicates(pairs: list[tuple[str, object]]) -> dict[str, object]: + result: dict[str, object] = {} + for key, value in pairs: + if key in result: + _error("status_duplicate_key") + result[key] = value + return result + + +def parse_status_bytes(payload: bytes) -> dict[str, object]: + if type(payload) is not bytes: + _error("status_bytes_invalid") + try: + value = json.loads(payload.decode("utf-8"), object_pairs_hook=_reject_duplicates) + except (UnicodeError, json.JSONDecodeError, RecursionError): + _error("status_bytes_invalid") + data = _validate_wire(value) + if serialize_status(data) != payload: + _error("status_noncanonical") + return data + + +def status_for_rows(payload: bytes, feed_records: Iterable[Mapping[str, object]]) -> dict[str, object]: + actual = parse_status_bytes(payload) + expected = build_status(feed_records) + if actual != expected: + _error("status_integrity_mismatch") + return actual diff --git a/src/political_event_tracking_research/rss_source_fetch.py b/src/political_event_tracking_research/rss_source_fetch.py index 0884243..964457e 100644 --- a/src/political_event_tracking_research/rss_source_fetch.py +++ b/src/political_event_tracking_research/rss_source_fetch.py @@ -5,17 +5,18 @@ import email.utils import hashlib import html -import json import re import urllib.request from collections.abc import Callable from dataclasses import dataclass from pathlib import Path +from types import MappingProxyType import defusedxml.ElementTree as ET from defusedxml.common import DefusedXmlException from .csv_utils import read_csv_rows, write_csv_rows +from .feed_primitives import PrimitiveRow, build_status, serialize_status USER_AGENT = ( @@ -55,6 +56,12 @@ class FeedXmlError(ValueError): """Sanitized producer-boundary XML failure.""" +@dataclass(frozen=True) +class ParsedFeed: + feed_kind: str + entries: tuple[MappingProxyType, ...] + + def load_feed_config(path: str | Path) -> list[FeedConfig]: feeds: list[FeedConfig] = [] for row in read_csv_rows(path): @@ -131,7 +138,21 @@ 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]]: +_RSS_FOREIGN_METADATA = frozenset( + { + "{http://www.w3.org/2005/Atom}link", + "{http://purl.org/dc/elements/1.1/}date", + "{http://purl.org/dc/elements/1.1/}creator", + "{http://purl.org/dc/elements/1.1/}subject", + "{http://purl.org/rss/1.0/modules/content/}encoded", + "{http://purl.org/rss/1.0/modules/syndication/}updatePeriod", + "{http://purl.org/rss/1.0/modules/syndication/}updateFrequency", + "{http://purl.org/rss/1.0/modules/syndication/}updateBase", + } +) + + +def parse_feed_snapshot(feed_bytes: bytes, feed: FeedConfig, *, max_items: int = 25) -> ParsedFeed: if type(feed_bytes) is not bytes: raise FeedXmlError("feed_xml_invalid") if len(feed_bytes) > MAX_XML_BYTES: @@ -140,67 +161,112 @@ def parse_feed_items(feed_bytes: bytes, feed: FeedConfig, *, max_items: int = 25 root = ET.fromstring(feed_bytes, forbid_dtd=True, forbid_entities=True, forbid_external=True) except (DefusedXmlException, ET.ParseError, LookupError, UnicodeError, ValueError, RecursionError): raise FeedXmlError("feed_xml_invalid") from None - rows: list[dict[str, str]] = [] + def rows_for(items: list[ET.Element], kind: str) -> ParsedFeed: + try: + rows: list[dict[str, str]] = [] + for item in items[:max_items]: + if kind == "rss2": + title = child_text(item, ("title",)) + link = rss_item_link(item) + published = child_text(item, ("pubDate", "{http://purl.org/dc/elements/1.1/}date")) + description = child_text(item, ("description", "{http://purl.org/rss/1.0/modules/content/}encoded")) + text = " ".join(part for part in (title, strip_html(description)) if part) + else: + title = child_text(item, ("{http://www.w3.org/2005/Atom}title", "title")) + link = atom_entry_link(item) + published = child_text( + item, + ( + "{http://www.w3.org/2005/Atom}published", + "{http://www.w3.org/2005/Atom}updated", + "published", + "updated", + ), + ) + summary = child_text( + item, + ("{http://www.w3.org/2005/Atom}summary", "{http://www.w3.org/2005/Atom}content"), + ) + text = " ".join(part for part in (title, strip_html(summary)) if part) + rows.append( + { + "item_id": stable_item_id(feed.feed_id, link, title), + "published_at": parse_datetime(published), + "source_type": feed.source_type, + "source_url": link, + "author": feed.author, + "text": text, + } + ) + return ParsedFeed(kind, tuple(MappingProxyType(row) for row in rows)) + except (TypeError, ValueError, UnicodeError, RecursionError): + raise FeedXmlError("feed_xml_invalid") from None + + if root.tag == "rss": + if root.attrib.get("version") != "2.0" or len(root) != 1 or root[0].tag != "channel": + raise FeedXmlError("feed_schema_invalid") + channel = root[0] + allowed = { + "title", + "link", + "description", + "language", + "copyright", + "managingEditor", + "webMaster", + "pubDate", + "lastBuildDate", + "category", + "generator", + "docs", + "cloud", + "ttl", + "image", + "rating", + "textInput", + "skipHours", + "skipDays", + "item", + } + if any(child.tag not in allowed and child.tag not in _RSS_FOREIGN_METADATA for child in channel): + raise FeedXmlError("feed_schema_invalid") + items = [child for child in channel if child.tag == "item"] + return rows_for(items, "rss2") + + atom_root = "{http://www.w3.org/2005/Atom}feed" + atom_entry = "{http://www.w3.org/2005/Atom}entry" + atom_metadata = { + "{http://www.w3.org/2005/Atom}id", + "{http://www.w3.org/2005/Atom}title", + "{http://www.w3.org/2005/Atom}updated", + "{http://www.w3.org/2005/Atom}subtitle", + "{http://www.w3.org/2005/Atom}link", + "{http://www.w3.org/2005/Atom}author", + "{http://www.w3.org/2005/Atom}category", + "{http://www.w3.org/2005/Atom}contributor", + "{http://www.w3.org/2005/Atom}generator", + "{http://www.w3.org/2005/Atom}icon", + "{http://www.w3.org/2005/Atom}logo", + "{http://www.w3.org/2005/Atom}rights", + } + if root.tag == atom_root and all(child.tag in atom_metadata or child.tag == atom_entry for child in root): + return rows_for([child for child in root if child.tag == atom_entry], "atom") + raise FeedXmlError("feed_schema_invalid") - rss_items = root.findall("./channel/item") - if rss_items: - for item in rss_items[:max_items]: - title = child_text(item, ("title",)) - link = rss_item_link(item) - published = child_text(item, ("pubDate", "{http://purl.org/dc/elements/1.1/}date")) - description = child_text(item, ("description", "{http://purl.org/rss/1.0/modules/content/}encoded")) - text = " ".join(part for part in (title, strip_html(description)) if part) - rows.append( - { - "item_id": stable_item_id(feed.feed_id, link, title), - "published_at": parse_datetime(published), - "source_type": feed.source_type, - "source_url": link, - "author": feed.author, - "text": text, - } - ) - return rows - - atom_entries = root.findall("{http://www.w3.org/2005/Atom}entry") - for entry in atom_entries[:max_items]: - title = child_text(entry, ("{http://www.w3.org/2005/Atom}title", "title")) - link = atom_entry_link(entry) - published = child_text( - entry, - ("{http://www.w3.org/2005/Atom}published", "{http://www.w3.org/2005/Atom}updated", "published", "updated"), - ) - summary = child_text(entry, ("{http://www.w3.org/2005/Atom}summary", "{http://www.w3.org/2005/Atom}content")) - text = " ".join(part for part in (title, strip_html(summary)) if part) - rows.append( - { - "item_id": stable_item_id(feed.feed_id, link, title), - "published_at": parse_datetime(published), - "source_type": feed.source_type, - "source_url": link, - "author": feed.author, - "text": text, - } - ) - return rows + +def parse_feed_items(feed_bytes: bytes, feed: FeedConfig, *, max_items: int = 25) -> list[dict[str, str]]: + return [dict(row) for row in parse_feed_snapshot(feed_bytes, feed, max_items=max_items).entries] 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,40 +279,49 @@ 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) + parsed = parse_feed_snapshot(fetcher(feed.feed_url), feed, max_items=max_items_per_feed) + feed_rows = [PrimitiveRow.from_mapping(row).to_mapping() for row in parsed.entries] + feed_records.append( + { + "feed_id": feed.feed_id, + "feed_url": feed.feed_url, + "kind": parsed.feed_kind, + "state": "accepted" if feed_rows else "quarantined", + "rows": feed_rows, + "error_code": None if feed_rows else "zero_entries", + } + ) + rows.extend(feed_rows) 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}", - ) + code = getattr(exc, "code", "fetch_failed") + if type(code) is not str or not re.fullmatch(r"[a-z][a-z0-9_]{0,63}", code): + code = "fetch_failed" + feed_records.append( + { + "feed_id": feed.feed_id, + "feed_url": feed.feed_url, + "kind": "unknown", + "state": "failed", + "rows": [], + "error_code": 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(record["state"] == "accepted" for record in feed_records): if status_output: - write_fetch_status(status_output, statuses, item_count=0) - raise RuntimeError("all configured RSS/Atom feeds failed") + write_fetch_status(status_output, feed_records) + if any(record["state"] == "failed" for record in feed_records): + raise RuntimeError("all configured RSS/Atom feeds failed") + write_csv_rows(output_path, ["item_id", "published_at", "source_type", "source_url", "author", "text"], []) + return [] 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 diff --git a/tests/test_feed_primitives.py b/tests/test_feed_primitives.py new file mode 100644 index 0000000..86a5af9 --- /dev/null +++ b/tests/test_feed_primitives.py @@ -0,0 +1,81 @@ +from __future__ import annotations + +import json + +import pytest + +from political_event_tracking_research.feed_primitives import ( + PrimitiveStatusError, + build_status, + parse_status_bytes, + serialize_status, +) + + +ROW = { + "item_id": "feed-a-1", + "published_at": "2026-05-01T12:30:00Z", + "source_type": "official", + "source_url": "https://example.test/item/1", + "author": "Example", + "text": "Policy mention", +} + + +def record( + feed_id: str, + *, + kind: str = "rss2", + 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": kind, + "state": state, + "rows": rows if rows is not None else [ROW], + "error_code": error_code, + } + + +def test_unknown_kind_is_only_valid_for_failed_feed() -> None: + failed = build_status([record("bad", kind="unknown", state="failed", rows=[], error_code="fetch_failed")]) + assert failed["feeds"][0]["kind"] == "unknown" + for state in ("accepted", "quarantined"): + with pytest.raises(PrimitiveStatusError, match="feed_kind_invalid|feed_state_invalid"): + build_status([record("bad", kind="unknown", state=state, rows=[], error_code="zero_entries")]) + + +def test_quarantined_feed_has_no_rows_or_digest_contribution() -> None: + quarantined = record("empty", state="quarantined", rows=[], error_code="zero_entries") + status = build_status([record("good"), quarantined]) + baseline = build_status([record("good")]) + assert status["accepted_row_count"] == baseline["accepted_row_count"] + assert status["aggregate_row_digest"] == baseline["aggregate_row_digest"] + assert status["publication_complete"] is False + assert status["eligible_for_live_publication"] is False + + +def test_wire_rejects_unknown_kind_for_accepted_or_quarantined() -> None: + status = build_status([record("good")]) + payload = json.loads(serialize_status(status)) + payload["feeds"][0]["kind"] = "unknown" + with pytest.raises(PrimitiveStatusError, match="feed_kind_invalid|feed_state_invalid"): + serialize_status(payload) + + +def test_wire_rejects_accepted_feed_with_error_code() -> None: + status = build_status([record("good")]) + payload = json.loads(serialize_status(status)) + payload["feeds"][0]["error_code"] = "unexpected" + with pytest.raises(PrimitiveStatusError, match="feed_state_invalid"): + serialize_status(payload) + + +def test_wire_roundtrip_is_canonical() -> None: + status = build_status([record("a"), record("empty", state="quarantined", rows=[], error_code="zero_entries")]) + wire = serialize_status(status) + assert parse_status_bytes(wire) == status + assert not wire.endswith(b"\n") diff --git a/tests/test_rss_source_fetch.py b/tests/test_rss_source_fetch.py index d1f324e..f675e59 100644 --- a/tests/test_rss_source_fetch.py +++ b/tests/test_rss_source_fetch.py @@ -7,7 +7,12 @@ import pytest from political_event_tracking_research import rss_source_fetch -from political_event_tracking_research.rss_source_fetch import FeedConfig, fetch_rss_sources, parse_feed_items +from political_event_tracking_research.rss_source_fetch import ( + FeedConfig, + fetch_rss_sources, + parse_feed_items, + parse_feed_snapshot, +) def test_parse_rss_feed_items_to_source_items() -> None: @@ -45,6 +50,12 @@ def test_parse_rss_feed_items_to_source_items() -> None: assert rows[0]["item_id"].startswith("whitehouse-test-") +def test_blank_optional_author_remains_valid() -> None: + payload = b"x" + rows = parse_feed_items(payload, FeedConfig("x", "https://example.test", "official", "")) + assert rows[0]["author"] == "" + + def test_parse_atom_feed_items_to_source_items() -> None: feed_xml = b""" @@ -70,6 +81,34 @@ def test_parse_atom_feed_items_to_source_items() -> None: assert "EVT2" in rows[0]["text"] +@pytest.mark.parametrize( + ("payload", "kind"), + [ + (b"", "rss2"), + (b"", "atom"), + ], +) +def test_empty_feed_preserves_kind_and_is_not_guessed_from_rows(payload: bytes, kind: str) -> None: + parsed = parse_feed_snapshot(payload, FeedConfig("x", "https://example.test", "official", "")) + assert parsed.feed_kind == kind + assert parsed.entries == () + + +@pytest.mark.parametrize( + "payload", + [ + b"", + b"", + b"", + b"", + b"", + ], +) +def test_extra_or_unknown_direct_children_fail_closed(payload: bytes) -> None: + with pytest.raises(ValueError, match="feed_schema_invalid"): + parse_feed_snapshot(payload, FeedConfig("x", "https://example.test", "official", "")) + + @pytest.mark.parametrize( "payload", [ @@ -171,8 +210,33 @@ def fake_fetch(url: str) -> bytes: payload = json.loads(status.read_text(encoding="utf-8")) 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"] + assert payload["feeds"][0]["feed_id"] == "bad" + assert payload["feeds"][0]["error_code"] == "fetch_failed" + + +def test_continue_on_feed_error_sanitizes_any_exception_but_not_base_exception(tmp_path: Path) -> None: + feeds_path = tmp_path / "feeds.csv" + feeds_path.write_text( + "feed_id,feed_url,source_type,author\n" + "bad,https://example.invalid/bad.xml,official_remarks,Example\n", + encoding="utf-8", + ) + with pytest.raises(RuntimeError, match="all configured"): + fetch_rss_sources( + feeds_path, + tmp_path / "source_items.csv", + continue_on_feed_error=True, + status_output=tmp_path / "status.json", + fetcher=lambda _url: (_ for _ in ()).throw(RuntimeError("do not leak")), + ) + assert "do not leak" not in (tmp_path / "status.json").read_text(encoding="utf-8") + with pytest.raises(KeyboardInterrupt): + fetch_rss_sources( + feeds_path, + tmp_path / "source_items.csv", + continue_on_feed_error=True, + fetcher=lambda _url: (_ for _ in ()).throw(KeyboardInterrupt()), + ) def test_fetch_rss_sources_fails_when_all_feeds_fail(tmp_path: Path) -> None: