diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 3597e75404a..9cc84400cf5 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -215,6 +215,34 @@ if MCP_AVAILABLE: _base_validate_and_normalize_mcp_server_payload(payload) _validate_mcp_server_name_fields(payload) + def stamp_omitted_oauth2_flow(payload: NewMCPServerRequest) -> None: + """Fallback only: fill in oauth2_flow when an oauth2 create omits it. + + An explicit oauth2_flow from the caller (the dashboard's flow selector, a REST + body, config.yaml) always wins and is never touched. The shape check below runs + solely for oauth2 creates that leave the field unset, so those rows still + persist a flow instead of relying on read-time inference. + + The create payload carries the plaintext credentials, so the M2M-vs-interactive + decision is reliable here in a way it is not at read time (credentials are + encrypted at rest and redacted in responses). The client_credentials shape + mirrors the legacy inference in MCPServerManager._resolve_oauth2_flow; every + other oauth2 configuration is the authorization_code grant, including + delegate_auth_to_upstream, where the client runs that grant upstream. + """ + if payload.auth_type != MCPAuth.oauth2: + return + if payload.oauth2_flow: + return + credentials = payload.credentials or {} + has_m2m_shape = bool( + payload.token_url + and credentials.get("client_id") + and credentials.get("client_secret") + and not payload.authorization_url + ) + payload.oauth2_flow = "client_credentials" if has_m2m_shape else "authorization_code" + _VALID_MCP_REQUIRED_FIELDS: frozenset = frozenset(NewMCPServerRequest.model_fields) def _validate_mcp_required_fields(payload: Any) -> None: @@ -1057,6 +1085,7 @@ if MCP_AVAILABLE: prisma_client = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") validate_and_normalize_mcp_server_payload(payload) + stamp_omitted_oauth2_flow(payload) _validate_mcp_required_fields(payload) payload.approval_status = MCPApprovalStatus.pending_review @@ -1322,6 +1351,7 @@ if MCP_AVAILABLE: # Validate and normalize payload fields validate_and_normalize_mcp_server_payload(payload) + stamp_omitted_oauth2_flow(payload) # AuthZ - restrict only proxy admins to create mcp servers if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role: @@ -1413,6 +1443,7 @@ if MCP_AVAILABLE: # Validate and normalize payload fields (alias/server name rules) validate_and_normalize_mcp_server_payload(payload) + stamp_omitted_oauth2_flow(payload) # Restrict to proxy admins similar to the persistent create endpoint if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role: diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index e223140b573..16508c8f2fd 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -5238,3 +5238,63 @@ class TestPerUserCredentialConfigServerResolution: _, _, _, updates, _ = merge_mock.await_args.args assert updates == {"CORP_USERNAME": "alice"} assert result.server_id == self.CONFIG_SERVER_ID + + +def _oauth2_create_payload(**overrides): + base = dict( + server_name="stamp_test_server", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type="oauth2", + ) + base.update(overrides) + return NewMCPServerRequest(**base) + + +def test_stamp_oauth2_flow_bare_oauth2_defaults_to_authorization_code(): + """A bare oauth2 create (no endpoints, no creds) is interactive: stamping it + authorization_code matches how needs_user_oauth_token treats a null flow.""" + payload = _oauth2_create_payload() + mgmt_endpoints.stamp_omitted_oauth2_flow(payload) + assert payload.oauth2_flow == "authorization_code" + + +def test_stamp_oauth2_flow_marks_m2m_shape_client_credentials(): + """token_url + full client credentials and no authorization_url is the M2M shape; + the stamp mirrors the legacy inference in _resolve_oauth2_flow so REST-created M2M + servers persist the flow instead of relying on read-time inference.""" + payload = _oauth2_create_payload( + token_url="https://idp.example.com/token", + credentials={"client_id": "cid", "client_secret": "csecret"}, + ) + mgmt_endpoints.stamp_omitted_oauth2_flow(payload) + assert payload.oauth2_flow == "client_credentials" + + +def test_stamp_oauth2_flow_authorization_url_wins_over_m2m_shape(): + """An authorization endpoint means interactive even when client creds + token_url + are present (GitHub Enterprise style); M2M never has an authorization endpoint.""" + payload = _oauth2_create_payload( + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + credentials={"client_id": "cid", "client_secret": "csecret"}, + ) + mgmt_endpoints.stamp_omitted_oauth2_flow(payload) + assert payload.oauth2_flow == "authorization_code" + + +def test_stamp_oauth2_flow_respects_explicit_value(): + """An explicit oauth2_flow from the caller must never be overridden by the stamp.""" + payload = _oauth2_create_payload( + oauth2_flow="authorization_code", + token_url="https://idp.example.com/token", + credentials={"client_id": "cid", "client_secret": "csecret"}, + ) + mgmt_endpoints.stamp_omitted_oauth2_flow(payload) + assert payload.oauth2_flow == "authorization_code" + + +def test_stamp_oauth2_flow_ignores_non_oauth2(): + payload = _oauth2_create_payload(auth_type="none") + mgmt_endpoints.stamp_omitted_oauth2_flow(payload) + assert payload.oauth2_flow is None diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx index 1245bcee3fa..db20bfb21b3 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx @@ -991,3 +991,99 @@ describe("CreateMCPServer", () => { }); }); }); + +describe("CreateMCPServer oauth2_flow persistence", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + const createdServer = { + server_id: "new-server-oauth", + server_name: "OAuth_Server", + alias: "OAuth_Server", + url: "https://example.com/mcp", + transport: "http", + auth_type: "oauth2", + created_at: "2024-01-01T00:00:00Z", + created_by: "user-1", + updated_at: "2024-01-01T00:00:00Z", + updated_by: "user-1", + }; + + async function setupHttpServerForm() { + render(); + await selectAntOption("Transport Type", "Streamable HTTP"); + await waitFor(() => { + expect(screen.getByPlaceholderText("https://your-mcp-server.com")).toBeInTheDocument(); + }); + const nameInput = document.getElementById("server_name") as HTMLInputElement; + await act(async () => { + fireEvent.change(nameInput, { target: { value: "OAuth_Server" } }); + }); + const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com"); + await act(async () => { + fireEvent.change(urlInput, { target: { value: "https://example.com/mcp" } }); + }); + } + + async function submitCreate() { + const submitButton = screen.getByRole("button", { name: "Add MCP Server" }); + await act(async () => { + fireEvent.click(submitButton); + }); + await waitFor(() => { + expect(networking.createMCPServer).toHaveBeenCalledTimes(1); + }); + const [, payload] = vi.mocked(networking.createMCPServer).mock.calls[0]; + return payload; + } + + it("persists authorization_code for an interactive OAuth create", async () => { + vi.mocked(networking.createMCPServer).mockResolvedValue(createdServer); + await setupHttpServerForm(); + await selectAntOption("Authentication", "OAuth"); + await waitFor(() => { + expect(screen.getByText("OAuth Flow Type")).toBeInTheDocument(); + }); + + const payload = await submitCreate(); + expect(payload.auth_type).toBe("oauth2"); + expect(payload.oauth2_flow).toBe("authorization_code"); + }); + + it("persists client_credentials for an M2M OAuth create", async () => { + vi.mocked(networking.createMCPServer).mockResolvedValue({ ...createdServer, oauth2_flow: "client_credentials" }); + await setupHttpServerForm(); + await selectAntOption("Authentication", "OAuth"); + await waitFor(() => { + expect(screen.getByText("OAuth Flow Type")).toBeInTheDocument(); + }); + await selectAntOption("OAuth Flow Type", "Machine-to-Machine (M2M)"); + await waitFor(() => { + expect(screen.getByPlaceholderText("Enter OAuth client ID")).toBeInTheDocument(); + }); + await act(async () => { + fireEvent.change(screen.getByPlaceholderText("Enter OAuth client ID"), { target: { value: "cid" } }); + }); + await act(async () => { + fireEvent.change(screen.getByPlaceholderText("Enter OAuth client secret"), { target: { value: "csecret" } }); + }); + await act(async () => { + fireEvent.change(screen.getByPlaceholderText("https://auth.example.com/oauth/token"), { + target: { value: "https://auth.example.com/oauth/token" }, + }); + }); + + const payload = await submitCreate(); + expect(payload.oauth2_flow).toBe("client_credentials"); + }); + + it("sends no oauth2_flow for a non-oauth2 create", async () => { + vi.mocked(networking.createMCPServer).mockResolvedValue({ ...createdServer, auth_type: "none" }); + await setupHttpServerForm(); + await selectAntOption("Authentication", "None"); + + const payload = await submitCreate(); + expect(payload.oauth2_flow).toBeUndefined(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx index 05a0696674a..b2bf16abe39 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx @@ -13,6 +13,7 @@ import { TRANSPORT, getMcpOAuthMode, MCP_OAUTH2_FLOW_M2M, + MCP_OAUTH2_FLOW_INTERACTIVE, } from "./types"; import OAuthFormFields from "./OAuthFormFields"; import MCPServerCostConfig from "./mcp_server_cost_config"; @@ -442,6 +443,12 @@ const CreateMCPServer: React.FC = ({ available_on_public_internet: Boolean(availableOnPublicInternetRaw), delegate_auth_to_upstream: Boolean(delegateAuthToUpstreamRaw), oauth_passthrough: Boolean(oauthPassthroughRaw), + ...(restValues.auth_type === AUTH_TYPE.OAUTH2 + ? { + oauth2_flow: + values.oauth_flow_type === OAUTH_FLOW.M2M ? MCP_OAUTH2_FLOW_M2M : MCP_OAUTH2_FLOW_INTERACTIVE, + } + : {}), static_headers: staticHeaders, env_vars: envVars, ...(tokenValidation !== null && { token_validation: tokenValidation }), diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx index 8da77120c09..44b2ba25f11 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx @@ -1003,3 +1003,59 @@ describe("MCPServerEdit (OAuth token persistence on save)", () => { expect(mockSetToken).not.toHaveBeenCalled(); }); }); + +describe("MCPServerEdit oauth2_flow preservation", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + async function saveAndGetPayload(server: Record) { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ ...interactiveOAuthServer }); + + render( + , + ); + + const saveButtons = screen.getAllByRole("button", { name: "Save Changes" }); + await act(async () => { + fireEvent.click(saveButtons[0]); + }); + + await waitFor(() => { + expect(networking.updateMCPServer).toHaveBeenCalledTimes(1); + }); + + const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0]; + return payload; + } + + it("never writes oauth2_flow for a legacy null-flow server with a token_url", async () => { + const payload = await saveAndGetPayload({ + token_url: "https://idp.example.com/oauth/token", + oauth2_flow: null, + }); + expect(payload).not.toHaveProperty("oauth2_flow"); + }); + + it("never writes oauth2_flow over an explicit client_credentials row", async () => { + const payload = await saveAndGetPayload({ + oauth2_flow: "client_credentials", + token_url: "https://idp.example.com/oauth/token", + }); + expect(payload).not.toHaveProperty("oauth2_flow"); + }); + + it("never writes oauth2_flow over the DCR authorization_code stamp", async () => { + const payload = await saveAndGetPayload({ + oauth2_flow: "authorization_code", + token_url: "https://idp.example.com/oauth/token", + }); + expect(payload).not.toHaveProperty("oauth2_flow"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx index 86e898809e5..bc9c3cfea07 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx @@ -219,7 +219,7 @@ const MCPServerEdit: React.FC = ({ static_headers: initialStaticHeaders, env_vars: initialEnvVars, extra_headers: mcpServer.extra_headers || [], - oauth_flow_type: mcpServer.token_url ? OAUTH_FLOW.M2M : OAUTH_FLOW.INTERACTIVE, + oauth_flow_type: mcpServer.oauth2_flow === MCP_OAUTH2_FLOW_M2M ? OAUTH_FLOW.M2M : OAUTH_FLOW.INTERACTIVE, token_validation_json: mcpServer.token_validation ? JSON.stringify(mcpServer.token_validation, null, 2) : undefined, @@ -1246,7 +1246,9 @@ const MCPServerEdit: React.FC = ({ transport: transportType ?? mcpServer.transport, auth_type: currentAuthType ?? mcpServer.auth_type, mcp_info: mcpServer.mcp_info, - oauth_flow_type: currentTokenUrl ?? mcpServer.token_url ? OAUTH_FLOW.M2M : OAUTH_FLOW.INTERACTIVE, + oauth_flow_type: + oauthFlowTypeValue ?? + (mcpServer.oauth2_flow === MCP_OAUTH2_FLOW_M2M ? OAUTH_FLOW.M2M : OAUTH_FLOW.INTERACTIVE), static_headers: currentStaticHeaders ?? mcpServer.static_headers, credentials: currentCredentials, authorization_url: currentAuthorizationUrl ?? mcpServer.authorization_url, diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index cd3ffcab5ec..ebf7c919b48 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -51,6 +51,8 @@ export const OAUTH_FLOW = { // from the UI-local OAUTH_FLOW.M2M ("m2m"); this is what the API actually returns. export const MCP_OAUTH2_FLOW_M2M = "client_credentials"; +export const MCP_OAUTH2_FLOW_INTERACTIVE = "authorization_code"; + export type McpOAuthMode = "m2m" | "passthrough" | "obo"; // Classify an OAuth2 MCP server into the mode that decides how the tool list is