mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(proxy): handle CRLF and CR SSE frame terminators and flush held tail in anthropic stream restamper
This commit is contained in:
parent
c02c81452c
commit
c21e895fe2
3 changed files with 144 additions and 26 deletions
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue