mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
fix(proxy): run pre-call hooks for Anthropic messages passthrough
Rebase the experimental Messages pass-through hook fix onto current litellm_internal_staging. The handler runs async_pre_call_hook and then async_pre_request_hook, and forwards the mutated request downstream
This commit is contained in:
parent
252c71c0b2
commit
d7694c26fd
3 changed files with 181 additions and 11 deletions
1
.github/workflows/test-unit-proxy-db.yml
vendored
1
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
133
tests/proxy_unit_tests/test_anthropic_messages_pre_call_hook.py
Normal file
133
tests/proxy_unit_tests/test_anthropic_messages_pre_call_hook.py
Normal file
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue