Skip to content

Commit a2057fd

Browse files
Pigbibicodex
andcommitted
Tighten primitive status canonical contract
Co-Authored-By: Codex <noreply@openai.com>
1 parent d3b0e66 commit a2057fd

4 files changed

Lines changed: 33 additions & 6 deletions

File tree

src/political_event_tracking_research/feed_primitives.py

Lines changed: 16 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010

1111
STATUS_VERSION = "pert.feed_primitives.v1"
1212
MAX_ROWS_PER_FEED = 10_000
13+
MAX_SAFE_JSON_INTEGER = 2**53 - 1
1314
_ROW_KEYS = frozenset({"item_id", "published_at", "source_type", "source_url", "author", "text"})
1415
_FEED_KEYS = frozenset({"feed_id", "feed_url", "kind", "state", "rows", "error_code"})
1516
_FEED_WIRE_KEYS = frozenset(
@@ -65,6 +66,12 @@ def _string(value: object, code: str, *, allow_empty: bool = True) -> str:
6566
return value
6667

6768

69+
def _safe_int(value: object, code: str) -> int:
70+
if type(value) is not int or value < 0 or value > MAX_SAFE_JSON_INTEGER:
71+
raise _fail(code)
72+
return value
73+
74+
6875
@dataclass(frozen=True, slots=True)
6976
class PrimitiveRow:
7077
item_id: str
@@ -180,6 +187,8 @@ def build_status(feed_records: Iterable[Mapping[str, object]]) -> dict[str, obje
180187
raise _fail("configured_feed_empty")
181188
if len({item["feed_id"] for item in records}) != len(records):
182189
raise _fail("feed_duplicate")
190+
if len(records) > MAX_SAFE_JSON_INTEGER:
191+
raise _fail("feed_count_overflow")
183192
records.sort(key=lambda item: (item["feed_id"], item["feed_url"]))
184193
feeds = [_feed_wire(item) for item in records]
185194
accepted = sum(item["state"] == "accepted" for item in records)
@@ -221,7 +230,7 @@ def _validate_wire(value: object) -> dict[str, object]:
221230
"accepted_row_count",
222231
"rejected_row_count",
223232
)
224-
if any(type(snapshot[key]) is not int or snapshot[key] < 0 for key in integer_keys):
233+
if any(_safe_int(snapshot[key], "status_counter_invalid") != snapshot[key] for key in integer_keys):
225234
raise _fail("status_counter_invalid")
226235
if any(type(snapshot[key]) is not bool for key in ("publication_complete", "eligible_for_live_publication")):
227236
raise _fail("status_counter_invalid")
@@ -239,6 +248,7 @@ def _validate_wire(value: object) -> dict[str, object]:
239248
"rejected_row_count": 0,
240249
}
241250
feed_ids: set[str] = set()
251+
feed_order: list[tuple[str, str]] = []
242252
for feed in snapshot["feeds"]:
243253
if not isinstance(feed, Mapping) or set(feed) != _FEED_WIRE_KEYS:
244254
raise _fail("feed_wire_shape_invalid")
@@ -251,11 +261,10 @@ def _validate_wire(value: object) -> dict[str, object]:
251261
if feed["feed_id"] in feed_ids:
252262
raise _fail("feed_duplicate")
253263
feed_ids.add(feed["feed_id"])
264+
feed_order.append((feed["feed_id"], feed["feed_url"]))
254265
if (
255-
type(feed["accepted_row_count"]) is not int
256-
or feed["accepted_row_count"] < 0
257-
or type(feed["rejected_row_count"]) is not int
258-
or feed["rejected_row_count"] < 0
266+
_safe_int(feed["accepted_row_count"], "feed_counter_invalid") != feed["accepted_row_count"]
267+
or _safe_int(feed["rejected_row_count"], "feed_counter_invalid") != feed["rejected_row_count"]
259268
):
260269
raise _fail("feed_counter_invalid")
261270
if not _is_digest(feed["row_digest"]):
@@ -275,6 +284,8 @@ def _validate_wire(value: object) -> dict[str, object]:
275284
expected_counts[f"{state if state != 'accepted' else 'successful'}_feed_count"] += 1
276285
expected_counts["accepted_row_count"] += feed["accepted_row_count"]
277286
expected_counts["rejected_row_count"] += feed["rejected_row_count"]
287+
if feed_order != sorted(feed_order):
288+
raise _fail("feed_order_invalid")
278289
if any(snapshot[key] != value for key, value in expected_counts.items()):
279290
raise _fail("status_counter_mismatch")
280291
complete = (

src/political_event_tracking_research/rss_source_fetch.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -173,7 +173,7 @@ def write_fetch_status(path: str | Path, feed_records: list[dict[str, object]])
173173
payload = serialize_status(build_status(feed_records))
174174
output_path = Path(path)
175175
output_path.parent.mkdir(parents=True, exist_ok=True)
176-
output_path.write_bytes(payload + b"\n")
176+
output_path.write_bytes(payload)
177177

178178

179179
def fetch_rss_sources(

tests/test_feed_primitives.py

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
import pytest
66

77
from political_event_tracking_research.feed_primitives import (
8+
MAX_SAFE_JSON_INTEGER,
89
PrimitiveStatusError,
910
build_status,
1011
parse_status_bytes,
@@ -131,3 +132,16 @@ def test_unknown_keys_duplicate_wire_and_unsafe_types_fail_closed() -> None:
131132
duplicate = serialize_status(status).replace(b'"status_version":', b'"status_version":')
132133
with pytest.raises(PrimitiveStatusError):
133134
parse_status_bytes(duplicate[:-1] + b',"status_version":"other"}')
135+
136+
137+
def test_safe_integer_bound_and_feed_order_are_strict() -> None:
138+
status = build_status([feed("a"), feed("b")])
139+
payload = json.loads(serialize_status(status))
140+
payload["feed_count"] = MAX_SAFE_JSON_INTEGER + 1
141+
with pytest.raises(PrimitiveStatusError, match="status_counter_invalid"):
142+
serialize_status(payload)
143+
144+
payload = json.loads(serialize_status(status))
145+
payload["feeds"] = list(reversed(payload["feeds"]))
146+
with pytest.raises(PrimitiveStatusError, match="feed_order_invalid"):
147+
serialize_status(payload)

tests/test_rss_source_fetch.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
import pytest
88

99
from political_event_tracking_research import rss_source_fetch
10+
from political_event_tracking_research.feed_primitives import parse_status_bytes
1011
from political_event_tracking_research.rss_source_fetch import FeedConfig, fetch_rss_sources, parse_feed_items
1112

1213

@@ -169,6 +170,7 @@ def fake_fetch(url: str) -> bytes:
169170

170171
assert len(rows) == 1
171172
payload = json.loads(status.read_text(encoding="utf-8"))
173+
parse_status_bytes(status.read_bytes())
172174
assert payload["successful_feed_count"] == 1
173175
assert payload["failed_feed_count"] == 1
174176
failed = next(item for item in payload["feeds"] if item["feed_id"] == "bad")

0 commit comments

Comments
 (0)