fix(mcp): keep the stored BYOK credential for catalog identity only on tools/list
Some checks are pending
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run

Listing used the resolved stored credential both to key the caller's catalog slot and as the
upstream transport header, so REST api_key and bearer_token listings sent the user's secret
instead of the server's static token and the MCPJWTSigner gate went quiet. The upstream client
and the signer gate now read the caller-supplied mcp_auth_header for every auth type, exactly
as before the catalog existed, and the stored credential only names the slot tools/call reads

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-10-01 01:59:51 +00:00
parent 0e4de07c65
commit 9c68ca4412
3 changed files with 72 additions and 31 deletions

View file

@ -1346,13 +1346,12 @@ async def _resolve_byok_mcp_auth_header(
return mcp_auth_header
async def _byok_listing_auth_header(
async def _byok_catalog_auth_header(
mcp_server: MCPServer,
user_api_key_auth: UserAPIKeyAuth | None,
mcp_auth_header: str | dict[str, str] | None,
) -> str | dict[str, str] | None:
"""The credential a tools/list may use: a supplied header forwards unchanged, and a missing one
falls back to the stored credential without the tool-call path's byok_auth_required raise."""
"""Keys the caller's catalog slot the way tools/call will look it up; never sent upstream."""
if not mcp_server.is_byok or mcp_auth_header is not None:
return mcp_auth_header
@ -1361,18 +1360,6 @@ 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,
@ -4511,13 +4498,9 @@ class MCPServerManager:
verbose_logger.info("_get_tools_from_server for %s...", server.name)
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,
mcp_auth_header=await _byok_catalog_auth_header(server, user_api_key_auth, mcp_auth_header),
raw_headers=raw_headers,
oauth2_headers=oauth2_headers,
)
@ -4564,7 +4547,7 @@ class MCPServerManager:
if (
get_mcp_jwt_signer() is not None
and not has_static_authorization
and not upstream_mcp_auth_header
and not mcp_auth_header
and not has_extra_authorization
):
extra_headers = await inject_mcp_jwt_headers_for_upstream(
@ -4587,7 +4570,7 @@ class MCPServerManager:
client = await self._create_mcp_client(
server=server,
mcp_auth_header=upstream_mcp_auth_header,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
stdio_env=stdio_env,
subject_token=subject_token,

View file

@ -245,6 +245,45 @@ def test_oauth2_byok_listing_sends_the_minted_token_not_the_users_stored_secret(
assert secret.encode() not in sent, "stored BYOK secret replaced the minted token on tools/list"
@pytest.mark.parametrize(("auth_type", "header", "shape"), STATIC_MODES[:2])
def test_byok_rest_listing_sends_the_servers_static_credential_not_the_users_stored_secret(
gateway: Gateway, auth_type: str, header: bytes, shape: str
) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
alias: Final = "byok" + uuid.uuid4().hex[:8]
static: Final = "static-" + uuid.uuid4().hex
identity: Final = register_mcp(
scenario, peer, alias, auth_type=auth_type, is_byok=True, credentials={"auth_value": static}
)
owner: Final = scenario.user()
owner_key: Final = scenario.key(user_id=owner, object_permission={"mcp_servers": [identity]})
secret: Final = "byok-" + uuid.uuid4().hex
stored: Final = gateway.client.post(
f"/v1/mcp/server/{identity}/user-credential",
json={"credential": secret},
headers={"x-litellm-api-key": owner_key},
)
assert stored.status_code in (200, 201), stored.text
scenario.cleanups.callback(
gateway.client.delete,
f"/v1/mcp/server/{identity}/user-credential",
headers={"x-litellm-api-key": owner_key},
)
peer.drain()
response: Final = gateway.client.get(
"/mcp-rest/tools/list", params={"server_id": identity}, headers={"x-litellm-api-key": owner_key}
)
assert response.status_code == 200, response.text
assert "add" in {tool["name"] for tool in response.json()["tools"]}, response.text
listings: Final = _listings(peer)
assert len(listings) == 1, listings
assert _header(listings[0], header) == shape.format(secret=static, basic="").encode(), listings[0]["headers"]
peer.drain()
called: Final = call_tool(gateway, owner_key, identity, f"{alias}-add", ADD)
assert called.status_code == 200, called.text
assert _header(_one_call(peer), header) == shape.format(secret=secret, basic="").encode()
def test_deprecated_string_x_mcp_auth_lists_a_byok_server_for_a_key_without_a_user(gateway: Gateway) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
alias: Final = "byok" + uuid.uuid4().hex[:8]

View file

@ -7380,10 +7380,32 @@ 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.parametrize(
"server_auth",
[
pytest.param(
{
"auth_type": MCPAuth.oauth2,
"client_id": "cid",
"client_secret": "csec",
"token_url": "http://cc1/token",
},
id="oauth2",
),
pytest.param({"auth_type": MCPAuth.api_key, "authentication_token": "STATIC-ADMIN-TOKEN"}, id="api_key"),
pytest.param(
{"auth_type": MCPAuth.bearer_token, "authentication_token": "STATIC-ADMIN-TOKEN"}, id="bearer_token"
),
pytest.param({"auth_type": MCPAuth.none}, id="none"),
],
)
@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."""
async def test_byok_listing_keys_the_catalog_by_the_stored_secret_but_never_sends_it_upstream(
self, server_auth: dict[str, object]
):
"""The stored BYOK secret keys the catalog slot tools/call reads, but tools/list sends upstream
exactly what the caller supplied (nothing here), so the static token, the M2M mint and
MCPJWTSigner all behave as they did before the catalog existed, whatever the auth_type."""
from litellm.proxy._experimental.mcp_server.byok_credential_cache import (
byok_credential_cache_key,
cache_byok_credential,
@ -7396,16 +7418,13 @@ class TestMCPServerManager:
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,
**server_auth,
)
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={})]
return_value=[MCPTool(name="echo", description="listed catalog", inputSchema={})]
)
signer_headers = AsyncMock(return_value={"Authorization": "Bearer signed-jwt"})
cache_byok_credential("alice", "cc1", "BYOK-ALICE-SECRET")
@ -7430,7 +7449,7 @@ class TestMCPServerManager:
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"
assert listed is not None and listed.description == "listed catalog"
@pytest.mark.parametrize(
("signer", "static_headers"),