mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): key the listed-tool slot by the caller's admission identity and forwarded bearer
The slot a tools/list records for a later tools/call was keyed by (user_id, api_key) only, so every team-only JWT caller shared one slot and one JWT user acting in two teams shared a slot; a tools/call then handed pre_mcp_call hooks a description another caller was served. The slot is now keyed by the hashed key, user, team and organization, plus the admission credential of a caller admitted with neither a key nor a user. The caller bearer split the slot only on client-forwarded-token and token-exchange servers; a legacy delegated oauth2 server (delegate_auth_to_upstream without client credentials) also forwards it upstream and served a different catalog per bearer into one slot. The bearer now splits the slot on every server whose egress forwards it (_consumes_caller_authorization) or exchanges it.
This commit is contained in:
parent
0219521559
commit
ddb672155f
2 changed files with 149 additions and 15 deletions
|
|
@ -1221,6 +1221,21 @@ def listed_tools_caller_for(
|
|||
)
|
||||
|
||||
|
||||
def _admission_identity(
|
||||
auth: UserAPIKeyAuth, raw_headers: Mapping[str, str] | None
|
||||
) -> tuple[str | None, str | None, str | None, str | None, str | None]:
|
||||
"""The admission identity the served catalog is shaped for: the hashed key, user, team and
|
||||
organization, plus the admission credential (``x-litellm-api-key``, else ``Authorization``) of a
|
||||
caller admitted with neither a key nor a user."""
|
||||
keyless: Final = auth.api_key is None and auth.user_id is None
|
||||
credential: Final = (
|
||||
_raw_header_value(raw_headers, "x-litellm-api-key") or _raw_header_value(raw_headers, "authorization")
|
||||
if keyless
|
||||
else None
|
||||
)
|
||||
return auth.api_key, auth.user_id, auth.team_id, auth.org_id, credential
|
||||
|
||||
|
||||
def _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str:
|
||||
"""Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection.
|
||||
|
||||
|
|
@ -4692,11 +4707,14 @@ class MCPServerManager:
|
|||
def _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None:
|
||||
"""Key the listed-tool cache by every request input that can change the served catalog.
|
||||
|
||||
The catalog is guardrail-shaped for the caller's own key (default-on guardrails, key or team
|
||||
selections and opt-outs), so every keyed caller gets its own slot, on OpenAPI servers too.
|
||||
Forwarded headers, header-driven stdio env, the caller bearer (forwarded as-is or exchanged as
|
||||
the OBO subject) and the server-specific auth header also reach upstream and split the slot
|
||||
further. Only unkeyed listings with none of those share the ``None`` slot.
|
||||
The catalog is guardrail-shaped for the caller's admission identity (default-on guardrails,
|
||||
key or team selections and opt-outs), so every admitted caller gets its own slot, keyed by
|
||||
``_admission_identity``: the hashed key, user, team and organization, plus the admission
|
||||
credential of a caller admitted with neither a key nor a user (a team-only JWT). Forwarded
|
||||
headers, header-driven stdio env, the caller bearer on every server whose egress forwards it
|
||||
(``_consumes_caller_authorization``) or exchanges it as the OBO subject, and the
|
||||
server-specific auth header also reach upstream and split the slot further. Only unkeyed
|
||||
listings with none of those share the ``None`` slot.
|
||||
"""
|
||||
if caller is None:
|
||||
return None
|
||||
|
|
@ -4706,19 +4724,18 @@ class MCPServerManager:
|
|||
stdio_env: Final = None if header_env == self._build_stdio_env(server) else header_env
|
||||
caller_bearer: Final = (
|
||||
self._extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth)
|
||||
if server.is_client_forwarded_token or server.auth_type == MCPAuth.oauth2_token_exchange
|
||||
if _consumes_caller_authorization(server) or server.auth_type == MCPAuth.oauth2_token_exchange
|
||||
else None
|
||||
)
|
||||
_, digest = self._discovery_key(
|
||||
server,
|
||||
auth,
|
||||
caller.mcp_auth_header,
|
||||
forwarded,
|
||||
stdio_env,
|
||||
caller_bearer,
|
||||
per_caller=auth is not None,
|
||||
identity: Final = None if auth is None else _admission_identity(auth, caller.raw_headers)
|
||||
if not (identity or caller.mcp_auth_header or forwarded or stdio_env or caller_bearer):
|
||||
return None
|
||||
material: Final = json.dumps(
|
||||
(identity, caller.mcp_auth_header, forwarded, stdio_env, caller_bearer),
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
return digest
|
||||
return hashlib.sha256(material.encode()).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _forwarded_header_values(
|
||||
|
|
|
|||
|
|
@ -7722,6 +7722,123 @@ class TestMCPServerManager:
|
|||
|
||||
assert listed is not None and listed.description == "slot a"
|
||||
|
||||
def test_listed_tools_slot_is_split_per_team_for_keyless_callers(self):
|
||||
"""A team-only JWT admits a caller with neither a key nor a user, so the team keys the slot."""
|
||||
manager: Final = MCPServerManager()
|
||||
server: Final = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
|
||||
team_one: Final = ListedToolsCaller(
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-one")
|
||||
)
|
||||
team_two: Final = ListedToolsCaller(
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-two")
|
||||
)
|
||||
manager._record_listed_tools(
|
||||
server, [MCPTool(name="foo", description="Fetch rows FLAGWORD", inputSchema={})], team_one
|
||||
)
|
||||
|
||||
assert manager.get_listed_tool(server, "foo", team_two) is None
|
||||
listed: Final = manager.get_listed_tool(server, "foo", team_one)
|
||||
assert listed is not None and listed.description == "Fetch rows FLAGWORD"
|
||||
|
||||
def test_listed_tools_slot_is_split_per_team_for_the_same_keyless_user(self):
|
||||
"""One JWT user acting in two teams is served two team-shaped catalogs, so each team is a slot."""
|
||||
manager: Final = MCPServerManager()
|
||||
server: Final = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
|
||||
alice_in_one: Final = ListedToolsCaller(
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id="alice", team_id="team-one")
|
||||
)
|
||||
alice_in_two: Final = ListedToolsCaller(
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id="alice", team_id="team-two")
|
||||
)
|
||||
manager._record_listed_tools(
|
||||
server, [MCPTool(name="foo", description="Fetch rows FLAGWORD", inputSchema={})], alice_in_one
|
||||
)
|
||||
|
||||
assert manager.get_listed_tool(server, "foo", alice_in_two) is None
|
||||
listed: Final = manager.get_listed_tool(server, "foo", alice_in_one)
|
||||
assert listed is not None and listed.description == "Fetch rows FLAGWORD"
|
||||
|
||||
def test_listed_tools_slot_is_split_by_the_admission_bearer_of_keyless_callers_without_a_user(self):
|
||||
"""Two team-only JWT callers of one team differ only in the JWT they were admitted with, so that
|
||||
credential keys the slot, on a server that never forwards it."""
|
||||
manager: Final = MCPServerManager()
|
||||
server: Final = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
|
||||
alice: Final = ListedToolsCaller(
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-one"),
|
||||
raw_headers={"authorization": "Bearer jwt-alice"},
|
||||
)
|
||||
bob: Final = ListedToolsCaller(
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-one"),
|
||||
raw_headers={"authorization": "Bearer jwt-bob"},
|
||||
)
|
||||
manager._record_listed_tools(server, [MCPTool(name="foo", description="alice view", inputSchema={})], alice)
|
||||
|
||||
assert manager.get_listed_tool(server, "foo", bob) is None
|
||||
listed: Final = manager.get_listed_tool(server, "foo", alice)
|
||||
assert listed is not None and listed.description == "alice view"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("server_kwargs", "forwards_bearer"),
|
||||
[
|
||||
pytest.param(
|
||||
{"auth_type": MCPAuth.oauth2, "delegate_auth_to_upstream": True, "oauth2_flow": "authorization_code"},
|
||||
True,
|
||||
id="oauth2-delegated-to-upstream",
|
||||
),
|
||||
pytest.param({"auth_type": MCPAuth.oauth_delegate}, True, id="oauth-delegate"),
|
||||
pytest.param({"auth_type": MCPAuth.true_passthrough}, True, id="true-passthrough"),
|
||||
pytest.param({"auth_type": MCPAuth.oauth2_token_exchange}, True, id="token-exchange"),
|
||||
pytest.param(
|
||||
{"auth_type": MCPAuth.none, "extra_headers": ["Authorization"], "oauth_passthrough": True},
|
||||
True,
|
||||
id="oauth-passthrough",
|
||||
),
|
||||
pytest.param({}, False, id="plain"),
|
||||
pytest.param(
|
||||
{"auth_type": MCPAuth.oauth2, "oauth2_flow": "authorization_code"},
|
||||
False,
|
||||
id="oauth2-gateway-managed",
|
||||
),
|
||||
pytest.param(
|
||||
{
|
||||
"auth_type": MCPAuth.oauth2,
|
||||
"oauth2_flow": "client_credentials",
|
||||
"delegate_auth_to_upstream": True,
|
||||
"client_id": "gateway",
|
||||
"client_secret": "secret",
|
||||
"token_url": "http://idp/token",
|
||||
},
|
||||
False,
|
||||
id="oauth2-client-credentials",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_listed_tools_slot_is_split_by_the_forwarded_bearer_on_servers_that_forward_it(
|
||||
self, server_kwargs: dict[str, object], forwards_bearer: bool
|
||||
):
|
||||
"""Two callers sharing one key but carrying different upstream bearers are served two upstream
|
||||
catalogs exactly on the servers whose egress forwards or exchanges that bearer."""
|
||||
manager: Final = MCPServerManager()
|
||||
server: Final = MCPServer(
|
||||
**{"server_id": "dg", "name": "dg", "transport": MCPTransport.http, "url": "http://dg", **server_kwargs}
|
||||
)
|
||||
caller_a: Final = ListedToolsCaller(
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-master"),
|
||||
raw_headers={"x-litellm-api-key": "Bearer sk-master", "authorization": "Bearer UP-A"},
|
||||
)
|
||||
caller_b: Final = ListedToolsCaller(
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-master"),
|
||||
raw_headers={"x-litellm-api-key": "Bearer sk-master", "authorization": "Bearer UP-B"},
|
||||
)
|
||||
manager._record_listed_tools(
|
||||
server, [MCPTool(name="lookup", description="Workspace A lookup FLAGWORD", inputSchema={})], caller_a
|
||||
)
|
||||
|
||||
for_b: Final = manager.get_listed_tool(server, "lookup", caller_b)
|
||||
assert (for_b is None) is forwards_bearer
|
||||
for_a: Final = manager.get_listed_tool(server, "lookup", caller_a)
|
||||
assert for_a is not None and for_a.description == "Workspace A lookup FLAGWORD"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_tool_hands_hooks_the_catalog_the_same_forwarded_headers_listed(self):
|
||||
manager = MCPServerManager()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue