mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): name the connected segment in sign-in challenge resource metadata
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
6a71a3b1a0
commit
fcfd7b4f14
6 changed files with 159 additions and 5 deletions
|
|
@ -531,7 +531,10 @@ def _resolve_mcp_server_by_name_or_id(lookup: str, client_ip: str | None) -> MCP
|
|||
by_name: Final = global_mcp_server_manager.get_mcp_server_by_name(lookup, client_ip=client_ip)
|
||||
if by_name is not None:
|
||||
return by_name
|
||||
return global_mcp_server_manager.get_mcp_server_by_id(lookup, client_ip=client_ip)
|
||||
by_id: Final = global_mcp_server_manager.get_mcp_server_by_id(lookup, client_ip=client_ip)
|
||||
if by_id is not None:
|
||||
return by_id
|
||||
return global_mcp_server_manager.get_mcp_server_answering_to(lookup, client_ip=client_ip)
|
||||
|
||||
|
||||
def _resolve_oauth2_server_for_root_endpoints(
|
||||
|
|
|
|||
|
|
@ -298,7 +298,7 @@ def raise_public(error: CredError) -> NoReturn:
|
|||
assert_never(error.tag)
|
||||
|
||||
|
||||
def oauth_protected_resource_path(root_path: str, server: MCPServer) -> str:
|
||||
def oauth_protected_resource_path(root_path: str, server: MCPServer, *, connected_as: str | None = None) -> str:
|
||||
"""The server's RFC 9728 Protected Resource Metadata path, the shared anchor of both challenges.
|
||||
|
||||
``root_path`` is the prefix the request was routed under, resolved by the caller (the imperative
|
||||
|
|
@ -320,7 +320,7 @@ def oauth_protected_resource_path(root_path: str, server: MCPServer) -> str:
|
|||
challenge would then disagree on where the resource metadata lives.
|
||||
"""
|
||||
prefix: Final = "" if root_path == "/" else root_path
|
||||
name: Final = server.alias or server.server_name or server.name or server.server_id
|
||||
name: Final = connected_as or server.alias or server.server_name or server.name or server.server_id
|
||||
scalar_env: Final = os.getenv("SERVER_ROOT_PATH", "").rstrip("/")
|
||||
if not prefix or (scalar_env and prefix == scalar_env):
|
||||
return f"/.well-known/oauth-protected-resource{prefix}/mcp/{name}"
|
||||
|
|
@ -348,6 +348,7 @@ def raise_token_exchange_challenge(
|
|||
*,
|
||||
root_path: str,
|
||||
claims: str | None = None,
|
||||
connected_as: str | None = None,
|
||||
) -> NoReturn:
|
||||
"""Raise the RFC 9728 / RFC 6750 challenge an OBO (``token_exchange``) server returns when the
|
||||
caller's subject token is missing or the IdP rejected it.
|
||||
|
|
@ -366,7 +367,7 @@ def raise_token_exchange_challenge(
|
|||
two literals) and the base64 claims draw from a fixed alphabet, so nothing from the IdP body
|
||||
reaches the header unescaped.
|
||||
"""
|
||||
resource_metadata: Final = oauth_protected_resource_path(root_path, server)
|
||||
resource_metadata: Final = oauth_protected_resource_path(root_path, server, connected_as=connected_as)
|
||||
encoded_claims: Final = base64.b64encode(claims.encode()).decode() if claims else None
|
||||
error: Final = "insufficient_claims" if encoded_claims else "invalid_token"
|
||||
error_description: Final = (
|
||||
|
|
|
|||
|
|
@ -1737,7 +1737,9 @@ if MCP_AVAILABLE:
|
|||
get_request_root_path,
|
||||
)
|
||||
|
||||
raise_token_exchange_challenge(server, root_path=get_request_root_path())
|
||||
raise_token_exchange_challenge(
|
||||
server, root_path=get_request_root_path(), connected_as=server_name
|
||||
)
|
||||
|
||||
# Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run
|
||||
# the exchange here at the transport edge, so a rejected subject raises the RFC 9728
|
||||
|
|
|
|||
|
|
@ -216,3 +216,46 @@ def test_agent_365_gated_server_challenges_at_connect_and_advertises_entra(gatew
|
|||
assert refused.status_code == 403, refused.text
|
||||
assert "www-authenticate" not in refused.headers
|
||||
assert api.drain() == ()
|
||||
|
||||
|
||||
def test_challenge_and_prm_resolve_the_connected_case_variant(gateway: Gateway, tmp_path: Path) -> None:
|
||||
def nothing(request: Request) -> Reply:
|
||||
return Reply(status=500)
|
||||
|
||||
with wire_server(nothing) as api:
|
||||
config: Final = _sign_in_config(
|
||||
{
|
||||
"guardrail": "agent_365",
|
||||
"mode": "pre_mcp_call",
|
||||
"default_on": True,
|
||||
"tenant_id": "00000000-0000-0000-0000-000000000000",
|
||||
"client_id": "22222222-2222-2222-2222-222222222222",
|
||||
"client_secret": "secret",
|
||||
"api_base": api.url,
|
||||
},
|
||||
tmp_path / "agent365-case.yaml",
|
||||
)
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config) as candidate,
|
||||
mcp_peer() as peer,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
alias: Final = "a365" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
granted: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
connected_as: Final = alias.upper()
|
||||
|
||||
challenged: Final = _rpc(candidate, f"/mcp/{connected_as}", granted, {})
|
||||
assert challenged.status_code == 401, challenged.text
|
||||
authenticate: Final = challenged.headers.get("www-authenticate", "")
|
||||
assert f'resource_metadata="/.well-known/oauth-protected-resource/mcp/{connected_as}"' in authenticate
|
||||
assert 'error="invalid_token"' in authenticate
|
||||
|
||||
discovery: Final = candidate.client.get(f"/.well-known/oauth-protected-resource/mcp/{connected_as}")
|
||||
assert discovery.status_code == 200, discovery.text
|
||||
document: Final = discovery.json()
|
||||
assert document["authorization_servers"] == [
|
||||
"https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0"
|
||||
]
|
||||
assert document["scopes_supported"] == ["api://22222222-2222-2222-2222-222222222222/access_as_user"]
|
||||
assert api.drain() == ()
|
||||
|
|
|
|||
|
|
@ -3743,6 +3743,56 @@ async def test_protected_resource_metadata_resolves_server_by_id_when_name_looku
|
|||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_protected_resource_metadata_resolves_the_connected_case_variant():
|
||||
"""The challenge points clients at the segment they connected with (``/mcp/CATALOG``), so the PRM
|
||||
route must resolve that same segment through the answering-to fallback."""
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
server = MCPServer(
|
||||
server_id="catalog-server-id-001",
|
||||
name="catalog",
|
||||
alias="catalog",
|
||||
server_name="catalog",
|
||||
url="https://catalog.test/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
mcp_info={"server_name": "catalog"},
|
||||
)
|
||||
sign_in: Final = CallerSignIn(
|
||||
issuers=("https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0",),
|
||||
scopes=("api://22222222-2222-2222-2222-222222222222/access_as_user",),
|
||||
)
|
||||
request = MagicMock(spec=Request)
|
||||
request.base_url = "https://llm.example.com/"
|
||||
request.headers = {}
|
||||
|
||||
with (
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None), # test-quality-ok: resolver seam
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=None), # test-quality-ok: resolver seam
|
||||
patch.object(
|
||||
global_mcp_server_manager, "get_filtered_registry", return_value={server.server_id: server}
|
||||
), # test-quality-ok: resolver seam
|
||||
patch.object(discoverable_endpoints, "caller_sign_in_for", return_value=sign_in), # test-quality-ok: provider seam
|
||||
):
|
||||
result = await discoverable_endpoints._build_oauth_protected_resource_response(
|
||||
request=request,
|
||||
mcp_server_name="CATALOG",
|
||||
use_standard_pattern=True,
|
||||
)
|
||||
|
||||
assert result["authorization_servers"] == [
|
||||
"https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0"
|
||||
]
|
||||
assert result["resource"] == "https://llm.example.com/mcp/CATALOG"
|
||||
|
||||
|
||||
def test_authorization_server_metadata_resolves_server_by_id_when_name_lookup_fails():
|
||||
from fastapi import Request
|
||||
|
||||
|
|
|
|||
|
|
@ -11063,11 +11063,22 @@ def _catalog_server() -> MCPServer:
|
|||
|
||||
|
||||
class _CallerSignInGuardrail(CustomGuardrail):
|
||||
def __init__(self, *args, preflight_result=None, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self._preflight_result = preflight_result
|
||||
self.preflight_calls = [] # mutable-ok: call recorder
|
||||
|
||||
def caller_sign_in(self, server, user_api_key_auth):
|
||||
from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn
|
||||
|
||||
return CallerSignIn(issuers=("https://idp.test",), scopes=("scope-a",))
|
||||
|
||||
async def preflight_caller_sign_in(self, server, user_api_key_auth, subject_token):
|
||||
from litellm.proxy._experimental.mcp_server.caller_sign_in import SignedIn
|
||||
|
||||
self.preflight_calls.append(subject_token)
|
||||
return self._preflight_result if self._preflight_result is not None else SignedIn()
|
||||
|
||||
|
||||
class TestConnectChallengeResolver:
|
||||
"""The connect-time sign-in challenge must resolve the server the same way the router resolves
|
||||
|
|
@ -11195,3 +11206,47 @@ class TestConnectChallengeResolver:
|
|||
'error="invalid_token", '
|
||||
'error_description="Missing or invalid subject token; authenticate with the IdP and retry"'
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"route_name",
|
||||
["catalog-server-id-001", "CATALOG"],
|
||||
ids=["server_id", "uppercase_name"],
|
||||
)
|
||||
async def test_challenge_resource_metadata_names_the_connected_segment(self, route_name):
|
||||
"""The PRM path in the challenge must name the segment the client connected with, or the
|
||||
client's follow-up metadata fetch 404s against the route it was pointed at."""
|
||||
from litellm.proxy._experimental.mcp_server import server as server_module
|
||||
|
||||
server = _catalog_server()
|
||||
guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub")
|
||||
litellm.logging_callback_manager.add_litellm_callback(guardrail)
|
||||
try:
|
||||
with (
|
||||
patch.object(
|
||||
mcp_operations.global_mcp_server_manager,
|
||||
"get_filtered_registry",
|
||||
return_value={server.server_id: server},
|
||||
),
|
||||
patch.object(
|
||||
mcp_operations,
|
||||
"_get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[server]),
|
||||
),
|
||||
pytest.raises(HTTPException) as exc,
|
||||
):
|
||||
await server_module._raise_preemptive_401_for_unauthenticated_servers(
|
||||
scope={"type": "http", "method": "POST", "path": f"/mcp/{route_name}", "headers": []},
|
||||
mcp_servers=[route_name],
|
||||
oauth2_headers=None,
|
||||
mcp_server_auth_headers=None,
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
|
||||
client_ip=None,
|
||||
)
|
||||
finally:
|
||||
litellm.logging_callback_manager.remove_callback_from_list_by_object(
|
||||
litellm.callbacks, guardrail, require_self=False
|
||||
)
|
||||
|
||||
authenticate: Final = (exc.value.headers or {}).get("WWW-Authenticate") or ""
|
||||
assert authenticate.startswith(f'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/{route_name}"')
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue