style(proxy): format anthropic messages hook fix

This commit is contained in:
JerryLee 2026-05-11 16:43:33 +10:00
parent d6abae7e6d
commit 19a8e3fdd6
2 changed files with 50 additions and 17 deletions

View file

@ -128,7 +128,8 @@ async def _execute_pre_request_hooks(
if (
"async_pre_call_hook" in vars(callback.__class__)
and callback.__class__.async_pre_call_hook != _CustomLogger.async_pre_call_hook
and callback.__class__.async_pre_call_hook
!= _CustomLogger.async_pre_call_hook
):
# Keep the experimental Messages path in parity with the proxy
# pre-call hook contract used by the standard /v1/messages path.
@ -233,7 +234,9 @@ async def anthropic_messages(
"""
Async: Make llm api request in Anthropic /messages API spec
"""
original_stream = stream or kwargs.get("_websearch_interception_converted_stream", False)
original_stream = stream or kwargs.get(
"_websearch_interception_converted_stream", False
)
# Execute pre-request hooks to allow CustomLoggers to modify request
request_kwargs = await _execute_pre_request_hooks(
@ -272,7 +275,9 @@ async def anthropic_messages(
# The litellm_params dict may have been overwritten by **kwargs in
# _execute_pre_request_hooks, so fall back to get_llm_provider() if needed.
if not custom_llm_provider:
custom_llm_provider = request_kwargs.get("litellm_params", {}).get("custom_llm_provider")
custom_llm_provider = request_kwargs.get("litellm_params", {}).get(
"custom_llm_provider"
)
if not custom_llm_provider:
try:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model)
@ -435,7 +440,9 @@ def anthropic_messages_handler(
# Check if stream was converted for WebSearch interception
# This is set in the async wrapper above when stream=True is converted to stream=False
if kwargs.get("_websearch_interception_converted_stream", False):
litellm_logging_obj.model_call_details["websearch_interception_converted_stream"] = True
litellm_logging_obj.model_call_details[
"websearch_interception_converted_stream"
] = True
if litellm_params.mock_response and isinstance(litellm_params.mock_response, str):
return mock_response(
@ -447,10 +454,14 @@ def anthropic_messages_handler(
anthropic_messages_provider_config: Optional[BaseAnthropicMessagesConfig] = None
if custom_llm_provider is not None and custom_llm_provider in [provider.value for provider in LlmProviders]:
anthropic_messages_provider_config = ProviderConfigManager.get_provider_anthropic_messages_config(
model=model,
provider=litellm.LlmProviders(custom_llm_provider),
if custom_llm_provider is not None and custom_llm_provider in [
provider.value for provider in LlmProviders
]:
anthropic_messages_provider_config = (
ProviderConfigManager.get_provider_anthropic_messages_config(
model=model,
provider=litellm.LlmProviders(custom_llm_provider),
)
)
if anthropic_messages_provider_config is None:
# Route to Responses API for OpenAI / Azure, chat/completions for everything else.
@ -476,8 +487,14 @@ def anthropic_messages_handler(
**kwargs,
)
if _should_route_to_responses_api(custom_llm_provider):
return LiteLLMMessagesToResponsesAPIHandler.anthropic_messages_handler(**_shared_kwargs)
return LiteLLMMessagesToCompletionTransformationHandler.anthropic_messages_handler(**_shared_kwargs)
return LiteLLMMessagesToResponsesAPIHandler.anthropic_messages_handler(
**_shared_kwargs
)
return (
LiteLLMMessagesToCompletionTransformationHandler.anthropic_messages_handler(
**_shared_kwargs
)
)
if custom_llm_provider is None:
raise ValueError(
@ -486,11 +503,16 @@ def anthropic_messages_handler(
local_vars.update(kwargs)
anthropic_messages_optional_request_params = (
AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(params=local_vars)
AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
params=local_vars
)
)
if is_reasoning_auto_summary_enabled():
thinking_param = anthropic_messages_optional_request_params.get("thinking")
if isinstance(thinking_param, dict) and thinking_param.get("type") != "disabled":
if (
isinstance(thinking_param, dict)
and thinking_param.get("type") != "disabled"
):
anthropic_messages_optional_request_params["thinking"] = {
**thinking_param,
"display": "summarized",
@ -500,7 +522,9 @@ def anthropic_messages_handler(
model=model,
messages=messages,
anthropic_messages_provider_config=anthropic_messages_provider_config,
anthropic_messages_optional_request_params=dict(anthropic_messages_optional_request_params),
anthropic_messages_optional_request_params=dict(
anthropic_messages_optional_request_params
),
_is_async=is_async,
client=client,
custom_llm_provider=custom_llm_provider,

View file

@ -50,11 +50,15 @@ def client_no_auth(fake_env_vars):
@mock_patch_anthropic_messages()
def test_anthropic_messages_runs_proxy_async_pre_call_hook(mock_anthropic_messages, client_no_auth):
def test_anthropic_messages_runs_proxy_async_pre_call_hook(
mock_anthropic_messages, client_no_auth
):
hook_calls = []
class AnthropicMessagesPreCallHook(CustomLogger):
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type, **kwargs):
async def async_pre_call_hook(
self, user_api_key_dict, cache, data, call_type, **kwargs
):
hook_calls.append(call_type)
data["metadata"] = {**(data.get("metadata") or {}), "source": "unit-test"}
return data
@ -76,7 +80,10 @@ def test_anthropic_messages_runs_proxy_async_pre_call_hook(mock_anthropic_messag
assert response.json()["content"][0]["text"] == "Hello from LiteLLM"
assert hook_calls == ["anthropic_messages"]
mock_anthropic_messages.assert_called_once()
assert mock_anthropic_messages.call_args.kwargs["metadata"]["source"] == "unit-test"
assert (
mock_anthropic_messages.call_args.kwargs["metadata"]["source"]
== "unit-test"
)
finally:
litellm.callbacks = original_callbacks
@ -90,7 +97,9 @@ async def test_experimental_anthropic_messages_runs_proxy_async_pre_call_hook():
hook_calls = []
class AnthropicMessagesPreCallHook(CustomLogger):
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type, **kwargs):
async def async_pre_call_hook(
self, user_api_key_dict, cache, data, call_type, **kwargs
):
hook_calls.append((call_type, data["model"]))
data["metadata"] = {
**(data.get("metadata") or {}),