mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge 7768fe9b00 into 02035120e4
This commit is contained in:
commit
e5615e834d
4 changed files with 106 additions and 53 deletions
|
|
@ -459,7 +459,10 @@ async def acompletion_with_mcp(
|
|||
)
|
||||
|
||||
# Make follow-up call with streaming
|
||||
follow_up_call_args: Final = dict(self.base_call_args)
|
||||
follow_up_call_args: Final = LiteLLM_Proxy_MCP_Handler.prepare_follow_up_call_params(
|
||||
self.base_call_args,
|
||||
original_stream_setting=True,
|
||||
)
|
||||
follow_up_call_args["messages"] = follow_up_messages
|
||||
follow_up_call_args["stream"] = True
|
||||
# Ensure follow-up call doesn't trigger MCP handler again
|
||||
|
|
@ -631,7 +634,10 @@ async def acompletion_with_mcp(
|
|||
)
|
||||
|
||||
# Make follow-up call with original stream setting
|
||||
follow_up_call_args: Final = dict(base_call_args)
|
||||
follow_up_call_args: Final = LiteLLM_Proxy_MCP_Handler.prepare_follow_up_call_params(
|
||||
base_call_args,
|
||||
original_stream_setting=stream,
|
||||
)
|
||||
follow_up_call_args["messages"] = follow_up_messages
|
||||
follow_up_call_args["stream"] = stream
|
||||
|
||||
|
|
|
|||
|
|
@ -71,6 +71,32 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
This handles when a user passes mcp server_url="litellm_proxy" in their tools.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def prepare_chained_call_params(
|
||||
params: Mapping[str, Any],
|
||||
) -> dict[str, Any]: # mutable-ok: returns a sanitized copy for the next provider call
|
||||
"""Copy request params without state owned by the previous LLM call.
|
||||
|
||||
MCP auto-execution keeps the trace identifier so chained rounds remain
|
||||
correlated, but each provider call must create its own logging object and
|
||||
call identifier. Reusing either makes success dispatch one-shot and drops
|
||||
spend rows for later rounds.
|
||||
"""
|
||||
follow_up_params = dict(params)
|
||||
follow_up_params.pop("litellm_logging_obj", None)
|
||||
follow_up_params.pop("litellm_call_id", None)
|
||||
if follow_up_params.get("web_search_options") is None:
|
||||
follow_up_params.pop("web_search_options", None)
|
||||
|
||||
nested_params = follow_up_params.get("litellm_params")
|
||||
if isinstance(nested_params, dict):
|
||||
nested_params = dict(nested_params)
|
||||
nested_params.pop("litellm_logging_obj", None)
|
||||
nested_params.pop("litellm_call_id", None)
|
||||
follow_up_params["litellm_params"] = nested_params
|
||||
|
||||
return follow_up_params
|
||||
|
||||
@staticmethod
|
||||
def _get_parent_request_tags(kwargs: dict[str, Any] | None) -> list[str]:
|
||||
"""Tags from the parent LLM request, using the same extraction logic as standard logging (incl. User-Agent)."""
|
||||
|
|
@ -702,13 +728,19 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
},
|
||||
}
|
||||
]
|
||||
tool_logging_call_id = litellm_call_id or str(uuid.uuid4())
|
||||
# SpendLogs.request_id is unique. The parent LLM call ID is
|
||||
# therefore metadata, not the tool execution's call ID: reusing
|
||||
# it causes all but the first MCP tool row to be skipped by the
|
||||
# database's duplicate protection.
|
||||
tool_logging_call_id = str(uuid.uuid4())
|
||||
logging_metadata: dict[str, object] = {
|
||||
"tool_call_id": tool_call_id,
|
||||
"tool_name": sanitized_tool_name,
|
||||
"server_name": server_name,
|
||||
"headers": logging_safe_headers,
|
||||
}
|
||||
if litellm_call_id:
|
||||
logging_metadata["parent_litellm_call_id"] = litellm_call_id
|
||||
logging_request_data = {
|
||||
"model": f"MCP: {tool_name}",
|
||||
"metadata": logging_metadata,
|
||||
|
|
@ -1229,14 +1261,14 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
return initial_params
|
||||
|
||||
@staticmethod
|
||||
def _prepare_follow_up_call_params(call_params: dict[str, Any], original_stream_setting: bool) -> dict[str, Any]:
|
||||
def prepare_follow_up_call_params(call_params: dict[str, Any], original_stream_setting: bool) -> dict[str, Any]:
|
||||
"""
|
||||
Prepare call parameters for the follow-up LLM call after tool execution.
|
||||
|
||||
Restores the original streaming setting and removes tool_choice since
|
||||
we're now providing tool results, not requesting tool calls.
|
||||
"""
|
||||
follow_up_params: Final = call_params.copy()
|
||||
follow_up_params: Final = LiteLLM_Proxy_MCP_Handler.prepare_chained_call_params(call_params)
|
||||
|
||||
# Restore original streaming setting for follow-up call
|
||||
follow_up_params["stream"] = original_stream_setting
|
||||
|
|
@ -1246,6 +1278,8 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
|
||||
return follow_up_params
|
||||
|
||||
_prepare_follow_up_call_params = prepare_follow_up_call_params
|
||||
|
||||
@staticmethod
|
||||
def _add_mcp_output_elements_to_response(
|
||||
response: ResponsesAPIResponse,
|
||||
|
|
|
|||
|
|
@ -788,7 +788,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
)
|
||||
|
||||
# Make follow-up call with streaming
|
||||
follow_up_params: Final = self.original_request_params.copy()
|
||||
follow_up_params: Final = LiteLLM_Proxy_MCP_Handler.prepare_chained_call_params(
|
||||
self.original_request_params
|
||||
)
|
||||
follow_up_params.update(
|
||||
{
|
||||
"input": follow_up_input,
|
||||
|
|
|
|||
|
|
@ -1,20 +1,20 @@
|
|||
import importlib
|
||||
import subprocess
|
||||
import sys
|
||||
import textwrap
|
||||
import types
|
||||
from typing import Any, cast
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
import importlib
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
|
||||
from litellm.types.responses.main import OutputFunctionToolCall
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
|
||||
LiteLLM_Proxy_MCP_Handler,
|
||||
)
|
||||
from typing import Any, cast
|
||||
from litellm.types.utils import ModelResponse
|
||||
from litellm.types.responses.main import OutputFunctionToolCall
|
||||
|
||||
|
||||
class _DummyMCPResult:
|
||||
|
|
@ -105,9 +105,7 @@ def test_extract_tool_calls_from_chat_response_handles_tool_calls():
|
|||
object="chat.completion",
|
||||
)
|
||||
|
||||
tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response(
|
||||
response
|
||||
)
|
||||
tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response(response)
|
||||
|
||||
assert len(tool_calls) == 1
|
||||
assert tool_calls[0]["function"]["name"] == "foo"
|
||||
|
|
@ -177,9 +175,7 @@ def test_transform_mcp_tools_to_openai_uses_chat_format(monkeypatch):
|
|||
fake_transform_responses,
|
||||
)
|
||||
|
||||
chat_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(
|
||||
["tool"], target_format="chat"
|
||||
)
|
||||
chat_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(["tool"], target_format="chat")
|
||||
resp_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(["tool"])
|
||||
|
||||
assert chat_tools == [{"chat": True}]
|
||||
|
|
@ -299,9 +295,7 @@ async def test_execute_tool_calls_strips_prefix_when_alias_differs_from_server_n
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server import mcp_server_manager as _msm
|
||||
|
||||
_msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock(
|
||||
return_value=fake_server
|
||||
)
|
||||
_msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock(return_value=fake_server)
|
||||
|
||||
tool_name = "my_deepwiki-read_wiki_structure"
|
||||
tool_calls = [
|
||||
|
|
@ -375,7 +369,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey
|
|||
|
||||
fake_manager = types.SimpleNamespace(
|
||||
get_registry=MagicMock(return_value={}),
|
||||
call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom"))
|
||||
call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom")),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
|
||||
|
|
@ -383,9 +377,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey
|
|||
)
|
||||
|
||||
tool_name = "deepwiki-read_wiki_structure"
|
||||
tool_calls = [
|
||||
{"id": "call-err", "function": {"name": tool_name, "arguments": "{}"}}
|
||||
]
|
||||
tool_calls = [{"id": "call-err", "function": {"name": tool_name, "arguments": "{}"}}]
|
||||
|
||||
user_auth = types.SimpleNamespace(api_key="test_key", user_id="test_user")
|
||||
|
||||
|
|
@ -403,14 +395,11 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey
|
|||
|
||||
post_call_failure_hook.assert_awaited_once()
|
||||
assert post_call_failure_hook.await_args is not None
|
||||
assert (
|
||||
post_call_failure_hook.await_args.kwargs.get("route")
|
||||
== "/responses/mcp/call_tool"
|
||||
)
|
||||
assert post_call_failure_hook.await_args.kwargs.get("route") == "/responses/mcp/call_tool"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_tool_calls_passes_litellm_call_id_and_trace_id_to_function_setup(
|
||||
async def test_execute_tool_calls_uses_unique_call_ids_and_preserves_parent_context(
|
||||
monkeypatch,
|
||||
):
|
||||
"""
|
||||
|
|
@ -420,22 +409,23 @@ async def test_execute_tool_calls_passes_litellm_call_id_and_trace_id_to_functio
|
|||
_setup_proxy_logging(monkeypatch)
|
||||
call_tool_mock = _setup_mcp_call_environment(monkeypatch)
|
||||
|
||||
captured = {}
|
||||
captured = []
|
||||
|
||||
def fake_function_setup(*_args, **kwargs):
|
||||
captured.update(kwargs)
|
||||
captured.append(kwargs)
|
||||
return None, None
|
||||
|
||||
# NOTE: Don't patch via dotted string path here because `litellm.responses`
|
||||
# is a function attribute on the `litellm` package (shadowing the submodule),
|
||||
# which breaks monkeypatch's importpath resolution.
|
||||
handler_module = importlib.import_module(
|
||||
"litellm.responses.mcp.litellm_proxy_mcp_handler"
|
||||
)
|
||||
handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler")
|
||||
monkeypatch.setattr(handler_module, "function_setup", fake_function_setup)
|
||||
|
||||
tool_name = "deepwiki-read_wiki_structure"
|
||||
tool_calls = [{"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}}]
|
||||
tool_calls = [
|
||||
{"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}},
|
||||
{"id": "call-2", "function": {"name": tool_name, "arguments": "{}"}},
|
||||
]
|
||||
|
||||
await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map={tool_name: "deepwiki"},
|
||||
|
|
@ -446,10 +436,41 @@ async def test_execute_tool_calls_passes_litellm_call_id_and_trace_id_to_functio
|
|||
)
|
||||
|
||||
# Ensure the tool call was attempted (sanity)
|
||||
assert call_tool_mock.await_count == 1
|
||||
assert call_tool_mock.await_count == 2
|
||||
assert len(captured) == 2
|
||||
assert captured[0]["litellm_call_id"] != captured[1]["litellm_call_id"]
|
||||
assert all(item["litellm_call_id"] != "cid" for item in captured)
|
||||
assert all(item["litellm_trace_id"] == "tid" for item in captured)
|
||||
assert all(item["metadata"]["parent_litellm_call_id"] == "cid" for item in captured)
|
||||
|
||||
assert captured.get("litellm_call_id") == "cid"
|
||||
assert captured.get("litellm_trace_id") == "tid"
|
||||
|
||||
def test_prepare_follow_up_call_params_resets_per_call_logging_state():
|
||||
original = {
|
||||
"model": "gpt-4",
|
||||
"litellm_call_id": "parent-call",
|
||||
"litellm_logging_obj": object(),
|
||||
"litellm_trace_id": "trace-1",
|
||||
"tool_choice": "auto",
|
||||
"web_search_options": None,
|
||||
"litellm_params": {
|
||||
"litellm_call_id": "nested-parent-call",
|
||||
"litellm_logging_obj": object(),
|
||||
"metadata": {"team": "legal"},
|
||||
},
|
||||
}
|
||||
|
||||
follow_up = LiteLLM_Proxy_MCP_Handler.prepare_follow_up_call_params(original, original_stream_setting=True)
|
||||
|
||||
assert "litellm_call_id" not in follow_up
|
||||
assert "litellm_logging_obj" not in follow_up
|
||||
assert "litellm_call_id" not in follow_up["litellm_params"]
|
||||
assert "litellm_logging_obj" not in follow_up["litellm_params"]
|
||||
assert follow_up["litellm_trace_id"] == "trace-1"
|
||||
assert follow_up["stream"] is True
|
||||
assert "tool_choice" not in follow_up
|
||||
assert "web_search_options" not in follow_up
|
||||
assert follow_up["litellm_params"]["metadata"] == {"team": "legal"}
|
||||
assert original["litellm_call_id"] == "parent-call"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -515,9 +536,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch
|
|||
user_auth = types.SimpleNamespace(api_key="test_key", user_id="test_user")
|
||||
tools, _server_names = await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager(
|
||||
user_api_key_auth=user_auth,
|
||||
mcp_tools_with_litellm_proxy=[
|
||||
{"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"}
|
||||
],
|
||||
mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"}],
|
||||
)
|
||||
|
||||
assert tools == []
|
||||
|
|
@ -528,9 +547,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch
|
|||
|
||||
|
||||
def test_get_parent_request_tags_from_metadata():
|
||||
tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(
|
||||
{"metadata": {"tags": ["team-a", "prod"]}}
|
||||
)
|
||||
tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags({"metadata": {"tags": ["team-a", "prod"]}})
|
||||
assert tags == ["team-a", "prod"]
|
||||
|
||||
|
||||
|
|
@ -567,9 +584,7 @@ async def test_get_mcp_tools_from_manager_forwards_request_tags(monkeypatch):
|
|||
|
||||
await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager(
|
||||
user_api_key_auth=types.SimpleNamespace(api_key="k", user_id="u"),
|
||||
mcp_tools_with_litellm_proxy=[
|
||||
{"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"}
|
||||
],
|
||||
mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"}],
|
||||
request_tags=["team-a"],
|
||||
)
|
||||
|
||||
|
|
@ -589,9 +604,7 @@ async def test_execute_tool_calls_exposes_sanitized_client_headers_to_logging(mo
|
|||
captured.update(kwargs)
|
||||
return None, None
|
||||
|
||||
handler_module = importlib.import_module(
|
||||
"litellm.responses.mcp.litellm_proxy_mcp_handler"
|
||||
)
|
||||
handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler")
|
||||
monkeypatch.setattr(handler_module, "function_setup", fake_function_setup)
|
||||
|
||||
tool_name = "deepwiki-read_wiki_structure"
|
||||
|
|
@ -617,9 +630,7 @@ async def test_execute_tool_calls_propagates_request_tags_to_function_setup(monk
|
|||
captured.update(kwargs)
|
||||
return None, None
|
||||
|
||||
handler_module = importlib.import_module(
|
||||
"litellm.responses.mcp.litellm_proxy_mcp_handler"
|
||||
)
|
||||
handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler")
|
||||
monkeypatch.setattr(handler_module, "function_setup", fake_function_setup)
|
||||
|
||||
tool_name = "deepwiki-read_wiki_structure"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue