mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge 83ae2a9bbe into eae8ed7f3c
This commit is contained in:
commit
e896ed68f6
6 changed files with 471 additions and 43 deletions
|
|
@ -3,7 +3,8 @@ import html as _html
|
|||
import json
|
||||
import secrets
|
||||
import time
|
||||
from collections.abc import Callable, Mapping
|
||||
from collections.abc import AsyncIterator, Callable, Mapping
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
|
||||
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
|
||||
|
|
@ -107,6 +108,14 @@ _OAUTH_METADATA_CACHE_MAX_SIZE: Final = 128
|
|||
# Per-(server_id, resource_url) async locks so concurrent discovery requests
|
||||
# coalesce onto a single upstream fetch instead of issuing N parallel calls.
|
||||
_OAUTH_METADATA_FETCH_LOCKS: Final[dict[tuple[str, str], asyncio.Lock]] = {}
|
||||
# Callers inside ``_oauth_metadata_fetch_slot`` per cache key, lock waiters included. ``Lock.locked()``
|
||||
# reads False between one holder's release and the next waiter's wake-up, so it cannot tell an
|
||||
# idle lock from one being handed off.
|
||||
_OAUTH_METADATA_FETCHERS: Final[dict[tuple[str, str], int]] = {}
|
||||
# Per-server_id generation, bumped on invalidation so a fetch that started before the server
|
||||
# definition changed cannot repopulate the cache with the stale reply. Only servers with a fetch
|
||||
# in flight carry an entry; the rest are pruned with the cache.
|
||||
_OAUTH_METADATA_GENERATIONS: Final[dict[str, int]] = {}
|
||||
|
||||
router: Final = APIRouter(
|
||||
tags=["mcp"],
|
||||
|
|
@ -130,13 +139,52 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None:
|
|||
for cache_key in cache_keys_by_expiry[:overflow]:
|
||||
_OAUTH_METADATA_CACHE.pop(cache_key, None)
|
||||
|
||||
# Drop locks whose cache entry has been evicted and that aren't currently
|
||||
# held; held locks stay so in-flight callers continue to coalesce.
|
||||
# Drop locks whose cache entry has been evicted and that nobody holds or
|
||||
# waits on; the rest stay so in-flight callers continue to coalesce.
|
||||
for cache_key in list(_OAUTH_METADATA_FETCH_LOCKS):
|
||||
if cache_key in _OAUTH_METADATA_CACHE:
|
||||
if cache_key in _OAUTH_METADATA_CACHE or not _oauth_metadata_lock_idle(cache_key):
|
||||
continue
|
||||
lock = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key)
|
||||
if lock is None or lock.locked():
|
||||
_OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None)
|
||||
|
||||
for server_id in [sid for sid in _OAUTH_METADATA_GENERATIONS if not _oauth_metadata_fetch_in_flight(sid)]:
|
||||
_OAUTH_METADATA_GENERATIONS.pop(server_id, None)
|
||||
|
||||
|
||||
def _oauth_metadata_fetch_in_flight(server_id: str) -> bool:
|
||||
return any(cache_key[0] == server_id for cache_key in _OAUTH_METADATA_FETCHERS)
|
||||
|
||||
|
||||
def _oauth_metadata_lock_idle(cache_key: tuple[str, str]) -> bool:
|
||||
if cache_key in _OAUTH_METADATA_FETCHERS:
|
||||
return False
|
||||
lock: Final = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key)
|
||||
return lock is None or not lock.locked()
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _oauth_metadata_fetch_slot(cache_key: tuple[str, str]) -> AsyncIterator[None]:
|
||||
_OAUTH_METADATA_FETCHERS[cache_key] = _OAUTH_METADATA_FETCHERS.get(cache_key, 0) + 1
|
||||
try:
|
||||
async with _OAUTH_METADATA_FETCH_LOCKS.setdefault(cache_key, asyncio.Lock()):
|
||||
yield
|
||||
finally:
|
||||
remaining: Final = _OAUTH_METADATA_FETCHERS.get(cache_key, 0) - 1
|
||||
if remaining > 0:
|
||||
_OAUTH_METADATA_FETCHERS[cache_key] = remaining
|
||||
else:
|
||||
_OAUTH_METADATA_FETCHERS.pop(cache_key, None)
|
||||
|
||||
|
||||
def invalidate_oauth_metadata_cache(server_id: str) -> None:
|
||||
"""Drop cached upstream IdP metadata for a server whose definition changed."""
|
||||
if _oauth_metadata_fetch_in_flight(server_id):
|
||||
_OAUTH_METADATA_GENERATIONS[server_id] = _OAUTH_METADATA_GENERATIONS.get(server_id, 0) + 1
|
||||
else:
|
||||
_OAUTH_METADATA_GENERATIONS.pop(server_id, None)
|
||||
for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]:
|
||||
del _OAUTH_METADATA_CACHE[cache_key]
|
||||
for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]:
|
||||
if not _oauth_metadata_lock_idle(cache_key):
|
||||
continue
|
||||
_OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None)
|
||||
|
||||
|
|
@ -2360,12 +2408,19 @@ async def fetch_upstream_oauth_protected_resource(
|
|||
if cached is not None and cached[0] > now:
|
||||
return cached[1]
|
||||
|
||||
lock: Final = _OAUTH_METADATA_FETCH_LOCKS.setdefault(cache_key, asyncio.Lock())
|
||||
async with lock:
|
||||
async with _oauth_metadata_fetch_slot(cache_key):
|
||||
now = time.time()
|
||||
cached = _OAUTH_METADATA_CACHE.get(cache_key)
|
||||
if cached is not None and cached[0] > now:
|
||||
return cached[1]
|
||||
generation: Final = _OAUTH_METADATA_GENERATIONS.get(mcp_server.server_id, 0)
|
||||
|
||||
def store(payload: dict | None, ttl_seconds: int) -> None:
|
||||
if _OAUTH_METADATA_GENERATIONS.get(mcp_server.server_id, 0) != generation:
|
||||
return
|
||||
stored_at: Final = time.time()
|
||||
_OAUTH_METADATA_CACHE[cache_key] = (stored_at + ttl_seconds, payload)
|
||||
_prune_oauth_metadata_cache(stored_at)
|
||||
|
||||
host_base: Final = f"{upstream.scheme}://{upstream.netloc}"
|
||||
candidates: Final = [f"{host_base}/.well-known/oauth-protected-resource"]
|
||||
|
|
@ -2407,12 +2462,7 @@ async def fetch_upstream_oauth_protected_resource(
|
|||
)
|
||||
continue
|
||||
if isinstance(payload, dict):
|
||||
now = time.time()
|
||||
_OAUTH_METADATA_CACHE[cache_key] = (
|
||||
now + _OAUTH_METADATA_CACHE_TTL_SECONDS,
|
||||
payload,
|
||||
)
|
||||
_prune_oauth_metadata_cache(now)
|
||||
store(payload, _OAUTH_METADATA_CACHE_TTL_SECONDS)
|
||||
return payload
|
||||
|
||||
if len(network_errors) == len(candidates):
|
||||
|
|
@ -2421,12 +2471,7 @@ async def fetch_upstream_oauth_protected_resource(
|
|||
# Negative-result caching: when no candidate yielded a usable payload,
|
||||
# remember that for a shorter TTL so we don't re-fetch on every
|
||||
# subsequent discovery request (and so the per-key lock can be pruned).
|
||||
now = time.time()
|
||||
_OAUTH_METADATA_CACHE[cache_key] = (
|
||||
now + _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS,
|
||||
None,
|
||||
)
|
||||
_prune_oauth_metadata_cache(now)
|
||||
store(None, _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS)
|
||||
return None
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -2664,7 +2664,7 @@ class MCPServerManager:
|
|||
self._assign_unique_short_prefix(new_server)
|
||||
_warn_legacy_delegate_auth_if_applicable(new_server, source="config")
|
||||
_warn_config_id_jag_server_outruns_sso(new_server)
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self._invalidate_server_definition_caches(server_id)
|
||||
self.config_mcp_servers[server_id] = new_server
|
||||
self._set_oauth_discovery_deferred(
|
||||
server_id,
|
||||
|
|
@ -2868,7 +2868,7 @@ class MCPServerManager:
|
|||
global_mcp_tool_registry,
|
||||
)
|
||||
|
||||
self._invalidate_discovery_lists(server.server_id)
|
||||
self._invalidate_server_definition_caches(server.server_id)
|
||||
prefix_root: Final = normalize_server_name(get_server_prefix(server))
|
||||
if server.spec_path and prefix_root:
|
||||
openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR
|
||||
|
|
@ -3275,7 +3275,7 @@ class MCPServerManager:
|
|||
# env_vars_are_encrypted=False.
|
||||
new_server: Final = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False)
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
self._invalidate_discovery_lists(mcp_server.server_id)
|
||||
self._invalidate_server_definition_caches(mcp_server.server_id)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
|
|
@ -3312,7 +3312,7 @@ class MCPServerManager:
|
|||
previous_server=self.registry[mcp_server.server_id],
|
||||
)
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
self._invalidate_discovery_lists(mcp_server.server_id)
|
||||
self._invalidate_server_definition_caches(mcp_server.server_id)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
|
|
@ -4492,24 +4492,16 @@ class MCPServerManager:
|
|||
if server.spec_path:
|
||||
# OpenAPI tools were stored in the registry under the prefix
|
||||
# active at registration time — fetch by that same prefix.
|
||||
_tools: Final = global_mcp_tool_registry.list_tools(tool_prefix=get_server_prefix(server))
|
||||
registry_prefix: Final = normalize_server_name(get_server_prefix(server)) + MCP_TOOL_PREFIX_SEPARATOR
|
||||
_tools: Final = global_mcp_tool_registry.list_tools(tool_prefix=registry_prefix)
|
||||
tools = global_mcp_tool_registry.convert_tools_to_mcp_sdk_tool_type(_tools)
|
||||
# OpenAPI tools are stored in the registry with their prefix already
|
||||
# applied (e.g. "test_petstore-getinventory"). Do NOT pass them
|
||||
# through _create_prefixed_tools — that would add the prefix a second
|
||||
# time producing "test_petstore-test_petstore-getinventory".
|
||||
if not add_prefix:
|
||||
prefix: Final = get_server_prefix(server)
|
||||
sep: Final = MCP_TOOL_PREFIX_SEPARATOR
|
||||
tools = [
|
||||
(
|
||||
t.model_copy(update={"name": t.name[len(prefix) + len(sep) :]})
|
||||
if t.name.startswith(f"{prefix}{sep}")
|
||||
else t
|
||||
)
|
||||
for t in tools
|
||||
]
|
||||
return tools
|
||||
if add_prefix:
|
||||
return tools
|
||||
return [t.model_copy(update=MappingProxyType({"name": t.name[len(registry_prefix) :]})) for t in tools]
|
||||
else:
|
||||
tools = await self._fetch_tools_with_timeout(client, server.name)
|
||||
self._remember_upstream_initialize_instructions(server, client)
|
||||
|
|
@ -4558,6 +4550,14 @@ class MCPServerManager:
|
|||
self._resource_discovery_cache.invalidate(server_id)
|
||||
self._template_discovery_cache.invalidate(server_id)
|
||||
|
||||
def _invalidate_server_definition_caches(self, server_id: str) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( # noqa: PLC0415 # lazy: discoverable_endpoints lazily imports this module's manager singleton
|
||||
invalidate_oauth_metadata_cache,
|
||||
)
|
||||
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
invalidate_oauth_metadata_cache(server_id)
|
||||
|
||||
def _discovery_key(
|
||||
self,
|
||||
server: MCPServer,
|
||||
|
|
@ -6704,7 +6704,7 @@ class MCPServerManager:
|
|||
|
||||
for server_id in previous_registry.keys() | registered_registry.keys():
|
||||
if previous_registry.get(server_id) != registered_registry.get(server_id):
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self._invalidate_server_definition_caches(server_id)
|
||||
self.registry = registered_registry
|
||||
_warn_on_shared_identifier_prefixes(registered_registry.values())
|
||||
# A discovery task may have published into ``previous_registry`` while
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import itertools
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
|
@ -7,10 +8,13 @@ import yaml
|
|||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.mcp import (
|
||||
McpCaller,
|
||||
McpPeer,
|
||||
call_tool,
|
||||
delete_mcp,
|
||||
forget_mcp,
|
||||
listed_tools,
|
||||
mcp_peer,
|
||||
openapi_peer,
|
||||
register_mcp,
|
||||
tool_calls,
|
||||
tool_names,
|
||||
|
|
@ -189,6 +193,54 @@ def test_duplicate_alias_is_rejected_so_tool_prefixes_cannot_collide(gateway: Ga
|
|||
scenario.cleanups.callback(forget_mcp, gateway, winner)
|
||||
|
||||
|
||||
def _openapi_server_lists_and_calls_only_its_own_tools(
|
||||
gateway: Gateway, key: str, peer: McpPeer, identity: str
|
||||
) -> None:
|
||||
listed: Final = set(listed_tools(gateway, key, identity))
|
||||
assert listed == {"getpet", "createpet"}, (identity, listed)
|
||||
peer.drain()
|
||||
called: Final = call_tool(gateway, key, identity, "getpet", {"petId": "7"})
|
||||
assert called.status_code == 200, called.text
|
||||
assert [(item["method"], item["path"]) for item in peer.drain()] == [("GET", "/pets/7")], identity
|
||||
|
||||
|
||||
def test_openapi_listing_is_scoped_to_the_exact_alias_when_aliases_overlap(gateway: Gateway) -> None:
|
||||
with openapi_peer() as short, openapi_peer() as long, gateway.scenario() as scenario:
|
||||
stem: Final = "pet" + uuid.uuid4().hex[:8]
|
||||
servers: Final = tuple(
|
||||
(peer, alias, register_mcp(scenario, peer, alias))
|
||||
for peer, alias in ((short, stem), (long, stem + "store"))
|
||||
)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity for _, _, identity in servers]})
|
||||
for peer, _, identity in servers:
|
||||
_openapi_server_lists_and_calls_only_its_own_tools(gateway, key, peer, identity)
|
||||
aggregate: Final = McpCaller(gateway, key, "mcp").list_tools()
|
||||
assert aggregate.ok, aggregate.raw
|
||||
assert sorted(aggregate.tools) == sorted(
|
||||
f"{prefix}-{tool}" for prefix, tool in itertools.product((stem, stem + "store"), ("getpet", "createpet"))
|
||||
), aggregate.tools
|
||||
assert all(peer.drain() == () for peer, _, _ in servers), "listing must not reach any OpenAPI upstream"
|
||||
|
||||
|
||||
def test_config_declared_openapi_server_with_a_space_in_its_name_lists_its_tools(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
with openapi_peer() as peer:
|
||||
config: Final = yaml.safe_load((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text())
|
||||
name: Final = "pet store " + uuid.uuid4().hex[:8]
|
||||
config["mcp_servers"] = {name: peer.registration()}
|
||||
path: Final = tmp_path / "openapi-space.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario:
|
||||
identity: Final = next(i for i, s in _servers(candidate).items() if s["server_name"] == name)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
_openapi_server_lists_and_calls_only_its_own_tools(candidate, key, peer, identity)
|
||||
aggregate: Final = McpCaller(candidate, key, "mcp").list_tools()
|
||||
assert aggregate.ok, aggregate.raw
|
||||
prefix: Final = name.replace(" ", "_")
|
||||
assert sorted(aggregate.tools) == [f"{prefix}-createpet", f"{prefix}-getpet"], aggregate.tools
|
||||
|
||||
|
||||
def test_invalid_registrations_are_rejected(gateway: Gateway) -> None:
|
||||
with mcp_peer() as peer, gateway.scenario() as scenario:
|
||||
alias: Final = "mgmt" + uuid.uuid4().hex[:8]
|
||||
|
|
|
|||
|
|
@ -1,17 +1,23 @@
|
|||
import json
|
||||
import queue
|
||||
import threading
|
||||
import uuid
|
||||
from urllib.parse import parse_qs, urlsplit
|
||||
from typing import Final, Literal
|
||||
from collections.abc import Callable
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Final, Literal
|
||||
from urllib.parse import parse_qs, urlsplit
|
||||
|
||||
import pytest
|
||||
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.client import Gateway, Scenario, eventually
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.mcp import McpPeer, call_tool, mcp_peer, register_mcp, tool_names
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
_Upstream = Callable[[Request], Reply]
|
||||
|
||||
|
||||
@pytest.mark.covers("other.mcp.oauth.discovery_cannot_erase_configured_authorization_endpoint")
|
||||
|
|
@ -104,6 +110,129 @@ def test_partial_discovery_and_unrelated_edit_keep_actual_authorization_destinat
|
|||
assert updated.status_code == 202, updated.text
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Hold:
|
||||
armed: threading.Event = field(default_factory=threading.Event)
|
||||
released: threading.Event = field(default_factory=threading.Event)
|
||||
|
||||
|
||||
def _idp_upstream(origin: Callable[[], str], moved: threading.Event, hold: _Hold | None = None) -> _Upstream:
|
||||
def issuer() -> str:
|
||||
return origin() + ("/idp-after" if moved.is_set() else "/idp-before")
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
if "oauth-authorization-server" in request.target or "openid-configuration" in request.target:
|
||||
current: Final = issuer()
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"issuer": current,
|
||||
"authorization_endpoint": current + "/authorize",
|
||||
"token_endpoint": current + "/token",
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
if request.target.startswith("/.well-known/oauth-protected-resource"):
|
||||
body: Final = json.dumps({"resource": origin() + "/mcp", "authorization_servers": [issuer()]}).encode()
|
||||
if hold is not None and hold.armed.is_set():
|
||||
assert hold.released.wait(timeout=15), "the held upstream metadata reply was never released"
|
||||
return Reply(body=body)
|
||||
return Reply(status=404, body=b'{"error":"unexpected"}')
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
def _register_pass_through(scenario: Scenario, wire: Wire, alias: str) -> str:
|
||||
return register_mcp(scenario, McpPeer(wire.url + "/mcp", queue.Queue()), alias, auth_type="true_passthrough")
|
||||
|
||||
|
||||
def _wire_requests(wire: Wire, seen: list[Request]) -> Callable[[], tuple[Request, ...]]:
|
||||
def observed() -> tuple[Request, ...]:
|
||||
seen.extend(wire.drain())
|
||||
return tuple(seen)
|
||||
|
||||
return observed
|
||||
|
||||
|
||||
def _registration_discovery_settled(requests: tuple[Request, ...]) -> bool:
|
||||
return any(
|
||||
"oauth-authorization-server" in item.target or "openid-configuration" in item.target for item in requests
|
||||
)
|
||||
|
||||
|
||||
def _advertised_authorization_servers(gateway: Gateway, alias: str) -> tuple[str, ...]:
|
||||
response: Final = gateway.client.get(f"/.well-known/oauth-protected-resource/{alias}/mcp")
|
||||
assert response.status_code == 200, response.text
|
||||
return tuple(TypeAdapter(list[str]).validate_python(response.json()["authorization_servers"]))
|
||||
|
||||
|
||||
def _eventually_advertises(gateway: Gateway, alias: str, issuer: str) -> None:
|
||||
eventually(
|
||||
lambda: gateway.client.get(f"/.well-known/oauth-protected-resource/{alias}/mcp"),
|
||||
lambda response: response.status_code == 200 and response.json()["authorization_servers"] == [issuer],
|
||||
seconds=40,
|
||||
)
|
||||
|
||||
|
||||
def test_saving_a_pass_through_server_refetches_its_upstream_oauth_metadata(gateway: Gateway) -> None:
|
||||
moved: Final = threading.Event()
|
||||
with wire_server(_idp_upstream(lambda: wire.url, moved)) as wire, gateway.scenario() as scenario:
|
||||
alias: Final = "pt" + uuid.uuid4().hex[:8]
|
||||
identity: Final = _register_pass_through(scenario, wire, alias)
|
||||
assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",)
|
||||
assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",)
|
||||
moved.set()
|
||||
wire.drain()
|
||||
saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"})
|
||||
assert saved.status_code == 202, saved.text
|
||||
assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",)
|
||||
assert any(request.target.startswith("/.well-known/oauth-protected-resource") for request in wire.drain()), (
|
||||
"the save must send protected-resource discovery back to the upstream"
|
||||
)
|
||||
|
||||
|
||||
def test_peer_worker_stops_advertising_the_old_idp_after_a_save_on_another_worker(
|
||||
gateway: Gateway, peer: Gateway
|
||||
) -> None:
|
||||
moved: Final = threading.Event()
|
||||
with wire_server(_idp_upstream(lambda: wire.url, moved)) as wire, gateway.scenario() as scenario:
|
||||
alias: Final = "pt" + uuid.uuid4().hex[:8]
|
||||
identity: Final = _register_pass_through(scenario, wire, alias)
|
||||
assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",)
|
||||
_eventually_advertises(peer, alias, wire.url + "/idp-before")
|
||||
moved.set()
|
||||
saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"})
|
||||
assert saved.status_code == 202, saved.text
|
||||
assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",)
|
||||
_eventually_advertises(peer, alias, wire.url + "/idp-after")
|
||||
|
||||
|
||||
def test_metadata_fetched_before_a_save_cannot_repopulate_the_cache_after_it(gateway: Gateway) -> None:
|
||||
moved: Final = threading.Event()
|
||||
hold: Final = _Hold()
|
||||
with (
|
||||
wire_server(_idp_upstream(lambda: wire.url, moved, hold)) as wire,
|
||||
gateway.scenario() as scenario,
|
||||
ThreadPoolExecutor(max_workers=1) as pool,
|
||||
):
|
||||
alias: Final = "pt" + uuid.uuid4().hex[:8]
|
||||
identity: Final = _register_pass_through(scenario, wire, alias)
|
||||
seen: Final[list[Request]] = []
|
||||
observed: Final = _wire_requests(wire, seen)
|
||||
eventually(observed, _registration_discovery_settled, seconds=10)
|
||||
settled: Final = len(seen)
|
||||
hold.armed.set()
|
||||
stale: Final = pool.submit(_advertised_authorization_servers, gateway, alias)
|
||||
eventually(observed, lambda requests: len(requests) > settled, seconds=10)
|
||||
assert seen[settled].target.startswith("/.well-known/oauth-protected-resource"), seen[settled:]
|
||||
moved.set()
|
||||
saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"})
|
||||
assert saved.status_code == 202, saved.text
|
||||
hold.released.set()
|
||||
assert stale.result(timeout=30) == (wire.url + "/idp-before",)
|
||||
assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",)
|
||||
|
||||
|
||||
@pytest.mark.covers("other.mcp.oauth.same_url_credentials_are_isolated_by_user_and_server")
|
||||
@pytest.mark.parametrize("transition", ("revoke", "expire"))
|
||||
def test_same_url_oauth_credentials_and_revocation_are_isolated_by_user_and_server(
|
||||
|
|
|
|||
|
|
@ -12611,3 +12611,151 @@ async def test_identity_bound_authorize_unrelated_bearer_uses_browser_session(
|
|||
proxy_server.prisma_client.db.litellm_mcpusercredentials.upsert.assert_not_called()
|
||||
proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called()
|
||||
proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_server_drops_cached_upstream_oauth_metadata():
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.proxy._types import LiteLLM_MCPServerTable, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="oauth-cache-server",
|
||||
name="oauth_cache_server",
|
||||
url="http://old-upstream/mcp",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
manager.registry[server.server_id] = server
|
||||
stale_key: Final = (server.server_id, server.url)
|
||||
other_key: Final = ("other-server", "http://other/mcp")
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE[stale_key] = (time.time() + 300, {"iss": "old-idp"})
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE[other_key] = (time.time() + 300, {"iss": "other"})
|
||||
try:
|
||||
await manager.update_server(
|
||||
LiteLLM_MCPServerTable(
|
||||
server_id=server.server_id,
|
||||
server_name=server.name,
|
||||
url="http://new-upstream/mcp",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
)
|
||||
assert stale_key not in discoverable_endpoints._OAUTH_METADATA_CACHE
|
||||
assert other_key in discoverable_endpoints._OAUTH_METADATA_CACHE
|
||||
finally:
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE.pop(stale_key, None)
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE.pop(other_key, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_metadata_fetched_before_invalidation_does_not_repopulate_the_cache():
|
||||
import asyncio
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
fetch_upstream_oauth_protected_resource,
|
||||
invalidate_oauth_metadata_cache,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
server = MCPServer(
|
||||
server_id="stale-write-server", name="stale_write", url="http://upstream/mcp", transport=MCPTransport.http
|
||||
)
|
||||
cache_key: Final = (server.server_id, server.url)
|
||||
started: Final = asyncio.Event()
|
||||
release: Final = asyncio.Event()
|
||||
|
||||
async def slow_get(url: str, headers: dict[str, str]) -> MagicMock:
|
||||
started.set()
|
||||
await release.wait()
|
||||
return MagicMock(status_code=200, json=MagicMock(return_value={"authorization_servers": ["old-idp"]}))
|
||||
|
||||
client = MagicMock()
|
||||
client.get = slow_get
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None)
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
|
||||
return_value=client,
|
||||
):
|
||||
in_flight: Final = asyncio.create_task(fetch_upstream_oauth_protected_resource(server))
|
||||
await started.wait()
|
||||
invalidate_oauth_metadata_cache(server.server_id)
|
||||
release.set()
|
||||
assert await in_flight == {"authorization_servers": ["old-idp"]}
|
||||
assert cache_key not in discoverable_endpoints._OAUTH_METADATA_CACHE
|
||||
discoverable_endpoints._prune_oauth_metadata_cache()
|
||||
assert server.server_id not in discoverable_endpoints._OAUTH_METADATA_GENERATIONS
|
||||
finally:
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None)
|
||||
discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server.server_id, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_waiting_on_a_lock_handoff_stays_tracked_through_invalidation():
|
||||
import asyncio
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
fetch_upstream_oauth_protected_resource,
|
||||
invalidate_oauth_metadata_cache,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
server = MCPServer(
|
||||
server_id="handoff-server", name="handoff", url="http://upstream/mcp", transport=MCPTransport.http
|
||||
)
|
||||
cache_key: Final = (server.server_id, server.url)
|
||||
started: Final = asyncio.Event()
|
||||
release: Final = asyncio.Event()
|
||||
|
||||
async def slow_get(url: str, headers: dict[str, str]) -> MagicMock:
|
||||
started.set()
|
||||
await release.wait()
|
||||
return MagicMock(status_code=200, json=MagicMock(return_value={"authorization_servers": ["pre-save-idp"]}))
|
||||
|
||||
client = MagicMock()
|
||||
client.get = slow_get
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None)
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
|
||||
return_value=client,
|
||||
):
|
||||
async with discoverable_endpoints._oauth_metadata_fetch_slot(cache_key):
|
||||
shared_lock: Final = discoverable_endpoints._OAUTH_METADATA_FETCH_LOCKS[cache_key]
|
||||
waiting: Final = asyncio.create_task(fetch_upstream_oauth_protected_resource(server))
|
||||
for _ in range(3):
|
||||
await asyncio.sleep(0)
|
||||
assert not started.is_set() and not waiting.done()
|
||||
invalidate_oauth_metadata_cache(server.server_id)
|
||||
assert discoverable_endpoints._OAUTH_METADATA_FETCH_LOCKS.get(cache_key) is shared_lock
|
||||
assert discoverable_endpoints._oauth_metadata_fetch_in_flight(server.server_id)
|
||||
await started.wait()
|
||||
invalidate_oauth_metadata_cache(server.server_id)
|
||||
release.set()
|
||||
assert await waiting == {"authorization_servers": ["pre-save-idp"]}
|
||||
assert cache_key not in discoverable_endpoints._OAUTH_METADATA_CACHE
|
||||
assert not discoverable_endpoints._oauth_metadata_fetch_in_flight(server.server_id)
|
||||
finally:
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None)
|
||||
discoverable_endpoints._OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None)
|
||||
discoverable_endpoints._OAUTH_METADATA_FETCHERS.pop(cache_key, None)
|
||||
discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server.server_id, None)
|
||||
|
||||
|
||||
def test_invalidating_an_idle_server_leaves_no_generation_behind():
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import invalidate_oauth_metadata_cache
|
||||
|
||||
server_ids: Final = tuple(f"churned-server-{i}" for i in range(50))
|
||||
try:
|
||||
for server_id in server_ids:
|
||||
invalidate_oauth_metadata_cache(server_id)
|
||||
assert not set(server_ids) & set(discoverable_endpoints._OAUTH_METADATA_GENERATIONS)
|
||||
finally:
|
||||
for server_id in server_ids:
|
||||
discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server_id, None)
|
||||
|
|
|
|||
|
|
@ -14062,6 +14062,60 @@ def test_discovery_cache_keys_isolate_user_dependent_auth(auth_type: MCPAuth) ->
|
|||
assert "second" not in str(second)
|
||||
|
||||
|
||||
def _register_local_tool(name: str, description: str) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
async def _handler(**kwargs):
|
||||
return None
|
||||
|
||||
global_mcp_tool_registry.register_tool(
|
||||
name=name, description=description, input_schema={"type": "object"}, handler=_handler
|
||||
)
|
||||
|
||||
|
||||
def _openapi_server(name: str) -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id=f"{name}-id", name=name, alias=name, transport=MCPTransport.http, url=None, spec_path="/spec.yaml"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openapi_listing_ignores_overlapping_server_prefix() -> None:
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
manager: Final = MCPServerManager()
|
||||
manager._create_mcp_client = AsyncMock(return_value=AsyncMock())
|
||||
for prefix in ("pet-", "petstore-"):
|
||||
global_mcp_tool_registry.unregister_tools_with_prefix(prefix)
|
||||
_register_local_tool("pet-list", "Local pet tool")
|
||||
_register_local_tool("petstore-list", "Foreign petstore tool")
|
||||
try:
|
||||
prefixed: Final = await manager._get_tools_from_server(server=_openapi_server("pet"), add_prefix=True)
|
||||
bare: Final = await manager._get_tools_from_server(server=_openapi_server("pet"), add_prefix=False)
|
||||
finally:
|
||||
for prefix in ("pet-", "petstore-"):
|
||||
global_mcp_tool_registry.unregister_tools_with_prefix(prefix)
|
||||
|
||||
assert [t.name for t in prefixed] == ["pet-list"]
|
||||
assert [t.name for t in bare] == ["list"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openapi_listing_finds_tools_registered_under_the_normalized_prefix() -> None:
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
manager: Final = MCPServerManager()
|
||||
manager._create_mcp_client = AsyncMock(return_value=AsyncMock())
|
||||
global_mcp_tool_registry.unregister_tools_with_prefix("pet_store-")
|
||||
_register_local_tool("pet_store-list", "Pet store tool")
|
||||
try:
|
||||
listed: Final = await manager._get_tools_from_server(server=_openapi_server("pet store"), add_prefix=False)
|
||||
finally:
|
||||
global_mcp_tool_registry.unregister_tools_with_prefix("pet_store-")
|
||||
|
||||
assert [t.name for t in listed] == ["list"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_cache_retries_cancelled_fetches() -> None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue