Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/anthropic_endpoints/streaming_model_restamp.py: 14%
115 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Restamp the public ``model`` on the Anthropic Messages ``message_start`` event, the only
3stream event carrying a model, so streamed responses report the requested model like
4non-streaming ones do.
6Chunks reach the serializer either as already-encoded SSE frames (``bytes``/``str``, the
7provider passthrough path) or as event dicts (fake-stream and agentic paths).
8"""
10import json
11import re
12from collections.abc import Mapping
13from typing import Final
15from pydantic import TypeAdapter, ValidationError
17_MESSAGE_START_EVENT: Final = "message_start"
18_MESSAGE_START_MARKER: Final = b"message_start"
19_SSE_DATA_FIELD: Final = "data:"
20_SSE_FRAME_END_PATTERN: Final = re.compile(rb"\r\n\r\n|\r\r|\n\n")
21_MAX_HELD_BYTES: Final = 65536
22_PING_MARKERS: Final = (b"event: ping", b'"type": "ping"', b'"type":"ping"')
24_EVENT_ADAPTER: Final = TypeAdapter(Mapping[str, object])
27def _restamped_event(event: Mapping[str, object], requested_model: str) -> Mapping[str, object] | None:
28 message: Final = event.get("message")
29 if event.get("type") != _MESSAGE_START_EVENT or not isinstance(message, dict):
30 return None
31 if message.get("model") == requested_model:
32 return None
33 return {**event, "message": {**message, "model": requested_model}} # mutable-ok: SSE payload, re-serialized as is
36def _restamped_data_line(line: str, requested_model: str) -> str | None:
37 stripped: Final = line.strip()
38 if not stripped.startswith(_SSE_DATA_FIELD):
39 return None
40 payload: Final = stripped[len(_SSE_DATA_FIELD) :].strip()
41 if not payload or payload == "[DONE]":
42 return None
43 try:
44 event: Final = _EVENT_ADAPTER.validate_json(payload)
45 except ValidationError:
46 return None
47 restamped: Final = _restamped_event(event, requested_model)
48 if restamped is None:
49 return None
50 terminator: Final = line[len(line.rstrip("\r\n")) :]
51 return f"data: {json.dumps(restamped, separators=(',', ':'))}{terminator}"
54def _restamped_frame(frame: str, requested_model: str) -> str | None:
55 lines: Final = frame.splitlines(keepends=True)
56 restamped: Final = tuple(_restamped_data_line(line, requested_model) for line in lines)
57 if all(line is None for line in restamped):
58 return None
59 return "".join(new if new is not None else old for new, old in zip(restamped, lines))
62def restamp_anthropic_stream_chunk_model(chunk: object, requested_model: str) -> object:
63 """
64 Return ``chunk`` with the ``message_start`` model replaced by ``requested_model``.
66 Chunks that carry no model are returned unchanged.
67 """
68 if isinstance(chunk, dict):
69 try:
70 event: Final = _EVENT_ADAPTER.validate_python(chunk)
71 except ValidationError:
72 return chunk
73 return _restamped_event(event, requested_model) or chunk
75 if isinstance(chunk, (bytes, bytearray)):
76 if _MESSAGE_START_EVENT.encode() not in chunk:
77 return chunk
78 restamped_bytes: Final = _restamped_frame(chunk.decode("utf-8", errors="ignore"), requested_model)
79 return chunk if restamped_bytes is None else restamped_bytes.encode("utf-8")
81 if isinstance(chunk, str):
82 if _MESSAGE_START_EVENT not in chunk:
83 return chunk
84 restamped_text: Final = _restamped_frame(chunk, requested_model)
85 return chunk if restamped_text is None else restamped_text
87 return chunk
90def _is_ping_frame(frame: bytes) -> bool:
91 return any(marker in frame for marker in _PING_MARKERS)
94class AnthropicStreamModelRestamper:
95 """
96 Per-stream restamper for the encoded passthrough path, where chunks are raw
97 transport reads: the ``message_start`` SSE frame can arrive split across
98 chunks or coalesced with later frames. Complete frames (``\\n\\n``,
99 ``\\r\\n\\r\\n``, or ``\\r\\r`` terminated) are emitted as their terminator
100 closes them and an incomplete tail is held until it completes, so the
101 restamp never misses a torn frame; ``flush`` returns whatever is still held
102 when the stream ends so no bytes are swallowed. Once ``message_start`` has
103 been handled, or the first real event proves the stream carries none, every
104 later chunk passes through untouched.
105 """
107 def __init__(self, requested_model: str) -> None:
108 self._requested_model: Final = requested_model
109 self._held = b""
110 self._armed = True
112 def process(self, chunk: object) -> object:
113 if not self._armed:
114 return chunk
115 if isinstance(chunk, (bytes, bytearray)):
116 return self._process_encoded(bytes(chunk))
117 if isinstance(chunk, str):
118 return self._process_encoded(chunk.encode("utf-8"))
119 restamped: Final = restamp_anthropic_stream_chunk_model(chunk, self._requested_model)
120 if isinstance(chunk, dict) and chunk.get("type") not in (None, "ping"):
121 self._armed = False
122 return restamped
124 def flush(self) -> bytes:
125 held: Final = self._held
126 self._held = b""
127 self._armed = False
128 if not held:
129 return b""
130 restamped: Final = restamp_anthropic_stream_chunk_model(held, self._requested_model)
131 return restamped if isinstance(restamped, bytes) else held
133 def _process_encoded(self, data: bytes) -> bytes:
134 combined: Final = self._held + data
135 boundaries: Final = tuple(match.end() for match in _SSE_FRAME_END_PATTERN.finditer(combined))
136 if not boundaries:
137 if len(combined) > _MAX_HELD_BYTES:
138 self._held = b""
139 self._armed = False
140 return combined
141 self._held = combined
142 return b""
143 emitted: Final = self._restamped_closed_block(combined[: boundaries[-1]])
144 tail: Final = combined[boundaries[-1] :]
145 if not self._armed:
146 self._held = b""
147 return emitted + tail
148 self._held = tail
149 return emitted
151 def _restamped_closed_block(self, closed: bytes) -> bytes:
152 boundaries: Final = tuple(match.end() for match in _SSE_FRAME_END_PATTERN.finditer(closed))
153 frames: Final = tuple(closed[start:end] for start, end in zip((0, *boundaries[:-1]), boundaries))
154 decider: Final = next(
155 (
156 index
157 for index, frame in enumerate(frames)
158 if _MESSAGE_START_MARKER in frame or (b"data:" in frame and not _is_ping_frame(frame))
159 ),
160 None,
161 )
162 if decider is None:
163 return closed
164 self._armed = False
165 if _MESSAGE_START_MARKER not in frames[decider]:
166 return closed
167 restamped_text: Final = _restamped_frame(
168 frames[decider].decode("utf-8", errors="ignore"), self._requested_model
169 )
170 if restamped_text is None:
171 return closed
172 return b"".join(
173 restamped_text.encode("utf-8") if index == decider else frame for index, frame in enumerate(frames)
174 )