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

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. 

5 

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""" 

9 

10import json 

11import re 

12from collections.abc import Mapping 

13from typing import Final 

14 

15from pydantic import TypeAdapter, ValidationError 

16 

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"') 

23 

24_EVENT_ADAPTER: Final = TypeAdapter(Mapping[str, object]) 

25 

26 

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 

34 

35 

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}" 

52 

53 

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)) 

60 

61 

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``. 

65 

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 

74 

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") 

80 

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 

86 

87 return chunk 

88 

89 

90def _is_ping_frame(frame: bytes) -> bool: 

91 return any(marker in frame for marker in _PING_MARKERS) 

92 

93 

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 """ 

106 

107 def __init__(self, requested_model: str) -> None: 

108 self._requested_model: Final = requested_model 

109 self._held = b"" 

110 self._armed = True 

111 

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 

123 

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 

132 

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 

150 

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 )