mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
8d0dc9294d
commit
fabe5c283a
15 changed files with 645 additions and 41 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue