mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): route a scoped name to the caller's granted server before the registry's pick
A key granted only the server named `docs` was refused with 403 on /mcp/docs when an ungranted server held `docs` as its alias, and a key granted only `github` was refused on /mcp/GITHUB when an ungranted `GitHub` existed: the scoped router took the registry-wide pick and answered "denied" whenever that pick was not among the caller's servers. The router now runs the registry's own pass order (exact alias, server_name, name, server_id; the same case-insensitively; prefix forms) over the caller's granted servers first, through get_mcp_server_answering_to(among=...), with IP hiding applied at the pass that found the server. A name hidden from the client IP stays denied before any grant lookup, and a name the registry places only on an ungranted server stays denied rather than being retried as an access group. The connect preflight reuses the router's selection for a single scoped name as the server it challenges, signs in and exchanges for, so a 401 names the granted server; an ungranted caller keeps the registry pick and the downstream 403, and the no-key path is unchanged. Unauthenticated RFC 9728 discovery has no grant list and stays on the registry pick. Tests: the `only_b` assertion in test_scoped_router_selects_the_server_the_connect_preflight_resolves and the no-IP `internal` assertion in test_scoped_router_hides_a_private_server_from_an_external_ip_like_the_connect_preflight now expect the granted server, which is what the base branch returned for both shapes; the unentitled-key exchange test's fixture becomes the empty selection the router returns for such a key; listing tests that stub the manager now point its lookup at a real empty manager so the `among` pass runs. New tests cover the grant-first router, the manager `among` pass order and IP hiding, and the connect preflight under an alias collision.
This commit is contained in:
parent
608fff7b08
commit
86c3b910a6
5 changed files with 270 additions and 36 deletions
|
|
@ -7199,12 +7199,18 @@ class MCPServerManager:
|
|||
return server
|
||||
return None
|
||||
|
||||
def get_mcp_server_answering_to(self, name: str, client_ip: str | None = None) -> MCPServer | None:
|
||||
def get_mcp_server_answering_to(
|
||||
self, name: str, client_ip: str | None = None, *, among: Sequence[MCPServer] | None = None
|
||||
) -> MCPServer | None:
|
||||
"""The one server a ``/mcp/{name}`` segment denotes, shared by the connect preflight, the scoped
|
||||
router, and RFC 9728 discovery so all three name the same server: the exact ``get_mcp_server_by_name``
|
||||
priority first, then the exact ``server_id``, then the name priority case-insensitively, then any prefix
|
||||
form routing accepts. A name that denotes a server hidden from ``client_ip`` resolves to ``None`` at the
|
||||
pass that found it: it never falls through to a looser pass that could name another server."""
|
||||
pass that found it: it never falls through to a looser pass that could name another server. ``among``
|
||||
runs the same passes over those servers alone instead of the registry, which is how the scoped router
|
||||
picks the caller's granted server answering to ``name``."""
|
||||
if among is not None:
|
||||
return self._server_among_answering_to(name, tuple(among), client_ip)
|
||||
exact: Final = self.get_mcp_server_by_name(name)
|
||||
if exact is not None:
|
||||
return exact if self._is_server_accessible_from_ip(exact, client_ip) else None
|
||||
|
|
@ -7230,6 +7236,33 @@ class MCPServerManager:
|
|||
None,
|
||||
)
|
||||
|
||||
def _server_among_answering_to(
|
||||
self, name: str, servers: Sequence[MCPServer], client_ip: str | None
|
||||
) -> MCPServer | None:
|
||||
"""``get_mcp_server_answering_to`` over ``servers`` instead of the registry: the same passes in the
|
||||
same order, with a server hidden from ``client_ip`` resolving to ``None`` at the pass that found it."""
|
||||
requested: Final = name.lower()
|
||||
passes: Final[tuple[Callable[[MCPServer], bool], ...]] = (
|
||||
lambda server: server.alias == name,
|
||||
lambda server: server.server_name == name,
|
||||
lambda server: server.name == name,
|
||||
lambda server: server.server_id == name,
|
||||
lambda server: (server.alias or "").lower() == requested,
|
||||
lambda server: (server.server_name or "").lower() == requested,
|
||||
lambda server: (server.name or "").lower() == requested,
|
||||
)
|
||||
for matches in passes:
|
||||
if (found := next((server for server in servers if matches(server)), None)) is not None:
|
||||
return found if self._is_server_accessible_from_ip(found, client_ip) else None
|
||||
return next(
|
||||
(
|
||||
server
|
||||
for server in servers
|
||||
if self._is_server_accessible_from_ip(server, client_ip) and server_answers_to_name(server, name)
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
def get_filtered_registry(self, client_ip: str | None = None) -> dict[str, MCPServer]:
|
||||
"""
|
||||
Get registry filtered by client IP access control.
|
||||
|
|
|
|||
|
|
@ -501,17 +501,22 @@ def _server_answers_to(server: MCPServer, name: str) -> bool:
|
|||
def _scoped_server(
|
||||
name: str, allowed_mcp_servers: Sequence[MCPServer], client_ip: str | None
|
||||
) -> MCPServer | Literal["denied"] | None:
|
||||
"""The granted server a scoped ``name`` selects: the registry's ``get_mcp_server_answering_to`` pick, made
|
||||
with the same ``client_ip`` the connect preflight and discovery use, when the caller holds it. ``"denied"``
|
||||
when the registry names a server the caller does not hold, or one hidden from ``client_ip``, so the name
|
||||
is neither rerouted to another granted server nor retried as an access group. ``None`` when the registry
|
||||
cannot place the name for any caller, after trying the granted servers answering to it."""
|
||||
"""The server a scoped ``name`` selects for the caller, in this order. ``"denied"`` when the registry's
|
||||
``get_mcp_server_answering_to`` pick, made with the same ``client_ip`` the connect preflight and discovery
|
||||
use, is a server hidden from that IP. Otherwise the granted server answering to ``name``: the registry's
|
||||
own pass order run over ``allowed_mcp_servers`` alone, so a granted server wins over an ungranted alias or
|
||||
case variant the registry would pick. With no granted server answering: ``"denied"`` when the registry
|
||||
places the name on a server the caller does not hold, so it is not retried as an access group; ``None``
|
||||
when the registry cannot place the name for any caller."""
|
||||
registry_pick: Final = global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip)
|
||||
if registry_pick is not None:
|
||||
return next((s for s in allowed_mcp_servers if s.server_id == registry_pick.server_id), "denied")
|
||||
if global_mcp_server_manager.get_mcp_server_answering_to(name) is not None:
|
||||
if registry_pick is None and global_mcp_server_manager.get_mcp_server_answering_to(name) is not None:
|
||||
return "denied"
|
||||
return next((s for s in allowed_mcp_servers if s and _server_answers_to(s, name)), None)
|
||||
granted: Final = global_mcp_server_manager.get_mcp_server_answering_to(
|
||||
name, client_ip=client_ip, among=allowed_mcp_servers
|
||||
)
|
||||
if granted is not None:
|
||||
return granted
|
||||
return "denied" if registry_pick is not None else None
|
||||
|
||||
|
||||
async def raise_denied_scoped_mcp_access(
|
||||
|
|
|
|||
|
|
@ -1603,7 +1603,24 @@ if MCP_AVAILABLE:
|
|||
a server it will be 403'd on immediately after authentication.
|
||||
"""
|
||||
for server_name in mcp_servers or []:
|
||||
server = operations.global_mcp_server_manager.get_mcp_server_answering_to(server_name, client_ip=client_ip)
|
||||
registry_pick = operations.global_mcp_server_manager.get_mcp_server_answering_to(
|
||||
server_name, client_ip=client_ip
|
||||
)
|
||||
obo_without_subject = (
|
||||
registry_pick is not None
|
||||
and registry_pick.auth_type == MCPAuth.oauth2_token_exchange
|
||||
and not oauth2_headers
|
||||
)
|
||||
allowed_single = (
|
||||
await operations._get_allowed_mcp_servers(
|
||||
user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
|
||||
)
|
||||
if registry_pick and not obo_without_subject and mcp_servers is not None and len(mcp_servers) == 1
|
||||
else ()
|
||||
)
|
||||
granted = next(iter(allowed_single), None)
|
||||
server = granted if granted is not None else registry_pick
|
||||
granted_single = granted is not None
|
||||
if server is not None and allowed_server_ids is not None and server.server_id not in allowed_server_ids:
|
||||
# Caller's narrowed scope excludes this server — skip the
|
||||
# preemptive challenge and let downstream authorization
|
||||
|
|
@ -1703,7 +1720,7 @@ if MCP_AVAILABLE:
|
|||
# JSON-RPC error and the WWW-Authenticate header is lost. OBO keeps its connect gate;
|
||||
# guardrail-only gates fire only on a single-server connect the key's grant admits, so a
|
||||
# key without access gets the grant's 403 instead of a sign-in it could not use. The one
|
||||
# admission lookup below serves the challenge, the sign-in preflight and the exchange.
|
||||
# admission lookup above serves the challenge, the sign-in preflight and the exchange.
|
||||
sign_in = caller_sign_in_for(server, user_api_key_auth) if server is not None else None
|
||||
subject_token = (
|
||||
operations.global_mcp_server_manager._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight
|
||||
|
|
@ -1712,19 +1729,6 @@ if MCP_AVAILABLE:
|
|||
if server is not None
|
||||
else None
|
||||
)
|
||||
obo_without_subject = (
|
||||
server is not None and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers
|
||||
)
|
||||
allowed_single = (
|
||||
await operations._get_allowed_mcp_servers(
|
||||
user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
|
||||
)
|
||||
if server and not obo_without_subject and mcp_servers is not None and len(mcp_servers) == 1
|
||||
else ()
|
||||
)
|
||||
granted_single = server is not None and any(
|
||||
allowed.server_id == server.server_id for allowed in allowed_single
|
||||
)
|
||||
if server and sign_in is not None and subject_token is None and (obo_without_subject or granted_single):
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph
|
||||
raise_token_exchange_challenge,
|
||||
|
|
|
|||
|
|
@ -924,7 +924,7 @@ async def test_get_tools_from_mcp_servers():
|
|||
mock_manager.get_mcp_server_by_id = lambda server_id: (
|
||||
mock_server_1 if server_id == "server1_id" else mock_server_2
|
||||
)
|
||||
mock_manager.get_mcp_server_answering_to = MagicMock(return_value=None)
|
||||
mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to
|
||||
mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1])
|
||||
# Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = MagicMock(
|
||||
|
|
|
|||
|
|
@ -32,7 +32,7 @@ import litellm
|
|||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
|
||||
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller, MCPServerManager
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_MCPServerTable,
|
||||
MCPTransport,
|
||||
|
|
@ -1279,7 +1279,7 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
|
|||
mock_manager.get_mcp_server_by_id = lambda server_id: (
|
||||
working_server if server_id == "working_server" else failing_server
|
||||
)
|
||||
mock_manager.get_mcp_server_answering_to = lambda name, client_ip=None: None
|
||||
mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to
|
||||
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
|
||||
server_ids,
|
||||
|
|
@ -6765,7 +6765,7 @@ async def test_list_tools_with_legacy_db_m2m_server_resolves_oauth2_flow():
|
|||
):
|
||||
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["legacy-m2m-id"])
|
||||
mock_manager.get_mcp_server_by_id = MagicMock(return_value=legacy_server)
|
||||
mock_manager.get_mcp_server_answering_to = MagicMock(return_value=None)
|
||||
mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to
|
||||
mock_manager.filter_server_ids_by_ip_with_info = MagicMock(return_value=(["legacy-m2m-id"], 0))
|
||||
mock_manager._get_tools_from_server = AsyncMock(side_effect=capture_extra_headers)
|
||||
|
||||
|
|
@ -8620,7 +8620,9 @@ async def test_scoped_router_selects_the_server_the_connect_preflight_resolves(a
|
|||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
assert only_b == [], "a name the registry gives to an ungranted server must not fall through to another"
|
||||
assert [s.server_id for s in only_b] == ["b-id"], (
|
||||
"the granted server answering to the name wins over the registry's ungranted alias holder"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -8691,10 +8693,202 @@ async def test_scoped_router_hides_a_private_server_from_an_external_ip_like_the
|
|||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
assert external == [], "a name the preflight hides from this IP must not reroute to a case variant"
|
||||
assert internal == []
|
||||
assert [s.server_id for s in internal] == ["u-id"], "with no IP hiding in play the granted case variant wins"
|
||||
assert [s.server_id for s in by_own_name] == ["u-id"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("registry", "scope", "granted", "expected"),
|
||||
[
|
||||
pytest.param(("a-id", "d-id"), "docs", ("d-id",), ["d-id"], id="alias-collision-alias-holder-listed-first"),
|
||||
pytest.param(("d-id", "a-id"), "docs", ("d-id",), ["d-id"], id="alias-collision-exact-name-listed-first"),
|
||||
pytest.param(("g1", "g2"), "GITHUB", ("g2",), ["g2"], id="case-variant-collision"),
|
||||
pytest.param(("p-id", "m-id"), "shared", ("m-id",), [], id="ungranted-only-name-stays-denied"),
|
||||
],
|
||||
)
|
||||
async def test_scoped_router_prefers_the_granted_server_answering_to_the_name(registry, scope, granted, expected):
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
from litellm.proxy._experimental.mcp_server.server import _get_allowed_mcp_servers_from_mcp_server_names
|
||||
|
||||
servers: Final = {
|
||||
"a-id": MCPServer(
|
||||
server_id="a-id", name="a_docs", server_name="a_docs", alias="docs", transport=MCPTransport.http
|
||||
),
|
||||
"d-id": MCPServer(server_id="d-id", name="docs", server_name="docs", transport=MCPTransport.http),
|
||||
"g1": MCPServer(server_id="g1", name="GitHub", server_name="GitHub", transport=MCPTransport.http),
|
||||
"g2": MCPServer(server_id="g2", name="github", server_name="github", transport=MCPTransport.http),
|
||||
"p-id": MCPServer(server_id="p-id", name="p", server_name="p", alias="shared", transport=MCPTransport.http),
|
||||
"m-id": MCPServer(server_id="m-id", name="m", server_name="m", transport=MCPTransport.http),
|
||||
}
|
||||
global_mcp_server_manager.registry.clear()
|
||||
global_mcp_server_manager.registry.update({server_id: servers[server_id] for server_id in registry})
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp."
|
||||
"MCPRequestHandler._get_mcp_servers_from_access_groups",
|
||||
new_callable=AsyncMock,
|
||||
return_value=["m-id"],
|
||||
) as groups:
|
||||
selected: Final = await _get_allowed_mcp_servers_from_mcp_server_names(
|
||||
mcp_servers=[scope], allowed_mcp_servers=[servers[server_id] for server_id in granted]
|
||||
)
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
assert [s.server_id for s in selected] == expected
|
||||
groups.assert_not_awaited()
|
||||
|
||||
|
||||
def test_get_mcp_server_answering_to_among_applies_the_registry_pass_order_and_ip_hiding():
|
||||
manager: Final = MCPServerManager()
|
||||
by_alias: Final = MCPServer(server_id="a-id", name="a", server_name="a", alias="svc", transport=MCPTransport.http)
|
||||
by_server_name: Final = MCPServer(server_id="b-id", name="b", server_name="svc", transport=MCPTransport.http)
|
||||
by_name: Final = MCPServer(server_id="c-id", name="svc", server_name="c", transport=MCPTransport.http)
|
||||
by_id: Final = MCPServer(server_id="svc", name="d", server_name="d", transport=MCPTransport.http)
|
||||
folded: Final = MCPServer(server_id="e-id", name="e", server_name="SVC", transport=MCPTransport.http)
|
||||
hidden: Final = MCPServer(
|
||||
server_id="h-id",
|
||||
name="h",
|
||||
server_name="h",
|
||||
alias="svc",
|
||||
transport=MCPTransport.http,
|
||||
available_on_public_internet=False,
|
||||
)
|
||||
|
||||
assert manager.get_mcp_server_answering_to("svc", among=[by_name, by_server_name, by_alias]) is by_alias
|
||||
assert manager.get_mcp_server_answering_to("svc", among=[by_name, by_server_name]) is by_server_name
|
||||
assert manager.get_mcp_server_answering_to("svc", among=[by_id, by_name]) is by_name
|
||||
assert manager.get_mcp_server_answering_to("svc", among=[folded, by_id]) is by_id
|
||||
assert manager.get_mcp_server_answering_to("svc", among=[folded]) is folded
|
||||
assert manager.get_mcp_server_answering_to("E-ID", among=[folded]) is folded
|
||||
assert manager.get_mcp_server_answering_to("svc", among=()) is None
|
||||
assert manager.get_mcp_server_answering_to("svc", client_ip="203.0.113.7", among=[hidden, folded]) is None
|
||||
assert manager.get_mcp_server_answering_to("svc", client_ip="10.0.0.7", among=[hidden, folded]) is hidden
|
||||
assert manager.get_mcp_server_answering_to("svc", client_ip="203.0.113.7", among=[folded]) is folded
|
||||
assert manager.get_mcp_server_answering_to("svc") is None
|
||||
|
||||
manager.registry = {"e-id": folded, "b-id": by_server_name, "a-id": by_alias}
|
||||
|
||||
assert manager.get_mcp_server_answering_to("svc") is by_alias
|
||||
assert manager.get_mcp_server_answering_to("svc") is manager.get_mcp_server_answering_to(
|
||||
"svc", among=tuple(manager.registry.values())
|
||||
)
|
||||
|
||||
|
||||
class _GrantedServerSignInGuardrail(CustomGuardrail):
|
||||
"""Requires caller sign-in on one server only and records every server it is asked about."""
|
||||
|
||||
def __init__(self, *args, gated_server_id: str, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self._gated_server_id = gated_server_id
|
||||
self.asked_about = [] # 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
|
||||
|
||||
self.asked_about.append(server.server_id)
|
||||
if server.server_id != self._gated_server_id:
|
||||
return None
|
||||
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
|
||||
|
||||
return SignedIn()
|
||||
|
||||
|
||||
class TestConnectPreflightRoutesLikeTheScopedRouter:
|
||||
"""A granted key connecting to ``/mcp/{name}`` is pre-flighted for the server the scoped router routes
|
||||
it to, so an ungranted server holding the name as an alias neither hides the granted server's sign-in
|
||||
challenge nor skips its connect-time exchange."""
|
||||
|
||||
@staticmethod
|
||||
def _register_alias_collision() -> None:
|
||||
alias_holder: Final = MCPServer(
|
||||
server_id="a-id",
|
||||
name="a_docs",
|
||||
server_name="a_docs",
|
||||
alias="docs",
|
||||
url="https://a.test/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
)
|
||||
granted: Final = MCPServer(
|
||||
server_id="d-id",
|
||||
name="docs",
|
||||
server_name="docs",
|
||||
url="https://d.test/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
)
|
||||
mcp_operations.global_mcp_server_manager.registry.update({"a-id": alias_holder, "d-id": granted})
|
||||
|
||||
@staticmethod
|
||||
async def _connect_to_docs() -> None:
|
||||
from litellm.proxy._experimental.mcp_server import server as server_module
|
||||
|
||||
await server_module._raise_preemptive_401_for_unauthenticated_servers(
|
||||
scope={"type": "http", "method": "POST", "path": "/mcp/docs", "headers": []},
|
||||
mcp_servers=["docs"],
|
||||
oauth2_headers=None,
|
||||
mcp_server_auth_headers=None,
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-granted-docs", user_id="u-1"),
|
||||
client_ip=None,
|
||||
raw_headers={"x-litellm-api-key": "sk-granted-docs"},
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_runs_for_the_granted_server_not_the_alias_holder(self):
|
||||
self._register_alias_collision()
|
||||
|
||||
async def refuse_exchange(server, **kwargs):
|
||||
raise HTTPException(
|
||||
status_code=401, detail=f"exchange refused for {server.server_id} as {kwargs['connected_as']}"
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it
|
||||
mcp_operations.global_mcp_server_manager, "get_allowed_mcp_servers", AsyncMock(return_value=["d-id"])
|
||||
),
|
||||
patch.object( # test-quality-ok: a real exchanger would call an IdP; this one reports which server it ran for
|
||||
mcp_operations.global_mcp_server_manager, "preflight_token_exchange", refuse_exchange
|
||||
),
|
||||
pytest.raises(HTTPException) as exc,
|
||||
):
|
||||
await self._connect_to_docs()
|
||||
|
||||
assert exc.value.status_code == 401
|
||||
assert exc.value.detail == "exchange refused for d-id as docs"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sign_in_challenge_names_the_granted_server_not_the_alias_holder(self, monkeypatch):
|
||||
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
|
||||
self._register_alias_collision()
|
||||
guardrail: Final = _GrantedServerSignInGuardrail(guardrail_name="sign-in-stub", gated_server_id="d-id")
|
||||
litellm.logging_callback_manager.add_litellm_callback(guardrail)
|
||||
try:
|
||||
with (
|
||||
patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it
|
||||
mcp_operations.global_mcp_server_manager,
|
||||
"get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=["d-id"]),
|
||||
),
|
||||
pytest.raises(HTTPException) as exc,
|
||||
):
|
||||
await self._connect_to_docs()
|
||||
finally:
|
||||
litellm.logging_callback_manager.remove_callback_from_list_by_object(
|
||||
litellm.callbacks, guardrail, require_self=False
|
||||
)
|
||||
|
||||
assert exc.value.status_code == 401
|
||||
assert ((exc.value.headers or {}).get("WWW-Authenticate") or "").startswith(
|
||||
'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/docs"'
|
||||
)
|
||||
assert guardrail.asked_about == ["d-id"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_allowed_mcp_servers_from_mcp_server_names_mixed_known_and_unknown():
|
||||
"""
|
||||
|
|
@ -9836,7 +10030,7 @@ async def test_aggregate_listing_reports_per_server_outcomes():
|
|||
mock_manager.get_mcp_server_by_id = lambda server_id: (
|
||||
working_server if server_id == "working_server" else broken_server
|
||||
)
|
||||
mock_manager.get_mcp_server_answering_to = lambda name, client_ip=None: None
|
||||
mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
|
||||
|
||||
async def mock_get_tools_from_server(server, **kwargs):
|
||||
|
|
@ -10351,9 +10545,7 @@ class TestOboPreflightScopedToAllowedServers:
|
|||
requested = _make_obo_server("obo_tools")
|
||||
key = UserAPIKeyAuth(api_key="sk-plain-only")
|
||||
|
||||
allowed_lookup, preflight = await self._run(
|
||||
requested, allowed=[_make_obo_server("plain_tools")], user_api_key_auth=key
|
||||
)
|
||||
allowed_lookup, preflight = await self._run(requested, allowed=[], user_api_key_auth=key)
|
||||
|
||||
preflight.assert_not_awaited()
|
||||
allowed_lookup.assert_awaited_once_with(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue