mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): keep the stored BYOK credential for catalog identity only on tools/list
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:
parent
0e4de07c65
commit
9c68ca4412
3 changed files with 72 additions and 31 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue