diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index f4bcc57366c..1f38742701e 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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, ) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 6d8635960a1..6391cadd96d 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -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( diff --git a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 0736391642b..bf21e3434ca 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -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 diff --git a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py index 44cc1e80b09..790df66b909 100644 --- a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -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)