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:
joshua-berri 2026-10-07 16:30:10 -07:00 • committed by GitHub
parent d910653b33
commit e66dbfc366
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 201 additions and 0 deletions

View file

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

View file

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

View file

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

View file

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