feat(mcp): record the delegated on-behalf-of user in MCP tool-call spend logs

When an MCP tool call runs under a validated user->agent delegation
(user_api_key_auth.delegated_user_id set), record that user in the tool-call
logging metadata alongside the agent key's existing attribution, so N users
behind one agent stay distinguishable in the logs. Injected once in
execute_mcp_tool, the chokepoint both the streamable and REST call paths share.
Non-delegated calls are unchanged. Part of LIT-4448.
This commit is contained in:
Tin Chi Lo 2026-07-17 15:31:40 -07:00
parent 490028b0ef
commit 54a00ff00b
3 changed files with 111 additions and 0 deletions

View file

@ -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

View file

@ -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):
"""

View file

@ -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