Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 26 additions & 2 deletions litellm/proxy/_experimental/mcp_server/rest_endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,13 +77,32 @@ def _connection_error_message(exc: BaseException) -> str:
ListMCPToolsRestAPIResponseObject,
MCPInfo,
MCPServer,
_fire_mcp_success_logging,
_tool_name_matches,
execute_mcp_tool,
filter_tools_by_allowed_tools,
)

########################################################
############ 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)
Comment thread
Sameerlite marked this conversation as resolved.

def _get_server_auth_header(
server,
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]],
Expand Down Expand Up @@ -798,7 +817,8 @@ async def call_tool_rest_api(
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,
Expand All @@ -809,6 +829,8 @@ async def call_tool_rest_api(
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")
Expand Down Expand Up @@ -876,11 +898,12 @@ async def call_tool_rest_api(
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()
Comment thread
Sameerlite marked this conversation as resolved.
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"),
Expand All @@ -889,6 +912,7 @@ async def call_tool_rest_api(
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(
Expand Down
33 changes: 22 additions & 11 deletions litellm/proxy/_experimental/mcp_server/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -1681,6 +1681,7 @@ async def _get_tools_from_mcp_servers(
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]:
"""
Expand Down Expand Up @@ -1724,6 +1725,7 @@ async def _get_tools_from_mcp_servers(
"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": [
Expand Down Expand Up @@ -1899,7 +1901,9 @@ async def _fetch_and_filter_server_tools(
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,
)
Expand Down Expand Up @@ -2741,6 +2745,22 @@ async def execute_mcp_tool(

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,
Expand Down Expand Up @@ -2812,16 +2832,7 @@ async def call_mcp_tool(
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(
Expand Down
19 changes: 19 additions & 0 deletions litellm/proxy/hooks/proxy_track_cost_callback.py
Original file line number Diff line number Diff line change
Expand Up @@ -190,6 +190,11 @@ async def _PROXY_track_cost_callback(
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))
Expand Down Expand Up @@ -388,6 +393,20 @@ def _should_track_errors_in_db():
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],
Expand Down
47 changes: 47 additions & 0 deletions litellm/proxy/spend_tracking/spend_management_endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -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[])
Comment thread
veria-ai[bot] marked this conversation as resolved.
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:
Expand Down
3 changes: 3 additions & 0 deletions litellm/responses/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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,
Expand Down
29 changes: 25 additions & 4 deletions litellm/responses/mcp/chat_completions_handler.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
"""Helpers for handling MCP-aware `/chat/completions` requests."""

import logging
from typing import (
Any,
List,
Expand Down Expand Up @@ -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)
(
Expand All @@ -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(
Expand Down Expand Up @@ -218,6 +221,7 @@ def __init__(
litellm_trace_id,
openai_tools,
base_call_args,
request_tags,
):
self.stream_wrapper = stream_wrapper
self.messages = messages
Expand All @@ -231,6 +235,7 @@ def __init__(
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
Expand Down Expand Up @@ -303,6 +308,17 @@ def _add_mcp_tool_metadata_to_final_chunk(self, chunk: ModelResponseStream) -> M

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:
Expand Down Expand Up @@ -332,15 +348,16 @@ async def __anext__(self):
)

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()
Comment thread
Sameerlite marked this conversation as resolved.

return chunk
except StopAsyncIteration:
Expand All @@ -354,6 +371,7 @@ async def __anext__(self):
# 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
Expand Down Expand Up @@ -426,6 +444,7 @@ async def _process_tool_calls(self):
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):
Expand Down Expand Up @@ -485,6 +504,7 @@ async def _prepare_follow_up_call(self):
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
Expand Down Expand Up @@ -596,6 +616,7 @@ def __next__(self):
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:
Expand Down
Loading
Loading