diff --git a/headroom/proxy/helpers.py b/headroom/proxy/helpers.py index b931b9339..d4b51f280 100644 --- a/headroom/proxy/helpers.py +++ b/headroom/proxy/helpers.py @@ -24,7 +24,7 @@ from typing import TYPE_CHECKING, Any, Literal, cast from headroom import paths as _paths from headroom._subprocess import run -from headroom.proxy import request_limit_policy, wire_debug_redaction_policy +from headroom.proxy import request_limit_policy, sse_byte_buffer_policy, wire_debug_redaction_policy from headroom.proxy.body_forwarding import ( BodyMutationTracker as BodyMutationTracker, # noqa: F401 - compatibility export ) @@ -533,22 +533,12 @@ def get_body_too_large_status() -> int: ) -# SSE byte-buffer helper supports LF and CRLF event separators. Per the SSE -# spec the default event name is "message"; we return ``None`` so callers can -# decide whether to apply that default. -_SSE_EVENT_TERMINATORS = (b"\n\n", b"\r\n\r\n") +_SSE_EVENT_TERMINATORS = sse_byte_buffer_policy.SSE_EVENT_TERMINATORS def _find_sse_event_terminator(buf: bytearray) -> tuple[int, int] | None: """Return the earliest complete SSE event terminator in ``buf``.""" - matches = [ - (idx, len(terminator)) - for terminator in _SSE_EVENT_TERMINATORS - if (idx := buf.find(terminator)) != -1 - ] - if not matches: - return None - return min(matches, key=lambda match: match[0]) + return sse_byte_buffer_policy.find_sse_event_terminator(buf) _SSE_EVENT_LINE_PREFIX = b"event:" @@ -595,36 +585,7 @@ def parse_sse_events_from_byte_buffer( ``decode("utf-8", errors="ignore")`` on a partial buffer; UTF-8 multi-byte characters split across TCP reads will corrupt content. """ - events: list[tuple[str | None, str]] = [] - while True: - terminator_match = _find_sse_event_terminator(buf) - if terminator_match is None: - break - idx, terminator_len = terminator_match - event_bytes = bytes(buf[:idx]) - # Drain the event + the trailing terminator from the buffer. - del buf[: idx + terminator_len] - # Decoding the COMPLETE event must succeed. If it doesn't, the - # upstream emitted invalid UTF-8 mid-stream — surface loudly. - event_text = event_bytes.decode("utf-8") - event_name: str | None = None - data_lines: list[str] = [] - for line in event_text.splitlines(): - if not line: - continue - # SSE comment line — ignored per spec. - if line.startswith(":"): - continue - if line.startswith("event:"): - event_name = line[len("event:") :].lstrip() - elif line.startswith("data:"): - data_lines.append(line[len("data:") :].lstrip()) - # Per SSE spec, multiple `data:` lines join with newline. We - # preserve that here even though OpenAI/Anthropic emit one - # `data:` per event. - if data_lines: - events.append((event_name, "\n".join(data_lines))) - return events + return sse_byte_buffer_policy.parse_sse_events_from_byte_buffer(buf) # Maximum message array length (prevents DoS from deeply nested payloads) diff --git a/headroom/proxy/sse_byte_buffer_policy.py b/headroom/proxy/sse_byte_buffer_policy.py new file mode 100644 index 000000000..c96a19df5 --- /dev/null +++ b/headroom/proxy/sse_byte_buffer_policy.py @@ -0,0 +1,60 @@ +"""Pure SSE byte-buffer parsing policy.""" + +from __future__ import annotations + +# SSE byte-buffer helper supports LF and CRLF event separators. Per the SSE +# spec the default event name is "message"; we return ``None`` so callers can +# decide whether to apply that default. +SSE_EVENT_TERMINATORS = (b"\n\n", b"\r\n\r\n") +SSE_EVENT_LINE_PREFIX = "event:" +SSE_DATA_LINE_PREFIX = "data:" + + +def find_sse_event_terminator(buf: bytearray) -> tuple[int, int] | None: + """Return the earliest complete SSE event terminator in ``buf``.""" + matches = [ + (idx, len(terminator)) + for terminator in SSE_EVENT_TERMINATORS + if (idx := buf.find(terminator)) != -1 + ] + if not matches: + return None + return min(matches, key=lambda match: match[0]) + + +def parse_sse_events_from_byte_buffer( + buf: bytearray, +) -> list[tuple[str | None, str]]: + """Drain complete ``event:`` + ``data:`` events from a bytes buffer. + + Returns list of ``(event_name, data_str)`` tuples for complete events. + Mutates ``buf`` in-place to leave only partial-event tail bytes. + + Operates on bytes; only decodes complete events as UTF-8 (raises if a + *complete* event has invalid UTF-8, which is an upstream protocol bug). + """ + events: list[tuple[str | None, str]] = [] + while True: + terminator_match = find_sse_event_terminator(buf) + if terminator_match is None: + break + idx, terminator_len = terminator_match + event_bytes = bytes(buf[:idx]) + del buf[: idx + terminator_len] + + event_text = event_bytes.decode("utf-8") + event_name: str | None = None + data_lines: list[str] = [] + for line in event_text.splitlines(): + if not line: + continue + if line.startswith(":"): + continue + if line.startswith(SSE_EVENT_LINE_PREFIX): + event_name = line[len(SSE_EVENT_LINE_PREFIX) :].lstrip() + elif line.startswith(SSE_DATA_LINE_PREFIX): + data_lines.append(line[len(SSE_DATA_LINE_PREFIX) :].lstrip()) + + if data_lines: + events.append((event_name, "\n".join(data_lines))) + return events diff --git a/tests/test_sse_byte_buffer_policy.py b/tests/test_sse_byte_buffer_policy.py new file mode 100644 index 000000000..d12466dfd --- /dev/null +++ b/tests/test_sse_byte_buffer_policy.py @@ -0,0 +1,33 @@ +from __future__ import annotations + +import pytest + +from headroom.proxy.sse_byte_buffer_policy import ( + find_sse_event_terminator, + parse_sse_events_from_byte_buffer, +) + + +def test_find_sse_event_terminator_returns_earliest_separator() -> None: + assert find_sse_event_terminator(bytearray(b"data: one\r\n\r\ndata: two\n\n")) == (9, 4) + + +def test_parse_sse_events_drains_complete_events_and_leaves_tail() -> None: + buf = bytearray(b": ignored\nevent: delta\ndata: one\ndata: two\n\npartial") + + assert parse_sse_events_from_byte_buffer(buf) == [("delta", "one\ntwo")] + assert bytes(buf) == b"partial" + + +def test_parse_sse_events_preserves_split_utf8_tail() -> None: + smile = "\U0001f642".encode() + buf = bytearray(b"data: hello " + smile[:2]) + + assert parse_sse_events_from_byte_buffer(buf) == [] + buf.extend(smile[2:] + b"\n\n") + assert parse_sse_events_from_byte_buffer(buf) == [(None, "hello \U0001f642")] + + +def test_parse_sse_events_raises_on_complete_invalid_utf8_event() -> None: + with pytest.raises(UnicodeDecodeError): + parse_sse_events_from_byte_buffer(bytearray(b"data: \xff\n\n"))