mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
2f0cdb35bf
commit
76eeaf2381
7 changed files with 256 additions and 2 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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 }),
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue