fix passthrough bugs

This commit is contained in:
mubashir1osmani 2026-06-07 20:35:53 -07:00
parent 3d22e0197f
commit 58b00309b3
No known key found for this signature in database
GPG key ID: AB055FF67D0B4D9A
4 changed files with 214 additions and 123 deletions

View file

@ -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")

View file

@ -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

View file

@ -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)

View file

@ -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"