mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(mcp): propagate model into model_call_details for passthrough tool calls (#30122)
* fix(mcp): propagate model into model_call_details for passthrough tool calls The @client decorator on call_mcp_tool creates the logging object via function_setup without a model kwarg, so model_call_details["model"] starts as None. execute_mcp_tool only set logging_obj.model as an instance attribute, which the spend-log writer never reads (it reads kwargs["model"] from model_call_details). MCP passthrough tools/call rows therefore persisted with model="" while list_tools rows showed "MCP: list_tools", degrading the Logs UI display and bucketing all MCP tool spend under an empty model in DailyUserSpend. Propagate the model into model_call_details alongside the existing attribute assignment so the StandardLoggingPayload and SpendLogs writer pick it up. Covers the /mcp passthrough, REST /mcp-rest/tools/call, and orchestrated paths (the latter already passed model into function_setup, so this is a no-op there). * test(mcp): trim regression test docstring
This commit is contained in:
parent
35dc441092
commit
abf04d03c3
2 changed files with 82 additions and 0 deletions
|
|
@ -2552,6 +2552,7 @@ if MCP_AVAILABLE:
|
|||
standard_logging_mcp_tool_call
|
||||
)
|
||||
litellm_logging_obj.model = f"MCP: {name}"
|
||||
litellm_logging_obj.model_call_details["model"] = f"MCP: {name}"
|
||||
# Resolve the MCP server early so BYOK checks and credential injection
|
||||
# apply to ALL dispatch paths (local tool registry AND managed MCP server).
|
||||
if mcp_server is None:
|
||||
|
|
|
|||
|
|
@ -5489,3 +5489,84 @@ async def test_create_mcp_client_sampling_enabled():
|
|||
|
||||
client = await manager._create_mcp_client(server=server)
|
||||
assert client._sampling_callback is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_sets_model_in_model_call_details():
|
||||
"""Regression test: MCP tools/call spend logs persisted with model="".
|
||||
|
||||
execute_mcp_tool set logging_obj.model only; the spend-log writer reads
|
||||
model_call_details["model"], which stays None when function_setup builds
|
||||
the logging object without a "model" kwarg.
|
||||
"""
|
||||
import uuid
|
||||
from datetime import timezone
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_module
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
from litellm.utils import Rules, function_setup
|
||||
|
||||
user = UserAPIKeyAuth(
|
||||
api_key="sk-user",
|
||||
user_id="alice",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER.value,
|
||||
)
|
||||
|
||||
fake_server = MagicMock()
|
||||
fake_server.name = "openapi-petstore"
|
||||
fake_server.is_byok = False
|
||||
fake_server.auth_type = None
|
||||
fake_server.mcp_info = None
|
||||
fake_server.server_id = "srv-1"
|
||||
fake_server.server_name = "openapi-petstore"
|
||||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "list_pets"
|
||||
|
||||
start_time = datetime.now(timezone.utc)
|
||||
litellm_logging_obj, _ = function_setup(
|
||||
original_function="call_mcp_tool",
|
||||
rules_obj=Rules(),
|
||||
start_time=start_time,
|
||||
litellm_call_id=str(uuid.uuid4()),
|
||||
name="list_pets",
|
||||
arguments={"limit": 10},
|
||||
)
|
||||
assert litellm_logging_obj.model_call_details.get("model") is None
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
mcp_module.global_mcp_server_manager,
|
||||
"_get_mcp_server_from_tool_name",
|
||||
return_value=fake_server,
|
||||
),
|
||||
patch.object(
|
||||
mcp_module.global_mcp_server_manager,
|
||||
"pre_call_tool_check",
|
||||
new=AsyncMock(return_value={}),
|
||||
),
|
||||
patch.object(
|
||||
mcp_module.global_mcp_tool_registry,
|
||||
"get_tool",
|
||||
return_value=fake_tool,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._handle_local_mcp_tool",
|
||||
new=AsyncMock(return_value=[]),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed",
|
||||
return_value=True,
|
||||
),
|
||||
):
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="list_pets",
|
||||
arguments={"limit": 10},
|
||||
allowed_mcp_servers=[fake_server],
|
||||
start_time=start_time,
|
||||
user_api_key_auth=user,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
||||
assert litellm_logging_obj.model_call_details["model"] == "MCP: list_pets"
|
||||
assert litellm_logging_obj.model == "MCP: list_pets"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue