fix(proxy): handle CRLF and CR SSE frame terminators and flush held tail in anthropic stream restamper

This commit is contained in:
mateo-berri 2026-09-01 12:16:27 -07:00
parent c02c81452c
commit c21e895fe2
3 changed files with 144 additions and 26 deletions

View file

@ -8,6 +8,7 @@ provider passthrough path) or as event dicts (fake-stream and agentic paths).
"""
import json
import re
from collections.abc import Mapping
from typing import Final
@ -16,7 +17,7 @@ from pydantic import TypeAdapter, ValidationError
_MESSAGE_START_EVENT: Final = "message_start"
_MESSAGE_START_MARKER: Final = b"message_start"
_SSE_DATA_FIELD: Final = "data:"
_SSE_FRAME_END: Final = b"\n\n"
_SSE_FRAME_END_PATTERN: Final = re.compile(rb"\r\n\r\n|\r\r|\n\n")
_MAX_HELD_BYTES: Final = 65536
_PING_MARKERS: Final = (b"event: ping", b'"type": "ping"', b'"type":"ping"')
@ -46,15 +47,16 @@ def _restamped_data_line(line: str, requested_model: str) -> str | None:
restamped: Final = _restamped_event(event, requested_model)
if restamped is None:
return None
return f"data: {json.dumps(restamped, separators=(',', ':'))}"
terminator: Final = line[len(line.rstrip("\r\n")) :]
return f"data: {json.dumps(restamped, separators=(',', ':'))}{terminator}"
def _restamped_frame(frame: str, requested_model: str) -> str | None:
lines: Final = frame.split("\n")
lines: Final = frame.splitlines(keepends=True)
restamped: Final = tuple(_restamped_data_line(line, requested_model) for line in lines)
if all(line is None for line in restamped):
return None
return "\n".join(new if new is not None else old for new, old in zip(restamped, lines))
return "".join(new if new is not None else old for new, old in zip(restamped, lines))
def restamp_anthropic_stream_chunk_model(chunk: object, requested_model: str) -> object:
@ -93,10 +95,12 @@ class AnthropicStreamModelRestamper:
"""
Per-stream restamper for the encoded passthrough path, where chunks are raw
transport reads: the ``message_start`` SSE frame can arrive split across
chunks or coalesced with later frames. Complete frames are emitted as their
terminator closes them and an incomplete tail is held until it completes,
so the restamp never misses a torn frame. Once ``message_start`` has been
handled, or the first real event proves the stream carries none, every
chunks or coalesced with later frames. Complete frames (``\\n\\n``,
``\\r\\n\\r\\n``, or ``\\r\\r`` terminated) are emitted as their terminator
closes them and an incomplete tail is held until it completes, so the
restamp never misses a torn frame; ``flush`` returns whatever is still held
when the stream ends so no bytes are swallowed. Once ``message_start`` has
been handled, or the first real event proves the stream carries none, every
later chunk passes through untouched.
"""
@ -117,17 +121,27 @@ class AnthropicStreamModelRestamper:
self._armed = False
return restamped
def flush(self) -> bytes:
held: Final = self._held
self._held = b""
self._armed = False
if not held:
return b""
restamped: Final = restamp_anthropic_stream_chunk_model(held, self._requested_model)
return restamped if isinstance(restamped, bytes) else held
def _process_encoded(self, data: bytes) -> bytes:
combined: Final = self._held + data
if _SSE_FRAME_END not in combined:
boundaries: Final = tuple(match.end() for match in _SSE_FRAME_END_PATTERN.finditer(combined))
if not boundaries:
if len(combined) > _MAX_HELD_BYTES:
self._held = b""
self._armed = False
return combined
self._held = combined
return b""
closed, _, tail = combined.rpartition(_SSE_FRAME_END)
emitted: Final = self._restamped_closed_block(closed + _SSE_FRAME_END)
emitted: Final = self._restamped_closed_block(combined[: boundaries[-1]])
tail: Final = combined[boundaries[-1] :]
if not self._armed:
self._held = b""
return emitted + tail
@ -135,7 +149,8 @@ class AnthropicStreamModelRestamper:
return emitted
def _restamped_closed_block(self, closed: bytes) -> bytes:
frames: Final = tuple(closed.split(_SSE_FRAME_END)[:-1])
boundaries: Final = tuple(match.end() for match in _SSE_FRAME_END_PATTERN.finditer(closed))
frames: Final = tuple(closed[start:end] for start, end in zip((0, *boundaries[:-1]), boundaries))
decider: Final = next(
(
index
@ -147,13 +162,13 @@ class AnthropicStreamModelRestamper:
if decider is None:
return closed
self._armed = False
decider_frame: Final = frames[decider] + _SSE_FRAME_END
if _MESSAGE_START_MARKER not in decider_frame:
if _MESSAGE_START_MARKER not in frames[decider]:
return closed
restamped_text: Final = _restamped_frame(decider_frame.decode("utf-8", errors="ignore"), self._requested_model)
restamped_text: Final = _restamped_frame(
frames[decider].decode("utf-8", errors="ignore"), self._requested_model
)
if restamped_text is None:
return closed
return b"".join(
restamped_text.encode("utf-8") if index == decider else frame + _SSE_FRAME_END
for index, frame in enumerate(frames)
restamped_text.encode("utf-8") if index == decider else frame for index, frame in enumerate(frames)
)

View file

@ -3410,12 +3410,10 @@ class ProxyBaseLLMRequestProcessing:
return chunk
@staticmethod
def _sse_chunk_serializer(restamp_model: str | None) -> StreamChunkSerializer:
if not restamp_model:
def _sse_chunk_serializer(restamper: AnthropicStreamModelRestamper | None) -> StreamChunkSerializer:
if restamper is None:
return ProxyBaseLLMRequestProcessing.return_sse_chunk
restamper: Final = AnthropicStreamModelRestamper(restamp_model)
def serialize(chunk: object) -> str:
return ProxyBaseLLMRequestProcessing.return_sse_chunk(restamper.process(chunk))
@ -3481,11 +3479,16 @@ class ProxyBaseLLMRequestProcessing:
serialize_chunk: StreamChunkSerializer,
serialize_error: StreamErrorSerializer,
request: Request | None = None,
flush_tail: Callable[[], bytes] | None = None,
) -> AsyncGenerator[str, None]:
"""
Shared streaming data generator: runs proxy iterator hook, per-chunk hook,
cost injection, then yields chunks via serialize_chunk; on exception runs
failure hook and yields via serialize_error. Use for SSE or NDJSON.
``flush_tail`` runs once after the upstream iterator completes cleanly and
its non-empty result is yielded, so a serializer that buffers bytes across
chunks can emit anything still held at end of stream.
"""
verbose_proxy_logger.debug("inside generator")
# Resolve per-stream (not per-chunk) whether the heavy per-chunk path
@ -3548,6 +3551,9 @@ class ProxyBaseLLMRequestProcessing:
# so it must not suppress that refund.
delivered_chunk = delivered_chunk or chunk != STREAM_SSE_KEEPALIVE_PING_BYTES
yield serialize_chunk(chunk)
held_tail: Final = flush_tail() if flush_tail is not None else b""
if held_tail:
yield serialize_chunk(held_tail)
stream_completed = True
except (asyncio.CancelledError, GeneratorExit):
# Client disconnected mid-stream. CancelledError / GeneratorExit
@ -3558,8 +3564,7 @@ class ProxyBaseLLMRequestProcessing:
# billing and release exactly once. This is the outermost generator
# Starlette closes on disconnect, so the nested iterator hook (which
# only sees GeneratorExit on GC) cannot own the refund.
if not stream_completed:
client_disconnected = True
client_disconnected = not stream_completed
if not delivered_chunk and not _withheld_provider_output(response):
from litellm.proxy.spend_tracking.budget_reservation import (
release_budget_reservation_on_cancel,
@ -3627,16 +3632,18 @@ class ProxyBaseLLMRequestProcessing:
event in place of the provider's model, matching what the non-streaming
response reports.
"""
restamper: Final = AnthropicStreamModelRestamper(restamp_model) if restamp_model else None
return ProxyBaseLLMRequestProcessing.async_streaming_data_generator(
response=response,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
proxy_logging_obj=proxy_logging_obj,
serialize_chunk=ProxyBaseLLMRequestProcessing._sse_chunk_serializer(restamp_model),
serialize_chunk=ProxyBaseLLMRequestProcessing._sse_chunk_serializer(restamper),
serialize_error=lambda proxy_exc: (
f"{STREAM_SSE_DATA_PREFIX}{json.dumps({'error': proxy_exc.to_dict()})}\n\n"
),
request=request,
flush_tail=None if restamper is None else restamper.flush,
)
@overload

View file

@ -14,12 +14,12 @@ from litellm.proxy.anthropic_endpoints.streaming_model_restamp import (
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
def _message_start_frame(model: str) -> bytes:
def _message_start_frame(model: str, line_end: str = "\n") -> bytes:
payload = {
"type": "message_start",
"message": {"id": "msg_1", "type": "message", "role": "assistant", "model": model, "content": []},
}
return f"event: message_start\ndata: {json.dumps(payload)}\n\n".encode()
return f"event: message_start{line_end}data: {json.dumps(payload)}{line_end}{line_end}".encode()
def _proxy_logging_obj_streaming(frames: list[bytes]) -> MagicMock:
@ -190,3 +190,99 @@ async def test_sse_generator_restamps_message_start_split_across_chunks():
joined = b"".join(chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in chunks)
assert _model_from_frame(joined) == "claude-auto-1"
def test_restamps_crlf_terminated_message_start_frame():
frame = _message_start_frame("claude-haiku-4-5-20251001", line_end="\r\n")
delta = b'event: content_block_delta\r\ndata: {"type":"content_block_delta","delta":{"text":"hi"}}\r\n\r\n'
restamper = AnthropicStreamModelRestamper("claude-auto-1")
emitted = restamper.process(frame)
assert isinstance(emitted, bytes)
assert _model_from_frame(emitted) == "claude-auto-1"
assert emitted.endswith(b"\r\n\r\n")
assert restamper.process(delta) == delta
def test_restamps_cr_terminated_message_start_frame():
frame = _message_start_frame("claude-haiku-4-5-20251001", line_end="\r")
restamper = AnthropicStreamModelRestamper("claude-auto-1")
emitted = restamper.process(frame)
assert isinstance(emitted, bytes)
assert b'"model":"claude-auto-1"' in emitted
assert emitted.endswith(b"\r\r")
def test_restamps_crlf_message_start_split_across_transport_chunks():
frame = _message_start_frame("claude-haiku-4-5-20251001", line_end="\r\n")
restamper = AnthropicStreamModelRestamper("claude-auto-1")
held = restamper.process(frame[:25])
emitted = restamper.process(frame[25:])
assert held == b""
assert isinstance(emitted, bytes)
assert _model_from_frame(emitted) == "claude-auto-1"
def test_flush_returns_restamped_held_tail():
unterminated = _message_start_frame("claude-haiku-4-5-20251001")[:-2]
restamper = AnthropicStreamModelRestamper("claude-auto-1")
assert restamper.process(unterminated) == b""
flushed = restamper.flush()
assert b'"model":"claude-auto-1"' in flushed
assert restamper.flush() == b""
def test_flush_disarms_the_restamper():
restamper = AnthropicStreamModelRestamper("claude-auto-1")
frame = _message_start_frame("claude-haiku-4-5-20251001")
assert restamper.flush() == b""
assert restamper.process(frame) == frame
@pytest.mark.asyncio
async def test_sse_generator_flushes_held_tail_at_end_of_stream():
unterminated = _message_start_frame("claude-haiku-4-5-20251001")[:-2]
proxy_logging_obj = _proxy_logging_obj_streaming([unterminated])
chunks = [
chunk
async for chunk in ProxyBaseLLMRequestProcessing.async_sse_data_generator(
response=MagicMock(),
user_api_key_dict=MagicMock(),
request_data={"model": "claude-auto-1"},
proxy_logging_obj=proxy_logging_obj,
restamp_model="claude-auto-1",
)
]
joined = b"".join(chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in chunks)
assert b'"model":"claude-auto-1"' in joined
@pytest.mark.asyncio
async def test_sse_generator_restamps_crlf_stream():
frame = _message_start_frame("claude-haiku-4-5-20251001", line_end="\r\n")
delta = b'event: content_block_delta\r\ndata: {"type":"content_block_delta","delta":{"text":"hi"}}\r\n\r\n'
proxy_logging_obj = _proxy_logging_obj_streaming([frame, delta])
chunks = [
chunk
async for chunk in ProxyBaseLLMRequestProcessing.async_sse_data_generator(
response=MagicMock(),
user_api_key_dict=MagicMock(),
request_data={"model": "claude-auto-1"},
proxy_logging_obj=proxy_logging_obj,
restamp_model="claude-auto-1",
)
]
assert _model_from_frame(chunks[0]) == "claude-auto-1"
assert chunks[1] == delta