Skip to content

Commit df84326

Browse files
Pigbibicodex
andcommitted
Harden fetch acceptance XML and feed states
Co-Authored-By: Codex <noreply@openai.com>
1 parent 76fa647 commit df84326

2 files changed

Lines changed: 76 additions & 5 deletions

File tree

src/political_event_tracking_research/fetch_acceptance.py

Lines changed: 45 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,12 @@
3131
_KINDS = frozenset({"rss", "atom", "unknown"})
3232
_SAFE_ERROR = re.compile(r"^[a-z][a-z0-9_]*$")
3333
_ATOM = "http://www.w3.org/2005/Atom"
34+
MAX_XML_BYTES = 1024 * 1024
35+
MAX_XML_DEPTH = 32
36+
MAX_XML_NODES = 10000
37+
MAX_XML_TEXT_BYTES = 256 * 1024
38+
MAX_XML_ATTRIBUTES = 128
39+
_FORBIDDEN_DECLARATION = re.compile(rb"<!\s*(?:doctype|entity|element|attlist|notation)\b|\b(?:system|public)\b", re.IGNORECASE)
3440

3541

3642
class FetchAcceptanceError(ValueError):
@@ -80,13 +86,41 @@ def _feed_result(feed_id: str, feed_url: str, kind: str, accepted: int, rejected
8086
return {"feed_id": feed_id, "feed_url": feed_url, "kind": kind, "accepted_row_count": accepted, "rejected_row_count": rejected, "state": state, "error_code": error_code}
8187

8288

83-
def classify_feed_payload(feed_id: str, feed_url: str, payload: bytes) -> dict[str, object]:
89+
def _parse_bounded_xml(payload: bytes) -> ET.Element:
8490
if type(payload) is not bytes:
8591
raise _fail("feed_payload_invalid")
92+
if len(payload) > MAX_XML_BYTES:
93+
raise _fail("xml_oversize")
94+
if _FORBIDDEN_DECLARATION.search(payload):
95+
raise _fail("xml_forbidden_declaration")
8696
try:
8797
root = ET.fromstring(payload)
88-
except (ET.ParseError, UnicodeError, ValueError):
89-
return _feed_result(feed_id, feed_url, "unknown", 0, 0, "failed", "xml_invalid")
98+
except (ET.ParseError, UnicodeError, ValueError, RecursionError):
99+
raise _fail("xml_invalid") from None
100+
nodes = 0
101+
text_bytes = 0
102+
stack: list[tuple[ET.Element, int]] = [(root, 1)]
103+
while stack:
104+
element, depth = stack.pop()
105+
nodes += 1
106+
if nodes > MAX_XML_NODES or depth > MAX_XML_DEPTH or len(element.attrib) > MAX_XML_ATTRIBUTES:
107+
raise _fail("xml_structure_over_limit")
108+
for text in (element.text, element.tail):
109+
if text is not None:
110+
text_bytes += len(text.encode("utf-8", errors="strict"))
111+
if text_bytes > MAX_XML_TEXT_BYTES:
112+
raise _fail("xml_structure_over_limit")
113+
stack.extend((child, depth + 1) for child in reversed(list(element)))
114+
return root
115+
116+
117+
def classify_feed_payload(feed_id: str, feed_url: str, payload: bytes) -> dict[str, object]:
118+
if type(payload) is not bytes:
119+
raise _fail("feed_payload_invalid")
120+
try:
121+
root = _parse_bounded_xml(payload)
122+
except FetchAcceptanceError as error:
123+
return _feed_result(feed_id, feed_url, "unknown", 0, 0, "failed", error.code)
90124
if root.tag == "rss":
91125
if root.attrib.get("version") not in {"2.0"}:
92126
return _feed_result(feed_id, feed_url, "unknown", 0, 0, "failed", "rss_schema_unsupported")
@@ -120,9 +154,15 @@ def _validate_feed_result(value: object) -> dict[str, object]:
120154
result = _feed_result(value["feed_id"], value["feed_url"], value["kind"], value["accepted_row_count"], value["rejected_row_count"], value["state"], value["error_code"])
121155
if result != dict(value):
122156
raise _fail("feed_result_noncanonical")
123-
if result["state"] == "accepted" and (result["accepted_row_count"] <= 0 or result["rejected_row_count"] != 0 or result["error_code"] is not None):
157+
state = result["state"]
158+
accepted_rows = result["accepted_row_count"]
159+
rejected_rows = result["rejected_row_count"]
160+
error_code = result["error_code"]
161+
if state == "accepted" and (accepted_rows <= 0 or rejected_rows != 0 or error_code is not None):
162+
raise _fail("feed_result_invalid")
163+
if state in {"failed", "stale", "missing"} and (accepted_rows != 0 or rejected_rows != 0 or not error_code):
124164
raise _fail("feed_result_invalid")
125-
if result["state"] == "quarantined" and not result["error_code"]:
165+
if state == "quarantined" and not error_code:
126166
raise _fail("feed_result_invalid")
127167
return result
128168

tests/test_fetch_acceptance.py

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -87,3 +87,34 @@ def test_invalid_entry_is_counted_not_accepted() -> None:
8787
assert result["accepted_row_count"] == 0
8888
assert result["rejected_row_count"] == 1
8989
assert result["state"] == "quarantined"
90+
91+
92+
@pytest.mark.parametrize("declaration", ["<!DOCTYPE rss>", "<!doctype rss>", "<! DOCTYPE rss>", "<!\nENTITY x \"x\">", "<!SYSTEM \"file:///tmp/x\">"])
93+
def test_forbidden_xml_declarations_are_sanitized(declaration: str) -> None:
94+
result = classify_feed_payload("feed", "https://example.test/feed", (declaration + "<rss version='2.0'><channel/></rss>").encode())
95+
assert result["state"] == "failed"
96+
assert result["error_code"] == "xml_forbidden_declaration"
97+
98+
99+
def test_xml_size_depth_nodes_text_and_attributes_are_bounded() -> None:
100+
assert classify_feed_payload("feed", "https://example.test/feed", RSS + b"x" * (1024 * 1024))["error_code"] == "xml_oversize"
101+
deep = "<rss version='2.0'><channel>" + "<x>" * 40 + "</x>" * 40 + "</channel></rss>"
102+
assert classify_feed_payload("feed", "https://example.test/feed", deep.encode())["error_code"] == "xml_structure_over_limit"
103+
attrs = "<rss version='2.0' " + " ".join(f"a{i}='x'" for i in range(130)) + "><channel/></rss>"
104+
assert classify_feed_payload("feed", "https://example.test/feed", attrs.encode())["error_code"] == "xml_structure_over_limit"
105+
106+
107+
@pytest.mark.parametrize(
108+
"feed",
109+
[
110+
{"feed_id": "x", "feed_url": "u", "kind": "rss", "accepted_row_count": 0, "rejected_row_count": 0, "state": "failed", "error_code": None},
111+
{"feed_id": "x", "feed_url": "u", "kind": "rss", "accepted_row_count": 1, "rejected_row_count": 0, "state": "failed", "error_code": "bad"},
112+
{"feed_id": "x", "feed_url": "u", "kind": "rss", "accepted_row_count": 0, "rejected_row_count": 0, "state": "stale", "error_code": ""},
113+
{"feed_id": "x", "feed_url": "u", "kind": "rss", "accepted_row_count": 0, "rejected_row_count": 0, "state": "missing", "error_code": None},
114+
{"feed_id": "x", "feed_url": "u", "kind": "rss", "accepted_row_count": 1, "rejected_row_count": 0, "state": "accepted", "error_code": "bad"},
115+
{"feed_id": "x", "feed_url": "u", "kind": "rss", "accepted_row_count": 0, "rejected_row_count": 0, "state": "quarantined", "error_code": None},
116+
],
117+
)
118+
def test_feed_state_row_and_error_invariants_are_strict(feed: dict[str, object]) -> None:
119+
with pytest.raises(FetchAcceptanceError):
120+
build_acceptance_status([feed])

0 commit comments

Comments
 (0)