mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
490028b0ef
commit
54a00ff00b
3 changed files with 111 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue