diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml index 3725e0f5805..f2bbd75682b 100644 --- a/.github/workflows/test-unit-proxy-db.yml +++ b/.github/workflows/test-unit-proxy-db.yml @@ -167,6 +167,7 @@ jobs: tests/proxy_unit_tests/test_proxy_setting_guardrails.py tests/proxy_unit_tests/test_banned_keyword_list.py tests/proxy_unit_tests/test_unit_test_proxy_hooks.py + tests/proxy_unit_tests/test_anthropic_messages_pre_call_hook.py workers: 4 dist: loadscope timeout: 15 diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py index 9d1e921cce4..6c38c91750a 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py @@ -9,7 +9,7 @@ import asyncio import contextvars from collections.abc import AsyncIterator, Coroutine, Iterator from functools import partial -from typing import Any, Final, cast +from typing import TYPE_CHECKING, Any, Final, cast import litellm from litellm.litellm_core_utils.exception_mapping_utils import exception_type @@ -40,6 +40,10 @@ from ..utils import is_reasoning_auto_summary_enabled from .interceptors import get_messages_interceptors from .utils import AnthropicMessagesRequestUtils, mock_response +if TYPE_CHECKING: + from litellm.caching.caching import DualCache + from litellm.proxy._types import UserAPIKeyAuth + # Providers that are routed directly to the OpenAI Responses API instead of # going through chat/completions. _RESPONSES_API_PROVIDERS: Final = frozenset({"openai"}) @@ -119,10 +123,10 @@ async def _execute_pre_request_hooks( **kwargs, ) -> dict: """ - Execute pre-request hooks from CustomLogger callbacks. + Execute pre-call and pre-request hooks from CustomLogger callbacks. Allows CustomLoggers to modify request parameters before the API call. - Used for WebSearch tool conversion, stream modification, etc. + Used for proxy guardrails, WebSearch tool conversion, stream modification, etc. Args: model: Model name @@ -145,6 +149,8 @@ async def _execute_pre_request_hooks( # Build complete request kwargs dict request_kwargs = { + "model": model, + "messages": messages, "tools": tools, "stream": stream, "litellm_params": { @@ -162,8 +168,28 @@ async def _execute_pre_request_hooks( if not isinstance(callback, _CustomLogger): continue - # Call the pre-request hook - modified_kwargs = await callback.async_pre_request_hook(model, messages, request_kwargs) + if ( + "async_pre_call_hook" in vars(callback.__class__) + and callback.__class__.async_pre_call_hook != _CustomLogger.async_pre_call_hook + ): + modified_kwargs = await callback.async_pre_call_hook( + user_api_key_dict=cast("UserAPIKeyAuth", request_kwargs.get("user_api_key_dict")), + cache=cast("DualCache", request_kwargs.get("cache")), + data=request_kwargs, + call_type="anthropic_messages", + ) + if isinstance(modified_kwargs, Exception): + raise modified_kwargs + if isinstance(modified_kwargs, str): + raise ValueError(modified_kwargs) + if modified_kwargs is not None: + request_kwargs = modified_kwargs + + modified_kwargs = await callback.async_pre_request_hook( + request_kwargs.get("model", model), + request_kwargs.get("messages", messages), + request_kwargs, + ) # If hook returned modified kwargs, use them if modified_kwargs is not None: @@ -286,7 +312,15 @@ async def anthropic_messages( tools=tools, stream=stream, custom_llm_provider=custom_llm_provider, + max_tokens=max_tokens, + metadata=metadata, + stop_sequences=stop_sequences, + system=system, + temperature=temperature, + thinking=thinking, tool_choice=tool_choice, + top_k=top_k, + top_p=top_p, **kwargs, ) @@ -294,6 +328,9 @@ async def anthropic_messages( # that we may forward explicitly downstream, so we (a) honor pre-request hook # overrides and (b) avoid duplicate-keyword conflicts when splatting `kwargs` # into call sites that already pass these as named arguments. + model = request_kwargs.pop("model", model) + messages = request_kwargs.pop("messages", messages) + max_tokens = request_kwargs.pop("max_tokens", max_tokens) tools = request_kwargs.pop("tools", tools) stream = request_kwargs.pop("stream", stream) metadata = request_kwargs.pop("metadata", metadata) @@ -382,12 +419,11 @@ async def anthropic_messages( api_base=api_base, client=client, custom_llm_provider=custom_llm_provider, - # messages were already empty-content-block sanitized at the top of this - # function and are NOT reassigned before this dispatch, so the handler - # can skip its (otherwise redundant) second full-messages scan. Passed - # explicitly (not via **kwargs) so it only affects this direct - # dispatch -- interceptor / sync entry points still sanitize. - _litellm_messages_presanitized=True, + # Messages are sanitized above. When no callback can edit them + # afterwards, skip the handler's second scan. Passed explicitly + # (not via **kwargs) so it only affects this direct dispatch -- + # interceptor / sync entry points still sanitize. + _litellm_messages_presanitized=not litellm.callbacks, **kwargs, ) ctx: Final = contextvars.copy_context() diff --git a/tests/proxy_unit_tests/test_anthropic_messages_pre_call_hook.py b/tests/proxy_unit_tests/test_anthropic_messages_pre_call_hook.py new file mode 100644 index 00000000000..4386545453d --- /dev/null +++ b/tests/proxy_unit_tests/test_anthropic_messages_pre_call_hook.py @@ -0,0 +1,133 @@ +import asyncio +import os +from unittest import mock + +import litellm +import pytest +from fastapi.testclient import TestClient + +from litellm.integrations.custom_logger import CustomLogger +from litellm.llms.anthropic.experimental_pass_through.messages.handler import ( + anthropic_messages, +) +from litellm.proxy.proxy_server import ( + app, + cleanup_router_config_variables, + initialize, +) + +EXAMPLE_ANTHROPIC_MESSAGES_RESULT = { + "id": "msg_test", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "Hello from LiteLLM"}], + "model": "gpt-3.5-turbo", + "stop_reason": "end_turn", + "usage": {"input_tokens": 5, "output_tokens": 5}, +} + + +def mock_patch_anthropic_messages(): + return mock.patch( + "litellm.proxy.proxy_server.llm_router.anthropic_messages", + return_value=EXAMPLE_ANTHROPIC_MESSAGES_RESULT, + ) + + +@pytest.fixture(scope="function") +def fake_env_vars(monkeypatch): + monkeypatch.setenv("OPENAI_API_KEY", "fake_openai_api_key") + monkeypatch.setenv("OPENAI_API_BASE", "http://fake-openai-api-base") + monkeypatch.setenv("AZURE_AI_API_BASE", "http://fake-azure-api-base") + monkeypatch.setenv("AZURE_OPENAI_API_KEY", "fake_azure_api_key") + monkeypatch.setenv("AZURE_SWEDEN_API_BASE", "http://fake-azure-sweden-api-base") + monkeypatch.setenv("REDIS_HOST", "localhost") + + +@pytest.fixture(scope="function") +def client_no_auth(fake_env_vars): + cleanup_router_config_variables() + test_dir = os.path.dirname(os.path.abspath(__file__)) + config_path = os.path.join(test_dir, "test_configs", "test_config_no_auth.yaml") + asyncio.run(initialize(config=config_path, debug=True)) + return TestClient(app) + + +@mock_patch_anthropic_messages() +def test_anthropic_messages_runs_proxy_async_pre_call_hook(mock_anthropic_messages, client_no_auth, monkeypatch): + hook_calls = [] + + class AnthropicMessagesPreCallHook(CustomLogger): + 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 + + monkeypatch.setattr(litellm, "callbacks", [AnthropicMessagesPreCallHook()]) + + response = client_no_auth.post( + "/v1/messages", + json={ + "model": "test_openai_models", + "max_tokens": 100, + "messages": [{"role": "user", "content": "hi"}], + }, + ) + + assert response.status_code == 200 + assert response.json()["content"][0]["text"] == "Hello from LiteLLM" + assert hook_calls == ["anthropic_messages"] + mock_anthropic_messages.assert_called_once() + metadata = mock_anthropic_messages.call_args.kwargs.get("metadata") + assert metadata == {"source": "unit-test"} + + +@pytest.mark.asyncio +async def test_experimental_anthropic_messages_runs_proxy_async_pre_call_hook( + monkeypatch, +): + hook_calls = [] + + class AnthropicMessagesPreCallHook(CustomLogger): + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type, **kwargs): + hook_calls.append(("pre_call", call_type, data["model"])) + data["metadata"] = { + **(data.get("metadata") or {}), + "source": "experimental-unit-test", + } + data["temperature"] = 0.2 + return data + + async def async_pre_request_hook(self, model, messages, kwargs): + hook_calls.append(("pre_request", model, kwargs.get("temperature"))) + updated = dict(kwargs) + updated["top_p"] = 0.4 + return updated + + monkeypatch.setattr(litellm, "callbacks", [AnthropicMessagesPreCallHook()]) + monkeypatch.setattr(litellm, "use_chat_completions_url_for_anthropic_messages", True) + + with mock.patch( + "litellm.llms.anthropic.experimental_pass_through.messages.handler.anthropic_messages_handler", + return_value=EXAMPLE_ANTHROPIC_MESSAGES_RESULT, + ) as mock_handler: + response = await anthropic_messages( + model="openai/gpt-4o-mini", + max_tokens=100, + messages=[{"role": "user", "content": "hi"}], + metadata={"existing": "keep"}, + custom_llm_provider="openai", + ) + + assert response == EXAMPLE_ANTHROPIC_MESSAGES_RESULT + assert hook_calls == [ + ("pre_call", "anthropic_messages", "openai/gpt-4o-mini"), + ("pre_request", "openai/gpt-4o-mini", 0.2), + ] + mock_handler.assert_called_once() + assert mock_handler.call_args.kwargs["metadata"] == { + "existing": "keep", + "source": "experimental-unit-test", + } + assert mock_handler.call_args.kwargs["temperature"] == 0.2 + assert mock_handler.call_args.kwargs["top_p"] == 0.4