mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): keep oauth2 listing on the minted or signed credential, not the stored BYOK secret
The listing helper that keys the per-caller catalog by the stored BYOK credential also handed that credential to the upstream client, which on an oauth2 server short-circuited the client_credentials mint and the MCPJWTSigner gate. Split the two: the catalog identity keeps the stored credential so tools/call finds the caller's slot, while an oauth2 server's tools/list sends only the per-request header, letting the M2M mint or signed JWT proceed as on main Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
c8870f0080
commit
de92a1103c
2 changed files with 69 additions and 2 deletions
|
|
@ -1361,6 +1361,18 @@ async def _byok_listing_auth_header(
|
|||
return await _get_byok_credential(mcp_server, user_api_key_auth)
|
||||
|
||||
|
||||
def _listing_upstream_auth_header(
|
||||
mcp_server: MCPServer,
|
||||
mcp_auth_header: str | dict[str, str] | None,
|
||||
byok_auth_header: str | dict[str, str] | None,
|
||||
) -> str | dict[str, str] | None:
|
||||
"""The credential a tools/list sends upstream: an oauth2 server's is minted or signed, never the
|
||||
stored BYOK one, the same rule ``_get_tools_from_mcp_servers`` applies to its own resolution."""
|
||||
if mcp_server.auth_type == MCPAuth.oauth2:
|
||||
return mcp_auth_header
|
||||
return byok_auth_header
|
||||
|
||||
|
||||
def _client_forwarded_authorization_headers(
|
||||
mcp_server: MCPServer,
|
||||
oauth2_headers: dict[str, str] | None,
|
||||
|
|
@ -4500,6 +4512,9 @@ class MCPServerManager:
|
|||
|
||||
client = None
|
||||
resolved_mcp_auth_header: Final = await _byok_listing_auth_header(server, user_api_key_auth, mcp_auth_header)
|
||||
upstream_mcp_auth_header: Final = _listing_upstream_auth_header(
|
||||
server, mcp_auth_header, resolved_mcp_auth_header
|
||||
)
|
||||
listed_caller: Final = ListedToolsCaller(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=resolved_mcp_auth_header,
|
||||
|
|
@ -4549,7 +4564,7 @@ class MCPServerManager:
|
|||
if (
|
||||
get_mcp_jwt_signer() is not None
|
||||
and not has_static_authorization
|
||||
and not resolved_mcp_auth_header
|
||||
and not upstream_mcp_auth_header
|
||||
and not has_extra_authorization
|
||||
):
|
||||
extra_headers = await inject_mcp_jwt_headers_for_upstream(
|
||||
|
|
@ -4572,7 +4587,7 @@ class MCPServerManager:
|
|||
|
||||
client = await self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=resolved_mcp_auth_header,
|
||||
mcp_auth_header=upstream_mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
|
|
|
|||
|
|
@ -7380,6 +7380,58 @@ class TestMCPServerManager:
|
|||
assert listed is not None and listed.description == "t"
|
||||
assert manager._create_mcp_client.await_args.kwargs["mcp_auth_header"] == "Bearer hdr"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth2_byok_listing_leaves_the_minted_token_and_signer_in_place(self):
|
||||
"""The stored BYOK secret keys the catalog slot tools/call reads, but never reaches the oauth2
|
||||
upstream on tools/list: the client mints its M2M token and MCPJWTSigner still signs."""
|
||||
from litellm.proxy._experimental.mcp_server.byok_credential_cache import (
|
||||
byok_credential_cache_key,
|
||||
cache_byok_credential,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.operations import byok_credential_cache
|
||||
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="cc1",
|
||||
name="cc1",
|
||||
transport=MCPTransport.http,
|
||||
url="http://cc1",
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="cid",
|
||||
client_secret="csec",
|
||||
token_url="http://cc1/token",
|
||||
is_byok=True,
|
||||
)
|
||||
alice = UserAPIKeyAuth(api_key="sk-alice", user_id="alice")
|
||||
manager._create_mcp_client = AsyncMock(return_value=AsyncMock())
|
||||
manager._fetch_tools_with_timeout = AsyncMock(
|
||||
return_value=[MCPTool(name="echo", description="m2m catalog", inputSchema={})]
|
||||
)
|
||||
signer_headers = AsyncMock(return_value={"Authorization": "Bearer signed-jwt"})
|
||||
cache_byok_credential("alice", "cc1", "BYOK-ALICE-SECRET")
|
||||
try:
|
||||
with (
|
||||
patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam
|
||||
"litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch( # test-quality-ok: same singleton's header injection, asserted on by call
|
||||
"litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.inject_mcp_jwt_headers_for_upstream",
|
||||
signer_headers,
|
||||
),
|
||||
):
|
||||
await manager._get_tools_from_server(server=server, user_api_key_auth=alice)
|
||||
finally:
|
||||
byok_credential_cache.delete_cache(byok_credential_cache_key("alice", "cc1"))
|
||||
|
||||
client_kwargs = manager._create_mcp_client.await_args.kwargs
|
||||
assert client_kwargs["mcp_auth_header"] is None, client_kwargs
|
||||
assert client_kwargs["extra_headers"] == {"Authorization": "Bearer signed-jwt"}
|
||||
signer_headers.assert_awaited_once()
|
||||
call_side = ListedToolsCaller(user_api_key_auth=alice, mcp_auth_header="BYOK-ALICE-SECRET")
|
||||
listed = manager.get_listed_tool(server, "echo", call_side)
|
||||
assert listed is not None and listed.description == "m2m catalog"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("signer", "static_headers"),
|
||||
[
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue