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..bf50b541744 --- /dev/null +++ b/tests/proxy_unit_tests/test_anthropic_messages_pre_call_hook.py @@ -0,0 +1,135 @@ +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( # test-quality-ok: hook metadata is only observable on the router call this route forwards + "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( # test-quality-ok: pre_call and pre_request mutations are only observable on the downstream handler kwargs + "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 diff --git a/uv.lock b/uv.lock index eb4cdef76f1..8ae0dcbc5bb 100644 --- a/uv.lock +++ b/uv.lock @@ -315,16 +315,16 @@ vertex = [ [[package]] name = "anyio" -version = "4.13.0" +version = "4.14.2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "exceptiongroup", marker = "python_full_version < '3.11'" }, { name = "idna" }, { name = "typing-extensions", marker = "python_full_version < '3.13'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/19/14/2c5dd9f512b66549ae92767a9c7b330ae88e1932ca57876909410251fe13/anyio-4.13.0.tar.gz", hash = "sha256:334b70e641fd2221c1505b3890c69882fe4a2df910cba14d97019b90b24439dc", size = 231622, upload-time = "2026-03-24T12:59:09.671Z" } +sdist = { url = "https://files.pythonhosted.org/packages/61/cc/a381afa6efea9f496eff839d4a6a1aed3bfafc7b3ab4b0d1b243a12573dd/anyio-4.14.2.tar.gz", hash = "sha256:cfa139f3ed1a23ee8f88a145ddb5ac7605b8bbfd8592baacd7ce3d8bb4313c7f", size = 260176, upload-time = "2026-07-12T20:29:07.082Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/da/42/e921fccf5015463e32a3cf6ee7f980a6ed0f395ceeaa45060b61d86486c2/anyio-4.13.0-py3-none-any.whl", hash = "sha256:08b310f9e24a9594186fd75b4f73f4a4152069e3853f1ed8bfbf58369f4ad708", size = 114353, upload-time = "2026-03-24T12:59:08.246Z" }, + { url = "https://files.pythonhosted.org/packages/da/35/f2287558c17e29fafc8ef3daf819bb9834061cfa43bff8014f7df7f63bdc/anyio-4.14.2-py3-none-any.whl", hash = "sha256:9f505dda5ac9f0c8309b5e8bd445a8c2bf7246f3ce950121e45ea15bc41d1494", size = 125813, upload-time = "2026-07-12T20:29:05.763Z" }, ] [[package]] @@ -9080,11 +9080,11 @@ wheels = [ [[package]] name = "soupsieve" -version = "2.8.4" +version = "2.9.2" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/47/2c/0a5f6f8ee0d5589e48c7640213ed5175d52cf540a06725b628cc1a45d6ce/soupsieve-2.8.4.tar.gz", hash = "sha256:e121fd02e975c695e4e9e8774a5ee35d74714b59307868dcc5319ad2d9e3328e", size = 121110, upload-time = "2026-05-24T13:55:57.154Z" } +sdist = { url = "https://files.pythonhosted.org/packages/69/99/a6ca3beb3ccacb41fb3321d8a60e5566f9e6467601ef8eba6a17e1b89778/soupsieve-2.9.2.tar.gz", hash = "sha256:4a55d8cf158a9c2e587fa4922f1bbb91d68ac829e2d6f25403a85747c71daf74", size = 122445, upload-time = "2026-08-07T00:57:24.801Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/5e/f5/0c41cb68dcae6b7de4fac4188a3a9589e21fb31df21ea3a2e888db95e6c9/soupsieve-2.8.4-py3-none-any.whl", hash = "sha256:e7e6b0769c8f51ed59acab6e994b00621096cfb1c640a7509295987388fbaf65", size = 37304, upload-time = "2026-05-24T13:55:55.406Z" }, + { url = "https://files.pythonhosted.org/packages/eb/dc/ad025c1ee131eba60c69f4dd5779b18fcf1e6b21a343e2162a84d5d133c7/soupsieve-2.9.2-py3-none-any.whl", hash = "sha256:8089a26fd974ca7a1f30276d3d8492ab266ab15af581642dfe8aa162e0c1c823", size = 37370, upload-time = "2026-08-07T00:57:23.524Z" }, ] [[package]]