diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index ccdcf9c8434..5cb41d27e0f 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2627,6 +2627,13 @@ if MCP_AVAILABLE: server_name=server_name, session_id=_mcp_session_id_from_headers(raw_headers), ) + # Dual attribution: on a validated delegation the agent key stays the primary + # spend/audit identity; record the delegated user it acted on behalf of so + # per-user visibility is not lost. This is the single chokepoint both the + # streamable and REST call paths funnel through. + delegated_user_id = getattr(user_api_key_auth, "delegated_user_id", None) if user_api_key_auth else None + if delegated_user_id: + standard_logging_mcp_tool_call["mcp_on_behalf_of_user_id"] = delegated_user_id litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj", None) if litellm_logging_obj: litellm_logging_obj.model_call_details["mcp_tool_call_metadata"] = standard_logging_mcp_tool_call diff --git a/litellm/types/utils.py b/litellm/types/utils.py index ec8a9336ca7..6767c93a10d 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2569,6 +2569,14 @@ class StandardLoggingMCPToolCall(TypedDict, total=False): Records which upstream received a relayed request; never a credential. """ + mcp_on_behalf_of_user_id: str | None + """ + When the call ran under a validated user->agent delegation, the user_id the calling agent + acted on behalf of. The agent key remains the primary spend/audit attribution; this records + the delegated (triggering) user so N users behind one agent stay distinguishable in the logs. + Absent on non-delegated calls. + """ + class StandardLoggingVectorStoreRequest(TypedDict, total=False): """ 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 7e59904a39c..dab4339c387 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 @@ -7573,3 +7573,99 @@ class TestPreemptive401ModeAware: await self._run(delegate, None, has_stored_token=False) assert exc.value.status_code == 401 await self._run(delegate, self.LITELLM_KEY_HEADERS, has_stored_token=False) + + +@pytest.mark.asyncio +async def test_execute_mcp_tool_records_delegated_user_in_logging_metadata(): + """A delegated call (user_api_key_auth.delegated_user_id set) records the + on-behalf-of user in mcp_tool_call_metadata, in addition to the agent key.""" + from mcp.types import TextContent + + from litellm.proxy._experimental.mcp_server import server as mcp_module + from litellm.proxy._types import UserAPIKeyAuth + + server = MCPServer( + server_id="srv", + name="echo", + server_name="echo", + url="http://127.0.0.1:5115/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + authentication_token="abc", + ) + + async def fake_handle_managed_mcp_tool(**kwargs): + return mcp_module.CallToolResult(content=[TextContent(type="text", text="ok")], isError=False) + + logging_obj = MagicMock() + logging_obj.model_call_details = {} + + agent_auth = UserAPIKeyAuth(api_key="agent-key", user_id="agent-svc-user", delegated_user_id="alice-id") + + with ( + patch.object(mcp_module.global_mcp_server_manager, "get_registry", return_value={server.server_id: server}), + patch.object(mcp_module.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=server), + patch.object(mcp_module, "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool), + patch.object(mcp_module.MCPRequestHandler, "is_tool_allowed", return_value=True), + patch.object(mcp_module.global_mcp_tool_registry, "get_tool", return_value=None), + ): + await mcp_module.execute_mcp_tool( + name="echo", + arguments={"message": "hi"}, + allowed_mcp_servers=[server], + start_time=datetime.now(), + requested_server_id=server.server_id, + user_api_key_auth=agent_auth, + litellm_logging_obj=logging_obj, + ) + + meta = logging_obj.model_call_details["mcp_tool_call_metadata"] + assert meta["mcp_on_behalf_of_user_id"] == "alice-id" + + +@pytest.mark.asyncio +async def test_execute_mcp_tool_omits_delegated_user_when_not_delegated(): + """No delegation -> the on-behalf-of field is absent (payload byte-identical + to today for non-delegated calls).""" + from mcp.types import TextContent + + from litellm.proxy._experimental.mcp_server import server as mcp_module + from litellm.proxy._types import UserAPIKeyAuth + + server = MCPServer( + server_id="srv", + name="echo", + server_name="echo", + url="http://127.0.0.1:5115/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + authentication_token="abc", + ) + + async def fake_handle_managed_mcp_tool(**kwargs): + return mcp_module.CallToolResult(content=[TextContent(type="text", text="ok")], isError=False) + + logging_obj = MagicMock() + logging_obj.model_call_details = {} + + plain_auth = UserAPIKeyAuth(api_key="user-key", user_id="alice-id") + + with ( + patch.object(mcp_module.global_mcp_server_manager, "get_registry", return_value={server.server_id: server}), + patch.object(mcp_module.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=server), + patch.object(mcp_module, "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool), + patch.object(mcp_module.MCPRequestHandler, "is_tool_allowed", return_value=True), + patch.object(mcp_module.global_mcp_tool_registry, "get_tool", return_value=None), + ): + await mcp_module.execute_mcp_tool( + name="echo", + arguments={"message": "hi"}, + allowed_mcp_servers=[server], + start_time=datetime.now(), + requested_server_id=server.server_id, + user_api_key_auth=plain_auth, + litellm_logging_obj=logging_obj, + ) + + meta = logging_obj.model_call_details["mcp_tool_call_metadata"] + assert "mcp_on_behalf_of_user_id" not in meta