From fabe5c283ac9d2e36ba0ca46a5ebc260a9c76fde Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 2 Jul 2026 20:46:39 +0530 Subject: [PATCH] fix(mcp): roll up MCP tool spend to user counters and usage UI (#31576) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * 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 * 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 * 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 * 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 * 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 Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .../mcp_server/rest_endpoints.py | 28 +- .../proxy/_experimental/mcp_server/server.py | 33 ++- .../proxy/hooks/proxy_track_cost_callback.py | 19 ++ .../spend_management_endpoints.py | 47 ++++ litellm/responses/main.py | 3 + .../responses/mcp/chat_completions_handler.py | 29 ++- .../mcp/litellm_proxy_mcp_handler.py | 38 ++- .../responses/mcp/mcp_streaming_iterator.py | 1 + .../mcp_server/test_mcp_server.py | 19 +- .../mcp_server/test_mcp_tool_search.py | 6 + .../mcp_server/test_rest_endpoints.py | 48 +++- .../hooks/test_proxy_track_cost_callback.py | 73 ++++++ .../test_spend_management_endpoints.py | 29 ++- .../mcp/test_chat_completions_handler.py | 239 ++++++++++++++++++ .../mcp/test_litellm_proxy_mcp_handler.py | 74 ++++++ 15 files changed, 645 insertions(+), 41 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 7ab7eb28147..a6067a60105 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index a65239b296f..57404793269 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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( diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 9c09231cd9f..b6342f4fa1a 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -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], diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 5624afcfa5e..9c152e30c52 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -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: diff --git a/litellm/responses/main.py b/litellm/responses/main.py index e8fe51ed484..8e3be2bc12d 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -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, diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index 10ff67f68d5..f2ccfd430ae 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -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: diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index e969208d1d9..999945b3823 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -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: diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index a961271e3f0..3c24a703b68 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 2b7c29ceff3..44f1d105093 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py index 9c20808df67..f2b74d65059 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 9e3862b43eb..46a51f8fe61 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -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( diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 0cbf308076c..813a0c5e38f 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -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" + ) diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index b37bb18744d..4acc94c737a 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -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"}) diff --git a/tests/test_litellm/responses/mcp/test_chat_completions_handler.py b/tests/test_litellm/responses/mcp/test_chat_completions_handler.py index fc5d2e5d382..ab4c5185057 100644 --- a/tests/test_litellm/responses/mcp/test_chat_completions_handler.py +++ b/tests/test_litellm/responses/mcp/test_chat_completions_handler.py @@ -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 diff --git a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py index 94712813783..35bfa6ea9e4 100644 --- a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -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"]