mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(mcp): preserve client application type during registration (#45159)
Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
parent
d910653b33
commit
e66dbfc366
4 changed files with 201 additions and 0 deletions
|
|
@ -1761,6 +1761,16 @@ def client_supplied_redirect_uris(value: object) -> list[str] | None:
|
|||
return uris if len(uris) == len(value) else None
|
||||
|
||||
|
||||
_CLIENT_APPLICATION_TYPE: Final = TypeAdapter(Literal["native", "web"] | None)
|
||||
|
||||
|
||||
def client_supplied_application_type(value: object) -> Literal["native", "web"] | None:
|
||||
try:
|
||||
return _CLIENT_APPLICATION_TYPE.validate_python(value)
|
||||
except ValidationError as exc:
|
||||
raise HTTPException(status_code=400, detail="application_type must be native or web") from exc
|
||||
|
||||
|
||||
async def _post_dcr_registration(
|
||||
registration_url: str,
|
||||
register_data: Mapping[str, object],
|
||||
|
|
@ -1925,6 +1935,7 @@ async def register_client_with_server(
|
|||
fallback_client_id: str | None = None,
|
||||
persist_credentials: bool = False,
|
||||
client_redirect_uris: list[str] | None = None,
|
||||
client_application_type: Literal["native", "web"] | None = None,
|
||||
):
|
||||
_raise_if_not_oauth2(mcp_server)
|
||||
request_base_url: Final = get_request_base_url(request)
|
||||
|
|
@ -1980,6 +1991,11 @@ async def register_client_with_server(
|
|||
)
|
||||
|
||||
register_data: Final = {
|
||||
**(
|
||||
{"application_type": client_application_type}
|
||||
if bridge_relay and client_application_type is not None
|
||||
else {}
|
||||
),
|
||||
"client_name": client_name,
|
||||
"redirect_uris": client_redirect_uris if bridge_relay else [current_redirect_uri],
|
||||
"grant_types": grant_types or (["authorization_code", "refresh_token"] if bridge_relay else []),
|
||||
|
|
@ -3094,6 +3110,7 @@ async def register_client(request: Request, mcp_server_name: str | None = None):
|
|||
return await register_aggregate_client(
|
||||
request=request, request_body=data, token_exchange_available=token_exchange_available()
|
||||
)
|
||||
client_application_type: Final = client_supplied_application_type(data.get("application_type"))
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
async with global_mcp_server_manager.catalog.operation():
|
||||
|
|
@ -3115,6 +3132,7 @@ async def register_client(request: Request, mcp_server_name: str | None = None):
|
|||
token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""),
|
||||
fallback_client_id=resolved.server_name or resolved.name,
|
||||
client_redirect_uris=client_redirect_uris,
|
||||
client_application_type=client_application_type,
|
||||
)
|
||||
return dummy_return
|
||||
|
||||
|
|
@ -3130,4 +3148,5 @@ async def register_client(request: Request, mcp_server_name: str | None = None):
|
|||
token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""),
|
||||
fallback_client_id=mcp_server_name,
|
||||
client_redirect_uris=client_redirect_uris,
|
||||
client_application_type=client_application_type,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -179,6 +179,7 @@ if MCP_AVAILABLE:
|
|||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_raise_if_not_oauth2,
|
||||
authorize_with_server,
|
||||
client_supplied_application_type,
|
||||
client_supplied_redirect_uris,
|
||||
exchange_token_with_server,
|
||||
get_request_base_url,
|
||||
|
|
@ -2426,6 +2427,7 @@ if MCP_AVAILABLE:
|
|||
request_data: Final = await _read_request_body(request=request)
|
||||
data: Final[Mapping[str, object]] = {**request_data}
|
||||
client_redirect_uris: Final = client_supplied_redirect_uris(data.get("redirect_uris"))
|
||||
client_application_type: Final = client_supplied_application_type(data.get("application_type"))
|
||||
|
||||
return await register_client_with_server(
|
||||
request=request,
|
||||
|
|
@ -2437,6 +2439,7 @@ if MCP_AVAILABLE:
|
|||
fallback_client_id=server_id,
|
||||
persist_credentials=_user_is_full_admin(user_api_key_dict),
|
||||
client_redirect_uris=client_redirect_uris,
|
||||
client_application_type=client_application_type,
|
||||
)
|
||||
|
||||
@router.delete(
|
||||
|
|
|
|||
|
|
@ -13247,3 +13247,118 @@ async def test_registration_losing_conditional_write_reuses_only_a_matching_winn
|
|||
assert result == ("reused" if winner_available else "failed")
|
||||
assert update.await_args.kwargs["expected_updated_at"] == row.updated_at
|
||||
assert server.client_id == ("winner-client" if winner_available else None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type", (MCPAuth.true_passthrough, MCPAuth.oauth_delegate, MCPAuth.oauth2))
|
||||
@pytest.mark.parametrize(
|
||||
"metadata", ({"application_type": "native"}, {"application_type": "web"}, {}, {"application_type": None})
|
||||
)
|
||||
async def test_register_preserves_client_application_type_only_for_bridge_relay(
|
||||
auth_type: MCPAuth, metadata: dict[str, object], monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
import httpx
|
||||
import respx
|
||||
from fastapi import FastAPI
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import router
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
server: Final = _bridge_server(auth_type=auth_type, server_id="application-client", alias="application-client")
|
||||
app: Final = FastAPI()
|
||||
app.include_router(router)
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
monkeypatch.setitem(global_mcp_server_manager.registry, server.server_id, server)
|
||||
client_redirect: Final = "http://127.0.0.1:53682/callback"
|
||||
with respx.mock as upstream:
|
||||
registration: Final = upstream.post(server.registration_url).respond(
|
||||
201, json={"client_id": "registered-client"}
|
||||
)
|
||||
async with httpx.AsyncClient(
|
||||
transport=httpx.ASGITransport(app=app), base_url="https://gateway.example"
|
||||
) as client:
|
||||
response: Final = await client.post(
|
||||
f"/{server.server_id}/register",
|
||||
json={"client_name": "Test client", "redirect_uris": [client_redirect], **metadata},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
assert response.json()["client_id"] == "registered-client"
|
||||
assert registration.call_count == 1
|
||||
posted: Final = json.loads(registration.calls[0].request.content)
|
||||
expected_type: Final = metadata.get("application_type") if auth_type != MCPAuth.oauth2 else None
|
||||
if expected_type is None:
|
||||
assert "application_type" not in posted
|
||||
else:
|
||||
assert posted["application_type"] == expected_type
|
||||
assert posted["redirect_uris"] == (
|
||||
["https://gateway.example/callback"] if auth_type == MCPAuth.oauth2 else [client_redirect]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("application_type", ("desktop", "", 1, ["native"], {"value": "native"}))
|
||||
async def test_register_rejects_invalid_application_type_before_upstream(
|
||||
application_type: object, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
import httpx
|
||||
import respx
|
||||
from fastapi import FastAPI
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import router
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
server: Final = _bridge_server(server_id="invalid-application-client", alias="invalid-application-client")
|
||||
app: Final = FastAPI()
|
||||
app.include_router(router)
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
monkeypatch.setitem(global_mcp_server_manager.registry, server.server_id, server)
|
||||
with respx.mock(assert_all_called=False) as upstream:
|
||||
registration: Final = upstream.post(server.registration_url).respond(
|
||||
201, json={"client_id": "must-not-register"}
|
||||
)
|
||||
async with httpx.AsyncClient(
|
||||
transport=httpx.ASGITransport(app=app), base_url="https://gateway.example"
|
||||
) as client:
|
||||
response: Final = await client.post(
|
||||
f"/{server.server_id}/register",
|
||||
json={"redirect_uris": ["http://127.0.0.1:53682/callback"], "application_type": application_type},
|
||||
)
|
||||
assert response.status_code == 400
|
||||
assert "application_type" in response.json()["detail"]
|
||||
assert registration.call_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("client_id", (None, "preconfigured-client"))
|
||||
async def test_register_application_type_keeps_no_registration_endpoint_fallback(
|
||||
client_id: str | None, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
import httpx
|
||||
import respx
|
||||
from fastapi import FastAPI
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import router
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
server: Final = _bridge_server(
|
||||
auth_type=MCPAuth.oauth2,
|
||||
server_id="static-client",
|
||||
alias="static-client",
|
||||
registration_url=None,
|
||||
client_id=client_id,
|
||||
)
|
||||
app: Final = FastAPI()
|
||||
app.include_router(router)
|
||||
monkeypatch.setitem(global_mcp_server_manager.registry, server.server_id, server)
|
||||
with respx.mock as upstream:
|
||||
async with httpx.AsyncClient(
|
||||
transport=httpx.ASGITransport(app=app), base_url="https://gateway.example"
|
||||
) as client:
|
||||
response: Final = await client.post(f"/{server.server_id}/register", json={"application_type": "native"})
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {
|
||||
"client_id": server.server_id,
|
||||
"client_secret": "dummy",
|
||||
"redirect_uris": ["https://gateway.example/callback"],
|
||||
}
|
||||
assert len(upstream.calls) == 0
|
||||
|
|
|
|||
|
|
@ -3751,6 +3751,7 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
fallback_client_id="server-1",
|
||||
persist_credentials=True,
|
||||
client_redirect_uris=None,
|
||||
client_application_type=None,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -11434,3 +11435,66 @@ def test_staged_issuer_edit_preserves_replacement_with_same_client_id(monkeypatc
|
|||
staged = management._inherit_credentials_from_existing_server(payload)
|
||||
assert staged.credentials == submitted
|
||||
assert saved.client_secret == "old-secret"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.respx(assert_all_called=False)
|
||||
@pytest.mark.parametrize("application_type", ("native", "web", None, "desktop"))
|
||||
async def test_mcp_register_application_type_reaches_upstream_or_is_rejected(
|
||||
application_type: str | None, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter
|
||||
) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="temporary-application-client",
|
||||
name="temporary-application-client",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.true_passthrough,
|
||||
dcr_bridge=True,
|
||||
authorization_url="https://provider.example/authorize",
|
||||
token_url="https://provider.example/token",
|
||||
registration_url="https://provider.example/register",
|
||||
)
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
mgmt_endpoints._cache_temporary_mcp_server(server, ttl_seconds=60)
|
||||
request: Final = Request(
|
||||
{
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"scheme": "https",
|
||||
"server": ("gateway.example", 443),
|
||||
"path": "/v1/mcp/server/oauth/temporary-application-client/register",
|
||||
"headers": [],
|
||||
},
|
||||
receive=AsyncMock(
|
||||
return_value={
|
||||
"type": "http.request",
|
||||
"body": json.dumps(
|
||||
{
|
||||
"redirect_uris": ["http://127.0.0.1:53682/callback"],
|
||||
"application_type": application_type,
|
||||
}
|
||||
).encode(),
|
||||
}
|
||||
),
|
||||
)
|
||||
registration: Final = respx_mock.post(server.registration_url).respond(201, json={"client_id": "registered-client"})
|
||||
try:
|
||||
if application_type == "desktop":
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await mgmt_endpoints.mcp_register(request, server.server_id, generate_mock_user_api_key_auth())
|
||||
assert exc.value.status_code == 400
|
||||
assert "application_type" in str(exc.value.detail)
|
||||
assert registration.call_count == 0
|
||||
return
|
||||
response: Final = await mgmt_endpoints.mcp_register(
|
||||
request, server.server_id, generate_mock_user_api_key_auth()
|
||||
)
|
||||
assert response.status_code == 200
|
||||
assert json.loads(response.body)["client_id"] == "registered-client"
|
||||
assert registration.call_count == 1
|
||||
posted: Final = json.loads(registration.calls[0].request.content)
|
||||
if application_type is None:
|
||||
assert "application_type" not in posted
|
||||
else:
|
||||
assert posted["application_type"] == application_type
|
||||
finally:
|
||||
mgmt_endpoints._temporary_mcp_servers.pop(server.server_id, None)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue