diff --git a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 8a0352e82a1..0736391642b 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -15,8 +15,11 @@ from litellm.proxy._types import LiteLLM_MCPServerTable from litellm.types.mcp import MCPAuth if TYPE_CHECKING: + from collections.abc import Iterator + import httpx from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey + from fastapi import APIRouter from litellm.proxy.auth.handle_jwt import JWTHandler @@ -11903,52 +11906,46 @@ def test_named_resource_discovery_follows_matching_authorization_issuer( assert authorization.json()["issuer"] == resource["authorization_servers"][0] -def test_static_root_path_authorization_discovery_preserves_issuer(monkeypatch, tmp_path): - import subprocess - import sys +@pytest.fixture +def _gateway_root_path_discovery_router(monkeypatch: pytest.MonkeyPatch) -> "Iterator[APIRouter]": + """The ``.well-known`` routes bake ``SERVER_ROOT_PATH`` into their paths when the module is + executed, so it is re-executed under the gateway env and its namespace restored afterwards + (a second reload would hand earlier importers a different ``router`` object).""" + import importlib + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + + snapshot: Final = dict(vars(discoverable_endpoints)) monkeypatch.setenv("SERVER_ROOT_PATH", "/gateway") - monkeypatch.setenv("PROXY_BASE_URL", "http://testserver/gateway") - monkeypatch.setenv("LITELLM_UI_PATH", str(tmp_path / "ui")) - result = subprocess.run( - [ - sys.executable, - "-c", - """ -import json -from fastapi import FastAPI -from fastapi.testclient import TestClient -from litellm.proxy._experimental.mcp_server.discoverable_endpoints import router -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 + monkeypatch.setenv("PROXY_BASE_URL", "https://llm.example.test/gateway") + try: + yield importlib.reload(discoverable_endpoints).router + finally: + added_keys: Final = [key for key in vars(discoverable_endpoints) if key not in snapshot] + for key in added_keys: + delattr(discoverable_endpoints, key) + vars(discoverable_endpoints).update(snapshot) -global_mcp_server_manager.registry['example'] = MCPServer( - server_id='example', name='example', server_name='example', alias='example', - transport=MCPTransport.http, auth_type=MCPAuth.oauth2, - authorization_url='https://idp.example.com/authorize', token_url='https://idp.example.com/token', -) -app = FastAPI(root_path='/gateway') -app.include_router(router) -with TestClient(app) as client: - responses = { - path: client.get('/.well-known/oauth-authorization-server/gateway/' + path) - for path in ('mcp/example', 'example/mcp', 'example', 'mcp') - } - print(json.dumps({path: {'status': response.status_code, 'body': response.json()} - for path, response in responses.items()})) -""", - ], - capture_output=True, - text=True, - check=True, - timeout=60, - ) - responses = json.loads(result.stdout) - for path in ("mcp/example", "example/mcp", "example", "mcp"): - assert responses[path]["status"] == 200, responses[path] - assert responses[path]["body"]["issuer"] == f"http://testserver/gateway/{path}" - assert responses["example/mcp"]["body"]["token_endpoint"] == "http://testserver/gateway/example/token" + +def test_static_root_path_authorization_discovery_preserves_issuer( + _gateway_root_path_discovery_router: "APIRouter", _isolated_mcp_registry: "dict[str, MCPServer]" +) -> None: + from fastapi import FastAPI + from fastapi.testclient import TestClient + + server: Final = _create_oauth2_server(server_id="example", name="example", server_name="example", alias="example") + _isolated_mcp_registry[server.server_id] = server # rebind-ok: the fixture hands the test an empty registry to fill + app: Final = FastAPI(root_path="/gateway") + app.include_router(_gateway_root_path_discovery_router) + with TestClient(app) as client: + responses: Final = { + path: client.get(f"/.well-known/oauth-authorization-server/gateway/{path}") + for path in ("mcp/example", "example/mcp", "example", "mcp") + } + for path, response in responses.items(): + assert response.status_code == 200, response.text + assert response.json()["issuer"] == f"https://llm.example.test/gateway/{path}" + assert responses["example/mcp"].json()["token_endpoint"] == "https://llm.example.test/gateway/example/token" @pytest.fixture