diff --git a/tests/mcp_tests/mcp_server.py b/tests/mcp_tests/mcp_server.py index bfbb29931c1..f6a18adbd26 100644 --- a/tests/mcp_tests/mcp_server.py +++ b/tests/mcp_tests/mcp_server.py @@ -27,6 +27,9 @@ _auth_mode = "none" _auth_secret: Optional[str] = None _client_id = DEFAULT_CLIENT_ID _client_secret = DEFAULT_CLIENT_SECRET +_expected_subject_token: Optional[str] = None +_expected_audience: Optional[str] = None +_expected_scope: Optional[str] = None def _request_is_authorized(headers) -> bool: @@ -42,7 +45,10 @@ def _request_is_authorized(headers) -> bool: return headers.get(DEFAULT_CUSTOM_HEADER) == _auth_secret if _auth_mode == "oauth2": auth = headers.get("authorization") or "" - return auth.startswith("Bearer ") and auth[len("Bearer ") :] in VALID_OAUTH_BEARER_TOKENS + return ( + auth.startswith("Bearer ") + and auth[len("Bearer ") :] in VALID_OAUTH_BEARER_TOKENS + ) return False @@ -59,7 +65,9 @@ class _AuthMiddleware: await self.app(scope, receive, send) return if not _request_is_authorized(request.headers): - response = JSONResponse({"error": "unauthorized", "auth_mode": _auth_mode}, status_code=401) + response = JSONResponse( + {"error": "unauthorized", "auth_mode": _auth_mode}, status_code=401 + ) await response(scope, receive, send) return await self.app(scope, receive, send) @@ -82,7 +90,10 @@ async def oauth_token(request: Request) -> JSONResponse: form = await request.form() grant_type = form.get("grant_type") - if form.get("client_id") != _client_id or form.get("client_secret") != _client_secret: + if ( + form.get("client_id") != _client_id + or form.get("client_secret") != _client_secret + ): return JSONResponse({"error": "invalid_client"}, status_code=401) if grant_type == "client_credentials": @@ -97,6 +108,15 @@ async def oauth_token(request: Request) -> JSONResponse: if grant_type == TOKEN_EXCHANGE_GRANT_TYPE: if not form.get("subject_token"): return JSONResponse({"error": "invalid_request"}, status_code=400) + if ( + _expected_subject_token + and form.get("subject_token") != _expected_subject_token + ): + return JSONResponse({"error": "invalid_subject_token"}, status_code=400) + if _expected_audience and form.get("audience") != _expected_audience: + return JSONResponse({"error": "invalid_audience"}, status_code=400) + if _expected_scope and form.get("scope") != _expected_scope: + return JSONResponse({"error": "invalid_scope"}, status_code=400) return JSONResponse( { "access_token": DEFAULT_OBO_ACCESS_TOKEN, @@ -130,11 +150,15 @@ def _parse_args() -> argparse.Namespace: parser.add_argument("--auth-secret", default=None) parser.add_argument("--client-id", default=DEFAULT_CLIENT_ID) parser.add_argument("--client-secret", default=DEFAULT_CLIENT_SECRET) + parser.add_argument("--expected-subject-token", default=None) + parser.add_argument("--expected-audience", default=None) + parser.add_argument("--expected-scope", default=None) return parser.parse_args() def main() -> None: global _auth_mode, _auth_secret, _client_id, _client_secret + global _expected_audience, _expected_scope, _expected_subject_token args = _parse_args() transport = (args.transport or "stdio").lower() @@ -142,6 +166,9 @@ def main() -> None: _auth_secret = args.auth_secret _client_id = args.client_id _client_secret = args.client_secret + _expected_subject_token = args.expected_subject_token + _expected_audience = args.expected_audience + _expected_scope = args.expected_scope if transport == "stdio": mcp.run(transport="stdio") diff --git a/tests/mcp_tests/test_configs/test_oauth2_mcp_config.yaml b/tests/mcp_tests/test_configs/test_oauth2_mcp_config.yaml index 87134ce01ea..b66e92d2e93 100644 --- a/tests/mcp_tests/test_configs/test_oauth2_mcp_config.yaml +++ b/tests/mcp_tests/test_configs/test_oauth2_mcp_config.yaml @@ -44,6 +44,7 @@ mcp_servers: url: "http://localhost:0/mcp" transport: "http" auth_type: "oauth2" + oauth2_flow: "client_credentials" client_id: "test-client" client_secret: "test-secret" token_url: "http://localhost:0/oauth/token" @@ -72,3 +73,9 @@ mcp_servers: - "mcp.tools.read" - "mcp.tools.execute" subject_token_type: "urn:ietf:params:oauth:token-type:access_token" + +# OAuth Passthrough (proxy forwards caller's Authorization to upstream) + math_oauth_passthrough: + url: "http://localhost:0/mcp" + transport: "http" + oauth_passthrough: true diff --git a/tests/mcp_tests/test_proxy_mcp_auth_e2e.py b/tests/mcp_tests/test_proxy_mcp_auth_e2e.py index 82072de9eac..a6ee4b84a64 100644 --- a/tests/mcp_tests/test_proxy_mcp_auth_e2e.py +++ b/tests/mcp_tests/test_proxy_mcp_auth_e2e.py @@ -53,7 +53,14 @@ SERVER_SPECS: dict[str, dict[str, typing.Optional[str]]] = { }, "test_oauth2_server": {"auth_mode": "oauth2", "auth_secret": None}, "test_pkce_server": {"auth_mode": "oauth2", "auth_secret": None}, - "internal_mcp_server": {"auth_mode": "oauth2", "auth_secret": None}, + "internal_mcp_server": { + "auth_mode": "oauth2", + "auth_secret": None, + "expected_subject_token": "user-subject-jwt", + "expected_audience": "api://internal-tools-mcp", + "expected_scope": "mcp.tools.read mcp.tools.execute", + }, + "math_oauth_passthrough": {"auth_mode": "oauth2", "auth_secret": None}, } @@ -100,7 +107,15 @@ def _reserve_port(host: str = "127.0.0.1") -> int: return sock.getsockname()[1] -def _start_mcp_server_process(*, auth_mode: str, port: int, auth_secret: typing.Optional[str]) -> subprocess.Popen: +def _start_mcp_server_process( + *, + auth_mode: str, + port: int, + auth_secret: typing.Optional[str], + expected_subject_token: typing.Optional[str] = None, + expected_audience: typing.Optional[str] = None, + expected_scope: typing.Optional[str] = None, +) -> subprocess.Popen: cmd = [ sys.executable, str(MCP_SERVER_SCRIPT), @@ -119,8 +134,16 @@ def _start_mcp_server_process(*, auth_mode: str, port: int, auth_secret: typing. ] if auth_secret is not None: cmd.extend(["--auth-secret", auth_secret]) + if expected_subject_token is not None: + cmd.extend(["--expected-subject-token", expected_subject_token]) + if expected_audience is not None: + cmd.extend(["--expected-audience", expected_audience]) + if expected_scope is not None: + cmd.extend(["--expected-scope", expected_scope]) - process = subprocess.Popen(cmd, cwd=str(PROJECT_ROOT), stdout=subprocess.PIPE, stderr=subprocess.PIPE) + process = subprocess.Popen( + cmd, cwd=str(PROJECT_ROOT), stdout=subprocess.PIPE, stderr=subprocess.PIPE + ) start_time = time.time() while True: @@ -136,7 +159,9 @@ def _start_mcp_server_process(*, auth_mode: str, port: int, auth_secret: typing. except OSError: if time.time() - start_time > PROXY_START_TIMEOUT: process.terminate() - raise TimeoutError(f"MCP server did not start in time (auth_mode={auth_mode})") + raise TimeoutError( + f"MCP server did not start in time (auth_mode={auth_mode})" + ) time.sleep(0.05) return process @@ -155,7 +180,9 @@ def _clear_proxy_database_env() -> typing.Iterator[None]: @pytest.fixture(scope="session") def mcp_auth_servers() -> typing.Iterator[dict[str, typing.Any]]: - servers = {name: {**spec, "port": _reserve_port()} for name, spec in SERVER_SPECS.items()} + servers = { + name: {**spec, "port": _reserve_port()} for name, spec in SERVER_SPECS.items() + } processes: list[subprocess.Popen] = [] try: @@ -164,6 +191,9 @@ def mcp_auth_servers() -> typing.Iterator[dict[str, typing.Any]]: auth_mode=spec["auth_mode"], port=spec["port"], auth_secret=spec["auth_secret"], + expected_subject_token=spec.get("expected_subject_token"), + expected_audience=spec.get("expected_audience"), + expected_scope=spec.get("expected_scope"), ) spec["process"] = process spec["base_url"] = f"http://127.0.0.1:{spec['port']}" @@ -188,7 +218,11 @@ def proxy_server_url( server_config = config["mcp_servers"][server_name] base_url = spec["base_url"] server_config["url"] = f"{base_url}/mcp" - for endpoint_key in ("token_url", "authorization_url", "token_exchange_endpoint"): + for endpoint_key in ( + "token_url", + "authorization_url", + "token_exchange_endpoint", + ): if endpoint_key in server_config: suffix = "authorize" if "authorization" in endpoint_key else "token" server_config[endpoint_key] = f"{base_url}/oauth/{suffix}" @@ -213,14 +247,15 @@ async def _call_add_tool( headers: typing.Optional[dict[str, str]] = None, ) -> typing.Optional[str]: request_headers = { - "Authorization": PROXY_AUTHORIZATION_HEADER, - "x-mcp-servers": server_name, + "x-litellm-api-key": "Bearer sk-1234" } if headers: request_headers.update(headers) async with asyncio.timeout(20): - async with streamablehttp_client(url=f"{proxy_server_url}/mcp", headers=request_headers) as ( + async with streamablehttp_client( + url=f"{proxy_server_url}/{server_name}/mcp", headers=request_headers + ) as ( read, write, _get_session_id, @@ -242,14 +277,15 @@ async def _list_tool_names( headers: typing.Optional[dict[str, str]] = None, ) -> list[str]: request_headers = { - "Authorization": PROXY_AUTHORIZATION_HEADER, - "x-mcp-servers": server_name, + "x-litellm-api-key": "Bearer sk-1234" } if headers: request_headers.update(headers) async with asyncio.timeout(20): - async with streamablehttp_client(url=f"{proxy_server_url}/mcp", headers=request_headers) as ( + async with streamablehttp_client( + url=f"{proxy_server_url}/{server_name}/mcp", headers=request_headers + ) as ( read, write, _get_session_id, @@ -272,16 +308,27 @@ class TestProxyMcpAuthE2E: ("test_oauth2_server", 9, 10, "19"), ], ) - async def test_proxy_forwards_configured_credential(self, proxy_server_url, server_name, a, b, expected) -> None: - result = await _call_add_tool(proxy_server_url=proxy_server_url, server_name=server_name, a=a, b=b) + async def test_proxy_forwards_configured_credential( + self, proxy_server_url, server_name, a, b, expected + ) -> None: + result = await _call_add_tool( + proxy_server_url=proxy_server_url, server_name=server_name, a=a, b=b + ) assert result == expected @pytest.mark.asyncio @pytest.mark.parametrize( "server_name", - ["math_api_key", "math_bearer_token", "math_authorization", "test_oauth2_server"], + [ + "math_api_key", + "math_bearer_token", + "math_authorization", + "test_oauth2_server", + ], ) - async def test_upstream_rejects_unauthenticated_request(self, mcp_auth_servers, server_name) -> None: + async def test_upstream_rejects_unauthenticated_request( + self, mcp_auth_servers, server_name + ) -> None: base_url = mcp_auth_servers[server_name]["base_url"] async with httpx.AsyncClient() as client: response = await client.post( @@ -291,7 +338,9 @@ class TestProxyMcpAuthE2E: assert response.status_code == 401 @pytest.mark.asyncio - async def test_oauth_m2m_ignores_caller_authorization(self, proxy_server_url) -> None: + async def test_oauth_m2m_ignores_caller_authorization( + self, proxy_server_url + ) -> None: """M2M servers must never forward the caller's Authorization; the proxy fetches its own client_credentials token. A bogus caller token must not break the call (proves the proxy substitutes its own upstream token).""" @@ -314,7 +363,9 @@ class TestProxyMcpAuthE2E: server_name="math_custom_header", a=4, b=5, - headers={f"x-mcp-math_custom_header-{DEFAULT_CUSTOM_HEADER}": DEFAULT_CUSTOM_HEADER_VALUE}, + headers={ + f"x-mcp-math_custom_header-{DEFAULT_CUSTOM_HEADER}": DEFAULT_CUSTOM_HEADER_VALUE + }, ) assert result == "9" @@ -322,7 +373,9 @@ class TestProxyMcpAuthE2E: async def test_custom_header_required_for_discovery(self, proxy_server_url) -> None: """Without the per-server custom header the upstream rejects the request, so its tools must not be discoverable through the proxy.""" - tool_names = await _list_tool_names(proxy_server_url=proxy_server_url, server_name="math_custom_header") + tool_names = await _list_tool_names( + proxy_server_url=proxy_server_url, server_name="math_custom_header" + ) assert not any(name.endswith("add") for name in tool_names) @pytest.mark.asyncio @@ -341,3 +394,34 @@ class TestProxyMcpAuthE2E: }, ) assert result == "13" + + @pytest.mark.asyncio + async def test_oauth_passthrough_forwards_caller_token( + self, proxy_server_url + ) -> None: + """Passthrough: the proxy forwards the caller's Authorization header + directly to the upstream server without any token exchange.""" + from tests.mcp_tests.mcp_server import DEFAULT_OAUTH_ACCESS_TOKEN + + result = await _call_add_tool( + proxy_server_url=proxy_server_url, + server_name="math_oauth_passthrough", + a=10, + b=11, + headers={ + "x-litellm-api-key": "Bearer sk-1234", + "Authorization": f"Bearer {DEFAULT_OAUTH_ACCESS_TOKEN}", + }, + ) + assert result == "21" + + @pytest.mark.asyncio + async def test_oauth_passthrough_rejects_without_token( + self, proxy_server_url + ) -> None: + """Passthrough without a caller token: the upstream rejects it.""" + tool_names = await _list_tool_names( + proxy_server_url=proxy_server_url, + server_name="math_oauth_passthrough", + ) + assert not any(name.endswith("add") for name in tool_names) diff --git a/tests/mcp_tests/test_proxy_mcp_byok_pkce_e2e.py b/tests/mcp_tests/test_proxy_mcp_byok_pkce_e2e.py index 4b8060af673..1531270cf8c 100644 --- a/tests/mcp_tests/test_proxy_mcp_byok_pkce_e2e.py +++ b/tests/mcp_tests/test_proxy_mcp_byok_pkce_e2e.py @@ -1,65 +1,71 @@ -"""DB-backed end-to-end test for the interactive PKCE (BYOK) MCP OAuth flow. +"""Interactive PKCE MCP e2e test against a real proxy at localhost:4000. -This drives LiteLLM acting as the OAuth 2.1 authorization server: - 1. mint a UI session cookie - 2. POST /v1/mcp/oauth/authorize with a PKCE code_challenge -> authorization code - 3. POST /v1/mcp/oauth/token with the code_verifier -> access token (and the - per-user credential is persisted to the DB) - 4. call the MCP server through the proxy using that access token; the proxy - loads the stored credential and forwards it upstream +Assumes CircleCI (or local dev) has: + - a litellm proxy running at localhost:4000 with DATABASE_URL and LITELLM_MASTER_KEY set + - a Postgres database accessible via DATABASE_URL -Requires a database (DATABASE_URL). Skipped otherwise. In CI this runs against -the Postgres service wired into the MCP workflow. +Flow: + 1. register a BYOK MCP server via the management API + 2. mint a UI session cookie (real litellm auth functions) + 3. POST /v1/mcp/oauth/authorize with a PKCE code_challenge -> authorization code + 4. POST /v1/mcp/oauth/token with the code_verifier -> access token + 5. use master key + access token to call mcp tools via the per-server path """ import asyncio -import base64 -import hashlib import os -import time import typing import uuid +from datetime import timedelta +from typing import cast from urllib.parse import parse_qs, urlparse import httpx import jwt +import litellm import pytest -import yaml from mcp import ClientSession from mcp.client.streamable_http import streamablehttp_client +from litellm.constants import LITELLM_PROXY_ADMIN_NAME +from litellm.proxy.auth.login_utils import LoginResult, create_ui_token_object +from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler from tests.mcp_tests.mcp_server import DEFAULT_OAUTH_ACCESS_TOKEN from tests.mcp_tests.test_proxy_mcp_auth_e2e import ( _reserve_port, _start_mcp_server_process, - _start_proxy_server, ) DATABASE_URL = os.getenv("DATABASE_URL") -pytestmark = pytest.mark.skipif(not DATABASE_URL, reason="BYOK PKCE e2e requires a database (DATABASE_URL)") +pytestmark = pytest.mark.skipif( + not DATABASE_URL, reason="BYOK PKCE e2e requires a database (DATABASE_URL)" +) -MASTER_KEY = "sk-1234" -BYOK_USER_ID = "byok-pkce-user" +MASTER_KEY = os.getenv("LITELLM_MASTER_KEY", "sk-1234") +PROXY_BASE_URL = os.getenv("LITELLM_PROXY_BASE_URL", "http://localhost:4000") REDIRECT_URI = "http://127.0.0.1:8765/callback" def _pkce_pair() -> tuple[str, str]: - verifier = base64.urlsafe_b64encode(os.urandom(40)).rstrip(b"=").decode() - digest = hashlib.sha256(verifier.encode()).digest() - challenge = base64.urlsafe_b64encode(digest).rstrip(b"=").decode() - return verifier, challenge + return SSOAuthenticationHandler.generate_pkce_params() def _session_cookie() -> str: - return jwt.encode( - { - "user_id": BYOK_USER_ID, - "login_method": "username_password", - "exp": int(time.time()) + 3600, - }, - MASTER_KEY, - algorithm="HS256", + login_result = LoginResult( + user_id=LITELLM_PROXY_ADMIN_NAME, + key="byok-pkce-e2e-ui-session", + user_email="byok-pkce-e2e@test.local", + user_role="proxy_admin", + login_method="username_password", ) + returned_ui_token_object = create_ui_token_object( + login_result=login_result, + general_settings={}, + premium_user=False, + ) + payload = dict(cast(dict, returned_ui_token_object)) + payload["exp"] = litellm.utils.get_utc_datetime() + timedelta(hours=1) + return jwt.encode(payload, MASTER_KEY, algorithm="HS256") @pytest.fixture(scope="module") @@ -78,67 +84,35 @@ def upstream_server() -> typing.Iterator[dict[str, typing.Any]]: @pytest.fixture(scope="module") def byok_server_id(upstream_server: dict[str, typing.Any]) -> str: - """Create a BYOK MCP server row directly in the DB and return its id.""" - from litellm.proxy._types import NewMCPServerRequest - from litellm.proxy._experimental.mcp_server.db import create_mcp_server - from litellm.proxy.utils import PrismaClient, ProxyLogging - from litellm.caching import DualCache - + """Register a BYOK MCP server via the proxy management API.""" server_id = str(uuid.uuid4()) + alias = f"byok_pkce_server_{server_id[:8]}" - async def _create() -> None: - prisma_client = PrismaClient( - database_url=DATABASE_URL, - proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), - ) - await prisma_client.connect() - try: - await create_mcp_server( - prisma_client, - NewMCPServerRequest( - server_id=server_id, - alias="byok_pkce_server", - url=f"{upstream_server['base_url']}/mcp", - transport="http", - auth_type="oauth2", - is_byok=True, - allow_all_keys=True, - ), - touched_by="byok-pkce-e2e", - ) - finally: - await prisma_client.disconnect() - - asyncio.run(_create()) - return server_id + response = httpx.post( + f"{PROXY_BASE_URL}/v1/mcp/server", + headers={"Authorization": f"Bearer {MASTER_KEY}"}, + json={ + "server_id": server_id, + "alias": alias, + "url": f"{upstream_server['base_url']}/mcp", + "transport": "http", + "auth_type": "oauth2", + "is_byok": True, + "allow_all_keys": True, + }, + ) + assert response.status_code == 201, ( + f"Failed to register BYOK server: {response.status_code} {response.text}" + ) + return alias @pytest.fixture(scope="module") -def proxy_server_url(tmp_path_factory: pytest.TempPathFactory, byok_server_id: str) -> typing.Iterator[str]: - os.environ["DATABASE_URL"] = DATABASE_URL # restore if a sibling test cleared it - os.environ["LITELLM_MASTER_KEY"] = MASTER_KEY - - config = { - "general_settings": {"master_key": MASTER_KEY}, - "model_list": [ - { - "model_name": "fake-model", - "litellm_params": {"model": "openai/fake", "api_key": "fake-key"}, - } - ], - } - config_path = tmp_path_factory.mktemp("byok_pkce") / "config.yaml" - config_path.write_text(yaml.safe_dump(config)) - - server_url, server, thread, sock = _start_proxy_server(str(config_path)) - yield server_url - - server.should_exit = True - thread.join(timeout=10) - sock.close() +def proxy_server_url() -> str: + return PROXY_BASE_URL -async def _complete_pkce_flow(proxy_server_url: str, server_id: str) -> str: +async def _complete_pkce_flow(proxy_server_url: str, server_alias: str) -> str: """Run authorize -> token and return the issued access token.""" verifier, challenge = _pkce_pair() cookies = {"token": _session_cookie()} @@ -151,7 +125,7 @@ async def _complete_pkce_flow(proxy_server_url: str, server_id: str) -> str: "code_challenge": challenge, "code_challenge_method": "S256", "state": "xyz", - "server_id": server_id, + "server_id": server_alias, "api_key": DEFAULT_OAUTH_ACCESS_TOKEN, "client_id": "byok-client", }, @@ -175,7 +149,9 @@ async def _complete_pkce_flow(proxy_server_url: str, server_id: str) -> str: @pytest.mark.asyncio -async def test_byok_pkce_authorize_rejects_wrong_verifier(proxy_server_url: str, byok_server_id: str) -> None: +async def test_byok_pkce_authorize_rejects_wrong_verifier( + proxy_server_url: str, byok_server_id: str +) -> None: """PKCE enforcement: a token request with a mismatched verifier is rejected.""" _verifier, challenge = _pkce_pair() cookies = {"token": _session_cookie()} @@ -212,7 +188,9 @@ async def test_byok_pkce_authorize_rejects_wrong_verifier(proxy_server_url: str, @pytest.mark.asyncio -async def test_byok_pkce_authorize_requires_session(proxy_server_url: str, byok_server_id: str) -> None: +async def test_byok_pkce_authorize_requires_session( + proxy_server_url: str, byok_server_id: str +) -> None: """Without a UI session cookie the authorize endpoint must refuse.""" _verifier, challenge = _pkce_pair() async with httpx.AsyncClient(follow_redirects=False) as client: @@ -231,27 +209,22 @@ async def test_byok_pkce_authorize_requires_session(proxy_server_url: str, byok_ @pytest.mark.asyncio -async def test_byok_pkce_end_to_end_tool_call(proxy_server_url: str, byok_server_id: str) -> None: - """Full interactive PKCE flow: authorize -> token -> call the MCP tool with - the issued access token. The proxy forwards the stored BYOK credential - upstream, so the tool call succeeds.""" +async def test_byok_pkce_end_to_end_tool_call( + proxy_server_url: str, byok_server_id: str +) -> None: access_token = await _complete_pkce_flow(proxy_server_url, byok_server_id) headers = { + "x-litellm-api-key": f"Bearer {MASTER_KEY}", "Authorization": f"Bearer {access_token}", - "x-mcp-servers": "byok_pkce_server", } async with asyncio.timeout(20): - async with streamablehttp_client(url=f"{proxy_server_url}/mcp", headers=headers) as ( - read, - write, - _get_session_id, - ): + async with streamablehttp_client( + url=f"{proxy_server_url}/{byok_server_id}/mcp", headers=headers + ) as (read, write, _get_session_id): async with ClientSession(read, write) as session: await session.initialize() - tools = await session.list_tools() - assert any(tool.name.endswith("add") for tool in tools.tools) - result = await session.call_tool("add", arguments={"a": 5, "b": 6}) assert result.content - assert getattr(result.content[0], "text", None) == "11" + text = getattr(result.content[0], "text", None) + assert text == "11"