mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
Merge pull request #40808 from BerriAI/litellm_fix_mcp_oauth_issuer_7078
fix(mcp): match per-server OAuth metadata issuers
This commit is contained in:
commit
e86adf98ac
2 changed files with 108 additions and 2 deletions
|
|
@ -2666,6 +2666,8 @@ async def oauth_protected_resource_mcp(request: Request, mcp_server_name: str |
|
|||
def _build_oauth_authorization_server_response(
|
||||
request: Request,
|
||||
mcp_server_name: str | None,
|
||||
*,
|
||||
issuer_path: str | None = None,
|
||||
) -> dict:
|
||||
"""Build OAuth authorization server metadata response (gateway-as-AS shape).
|
||||
|
||||
|
|
@ -2694,7 +2696,13 @@ def _build_oauth_authorization_server_response(
|
|||
|
||||
_raise_unless_oauth2_discovery_server(mcp_server, mcp_server_name, "not an OAuth authorization server")
|
||||
|
||||
issuer: Final = f"{request_base_url}/{mcp_server_name}" if explicitly_named else request_base_url
|
||||
issuer: Final = (
|
||||
f"{request_base_url}/{issuer_path}"
|
||||
if issuer_path is not None
|
||||
else f"{request_base_url}/{mcp_server_name}"
|
||||
if explicitly_named
|
||||
else request_base_url
|
||||
)
|
||||
|
||||
return {
|
||||
"issuer": issuer,
|
||||
|
|
@ -2724,6 +2732,7 @@ async def oauth_authorization_server_mcp_standard(request: Request, mcp_server_n
|
|||
return _build_oauth_authorization_server_response(
|
||||
request=request,
|
||||
mcp_server_name=mcp_server_name,
|
||||
issuer_path=f"mcp/{mcp_server_name}",
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -2802,7 +2811,7 @@ async def jwks_json(request: Request):
|
|||
|
||||
|
||||
# Additional legacy pattern support
|
||||
@router.get("/.well-known/oauth-authorization-server/{mcp_server_name}/mcp")
|
||||
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/{{mcp_server_name}}/mcp")
|
||||
async def oauth_authorization_server_legacy(request: Request, mcp_server_name: str):
|
||||
"""
|
||||
OAuth authorization server discovery for legacy /{server_name}/mcp pattern.
|
||||
|
|
@ -2810,6 +2819,7 @@ async def oauth_authorization_server_legacy(request: Request, mcp_server_name: s
|
|||
return _build_oauth_authorization_server_response(
|
||||
request=request,
|
||||
mcp_server_name=mcp_server_name,
|
||||
issuer_path=f"{mcp_server_name}/mcp",
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -11186,3 +11186,99 @@ async def test_dcr_refusal_is_actionable_without_upstream_body(
|
|||
assert f"HTTP {upstream_status}" in str(exc.value.detail)
|
||||
assert "pre-registered OAuth client" in str(exc.value.detail)
|
||||
assert "private upstream details" not in str(exc.value.detail)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("prefix", ["", "/tenant-a", "/tenant-b"])
|
||||
@pytest.mark.parametrize(
|
||||
("server_name", "pattern"),
|
||||
[
|
||||
("issuer_test", "mcp/{server}"),
|
||||
("issuer_test", "{server}/mcp"),
|
||||
("issuer_test", "{server}"),
|
||||
("mcp", "mcp/{server}"),
|
||||
("mcp", "{server}/mcp"),
|
||||
],
|
||||
)
|
||||
def test_per_server_authorization_metadata_issuer_matches_discovery_path(
|
||||
_no_proxy_base_url, _isolated_mcp_registry, prefix, server_name, pattern
|
||||
):
|
||||
server = _create_oauth2_server(server_id=server_name, name=server_name, server_name=server_name, alias=server_name)
|
||||
_isolated_mcp_registry[server.server_id] = server
|
||||
client = _prefixed_discovery_client(["/tenant-a", "/tenant-b"])
|
||||
path = pattern.format(server=server_name)
|
||||
response = client.get(f"{prefix}/.well-known/oauth-authorization-server/{path}")
|
||||
assert response.status_code == 200
|
||||
metadata = response.json()
|
||||
assert metadata["issuer"] == f"http://testserver{prefix}/{path}"
|
||||
assert metadata["authorization_endpoint"] == f"http://testserver{prefix}/{server_name}/authorize"
|
||||
assert metadata["token_endpoint"] == f"http://testserver{prefix}/{server_name}/token"
|
||||
assert metadata["registration_endpoint"] == f"http://testserver{prefix}/{server_name}/register"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("prefix", ["", "/tenant-a"])
|
||||
@pytest.mark.parametrize("relay", [False, True])
|
||||
@pytest.mark.parametrize("pattern", ["mcp/{server}", "{server}/mcp"])
|
||||
def test_named_resource_discovery_follows_matching_authorization_issuer(
|
||||
_no_proxy_base_url, _isolated_mcp_registry, prefix, relay, pattern
|
||||
):
|
||||
server = _create_oauth2_server().model_copy(update={"per_server_oauth_discovery": relay})
|
||||
_isolated_mcp_registry[server.server_id] = server
|
||||
client = _prefixed_discovery_client(["/tenant-a"])
|
||||
path = pattern.format(server=server.server_name)
|
||||
response = client.get(f"{prefix}/.well-known/oauth-protected-resource/{path}")
|
||||
assert response.status_code == 200
|
||||
resource = response.json()
|
||||
issuer_path = server.server_name if relay else "mcp"
|
||||
assert resource["resource"] == f"http://testserver{prefix}/{path}"
|
||||
assert resource["authorization_servers"] == [f"http://testserver{prefix}/{issuer_path}"]
|
||||
authorization = client.get(f"{prefix}/.well-known/oauth-authorization-server/{issuer_path}")
|
||||
assert authorization.status_code == 200
|
||||
assert authorization.json()["issuer"] == resource["authorization_servers"][0]
|
||||
|
||||
|
||||
def test_static_root_path_authorization_discovery_preserves_issuer(monkeypatch, tmp_path):
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
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
|
||||
|
||||
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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue