fix(mcp): roll up MCP tool spend to user counters and usage UI (#31576)

* fix(mcp): roll up MCP tool spend to user counters and usage UI

Direct REST MCP tool calls now fire success logging so spend_logs and
user/team rollups include configured mcp_server_cost_info charges.

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(mcp): gate key-info enrichment to requests missing user_id; fix import order

- Only call _enrich_failure_metadata_with_key_info when user_api_key_user_id is
  absent, avoiding a cache/DB lookup on every normal LLM request.
- Move LiteLLMProxyRequestSetup import to correct alphabetical position (I001).

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(mcp): scope MCP spend aggregate by api_key to prevent cross-tenant disclosure

Add api_key = ANY($2) to the MCP session aggregate query so it is
bounded by the same ownership already applied to the main page query.

Co-authored-by: Cursor <cursoragent@cursor.com>

* Fix spend logs for call and list mcp tools

* Add tags in mcp logging

* Fix ruff

* fix(lint): replace List/Dict with list/dict in new annotations (UP006)

Replace the 8 new UP006 violations introduced by the mcp-tags changes:
- Optional[List[str]] → Optional[list[str]] for request_tags params
- List[str] return type → list[str] in _get_parent_request_tags
- Dict[str, Dict[...]] → dict[str, dict[...]] for mcp_spend_map annotation

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(lint): keep call_tool_rest_api within complexity budget and narrow MCP spend enrichment except to PrismaError

* fix(mcp): keep final streaming chunk when draining inner stream fails

* fix: handle MCP logging edge cases

* fix: propagate MCP logging cancellation

---------

Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
Sameer Kankute 2026-07-02 20:46:39 +05:30 • committed by GitHub
parent 8d0dc9294d
commit fabe5c283a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
15 changed files with 645 additions and 41 deletions

View file

@ -77,6 +77,7 @@ if MCP_AVAILABLE:
ListMCPToolsRestAPIResponseObject,
MCPInfo,
MCPServer,
_fire_mcp_success_logging,
_tool_name_matches,
execute_mcp_tool,
filter_tools_by_allowed_tools,
@ -84,6 +85,24 @@ if MCP_AVAILABLE:
########################################################
############ MCP Server REST API Routes #################
async def _safe_fire_mcp_success_logging(
logging_obj: Optional[Any],
result: Any,
start_time: datetime,
end_time: datetime,
) -> None:
if logging_obj is None:
return
logging_results = await asyncio.gather(
_fire_mcp_success_logging(logging_obj, result, start_time, end_time),
return_exceptions=True,
)
logging_error = logging_results[0]
if isinstance(logging_error, asyncio.CancelledError):
raise logging_error
if isinstance(logging_error, BaseException):
verbose_logger.warning("MCP tool success logging failed (continuing): %s", logging_error)
def _get_server_auth_header(
server,
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]],
@ -798,7 +817,8 @@ if MCP_AVAILABLE:
proxy_logging_obj=proxy_logging_obj,
general_settings=general_settings,
)
return await handle_mcp_tool_call(
_tool_start_time = datetime.now()
result = await handle_mcp_tool_call(
tool_name=tool_arguments.get("tool_name", ""),
arguments=tool_arguments.get("arguments") or {},
user_api_key_dict=user_api_key_dict,
@ -809,6 +829,8 @@ if MCP_AVAILABLE:
raw_headers=virtual_raw_headers,
litellm_logging_obj=virtual_logging_obj,
)
await _safe_fire_mcp_success_logging(virtual_logging_obj, result, _tool_start_time, datetime.now())
return result
# Validate required parameters early
server_id = data.get("server_id")
@ -876,11 +898,12 @@ if MCP_AVAILABLE:
user_oauth_extra_headers = await _get_user_oauth_extra_headers(target_server, user_api_key_dict)
# Call execute_mcp_tool directly (permission checks already done)
_tool_start_time = datetime.now()
result = await execute_mcp_tool(
name=tool_name,
arguments=tool_arguments,
allowed_mcp_servers=allowed_mcp_servers,
start_time=datetime.now(),
start_time=_tool_start_time,
user_api_key_auth=data.get("user_api_key_auth"),
mcp_auth_header=data.get("mcp_auth_header"),
mcp_server_auth_headers=data.get("mcp_server_auth_headers"),
@ -889,6 +912,7 @@ if MCP_AVAILABLE:
litellm_logging_obj=data.get("litellm_logging_obj"),
requested_server_id=canonical_server_id,
)
await _safe_fire_mcp_success_logging(logging_obj, result, _tool_start_time, datetime.now())
return result
except MCPMissingUserEnvVarsError as e:
verbose_logger.info(

View file

@ -1681,6 +1681,7 @@ if MCP_AVAILABLE:
log_list_tools_to_spendlogs: bool = False,
list_tools_log_source: Optional[str] = None,
litellm_trace_id: Optional[str] = None,
request_tags: Optional[list[str]] = None,
client_ip: Optional[str] = None,
) -> List[MCPTool]:
"""
@ -1724,6 +1725,7 @@ if MCP_AVAILABLE:
"litellm_trace_id": effective_litellm_trace_id,
"metadata": {
"spend_logs_metadata": spend_logs_metadata,
**({"tags": request_tags} if request_tags else {}),
},
# Provide a small input payload for standard logging
"input": [
@ -1899,7 +1901,9 @@ if MCP_AVAILABLE:
end_time = datetime.now()
try:
await litellm_logging_obj.async_success_handler(
result=all_tools,
result=[
tool.model_dump(mode="json") if isinstance(tool, MCPTool) else tool for tool in all_tools
],
start_time=list_tools_start_time,
end_time=end_time,
)
@ -2741,6 +2745,22 @@ if MCP_AVAILABLE:
return response
async def _fire_mcp_success_logging(
logging_obj: LiteLLMLoggingObj,
result: Any,
start_time: datetime,
end_time: datetime,
) -> None:
logging_obj.post_call(original_response=result)
await logging_obj.async_post_mcp_tool_call_hook(
kwargs=logging_obj.model_call_details,
response_obj=result,
start_time=start_time,
end_time=end_time,
)
logging_obj.call_type = CallTypes.call_mcp_tool.value
await logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time)
@client
async def call_mcp_tool(
name: str,
@ -2812,16 +2832,7 @@ if MCP_AVAILABLE:
raise
if litellm_logging_obj:
litellm_logging_obj.post_call(original_response=response)
end_time = datetime.now()
await litellm_logging_obj.async_post_mcp_tool_call_hook(
kwargs=litellm_logging_obj.model_call_details,
response_obj=response,
start_time=start_time,
end_time=end_time,
)
litellm_logging_obj.call_type = CallTypes.call_mcp_tool.value
await litellm_logging_obj.async_success_handler(result=response, start_time=start_time, end_time=end_time)
await _fire_mcp_success_logging(litellm_logging_obj, response, start_time, datetime.now())
return response
async def mcp_get_prompt(

View file

@ -190,6 +190,11 @@ class _ProxyDBLogger(CustomLogger):
litellm_params = kwargs.get("litellm_params", {}) or {}
end_user_id = get_end_user_id_for_cost_tracking(litellm_params)
metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs)
# Only fetch key details when user_id wasn't already populated (e.g. direct MCP REST calls).
# Avoids a cache/DB lookup on every normal LLM request.
if metadata.get("user_api_key") and not metadata.get("user_api_key_user_id"):
metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata=metadata)
_write_spend_metadata_to_kwargs(kwargs=kwargs, metadata=metadata)
budget_reservation = _get_budget_reservation_from_metadata(metadata=metadata)
user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None))
team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None))
@ -388,6 +393,20 @@ class _ProxyDBLogger(CustomLogger):
return
def _write_spend_metadata_to_kwargs(kwargs: dict, metadata: dict) -> None:
patch = {k: v for k, v in metadata.items() if (k.startswith("user_api_key") or k == "tags") and v is not None}
if not patch:
return
litellm_params = kwargs.setdefault("litellm_params", {})
for bucket_name in ("litellm_metadata", "metadata"):
bucket = litellm_params.get(bucket_name)
if isinstance(bucket, dict):
for key, value in patch.items():
if bucket.get(key) is None:
bucket[key] = value
def _should_track_cost_callback(
user_api_key: Optional[str],
user_id: Optional[str],

View file

@ -3444,12 +3444,59 @@ async def _build_ui_spend_logs_response(
)
count_map = {r["session_id"]: r["_count"]["session_id"] for r in counts if r.get("session_id")}
mcp_spend_map: dict[str, dict[str, Union[int, float]]] = {}
if enrich_session_counts and session_ids:
from prisma.errors import PrismaError
try:
# Collect api_keys already present in the authorized page rows so the
# aggregate is scoped to the same ownership as the main query — prevents
# cross-tenant disclosure via a colliding session_id.
authorized_api_keys = list(
{
(row.get("api_key") if isinstance(row, dict) else getattr(row, "api_key", None))
for row in data
if (row.get("api_key") if isinstance(row, dict) else getattr(row, "api_key", None))
}
)
rows = await prisma_client.db.query_raw(
"""
SELECT session_id,
COUNT(*)::int AS mcp_tool_call_count,
COALESCE(SUM(spend), 0)::double precision AS mcp_tool_call_spend
FROM "LiteLLM_SpendLogs"
WHERE session_id = ANY($1::text[])
AND api_key = ANY($2::text[])
AND call_type IN ('call_mcp_tool', 'list_mcp_tools')
GROUP BY session_id
""",
session_ids,
authorized_api_keys,
)
mcp_spend_map = {
row["session_id"]: {
"mcp_tool_call_count": int(row.get("mcp_tool_call_count") or 0),
"mcp_tool_call_spend": float(row.get("mcp_tool_call_spend") or 0.0),
}
for row in rows
if row.get("session_id")
}
except PrismaError:
verbose_proxy_logger.debug(
"Failed to enrich MCP session spend aggregates for spend logs UI",
exc_info=True,
)
if enrich_session_counts:
enriched: List[dict] = []
for row in data:
row_dict = dict(row) if isinstance(row, dict) else row.model_dump()
sid = row_dict.get("session_id")
row_dict["session_total_count"] = count_map.get(sid, 1) if sid else 1
mcp_stats = mcp_spend_map.get(sid) if sid else None
if mcp_stats:
row_dict["mcp_tool_call_count"] = mcp_stats["mcp_tool_call_count"]
row_dict["mcp_tool_call_spend"] = mcp_stats["mcp_tool_call_spend"]
enriched.append(row_dict)
response_data: list = enriched
else:

View file

@ -212,6 +212,7 @@ async def aresponses_api_with_mcp(
litellm_trace_id=kwargs.get("litellm_trace_id"),
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
request_tags=LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(kwargs),
)
openai_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(original_mcp_tools)
@ -327,6 +328,7 @@ async def aresponses_api_with_mcp(
raw_headers=raw_headers_from_request,
litellm_call_id=kwargs.get("litellm_call_id"),
litellm_trace_id=kwargs.get("litellm_trace_id"),
request_tags=LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(kwargs),
)
if tool_results:
@ -382,6 +384,7 @@ async def aresponses_api_with_mcp(
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
request_tags=LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(kwargs),
)
final_response = LiteLLM_Proxy_MCP_Handler._add_mcp_output_elements_to_response(
response=final_response,

View file

@ -1,5 +1,6 @@
"""Helpers for handling MCP-aware `/chat/completions` requests."""
import logging
from typing import (
Any,
List,
@ -115,6 +116,7 @@ async def acompletion_with_mcp(
# Extract user_api_key_auth from metadata or kwargs
user_api_key_auth = kwargs.get("user_api_key_auth") or ((kwargs.get("metadata", {}) or {}).get("user_api_key_auth"))
request_tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(kwargs)
# Extract MCP auth headers before fetching tools (needed for dynamic auth)
(
@ -137,6 +139,7 @@ async def acompletion_with_mcp(
litellm_trace_id=kwargs.get("litellm_trace_id"),
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
request_tags=request_tags,
)
openai_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(
@ -218,6 +221,7 @@ async def acompletion_with_mcp(
litellm_trace_id,
openai_tools,
base_call_args,
request_tags,
):
self.stream_wrapper = stream_wrapper
self.messages = messages
@ -231,6 +235,7 @@ async def acompletion_with_mcp(
self.litellm_trace_id = litellm_trace_id
self.openai_tools = openai_tools
self.base_call_args = base_call_args
self.request_tags = request_tags
self.collected_chunks: List[ModelResponseStream] = []
self.tool_calls: Optional[List] = None
self.tool_results: Optional[List] = None
@ -303,6 +308,17 @@ async def acompletion_with_mcp(
return chunk
async def _drain_inner_stream(self):
try:
while True:
await self._stream_iterator.__anext__()
except StopAsyncIteration:
pass
except Exception:
logging.getLogger("LiteLLM").exception(
"Error draining inner MCP stream after final chunk; spend logging may be incomplete"
)
async def __anext__(self):
# Phase 1: Collect and yield initial stream chunks
if not self.stream_exhausted:
@ -332,15 +348,16 @@ async def acompletion_with_mcp(
)
if is_final:
# This is the final chunk, mark stream as exhausted
self.stream_exhausted = True
# Process tool calls after we've collected all chunks
await self._process_tool_calls()
# Apply MCP metadata (tool_calls and tool_results) to final chunk
chunk = self._add_mcp_tool_metadata_to_final_chunk(chunk)
# If we have tool results, prepare follow-up call immediately
if self.tool_results and self.complete_response:
await self._prepare_follow_up_call()
# Drain inner stream so CustomStreamWrapper fires its
# end-of-stream handler (dispatch_success_handlers →
# _ProxyDBLogger → LiteLLM_SpendLogs). The CSW may
# yield one usage chunk before raising StopAsyncIteration.
await self._drain_inner_stream()
return chunk
except StopAsyncIteration:
@ -354,6 +371,7 @@ async def acompletion_with_mcp(
# If we have tool results, prepare follow-up call
if self.tool_results and self.complete_response:
await self._prepare_follow_up_call()
await self._drain_inner_stream()
return final_chunk
# Phase 2: Yield follow-up stream chunks if available
@ -426,6 +444,7 @@ async def acompletion_with_mcp(
raw_headers=self.raw_headers,
litellm_call_id=self.litellm_call_id,
litellm_trace_id=self.litellm_trace_id,
request_tags=self.request_tags,
)
async def _prepare_follow_up_call(self):
@ -485,6 +504,7 @@ async def acompletion_with_mcp(
litellm_trace_id=kwargs.get("litellm_trace_id"),
openai_tools=openai_tools,
base_call_args=base_call_args,
request_tags=request_tags,
)
# Create a wrapper class that delegates to our custom iterator
@ -596,6 +616,7 @@ async def acompletion_with_mcp(
raw_headers=raw_headers,
litellm_call_id=kwargs.get("litellm_call_id"),
litellm_trace_id=kwargs.get("litellm_trace_id"),
request_tags=request_tags,
)
if not tool_results:

View file

@ -17,6 +17,7 @@ from litellm._logging import verbose_logger
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._experimental.mcp_server.utils import split_server_prefix_from_name
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.responses.main import aresponses
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
from litellm.types.llms.openai import ResponsesAPIResponse
@ -59,6 +60,20 @@ class LiteLLM_Proxy_MCP_Handler:
This handles when a user passes mcp server_url="litellm_proxy" in their tools.
"""
@staticmethod
def _get_parent_request_tags(kwargs: Optional[dict[str, Any]]) -> list[str]:
"""Tags from the parent LLM request, using the same extraction logic as standard logging (incl. User-Agent)."""
if not kwargs:
return []
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
litellm_params = kwargs.get("litellm_params") or kwargs
proxy_server_request = litellm_params.get("proxy_server_request") or kwargs.get("proxy_server_request") or {}
return StandardLoggingPayloadSetup._get_request_tags(
litellm_params=litellm_params,
proxy_server_request=proxy_server_request,
)
@staticmethod
def _should_use_litellm_mcp_gateway(tools: Optional[Iterable[ToolParam]]) -> bool:
"""
@ -162,6 +177,7 @@ class LiteLLM_Proxy_MCP_Handler:
litellm_trace_id: Optional[str] = None,
mcp_auth_header: Optional[str] = None,
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None,
request_tags: Optional[list[str]] = None,
) -> tuple[List[MCPTool], List[str]]:
"""
Get available tools from the MCP server manager.
@ -250,6 +266,7 @@ class LiteLLM_Proxy_MCP_Handler:
log_list_tools_to_spendlogs=True,
list_tools_log_source="responses",
litellm_trace_id=litellm_trace_id,
request_tags=request_tags,
)
allowed_mcp_server_ids = await global_mcp_server_manager.get_allowed_mcp_servers(user_api_key_auth)
@ -351,6 +368,7 @@ class LiteLLM_Proxy_MCP_Handler:
user_api_key_auth: Any,
mcp_tools_with_litellm_proxy: List[ToolParam],
litellm_trace_id: Optional[str] = None,
request_tags: Optional[list[str]] = None,
) -> tuple[List[Any], dict[str, str]]:
"""
Centralized method to process MCP tools through the complete pipeline.
@ -371,6 +389,7 @@ class LiteLLM_Proxy_MCP_Handler:
user_api_key_auth,
mcp_tools_with_litellm_proxy,
litellm_trace_id=litellm_trace_id,
request_tags=request_tags,
)
openai_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(deduplicated_mcp_tools)
@ -384,6 +403,7 @@ class LiteLLM_Proxy_MCP_Handler:
litellm_trace_id: Optional[str] = None,
mcp_auth_header: Optional[str] = None,
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None,
request_tags: Optional[list[str]] = None,
) -> tuple[List[Any], dict[str, str]]:
"""
Process MCP tools through filtering and deduplication pipeline without OpenAI transformation.
@ -411,6 +431,7 @@ class LiteLLM_Proxy_MCP_Handler:
litellm_trace_id=litellm_trace_id,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
request_tags=request_tags,
)
# Step 2: Filter tools based on allowed_tools parameter
@ -597,6 +618,7 @@ class LiteLLM_Proxy_MCP_Handler:
raw_headers: Optional[Dict[str, str]] = None,
litellm_call_id: Optional[str] = None,
litellm_trace_id: Optional[str] = None,
request_tags: Optional[list[str]] = None,
) -> List[Dict[str, Any]]:
"""Execute tool calls and return results."""
from fastapi import HTTPException
@ -672,17 +694,19 @@ class LiteLLM_Proxy_MCP_Handler:
}
if litellm_trace_id:
logging_request_data["litellm_trace_id"] = litellm_trace_id
user_identifier = None
if request_tags:
logging_request_data["metadata"]["tags"] = request_tags
if user_api_key_auth is not None:
user_api_key = getattr(user_api_key_auth, "api_key", None)
if user_api_key:
logging_request_data["metadata"]["user_api_key"] = user_api_key
LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata(
data=logging_request_data,
user_api_key_dict=user_api_key_auth,
_metadata_variable_name="metadata",
)
user_identifier = getattr(user_api_key_auth, "end_user_id", None) or getattr(
user_api_key_auth, "user_id", None
)
if user_identifier:
logging_request_data["user"] = user_identifier
if user_identifier:
logging_request_data["user"] = user_identifier
litellm_logging_obj: Optional[LiteLLMLoggingObj] = None
try:

View file

@ -630,6 +630,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
raw_headers=self.raw_headers,
litellm_call_id=self.litellm_call_id,
litellm_trace_id=self.litellm_trace_id,
request_tags=LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(self.original_request_params),
)
# Create completion events and output_item.done events for tool execution

View file

@ -4164,6 +4164,7 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab
_get_tools_from_mcp_servers,
)
from litellm.proxy._types import UserAPIKeyAuth
from mcp.types import Tool as MCPTool
except ImportError:
pytest.skip("MCP server not available")
@ -4177,12 +4178,20 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab
server_a.auth_type = None
server_a.extra_headers = None
tool_1 = MagicMock()
tool_1.name = "server_a-tool_1"
tool_1 = MCPTool(
name="server_a-tool_1",
description="test tool",
inputSchema={"type": "object"},
)
dummy_logging_obj = MagicMock()
dummy_logging_obj.model_call_details = {"metadata": {"spend_logs_metadata": {}}}
dummy_logging_obj.async_success_handler = AsyncMock()
function_setup_kwargs = {}
def _capture_function_setup(*_args, **kwargs):
function_setup_kwargs.update(kwargs)
return dummy_logging_obj, None
with (
patch(
@ -4206,7 +4215,7 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab
),
patch(
"litellm.proxy._experimental.mcp_server.server.function_setup",
return_value=(dummy_logging_obj, None),
side_effect=_capture_function_setup,
),
):
mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1])
@ -4218,13 +4227,15 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab
mcp_server_auth_headers=None,
log_list_tools_to_spendlogs=True,
list_tools_log_source="mcp_protocol",
request_tags=["team-a"],
)
assert tools == [tool_1]
dummy_logging_obj.async_success_handler.assert_awaited_once()
assert dummy_logging_obj.async_success_handler.await_args.kwargs["result"] == [
tool_1
tool_1.model_dump(mode="json")
]
assert function_setup_kwargs["metadata"]["tags"] == ["team-a"]
spend_meta = dummy_logging_obj.model_call_details["metadata"]["spend_logs_metadata"]
assert spend_meta["tool_count_total"] == 1

View file

@ -432,6 +432,11 @@ class TestCallToolRestApiVirtualTools:
new_callable=AsyncMock,
return_value=fake_result,
) as mock_execute,
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._fire_mcp_success_logging",
new_callable=AsyncMock,
side_effect=RuntimeError("logging failed"),
) as mock_fire_logging,
):
result = await self._get_call_fn()(
request=request,
@ -439,6 +444,7 @@ class TestCallToolRestApiVirtualTools:
)
mock_execute.assert_awaited_once()
mock_fire_logging.assert_awaited_once()
assert mock_execute.await_args.kwargs["name"] == "github-create_issue"
assert result.isError is False

View file

@ -1,6 +1,8 @@
import asyncio
import json
from datetime import datetime
from typing import Any, Dict, Optional
from unittest.mock import MagicMock
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
@ -330,7 +332,8 @@ class TestExecuteWithMcpClient:
@pytest.mark.asyncio
async def test_m2m_does_not_build_presented_store(self, monkeypatch):
"""M2M (client_credentials): to_server_spec returns None, so no presented provider is built;
the auto-fetch path is unchanged (no cred_provider, the incoming header dropped as before)."""
the auto-fetch path is unchanged (no cred_provider, the incoming header dropped as before).
"""
captured: dict = {}
def fake_build_stdio_env(server, raw_headers):
@ -377,7 +380,8 @@ class TestExecuteWithMcpClient:
@pytest.mark.asyncio
async def test_token_exchange_does_not_build_presented_store(self, monkeypatch):
"""OBO / token-exchange (auth_type oauth2_token_exchange, not oauth2): excluded by the
auth_type == oauth2 guard, so no presented provider is built and the v1 exchange path runs."""
auth_type == oauth2 guard, so no presented provider is built and the v1 exchange path runs.
"""
captured: dict = {}
def fake_build_stdio_env(server, raw_headers):
@ -1196,6 +1200,13 @@ class TestCallToolRestAPI:
fake_execute_mcp_tool,
raising=False,
)
fire_logging = AsyncMock(side_effect=RuntimeError("logging failed"))
monkeypatch.setattr(
rest_endpoints,
"_fire_mcp_success_logging",
fire_logging,
raising=False,
)
request_payload = {
"server_id": "server-1",
@ -1217,6 +1228,23 @@ class TestCallToolRestAPI:
assert captured["name"] == "demo-tool"
assert captured["arguments"] == {"foo": "bar"}
assert captured["allowed_mcp_servers"] == [stub_server]
fire_logging.assert_awaited_once()
async def test_success_logging_cancellation_propagates(self, monkeypatch):
fire_logging = AsyncMock(side_effect=asyncio.CancelledError())
monkeypatch.setattr(
rest_endpoints,
"_fire_mcp_success_logging",
fire_logging,
raising=False,
)
with pytest.raises(asyncio.CancelledError):
await rest_endpoints._safe_fire_mcp_success_logging(
object(), {"result": "ok"}, datetime.now(), datetime.now()
)
fire_logging.assert_awaited_once()
class TestGetToolsForSingleServer:
@ -1809,9 +1837,9 @@ class TestPreviewOpenAPITools:
names = [t["name"] for t in result["tools"]]
anthropic_re = re.compile(r"^[a-zA-Z0-9_-]{1,128}$")
for name in names:
assert anthropic_re.match(name), (
f"preview tool name {name!r} violates ^[a-zA-Z0-9_-]+$"
)
assert anthropic_re.match(
name
), f"preview tool name {name!r} violates ^[a-zA-Z0-9_-]+$"
assert "actions_download-job-logs-for-workflow-run" in names
assert "pulls_list-files" in names
@ -1868,7 +1896,9 @@ class TestPreviewOpenAPITools:
registered_summary_to_name: dict = {}
def fake_create_tool_function(path, method, operation, base_url): # noqa: ANN001
def fake_create_tool_function(
path, method, operation, base_url
): # noqa: ANN001
def _f():
return None
@ -1881,7 +1911,9 @@ class TestPreviewOpenAPITools:
)
class _StubRegistry:
def register_tool(self, name, description, input_schema, handler): # noqa: ANN001
def register_tool(
self, name, description, input_schema, handler
): # noqa: ANN001
registered_summary_to_name[description] = name
monkeypatch.setattr(

View file

@ -1104,3 +1104,76 @@ async def test_async_post_call_failure_hook_records_recovered_partial_spend():
mock_update_database.assert_called_once()
assert mock_update_database.call_args[1]["response_cost"] == 3.5e-05
@pytest.mark.asyncio
async def test_track_cost_callback_enriches_user_id_for_mcp_style_metadata():
"""MCP tool calls may only carry user_api_key; user/team rollups still need user_id."""
from litellm.proxy._types import UserAPIKeyAuth
logger = _ProxyDBLogger()
key_obj = UserAPIKeyAuth(
api_key="hashed-key",
user_id="mcp-user@example.com",
team_id="team-123",
org_id="org-456",
key_alias="mcp-key",
)
kwargs = {
"call_type": "call_mcp_tool",
"model": "MCP: echo",
"litellm_params": {
"metadata": {
"user_api_key": "hashed-key",
}
},
"standard_logging_object": {
"response_cost": 10.0,
"request_tags": [],
"metadata": {},
},
}
with (
patch(
"litellm.proxy.hooks.proxy_track_cost_callback.get_key_object",
new_callable=AsyncMock,
return_value=key_obj,
),
patch(
"litellm.proxy.proxy_server.increment_spend_counters",
new_callable=AsyncMock,
) as mock_increment,
patch(
"litellm.proxy.proxy_server.update_cache",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.proxy_server.proxy_logging_obj",
) as mock_proxy_logging,
):
mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock()
mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock()
await logger._PROXY_track_cost_callback(
kwargs=kwargs,
completion_response={"id": "mcp-call-1"},
start_time=datetime.now(),
end_time=datetime.now(),
)
mock_increment.assert_awaited_once()
assert mock_increment.call_args.kwargs["user_id"] == "mcp-user@example.com"
assert mock_increment.call_args.kwargs["team_id"] == "team-123"
assert mock_increment.call_args.kwargs["org_id"] == "org-456"
update_kwargs = (
mock_proxy_logging.db_spend_update_writer.update_database.await_args.kwargs
)
assert update_kwargs["user_id"] == "mcp-user@example.com"
assert update_kwargs["team_id"] == "team-123"
assert (
kwargs["litellm_params"]["metadata"]["user_api_key_user_id"]
== "mcp-user@example.com"
)

View file

@ -1051,7 +1051,10 @@ async def test_ui_view_spend_logs_sort_by_ttft_ms(client, monkeypatch):
page_size = params[-2] if len(params) >= 2 else 50
skip = params[-1] if len(params) >= 1 else 0
return [
{**{k: v for k, v in row.items() if k != "_ttft_ms"}, "total_count": len(base_logs)}
{
**{k: v for k, v in row.items() if k != "_ttft_ms"},
"total_count": len(base_logs),
}
for row in sorted_logs[skip : skip + page_size]
]
@ -2917,10 +2920,11 @@ async def test_build_ui_spend_logs_response_dict_rows_session_counts():
)
session_id = "sess-abc-123"
api_key = "hashed-key-xyz"
dict_rows = [
{"request_id": "req-1", "session_id": session_id, "call_type": "completion"},
{"request_id": "req-2", "session_id": session_id, "call_type": "mcp_tool_call"},
{"request_id": "req-3", "session_id": None, "call_type": "completion"},
{"request_id": "req-1", "session_id": session_id, "call_type": "completion", "api_key": api_key},
{"request_id": "req-2", "session_id": session_id, "call_type": "mcp_tool_call", "api_key": api_key},
{"request_id": "req-3", "session_id": None, "call_type": "completion", "api_key": api_key},
]
mock_prisma = MagicMock()
@ -2929,6 +2933,15 @@ async def test_build_ui_spend_logs_response_dict_rows_session_counts():
{"session_id": session_id, "_count": {"session_id": 2}},
]
)
mock_prisma.db.query_raw = AsyncMock(
return_value=[
{
"session_id": session_id,
"mcp_tool_call_count": 1,
"mcp_tool_call_spend": 10.0,
}
]
)
result = await _build_ui_spend_logs_response(
prisma_client=mock_prisma,
@ -2946,6 +2959,10 @@ async def test_build_ui_spend_logs_response_dict_rows_session_counts():
# Rows with the shared session_id should have session_total_count=2
assert rows[0]["session_total_count"] == 2
assert rows[1]["session_total_count"] == 2
assert rows[0]["mcp_tool_call_count"] == 1
assert rows[0]["mcp_tool_call_spend"] == 10.0
assert rows[1]["mcp_tool_call_count"] == 1
assert rows[1]["mcp_tool_call_spend"] == 10.0
# Row without a session_id defaults to 1
assert rows[2]["session_total_count"] == 1
@ -4104,7 +4121,9 @@ async def test_cold_storage_handler_returns_none_when_no_logger_configured(monke
@pytest.mark.asyncio
async def test_cold_storage_handler_resolves_configured_logger_from_registry(monkeypatch):
async def test_cold_storage_handler_resolves_configured_logger_from_registry(
monkeypatch,
):
from litellm.proxy.spend_tracking.cold_storage_handler import ColdStorageHandler
logger = _FakeColdStorageLogger({"messages": "from-registry"})

View file

@ -1105,3 +1105,242 @@ async def test_execute_tool_calls_sets_proxy_server_request_arguments(monkeypatc
"param1": "value1",
"param2": 123,
}, "arguments should be parsed correctly"
@pytest.mark.asyncio
async def test_acompletion_with_mcp_streaming_drain_error_does_not_drop_final_chunk(monkeypatch):
"""
Regression test: after yielding the final chunk, MCPStreamingIterator drains
the inner CustomStreamWrapper to fire end-of-stream spend logging. If the
inner stream raises a non-StopAsyncIteration error during that drain (e.g.
a transient APIError on the trailing usage chunk), the error must not
escape __anext__ and drop the already-assembled final chunk.
"""
from unittest.mock import MagicMock
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
from litellm.utils import CustomStreamWrapper
tools = [{"type": "mcp", "server_url": "litellm_proxy/mcp/local"}]
openai_tools = [{"type": "function", "function": {"name": "local_search"}}]
def create_chunk(content, finish_reason=None):
return ModelResponseStream(
id="test-stream",
model="test-model",
created=1234567890,
object="chat.completion.chunk",
choices=[
StreamingChoices(
index=0,
delta=Delta(content=content, role="assistant"),
finish_reason=finish_reason,
)
],
)
chunks = [
create_chunk("Hello"),
create_chunk(" world", finish_reason="stop"),
]
logging_obj = MagicMock()
logging_obj.model_call_details = {}
class DrainErrorStreamingResponse(CustomStreamWrapper):
def __init__(self):
super().__init__(
completion_stream=None,
model="test-model",
logging_obj=logging_obj,
)
self.chunks = chunks
self._index = 0
def __aiter__(self):
return self
async def __anext__(self):
if self._index < len(self.chunks):
chunk = self.chunks[self._index]
self._index += 1
return chunk
if self._index == len(self.chunks):
self._index += 1
raise RuntimeError("connection dropped on trailing usage chunk")
raise StopAsyncIteration
mock_acompletion = AsyncMock(return_value=DrainErrorStreamingResponse())
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_should_use_litellm_mcp_gateway",
staticmethod(lambda tools: True),
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_parse_mcp_tools",
staticmethod(lambda tools: (tools, [])),
)
async def mock_process(**_):
return (tools, {"local_search": "local"})
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_process_mcp_tools_without_openai_transform",
mock_process,
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_transform_mcp_tools_to_openai",
staticmethod(lambda *_, **__: openai_tools),
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_should_auto_execute_tools",
staticmethod(lambda **_: True),
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_extract_tool_calls_from_chat_response",
staticmethod(lambda **_: []),
)
monkeypatch.setattr(
ResponsesAPIRequestUtils,
"extract_mcp_headers_from_request",
staticmethod(lambda **_: (None, None, None, None)),
)
with patch("litellm.acompletion", mock_acompletion):
result = await acompletion_with_mcp(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
tools=tools,
stream=True,
)
all_chunks = []
async for chunk in result:
all_chunks.append(chunk)
final_chunks = [
chunk
for chunk in all_chunks
if chunk.choices and chunk.choices[0].finish_reason == "stop"
]
assert len(final_chunks) == 1, f"Final chunk must survive a drain error. Got chunks: {all_chunks}"
assert all_chunks[-1].choices[0].finish_reason == "stop"
@pytest.mark.asyncio
async def test_acompletion_with_mcp_streaming_drains_inner_stream_after_exhaustion(monkeypatch):
from unittest.mock import MagicMock
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
from litellm.utils import CustomStreamWrapper
tools = [{"type": "mcp", "server_url": "litellm_proxy/mcp/local"}]
openai_tools = [{"type": "function", "function": {"name": "local_search"}}]
def create_chunk(content):
return ModelResponseStream(
id="test-stream",
model="test-model",
created=1234567890,
object="chat.completion.chunk",
choices=[
StreamingChoices(
index=0,
delta=Delta(content=content, role="assistant"),
finish_reason=None,
)
],
)
chunks = [create_chunk("Hello"), create_chunk(" world")]
logging_obj = MagicMock()
logging_obj.model_call_details = {}
class ExhaustingStreamingResponse(CustomStreamWrapper):
def __init__(self):
super().__init__(
completion_stream=None,
model="test-model",
logging_obj=logging_obj,
)
self.chunks = chunks
self._index = 0
self.drained_after_exhaustion = False
def __aiter__(self):
return self
async def __anext__(self):
if self._index < len(self.chunks):
chunk = self.chunks[self._index]
self._index += 1
return chunk
if self._index == len(self.chunks):
self._index += 1
raise StopAsyncIteration
self.drained_after_exhaustion = True
raise StopAsyncIteration
initial_stream = ExhaustingStreamingResponse()
mock_acompletion = AsyncMock(return_value=initial_stream)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_should_use_litellm_mcp_gateway",
staticmethod(lambda tools: True),
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_parse_mcp_tools",
staticmethod(lambda tools: (tools, [])),
)
async def mock_process(**_):
return (tools, {"local_search": "local"})
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_process_mcp_tools_without_openai_transform",
mock_process,
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_transform_mcp_tools_to_openai",
staticmethod(lambda *_, **__: openai_tools),
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_should_auto_execute_tools",
staticmethod(lambda **_: True),
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_extract_tool_calls_from_chat_response",
staticmethod(lambda **_: []),
)
monkeypatch.setattr(
ResponsesAPIRequestUtils,
"extract_mcp_headers_from_request",
staticmethod(lambda **_: (None, None, None, None)),
)
with patch("litellm.acompletion", mock_acompletion):
result = await acompletion_with_mcp(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
tools=tools,
stream=True,
)
all_chunks = []
async for chunk in result:
all_chunks.append(chunk)
assert len(all_chunks) == 3
assert initial_stream.drained_after_exhaustion is True

View file

@ -401,3 +401,77 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch
assert mock_get_tools.await_args is not None
assert mock_get_tools.await_args.kwargs["log_list_tools_to_spendlogs"] is True
assert mock_get_tools.await_args.kwargs["list_tools_log_source"] == "responses"
def test_get_parent_request_tags_from_metadata():
tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(
{"metadata": {"tags": ["team-a", "prod"]}}
)
assert tags == ["team-a", "prod"]
def test_get_parent_request_tags_from_nested_litellm_params():
tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(
{
"metadata": {"tags": ["top-level"]},
"litellm_params": {
"metadata": {"tags": ["nested"]},
"proxy_server_request": {"headers": {"user-agent": "client/1.0"}},
},
}
)
assert tags == ["nested", "User-Agent: client", "User-Agent: client/1.0"]
@pytest.mark.asyncio
async def test_get_mcp_tools_from_manager_forwards_request_tags(monkeypatch):
mock_get_tools = AsyncMock(return_value=[])
monkeypatch.setattr(
"litellm.proxy._experimental.mcp_server.server._get_tools_from_mcp_servers",
mock_get_tools,
)
fake_manager = types.SimpleNamespace(
get_allowed_mcp_servers=AsyncMock(return_value=[]),
get_mcp_servers_from_ids=MagicMock(return_value=[]),
get_mcp_server_by_name=MagicMock(return_value=None),
)
monkeypatch.setattr(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
fake_manager,
)
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"}
],
request_tags=["team-a"],
)
assert mock_get_tools.await_args.kwargs["request_tags"] == ["team-a"]
@pytest.mark.asyncio
async def test_execute_tool_calls_propagates_request_tags_to_function_setup(monkeypatch):
_setup_proxy_logging(monkeypatch)
_setup_mcp_call_environment(monkeypatch)
captured = {}
def fake_function_setup(*_args, **kwargs):
captured.update(kwargs)
return None, None
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"
await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_server_map={tool_name: "deepwiki"},
tool_calls=[{"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}}],
user_api_key_auth=None,
request_tags=["team-a", "prod"],
)
assert captured["metadata"]["tags"] == ["team-a", "prod"]