mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix passthrough bugs
This commit is contained in:
parent
3d22e0197f
commit
58b00309b3
4 changed files with 214 additions and 123 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue