mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
commit
930ec9643a
4 changed files with 83 additions and 6 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue