feat(mcp): persist oauth2_flow explicitly on create instead of inferring it at read time (#32288)

* feat(mcp): persist oauth2_flow explicitly on create instead of inferring it at read time

The UI create payload never carried oauth2_flow, so every UI-created oauth2 server
persisted a null flow and relied on _resolve_oauth2_flow's field-shape inference at
registry build. That inference cannot tell a DCR-registered interactive server
(client creds + token_url, no persisted authorization_url) from an M2M server unless
endpoint discovery succeeds first, and the dashboard cannot reproduce it at all
because credentials are redacted in responses

The create form now persists the selected flow for oauth2 servers: authorization_code
for Interactive (PKCE), client_credentials for M2M. The REST create endpoints stamp an
omitted oauth2_flow server-side with the same discriminator the legacy inference uses,
run at write time where the payload carries plaintext credentials, so the decision is
made once with full information and stored. Applied to the admin create, the BYOM
submission, and the temporary session-server endpoints

The edit form derives its flow display from oauth2_flow instead of token_url presence
(token_url is present on authorization_code servers too, so it cannot distinguish M2M)
and deliberately never writes oauth2_flow: it has no flow selector, so a write from
edit could only erase an explicit value, including the authorization_code stamp the
DCR flow persists. Regression tests pin all of this down

Second step of persisting oauth2_flow at every write site so the legacy inference can
eventually be deleted; the backfill for existing null rows lands next

* refactor(mcp): name the create-time flow stamp for its fallback-only contract

stamp_omitted_oauth2_flow with a dedicated explicit-value early return and the shape
check renamed to has_m2m_shape, so the precedence (caller's oauth2_flow always wins,
inference only fills an omitted field) reads directly off the code
This commit is contained in:
tin-berri 2026-07-06 17:53:17 -07:00 • committed by GitHub
parent 2f0cdb35bf
commit 76eeaf2381
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 256 additions and 2 deletions

View file

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

View file

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

View file

@ -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(<CreateMCPServer {...defaultProps} />);
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();
});
});

View file

@ -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<CreateMCPServerProps> = ({
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 }),

View file

@ -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<string, unknown>) {
vi.mocked(networking.updateMCPServer).mockResolvedValue({ ...interactiveOAuthServer });
render(
<MCPServerEdit
mcpServer={{ ...interactiveOAuthServer, ...server }}
accessToken="access-token"
onCancel={vi.fn()}
onSuccess={vi.fn()}
availableAccessGroups={[]}
/>,
);
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");
});
});

View file

@ -219,7 +219,7 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
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<MCPServerEditProps> = ({
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,

View file

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