Merge pull request #41446 from BerriAI/litellm_fix_anthropic_stream_served_model

fix(anthropic): carry the served model from message_start onto stream chunks
This commit is contained in:
kerry-berri 2026-09-16 14:14:54 -07:00 • committed by GitHub
commit 930ec9643a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 83 additions and 6 deletions

View file

@ -632,6 +632,7 @@ class ModelResponseIterator:
self.tool_name_reverse_map: dict[str, str] = tool_name_reverse_map or {}
# Generate response ID once per stream to match OpenAI-compatible behavior
self.response_id = _generate_id()
self.served_model: str | None = None
# Track if we're currently streaming a response_format tool
self.is_response_format_tool: bool = False
@ -1067,6 +1068,9 @@ class ModelResponseIterator:
}
"""
message_start_block: Final = MessageStartBlock(**chunk)
start_message: Final = message_start_block["message"]
if "model" in start_message:
self.served_model = start_message["model"]
if "usage" in message_start_block["message"]:
usage = self._handle_usage(anthropic_usage_chunk=message_start_block["message"]["usage"])
elif type_chunk == "error":
@ -1098,6 +1102,7 @@ class ModelResponseIterator:
],
usage=usage,
id=self.response_id,
model=self.served_model,
)
return returned_chunk

View file

@ -2719,3 +2719,74 @@ class TestRustChatCompletionsHook:
"model": "m",
"messages": [],
}
def _served_model_stream_chunks(model: str | None) -> list[dict[str, object]]:
return [
{
"type": "message_start",
"message": {
"id": "msg_served",
"type": "message",
"role": "assistant",
"content": [],
"usage": {"input_tokens": 10, "output_tokens": 1},
**({"model": model} if model is not None else {}),
},
},
{
"type": "content_block_start",
"index": 0,
"content_block": {"type": "text", "text": ""},
},
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "text_delta", "text": "Hello"},
},
{"type": "content_block_stop", "index": 0},
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"output_tokens": 2},
},
{"type": "message_stop"},
]
def test_message_start_model_is_carried_on_stream_chunks():
iterator: Final = ModelResponseIterator(None, sync_stream=True)
parsed: Final = [iterator.chunk_parser(chunk) for chunk in _served_model_stream_chunks("claude-served-1")]
assert all(chunk.model == "claude-served-1" for chunk in parsed)
def test_message_start_without_model_leaves_chunk_model_unset():
iterator: Final = ModelResponseIterator(None, sync_stream=True)
parsed: Final = [iterator.chunk_parser(chunk) for chunk in _served_model_stream_chunks(None)]
assert all(chunk.model is None for chunk in parsed)
def test_served_model_reaches_assembled_stream_through_custom_stream_wrapper():
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
served_model: Final = "claude-served-1"
sse_lines: Final = [f"data: {json.dumps(chunk)}\n".encode() for chunk in _served_model_stream_chunks(served_model)]
iterator: Final = ModelResponseIterator(iter(sse_lines), sync_stream=True)
wrapper: Final = CustomStreamWrapper(
completion_stream=iter(iterator),
model="anthropic/claude-requested",
custom_llm_provider="anthropic",
logging_obj=MagicMock(),
)
chunks: Final = list(wrapper)
assert len(chunks) > 1
for chunk in chunks[1:]:
assert chunk._hidden_params["provider_response_model"] == served_model
assembled: Final = litellm.stream_chunk_builder(chunks=list(chunks), messages=[{"role": "user", "content": "hi"}])
assert assembled._hidden_params["provider_response_model"] == served_model

View file

@ -1291,8 +1291,9 @@ class TestToolPermissionGuardrailAnthropicMessages:
async def test_rewrite_mode_keeps_the_stream_identity_it_had_before_the_shared_helper(self):
"""Well-formed SSE must round-trip exactly as it did before the helpers were shared.
The shared module can stamp the upstream message id and model onto the assembled response
for callers that ask for it; this path never did, and a client reads those bytes.
The shared module can stamp the upstream message id onto the assembled response for
callers that ask for it; this path never did, and a client reads those bytes. The model,
though, is now the upstream's, matching what the untouched passthrough shows clients.
"""
with patch.object(self.rewriting, "should_run_guardrail", return_value=True):
out = await self._drain(self.rewriting, self._sse_chunks("Read"))
@ -1304,7 +1305,7 @@ class TestToolPermissionGuardrailAnthropicMessages:
if line.startswith("data: ") and json.loads(line[6:]).get("type") == "message_start"
)["message"]
assert message_start["id"].startswith("chatcmpl-"), "the rewritten stream must not adopt the upstream message id"
assert message_start["model"] == "unknown-model", "the rewritten stream must not adopt the upstream model"
assert message_start["model"] == "claude-sonnet-4-5", "the rewritten stream reports the model the upstream served"
@pytest.mark.asyncio
async def test_message_start_without_a_dict_message_fails_closed(self):

View file

@ -2541,7 +2541,7 @@ class TestRecordPartialUsageForFailure:
function_id="test-partial-usage-failure",
)
def _interrupted_chunks(self):
def _interrupted_chunks(self, *, model: str = "claude-sonnet-5"):
return [
self._sse(
"message_start",
@ -2551,7 +2551,7 @@ class TestRecordPartialUsageForFailure:
"id": "msg_abc",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-5",
"model": model,
"content": [],
"stop_reason": None,
"stop_sequence": None,
@ -2588,7 +2588,7 @@ class TestRecordPartialUsageForFailure:
AnthropicPassthroughLoggingHandler.record_partial_usage_for_failure(
litellm_logging_obj=logging_obj,
request_body={"model": "claude-unpriced-test-model", "stream": True},
all_chunks=self._interrupted_chunks(),
all_chunks=self._interrupted_chunks(model="claude-unpriced-test-model"),
)
usage = logging_obj.model_call_details["combined_usage_object"]