diff --git a/litellm/proxy/anthropic_endpoints/streaming_model_restamp.py b/litellm/proxy/anthropic_endpoints/streaming_model_restamp.py index e8d54f03949..7da5e5099fc 100644 --- a/litellm/proxy/anthropic_endpoints/streaming_model_restamp.py +++ b/litellm/proxy/anthropic_endpoints/streaming_model_restamp.py @@ -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) ) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 989cc7c18fb..eda7ebfa624 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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 diff --git a/tests/test_litellm/proxy/anthropic_endpoints/test_streaming_model_restamp.py b/tests/test_litellm/proxy/anthropic_endpoints/test_streaming_model_restamp.py index 385173b24a9..b7bc670c7f8 100644 --- a/tests/test_litellm/proxy/anthropic_endpoints/test_streaming_model_restamp.py +++ b/tests/test_litellm/proxy/anthropic_endpoints/test_streaming_model_restamp.py @@ -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