Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -0,0 +1,18 @@
-- Connector OAuth authorizations in flight. They used to live in an in-memory
-- map, which lost every pending consent on a restart or blue/green deploy and
-- never recorded more than which user started the flow. See the model comment.
CREATE TABLE "connector_oauth_attempts" (
"id" TEXT NOT NULL,
"state_hash" TEXT NOT NULL,
"user_id" TEXT NOT NULL,
"connector_id" TEXT NOT NULL,
"payload" TEXT NOT NULL,
"return_to" TEXT,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"expires_at" TIMESTAMP(3) NOT NULL,

CONSTRAINT "connector_oauth_attempts_pkey" PRIMARY KEY ("id")
);

CREATE UNIQUE INDEX "connector_oauth_attempts_state_hash_key" ON "connector_oauth_attempts"("state_hash");
CREATE INDEX "connector_oauth_attempts_expires_at_idx" ON "connector_oauth_attempts"("expires_at");
26 changes: 26 additions & 0 deletions packages/backend/prisma/schema.prisma
Original file line number Diff line number Diff line change
Expand Up @@ -991,6 +991,32 @@ model UserIdentity {
/// Everything security-relevant lives here rather than in a cookie that has to
/// survive the hop to the IdP and back: the CSRF `state`, the replay-guarding
/// `nonce`, and the PKCE verifier. Consumed atomically exactly once.
/// A connector OAuth authorization in flight ("Authorize with Provider"):
/// created when the user starts it, consumed once when the code comes back.
///
/// Persisted rather than held in memory so a restart or a blue/green deploy
/// between consent and callback does not lose it. The state is stored as a
/// SHA-256 hash and the rest (PKCE verifier, client credentials, signing
/// settings) encrypted with ENCRYPTION_KEY. `userId` is who started it: the
/// code is only exchanged in an authenticated request by that same user, so a
/// consent link started on someone else's connector cannot be completed by a
/// victim on the attacker's behalf.
model ConnectorOAuthAttempt {
id String @id @default(cuid())
stateHash String @unique @map("state_hash")
userId String @map("user_id")
connectorId String @map("connector_id")
/// Encrypted JSON of the pending flow.
payload String
/// Internal path to land on afterwards; validated when written and when read.
returnTo String? @map("return_to")
createdAt DateTime @default(now()) @map("created_at")
expiresAt DateTime @map("expires_at")

@@index([expiresAt])
@@map("connector_oauth_attempts")
}

model SsoLoginAttempt {
id String @id @default(cuid())

Expand Down
28 changes: 24 additions & 4 deletions packages/backend/src/connectors/connectors.controller.ts
Original file line number Diff line number Diff line change
Expand Up @@ -1096,6 +1096,22 @@ export class ConnectorsController {
return this.connectorsService.testConnection(id);
}

@Get('oauth/redirect-uri')
@ApiOperation({
summary: 'The OAuth redirect URI to register in a provider app',
description:
'Computed by the server from SERVER_URL, so it matches what the authorization request sends.',
})
oauthRedirectUri() {
return { redirectUri: this.oauthCallbackUrl() };
}

/** Where providers send the browser back; must equal what the user registered. */
private oauthCallbackUrl(): string {
const base = (this.configService.get<string>('SERVER_URL') || 'http://localhost:4000').replace(/\/+$/, '');
return `${base}/api/mcp-oauth/callback`;
}

@Post(':id/oauth/authorize')
@ApiOperation({
summary: 'Initiate OAuth2 authorization for a connector',
Expand All @@ -1104,7 +1120,11 @@ export class ConnectorsController {
'For REST/GraphQL connectors: uses authorizationUrl and tokenUrl from authConfig. ' +
'Returns an authorization URL for the user to visit.',
})
async initiateOAuth(@Req() req: any, @Param('id') id: string) {
async initiateOAuth(
@Req() req: any,
@Param('id') id: string,
@Body() body?: { returnTo?: string },
) {
const connector = await this.connectorsService.findById(id);
this.assertCanWrite(connector, req);

Expand All @@ -1113,7 +1133,7 @@ export class ConnectorsController {
}

try {
const callbackUrl = `${this.configService.get('SERVER_URL') || 'http://localhost:4000'}/api/mcp-oauth/callback`;
const callbackUrl = this.oauthCallbackUrl();
const authConfig = connector.authConfig
? JSON.parse(decrypt(connector.authConfig, this.encryptionKey))
: {};
Expand Down Expand Up @@ -1200,7 +1220,7 @@ export class ConnectorsController {
const state = this.mcpOAuthService.generateState();

// Store pending flow
this.mcpOAuthService.storePendingFlow(state, {
await this.mcpOAuthService.storePendingFlow(state, {
codeVerifier,
connectorId: connector.id,
userId: req.user.sub,
Expand All @@ -1212,7 +1232,7 @@ export class ConnectorsController {
clientAssertion,
persistAuthConfig,
createdAt: Date.now(),
});
}, { returnTo: body?.returnTo });

// Build authorization URL
const authorizationUrl = this.mcpOAuthService.buildAuthorizationUrl({
Expand Down
135 changes: 100 additions & 35 deletions packages/backend/src/connectors/mcp-oauth-callback.controller.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -11,28 +11,32 @@ function makeController(overrides: {
connectorType?: string;
remoteTools?: Array<{ name: string }>;
flow?: Record<string, unknown>;
noFlow?: boolean;
returnTo?: string;
} = {}) {
const reloadConnectorTools = jest.fn().mockResolvedValue(undefined);
const updateAuthConfigMerge = jest.fn().mockResolvedValue(undefined);
const deletePendingFlow = jest.fn();
const flow = {
connectorId: 'conn-1',
userId: 'user-1',
tokenUrl: 'https://sandbox-api.datev.de/token',
redirectUri: 'https://cloud.example.com/api/mcp-oauth/callback',
clientId: 'cid',
clientSecret: 'sec',
codeVerifier: 'verifier',
tokenAuthMethod: 'basic',
...overrides.flow,
};
const record = overrides.noFlow ? undefined : { flow, returnTo: overrides.returnTo };

const mcpOAuthService: any = {
getPendingFlow: jest.fn().mockReturnValue({
connectorId: 'conn-1',
tokenUrl: 'https://sandbox-api.datev.de/token',
redirectUri: 'https://cloud.example.com/api/mcp-oauth/callback',
clientId: 'cid',
clientSecret: 'sec',
codeVerifier: 'verifier',
tokenAuthMethod: 'basic',
...overrides.flow,
}),
getPendingFlow: jest.fn().mockResolvedValue(record),
takePendingFlow: jest.fn().mockResolvedValue(record),
exchangeCodeForTokens: jest.fn().mockResolvedValue({
accessToken: 'AT',
refreshToken: 'RT',
expiresIn: 3600,
}),
deletePendingFlow,
};
const connectorsService: any = {
updateAuthConfigMerge,
Expand Down Expand Up @@ -73,13 +77,88 @@ function makeRes() {
return { redirect: jest.fn() } as any;
}

describe('McpOAuthCallbackController', () => {
const asUser = (sub: string) => ({ user: { sub } });

describe('McpOAuthCallbackController — provider redirect', () => {
it('forwards code and state to the dashboard and exchanges nothing itself', async () => {
const { controller, mcpOAuthService, updateAuthConfigMerge } = makeController();
const res = makeRes();
await controller.oauthCallback('the-code', 'the-state', undefined, undefined, res);
expect(res.redirect).toHaveBeenCalledWith(
'https://cloud.example.com/connectors/oauth/complete?state=the-state&code=the-code',
);
expect(mcpOAuthService.exchangeCodeForTokens).not.toHaveBeenCalled();
expect(mcpOAuthService.takePendingFlow).not.toHaveBeenCalled();
expect(updateAuthConfigMerge).not.toHaveBeenCalled();
});

it('does not forward a state it never issued', async () => {
const { controller } = makeController({ noFlow: true });
const res = makeRes();
await controller.oauthCallback('the-code', 'forged', undefined, undefined, res);
expect(res.redirect.mock.calls[0][0]).toMatch(/complete\?error=/);
expect(res.redirect.mock.calls[0][0]).not.toContain('code=');
});

it('redirects with an error when code/state are missing', async () => {
const { controller, reloadConnectorTools } = makeController();
const res = makeRes();
await controller.oauthCallback('', '', undefined, undefined, res);
expect(reloadConnectorTools).not.toHaveBeenCalled();
expect(res.redirect).toHaveBeenCalledWith(expect.stringContaining('error='));
});

it('treats repeated parameters as missing, never as arrays', async () => {
const { controller, mcpOAuthService } = makeController();
const res = makeRes();
await controller.oauthCallback(['a', 'b'], ['s1', 's2'], ['access_denied', 'x'], ['d'], res);
expect(mcpOAuthService.takePendingFlow).not.toHaveBeenCalled();
expect(mcpOAuthService.getPendingFlow).not.toHaveBeenCalled();
expect(res.redirect.mock.calls[0][0]).toMatch(/complete\?error=/);
expect(res.redirect.mock.calls[0][0]).not.toContain('code=');
});

it('spends the attempt and explains a refusal at the provider', async () => {
const { controller, mcpOAuthService } = makeController();
const res = makeRes();
await controller.oauthCallback('', 'the-state', 'access_denied', 'User said no', res);
expect(mcpOAuthService.takePendingFlow).toHaveBeenCalledWith('the-state');
const url = new URL(res.redirect.mock.calls[0][0]);
expect(url.searchParams.get('error')).toMatch(/cancelled at the provider/);
expect(url.searchParams.get('connectorId')).toBe('conn-1');
});
});

describe('McpOAuthCallbackController — completion by the dashboard', () => {
it('refuses a user other than the one who started the flow, and kills the attempt', async () => {
const { controller, mcpOAuthService, updateAuthConfigMerge } = makeController();
await expect(
controller.complete(asUser('attacker'), { state: 'the-state', code: 'the-code' }),
).rejects.toThrow(/started by another account/);
expect(mcpOAuthService.takePendingFlow).toHaveBeenCalledWith('the-state');
expect(mcpOAuthService.exchangeCodeForTokens).not.toHaveBeenCalled();
expect(updateAuthConfigMerge).not.toHaveBeenCalled();
});

it('answers 410 for an expired or reused state', async () => {
const { controller } = makeController({ noFlow: true });
await expect(
controller.complete(asUser('user-1'), { state: 'the-state', code: 'the-code' }),
).rejects.toThrow(/expired or was already used/);
});

it('returns where the dashboard should land', async () => {
const { controller } = makeController({ returnTo: '/connectors/setup/etsy?step=done' });
await expect(
controller.complete(asUser('user-1'), { state: 'the-state', code: 'the-code' }),
).resolves.toEqual({ connectorId: 'conn-1', toolsImported: 0, returnTo: '/connectors/setup/etsy?step=done' });
});

it('reloads connector tools after storing the token even when MCP discovery throws (REST connector)', async () => {
const { controller, reloadConnectorTools, updateAuthConfigMerge } =
makeController({ listToolsThrows: true });
const res = makeRes();

await controller.oauthCallback('the-code', 'the-state', res);
await controller.complete(asUser('user-1'), { state: 'the-state', code: 'the-code' });

// Token was persisted via a MERGE (preserves authorizationUrl/scopes)...
expect(updateAuthConfigMerge).toHaveBeenCalledWith(
Expand All @@ -88,10 +167,6 @@ describe('McpOAuthCallbackController', () => {
);
// ...and the registry was reloaded despite discovery throwing.
expect(reloadConnectorTools).toHaveBeenCalledWith('conn-1');
// Redirects to success.
expect(res.redirect).toHaveBeenCalledWith(
expect.stringContaining('oauth=success'),
);
});

it('does not import MCP tools into a REST connector whose host also speaks MCP', async () => {
Expand All @@ -102,35 +177,25 @@ describe('McpOAuthCallbackController', () => {
connectorType: 'REST',
remoteTools: [{ name: 'get' }, { name: 'query' }, { name: 'sites_list' }],
});
const res = makeRes();

await controller.oauthCallback('the-code', 'the-state', res);
const out = await controller.complete(asUser('user-1'), { state: 'the-state', code: 'the-code' });

expect(mcpClientEngine.listTools).not.toHaveBeenCalled();
expect(prisma.mcpTool.create).not.toHaveBeenCalled();
expect(res.redirect).toHaveBeenCalledWith(expect.stringContaining('tools=0'));
expect(out.toolsImported).toBe(0);
});

it('still discovers tools for an MCP connector', async () => {
const { controller, mcpClientEngine, prisma } = makeController({
connectorType: 'MCP',
remoteTools: [{ name: 'search' }, { name: 'fetch' }],
});
const res = makeRes();

await controller.oauthCallback('the-code', 'the-state', res);
const out = await controller.complete(asUser('user-1'), { state: 'the-state', code: 'the-code' });

expect(mcpClientEngine.listTools).toHaveBeenCalled();
expect(prisma.mcpTool.create).toHaveBeenCalledTimes(2);
expect(res.redirect).toHaveBeenCalledWith(expect.stringContaining('tools=2'));
});

it('redirects with an error when code/state are missing', async () => {
const { controller, reloadConnectorTools } = makeController();
const res = makeRes();
await controller.oauthCallback('', '', res);
expect(reloadConnectorTools).not.toHaveBeenCalled();
expect(res.redirect).toHaveBeenCalledWith(expect.stringContaining('error='));
expect(out.toolsImported).toBe(2);
});

it('writes what the flow took from the catalog next to the tokens, and not the resolved client', async () => {
Expand All @@ -150,7 +215,7 @@ describe('McpOAuthCallbackController', () => {
},
});

await controller.oauthCallback('the-code', 'the-state', makeRes());
await controller.complete(asUser('user-1'), { state: 'the-state', code: 'the-code' });

const patch = updateAuthConfigMerge.mock.calls[0][1];
expect(patch).toMatchObject({
Expand All @@ -166,7 +231,7 @@ describe('McpOAuthCallbackController', () => {

it('still writes the client settings when the flow does not say otherwise (MCP)', async () => {
const { controller, updateAuthConfigMerge } = makeController({ connectorType: 'MCP' });
await controller.oauthCallback('the-code', 'the-state', makeRes());
await controller.complete(asUser('user-1'), { state: 'the-state', code: 'the-code' });
expect(updateAuthConfigMerge.mock.calls[0][1]).toMatchObject({
clientId: 'cid',
clientSecret: 'sec',
Expand Down
Loading
Loading