fix(vertex-ai): preserve provider metadata in Claude responses

This commit is contained in:
林SO 2026-08-05 22:05:58 +08:00
parent f4308bc124
commit e3638be2e0
4 changed files with 115 additions and 8 deletions

View file

@ -206,6 +206,7 @@ class AnthropicChatCompletion(BaseLLM):
model: str,
messages: list,
api_base: str,
custom_llm_provider: str,
custom_prompt_dict: dict,
model_response: ModelResponse,
print_verbose: Callable,
@ -245,7 +246,7 @@ class AnthropicChatCompletion(BaseLLM):
streamwrapper: Final = CustomStreamWrapper(
completion_stream=completion_stream,
model=model,
custom_llm_provider="anthropic",
custom_llm_provider=custom_llm_provider,
logging_obj=logging_obj,
_response_headers=process_anthropic_headers(headers),
)
@ -413,6 +414,7 @@ class AnthropicChatCompletion(BaseLLM):
messages=messages,
data=data,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
custom_prompt_dict=custom_prompt_dict,
model_response=model_response,
print_verbose=print_verbose,
@ -489,7 +491,7 @@ class AnthropicChatCompletion(BaseLLM):
return CustomStreamWrapper(
completion_stream=completion_stream,
model=model,
custom_llm_provider="anthropic",
custom_llm_provider=custom_llm_provider,
logging_obj=logging_obj,
_response_headers=process_anthropic_headers(headers),
)

View file

@ -2593,8 +2593,6 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
speed: str | None = None,
tool_name_reverse_map: dict[str, str] | None = None,
):
_hidden_params: Final[dict] = {}
_hidden_params["additional_headers"] = process_anthropic_headers(dict(raw_response.headers))
if "error" in completion_response:
response_headers: Final = getattr(raw_response, "headers", None)
raise AnthropicError(
@ -2668,7 +2666,6 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
_message = json_mode_message
model_response.choices[0].message = _message
model_response._hidden_params["original_response"] = completion_response["content"]
model_response.choices[0].finish_reason = cast(
OpenAIChatCompletionFinishReason,
map_finish_reason(completion_response["stop_reason"]),
@ -2685,8 +2682,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
model_response.created = int(time.time())
model_response.model = completion_response["model"]
_hidden_params["provider_specific_fields"] = provider_specific_fields
model_response._hidden_params = _hidden_params
model_response._hidden_params = {
**model_response._hidden_params,
"additional_headers": process_anthropic_headers(dict(raw_response.headers)),
"original_response": completion_response["content"],
"provider_specific_fields": provider_specific_fields,
}
return model_response
def get_prefix_prompt(self, messages: list[AllMessageValues]) -> str | None:

View file

@ -10,13 +10,18 @@ import pytest
import litellm
from litellm._uuid import uuid
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
from litellm.llms.anthropic.chat.handler import ModelResponseIterator, make_call
from litellm.llms.anthropic.chat.handler import (
AnthropicChatCompletion,
ModelResponseIterator,
make_call,
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.types.llms.openai import (
ChatCompletionToolCallChunk,
ChatCompletionToolCallFunctionChunk,
)
from litellm.types.responses.main import OutputCodeInterpreterCall
from litellm.types.utils import ModelResponse
@pytest.mark.asyncio
@ -2452,3 +2457,71 @@ def test_served_model_reaches_assembled_stream_through_custom_stream_wrapper():
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
def test_streaming_completion_preserves_custom_llm_provider():
response = MagicMock()
response.status_code = 200
response.headers = {}
response.iter_lines.return_value = iter(())
client = MagicMock(spec=HTTPHandler)
client.post.return_value = response
logging_obj = MagicMock()
logging_obj.model_call_details = {
"custom_llm_provider": "vertex_ai",
"litellm_params": {},
}
result = AnthropicChatCompletion().completion(
model="claude-sonnet-4-5@20250929",
messages=[{"role": "user", "content": "Hello"}],
api_base="https://example.com/v1/messages",
custom_llm_provider="vertex_ai",
custom_prompt_dict={},
model_response=ModelResponse(),
print_verbose=MagicMock(),
encoding=None,
api_key="test-key",
logging_obj=logging_obj,
optional_params={"stream": True, "max_tokens": 16, "is_vertex_request": True},
timeout=60.0,
litellm_params={},
client=client,
)
assert result.custom_llm_provider == "vertex_ai"
@pytest.mark.asyncio
async def test_async_streaming_completion_preserves_custom_llm_provider():
response = MagicMock()
response.headers = {}
response.aiter_lines.return_value = iter(())
client = MagicMock(spec=AsyncHTTPHandler)
client.post = AsyncMock(return_value=response)
logging_obj = MagicMock()
logging_obj.model_call_details = {
"custom_llm_provider": "vertex_ai",
"litellm_params": {},
}
completion = AnthropicChatCompletion().completion(
model="claude-sonnet-4-5@20250929",
messages=[{"role": "user", "content": "Hello"}],
api_base="https://example.com/v1/messages",
custom_llm_provider="vertex_ai",
custom_prompt_dict={},
model_response=ModelResponse(),
print_verbose=MagicMock(),
encoding=None,
api_key="test-key",
logging_obj=logging_obj,
optional_params={"stream": True, "max_tokens": 16, "is_vertex_request": True},
timeout=60.0,
litellm_params={},
acompletion=True,
client=client,
)
result = await completion
assert result.custom_llm_provider == "vertex_ai"

View file

@ -6798,3 +6798,34 @@ def test_chat_dummy_tool_result_for_an_orphaned_tool_call_replays_a_byte_identic
_assert_prefix_stable(requests)
assert [m["role"] for m in requests[0]["messages"]] == ["user", "assistant", "user"]
assert requests[0]["messages"][2]["content"][0]["type"] == "tool_result"
def test_transform_parsed_response_preserves_existing_hidden_params():
from litellm.types.utils import ModelResponse
config = AnthropicConfig()
raw_response = MagicMock()
raw_response.headers = {"request-id": "req_vertex"}
raw_response.status_code = 200
completion_response = {
"id": "msg_vertex",
"model": "claude-sonnet-4-5@20250929",
"stop_reason": "end_turn",
"usage": {"input_tokens": 10, "output_tokens": 5},
"content": [{"type": "text", "text": "Hello"}],
}
model_response = ModelResponse()
model_response._hidden_params = {
"custom_llm_provider": "vertex_ai",
"region_name": "us-east5",
}
result = config.transform_parsed_response(
completion_response=completion_response,
raw_response=raw_response,
model_response=model_response,
)
assert result._hidden_params["custom_llm_provider"] == "vertex_ai"
assert result._hidden_params["region_name"] == "us-east5"
assert result._hidden_params["original_response"] == completion_response["content"]