diff --git a/packages/backend/prisma/migrations/20261003100000_connector_oauth_attempts/migration.sql b/packages/backend/prisma/migrations/20261003100000_connector_oauth_attempts/migration.sql new file mode 100644 index 00000000..4cc79f05 --- /dev/null +++ b/packages/backend/prisma/migrations/20261003100000_connector_oauth_attempts/migration.sql @@ -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"); diff --git a/packages/backend/prisma/schema.prisma b/packages/backend/prisma/schema.prisma index aa494f63..bba2b325 100644 --- a/packages/backend/prisma/schema.prisma +++ b/packages/backend/prisma/schema.prisma @@ -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()) diff --git a/packages/backend/src/connectors/connectors.controller.ts b/packages/backend/src/connectors/connectors.controller.ts index ae9dafd9..771af71e 100644 --- a/packages/backend/src/connectors/connectors.controller.ts +++ b/packages/backend/src/connectors/connectors.controller.ts @@ -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('SERVER_URL') || 'http://localhost:4000').replace(/\/+$/, ''); + return `${base}/api/mcp-oauth/callback`; + } + @Post(':id/oauth/authorize') @ApiOperation({ summary: 'Initiate OAuth2 authorization for a connector', @@ -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); @@ -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)) : {}; @@ -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, @@ -1212,7 +1232,7 @@ export class ConnectorsController { clientAssertion, persistAuthConfig, createdAt: Date.now(), - }); + }, { returnTo: body?.returnTo }); // Build authorization URL const authorizationUrl = this.mcpOAuthService.buildAuthorizationUrl({ diff --git a/packages/backend/src/connectors/mcp-oauth-callback.controller.spec.ts b/packages/backend/src/connectors/mcp-oauth-callback.controller.spec.ts index 2d91c5ec..a79e44d0 100644 --- a/packages/backend/src/connectors/mcp-oauth-callback.controller.spec.ts +++ b/packages/backend/src/connectors/mcp-oauth-callback.controller.spec.ts @@ -11,28 +11,32 @@ function makeController(overrides: { connectorType?: string; remoteTools?: Array<{ name: string }>; flow?: Record; + 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, @@ -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( @@ -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 () => { @@ -102,13 +177,12 @@ 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 () => { @@ -116,21 +190,12 @@ describe('McpOAuthCallbackController', () => { 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 () => { @@ -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({ @@ -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', diff --git a/packages/backend/src/connectors/mcp-oauth-callback.controller.ts b/packages/backend/src/connectors/mcp-oauth-callback.controller.ts index 5438b68d..d6de7725 100644 --- a/packages/backend/src/connectors/mcp-oauth-callback.controller.ts +++ b/packages/backend/src/connectors/mcp-oauth-callback.controller.ts @@ -1,16 +1,38 @@ -import { Controller, Get, Query, Res, Logger } from '@nestjs/common'; +import { + Body, + Controller, + ForbiddenException, + Get, + GoneException, + HttpCode, + Logger, + Post, + Query, + Req, + Res, + UseGuards, +} from '@nestjs/common'; +import { AuthGuard } from '@nestjs/passport'; import { ApiTags, ApiOperation } from '@nestjs/swagger'; import { ConfigService } from '@nestjs/config'; import type { Response } from 'express'; -import { McpOAuthService } from './mcp-oauth.service'; +import { McpOAuthService, PendingOAuthFlow } from './mcp-oauth.service'; import { ConnectorsService } from './connectors.service'; import { McpClientEngine } from './engines/mcp-client.engine'; import { PrismaService } from '../common/prisma.service'; import { McpServerService } from '../mcp-server/mcp-server.service'; /** - * Separate controller for the OAuth2 callback — no JWT guard. - * The remote MCP server redirects the user's browser here after login. + * Connector OAuth: where the provider sends the browser back, and where the + * dashboard completes the authorization. + * + * The two steps are split on purpose. The provider's redirect carries no + * proof of who is at the keyboard, so the callback only checks that the state + * is one we issued and forwards code + state to the dashboard. The dashboard + * then posts them in an authenticated request, and the code is exchanged only + * if that user is the one who started the flow. Exchanging in the callback, + * as before, let anyone who started an authorization on their own connector + * send the consent link to someone else and receive that person's tokens. */ @ApiTags('MCP OAuth') @Controller('api/mcp-oauth') @@ -26,125 +48,170 @@ export class McpOAuthCallbackController { private readonly configService: ConfigService, ) {} + private frontendUrl(): string { + return ( + this.configService.get('FRONTEND_URL') || 'http://localhost:3000' + ).replace(/\/+$/, ''); + } + + private completePage(params: Record): string { + return `${this.frontendUrl()}/connectors/oauth/complete?${new URLSearchParams(params).toString()}`; + } + @Get('callback') @ApiOperation({ - summary: 'OAuth2 callback handler for MCP connector authorization', + summary: 'OAuth2 redirect target for connector authorization', description: - 'Handles the redirect from a remote MCP server after user authorization. ' + - 'Exchanges the auth code for tokens and auto-discovers MCP tools.', + 'Checks the state and forwards the code to the dashboard, which completes ' + + 'the authorization with POST /api/mcp-oauth/complete.', }) async oauthCallback( - @Query('code') code: string, - @Query('state') state: string, + @Query('code') rawCode: unknown, + @Query('state') rawState: unknown, + @Query('error') rawError: unknown, + @Query('error_description') rawErrorDescription: unknown, @Res() res: Response, ) { - const frontendUrl = - this.configService.get('FRONTEND_URL') || 'http://localhost:3000'; + // A repeated parameter (?state=a&state=b) arrives as an array: take + // strings only, so nothing below is fed a value of the wrong type. + const code = singleQueryValue(rawCode); + const state = singleQueryValue(rawState); + const providerError = singleQueryValue(rawError); + const providerErrorDescription = singleQueryValue(rawErrorDescription); + if (providerError) { + // The user declined, or the provider refused the request. The attempt + // is spent either way. + const record = state ? await this.mcpOAuthService.takePendingFlow(state) : undefined; + return res.redirect( + this.completePage({ + error: describeProviderError(providerError, providerErrorDescription), + ...(record ? { connectorId: record.flow.connectorId } : {}), + }), + ); + } if (!code || !state) { return res.redirect( - `${frontendUrl}/connectors?error=${encodeURIComponent('Missing code or state in OAuth callback')}`, + this.completePage({ error: 'The provider came back without an authorization code. Start the authorization again.' }), ); } - const flow = this.mcpOAuthService.getPendingFlow(state); - if (!flow) { - this.logger.warn(`OAuth callback with unknown state: ${state}`); + const record = await this.mcpOAuthService.getPendingFlow(state); + if (!record) { + this.logger.warn('OAuth callback with an unknown or expired state'); return res.redirect( - `${frontendUrl}/connectors?error=${encodeURIComponent('OAuth session expired or invalid state')}`, + this.completePage({ error: 'This authorization expired or was already used. Start it again from the connector.' }), ); } - try { - // 1. Exchange auth code for tokens - const tokens = await this.mcpOAuthService.exchangeCodeForTokens({ - tokenUrl: flow.tokenUrl, - code, - redirectUri: flow.redirectUri, - clientId: flow.clientId, - clientSecret: flow.clientSecret, - codeVerifier: flow.codeVerifier, - tokenAuthMethod: flow.tokenAuthMethod, - clientAssertion: flow.clientAssertion, - }); - - this.logger.log( - `OAuth tokens obtained for connector ${flow.connectorId}`, + return res.redirect(this.completePage({ state, code })); + } + + @Post('complete') + @UseGuards(AuthGuard('jwt')) + @HttpCode(200) + @ApiOperation({ + summary: 'Complete a connector authorization (dashboard, authenticated)', + description: + 'Exchanges the authorization code, but only for the user who started the flow.', + }) + async complete( + @Req() req: any, + @Body() body: { state?: string; code?: string }, + ): Promise<{ connectorId: string; toolsImported: number; returnTo?: string }> { + const record = await this.mcpOAuthService.takePendingFlow(String(body?.state || '')); + if (!record) { + throw new GoneException('This authorization expired or was already used. Start it again from the connector.'); + } + const { flow, returnTo } = record; + if (flow.userId !== req.user?.sub) { + // Consumed above on purpose: a link someone else started is dead now. + this.logger.warn( + `Refused OAuth completion: connector ${flow.connectorId} was authorized by another user`, + ); + throw new ForbiddenException( + 'This authorization was started by another account. Sign in as that account, or start it again yourself.', ); + } + if (!body?.code) { + throw new GoneException('The provider came back without an authorization code. Start the authorization again.'); + } - // 2. Store tokens (encrypted) in the connector's authConfig. Merge, don't - // replace — preserves static config (authorizationUrl, scopes) needed for - // later re-authorization. - const clientSettings = flow.persistAuthConfig ?? { - tokenUrl: flow.tokenUrl, - clientId: flow.clientId, - clientSecret: flow.clientSecret, - tokenAuthMethod: flow.tokenAuthMethod, - }; - await this.connectorsService.updateAuthConfigMerge(flow.connectorId, { - ...clientSettings, - accessToken: tokens.accessToken, - refreshToken: tokens.refreshToken, - expiresIn: tokens.expiresIn, - expiresAt: Date.now() + (tokens.expiresIn || 3600) * 1000, - authorizedAt: new Date().toISOString(), - }); - - // Reload the connector's tools into the in-memory MCP registry so the - // freshly-stored access token takes effect immediately. The registry - // caches a snapshot of authConfig (incl. the token) per tool, so without - // this a just-authorized connector would keep serving with the stale - // (token-less) snapshot. For REST/GraphQL OAuth connectors this is the - // ONLY reload — the MCP auto-discovery block below throws for non-MCP - // servers and never reaches its own reloadConnectorTools() call. - try { - await this.mcpServer.reloadConnectorTools(flow.connectorId); - } catch (reloadErr: any) { - this.logger.warn( - `Failed to reload tools after OAuth for connector ${flow.connectorId}: ${reloadErr.message}`, - ); - } + const toolsImported = await this.exchangeAndStore(flow, String(body.code)); + return { connectorId: flow.connectorId, toolsImported, ...(returnTo ? { returnTo } : {}) }; + } - // 3. Auto-discover tools from the remote MCP server (MCP connectors only) - let toolsImported = 0; - try { - const connector = await this.connectorsService.findByIdInternal( - flow.connectorId, - ); + /** Exchange the code, store the tokens, reload the tools. Throws on failure. */ + private async exchangeAndStore(flow: PendingOAuthFlow, code: string): Promise { + // 1. Exchange auth code for tokens + const tokens = await this.mcpOAuthService.exchangeCodeForTokens({ + tokenUrl: flow.tokenUrl, + code, + redirectUri: flow.redirectUri, + clientId: flow.clientId, + clientSecret: flow.clientSecret, + codeVerifier: flow.codeVerifier, + tokenAuthMethod: flow.tokenAuthMethod, + clientAssertion: flow.clientAssertion, + }); - // A REST or GraphQL connector already has its tools. Discovery used - // to run for them too and relied on the host not speaking MCP; Google - // does (searchconsole.googleapis.com), so authorising the Search - // Console connector added three MCP tools mapped as REST calls. - if (connector.type === 'MCP') { - toolsImported = await this.importRemoteTools( - flow.connectorId, - connector, - tokens.accessToken, - ); - } - } catch (discoverErr: any) { - this.logger.warn( - `Tool discovery failed after OAuth (will proceed anyway): ${discoverErr.message}`, - ); - } + this.logger.log(`OAuth tokens obtained for connector ${flow.connectorId}`); - // 4. Clean up - this.mcpOAuthService.deletePendingFlow(state); + // 2. Store tokens (encrypted) in the connector's authConfig. Merge, don't + // replace — preserves static config (authorizationUrl, scopes) needed for + // later re-authorization. + const clientSettings = flow.persistAuthConfig ?? { + tokenUrl: flow.tokenUrl, + clientId: flow.clientId, + clientSecret: flow.clientSecret, + tokenAuthMethod: flow.tokenAuthMethod, + }; + await this.connectorsService.updateAuthConfigMerge(flow.connectorId, { + ...clientSettings, + accessToken: tokens.accessToken, + refreshToken: tokens.refreshToken, + expiresIn: tokens.expiresIn, + expiresAt: Date.now() + (tokens.expiresIn || 3600) * 1000, + authorizedAt: new Date().toISOString(), + }); - // 5. Redirect to frontend - return res.redirect( - `${frontendUrl}/connectors/${flow.connectorId}?oauth=success&tools=${toolsImported}`, - ); - } catch (error: any) { - this.logger.error( - `OAuth callback failed for connector ${flow.connectorId}: ${error.message}`, + // Reload the connector's tools into the in-memory MCP registry so the + // freshly-stored access token takes effect immediately. The registry + // caches a snapshot of authConfig (incl. the token) per tool, so without + // this a just-authorized connector would keep serving with the stale + // (token-less) snapshot. For REST/GraphQL OAuth connectors this is the + // ONLY reload — the MCP auto-discovery block below throws for non-MCP + // servers and never reaches its own reloadConnectorTools() call. + try { + await this.mcpServer.reloadConnectorTools(flow.connectorId); + } catch (reloadErr: any) { + this.logger.warn( + `Failed to reload tools after OAuth for connector ${flow.connectorId}: ${reloadErr.message}`, ); - this.mcpOAuthService.deletePendingFlow(state); - return res.redirect( - `${frontendUrl}/connectors/${flow.connectorId}?oauth=error&message=${encodeURIComponent(error.message)}`, + } + + // 3. Auto-discover tools from the remote MCP server (MCP connectors only) + let toolsImported = 0; + try { + const connector = await this.connectorsService.findByIdInternal(flow.connectorId); + + // A REST or GraphQL connector already has its tools. Discovery used + // to run for them too and relied on the host not speaking MCP; Google + // does (searchconsole.googleapis.com), so authorising the Search + // Console connector added three MCP tools mapped as REST calls. + if (connector.type === 'MCP') { + toolsImported = await this.importRemoteTools( + flow.connectorId, + connector, + tokens.accessToken, + ); + } + } catch (discoverErr: any) { + this.logger.warn( + `Tool discovery failed after OAuth (will proceed anyway): ${discoverErr.message}`, ); } + return toolsImported; } /** Import the tools a remote MCP server lists, skipping ones already present. */ @@ -200,3 +267,22 @@ export class McpOAuthCallbackController { return toolsImported; } } + +function singleQueryValue(value: unknown): string | undefined { + return typeof value === 'string' ? value : undefined; +} + +/** The provider's refusal in words a user can act on. */ +function describeProviderError(code: string, description?: string): string { + const detail = description ? ` (${description.slice(0, 200)})` : ''; + if (code === 'access_denied') { + return `The authorization was cancelled at the provider${detail}. Nothing was changed; start it again when you are ready.`; + } + if (code === 'invalid_scope') { + return `The provider refused the requested permissions${detail}. Check the scopes in the connector's OAuth settings.`; + } + if (code === 'unauthorized_client' || code === 'invalid_client') { + return `The provider does not accept this app${detail}. Check the client ID and that the redirect URI shown in AnythingMCP is registered in the app.`; + } + return `The provider returned an error: ${code.slice(0, 80)}${detail}.`; +} diff --git a/packages/backend/src/connectors/mcp-oauth.service.spec.ts b/packages/backend/src/connectors/mcp-oauth.service.spec.ts index 76c7a2c3..c80f590e 100644 --- a/packages/backend/src/connectors/mcp-oauth.service.spec.ts +++ b/packages/backend/src/connectors/mcp-oauth.service.spec.ts @@ -324,7 +324,7 @@ describe('McpOAuthService PKCE', () => { expect(params.get('client_id')).toBe('keystring'); }); - it('keeps the verifier with its state, for ten minutes', () => { + it('keeps the verifier with its state, until it expires or is taken', async () => { const s = new McpOAuthService(); const flow = { codeVerifier: 'v', @@ -335,10 +335,91 @@ describe('McpOAuthService PKCE', () => { tokenUrl: 't', createdAt: Date.now(), }; - s.storePendingFlow('state-a', flow); - expect(s.getPendingFlow('state-a')?.codeVerifier).toBe('v'); - expect(s.getPendingFlow('state-b')).toBeUndefined(); - s.storePendingFlow('state-old', { ...flow, createdAt: Date.now() - 11 * 60 * 1000 }); - expect(s.getPendingFlow('state-old')).toBeUndefined(); + await s.storePendingFlow('state-a', flow, { returnTo: '/connectors/c' }); + expect((await s.getPendingFlow('state-a'))?.flow.codeVerifier).toBe('v'); + expect((await s.getPendingFlow('state-a'))?.returnTo).toBe('/connectors/c'); + expect(await s.getPendingFlow('state-b')).toBeUndefined(); + // Taking it consumes it. + expect((await s.takePendingFlow('state-a'))?.flow.codeVerifier).toBe('v'); + expect(await s.takePendingFlow('state-a')).toBeUndefined(); + }); + +}); + +describe('McpOAuthService pending flows in the database', () => { + const flow = { + codeVerifier: 'the-verifier', + connectorId: 'conn-1', + userId: 'user-1', + redirectUri: 'https://cloud.example.com/api/mcp-oauth/callback', + clientId: 'cid', + clientSecret: 'very-secret', + tokenUrl: 'https://example.com/token', + createdAt: Date.now(), + }; + const saved = process.env.ENCRYPTION_KEY; + beforeAll(() => { + process.env.ENCRYPTION_KEY = 'k'.repeat(32); + }); + afterAll(() => { + if (saved === undefined) delete process.env.ENCRYPTION_KEY; + else process.env.ENCRYPTION_KEY = saved; + }); + + function fakePrisma() { + const rows = new Map(); + return { + rows, + connectorOAuthAttempt: { + deleteMany: jest.fn().mockResolvedValue({ count: 0 }), + create: jest.fn(async ({ data }: any) => { + rows.set(data.stateHash, data); + return data; + }), + findUnique: jest.fn(async ({ where }: any) => rows.get(where.stateHash) ?? null), + delete: jest.fn(async ({ where }: any) => { + const row = rows.get(where.stateHash); + if (!row) throw new Error('P2025'); + rows.delete(where.stateHash); + return row; + }), + }, + }; + } + + it('stores the state hashed and the secrets encrypted', async () => { + const prisma = fakePrisma(); + const s = new McpOAuthService(prisma as any); + await s.storePendingFlow('the-state', flow); + const [row] = [...prisma.rows.values()]; + expect(row.stateHash).not.toContain('the-state'); + expect(row.payload).not.toContain('very-secret'); + expect(row.payload).not.toContain('the-verifier'); + expect(row.userId).toBe('user-1'); + expect((await s.getPendingFlow('the-state'))?.flow.clientSecret).toBe('very-secret'); + }); + + it('can be taken only once', async () => { + const s = new McpOAuthService(fakePrisma() as any); + await s.storePendingFlow('the-state', flow); + expect((await s.takePendingFlow('the-state'))?.flow.userId).toBe('user-1'); + expect(await s.takePendingFlow('the-state')).toBeUndefined(); + }); + + it('ignores an expired attempt', async () => { + const prisma = fakePrisma(); + const s = new McpOAuthService(prisma as any); + await s.storePendingFlow('the-state', flow); + for (const row of prisma.rows.values()) row.expiresAt = new Date(Date.now() - 1000); + expect(await s.getPendingFlow('the-state')).toBeUndefined(); + expect(await s.takePendingFlow('the-state')).toBeUndefined(); + }); + + it('drops a returnTo that would leave the dashboard', async () => { + const s = new McpOAuthService(fakePrisma() as any); + for (const bad of ['https://evil.example/x', '//evil.example', '/\\evil.example', 'connectors/1']) { + await s.storePendingFlow(`st-${bad}`, flow, { returnTo: bad }); + expect((await s.getPendingFlow(`st-${bad}`))?.returnTo).toBeUndefined(); + } }); }); diff --git a/packages/backend/src/connectors/mcp-oauth.service.ts b/packages/backend/src/connectors/mcp-oauth.service.ts index cb14c1a5..0069e216 100644 --- a/packages/backend/src/connectors/mcp-oauth.service.ts +++ b/packages/backend/src/connectors/mcp-oauth.service.ts @@ -1,5 +1,7 @@ -import { Injectable, Logger } from '@nestjs/common'; +import { Injectable, Logger, Optional } from '@nestjs/common'; import { createHash, randomBytes } from 'crypto'; +import { PrismaService } from '../common/prisma.service'; +import { decrypt, encrypt } from '../common/crypto/encryption.util'; import axios from 'axios'; import { assertSafeOutboundUrl } from '../common/ssrf.util'; import { @@ -18,7 +20,7 @@ interface OAuthMetadata { code_challenge_methods_supported?: string[]; } -interface PendingOAuthFlow { +export interface PendingOAuthFlow { codeVerifier: string; connectorId: string; userId: string; @@ -55,9 +57,13 @@ interface PendingOAuthFlow { export class McpOAuthService { private readonly logger = new Logger(McpOAuthService.name); - // In-memory store for pending OAuth flows, keyed by state. - // Entries auto-expire after 10 minutes. - private pendingFlows = new Map(); + /** + * Pending flows when no database is wired (unit tests). In the application + * they live in `connector_oauth_attempts`, see storePendingFlow. + */ + private memoryFlows = new Map(); + + constructor(@Optional() private readonly prisma?: PrismaService) {} /** * Discover the OAuth metadata of a remote MCP server. @@ -412,33 +418,115 @@ export class McpOAuthService { } // --- Pending Flow Storage --- + // + // An authorization in flight is stored in the database (state hashed, the + // rest encrypted), so it survives a restart or a blue/green deploy between + // consent and callback. It is read twice: peeked by the provider callback, + // which only forwards the code to the dashboard, and taken (deleted) by the + // authenticated request that exchanges the code, after checking that the + // same user started it. + + private stateHash(state: string): string { + return createHash('sha256').update(state).digest('hex'); + } - storePendingFlow(state: string, data: PendingOAuthFlow): void { - // Clean up expired entries (>10 min) - const now = Date.now(); - for (const [key, flow] of this.pendingFlows) { - if (now - flow.createdAt > 10 * 60 * 1000) { - this.pendingFlows.delete(key); - } - } + private encryptionKey(): string { + const key = process.env.ENCRYPTION_KEY; + if (!key) throw new Error('ENCRYPTION_KEY is not set'); + return key; + } - this.pendingFlows.set(state, data); + async storePendingFlow( + state: string, + data: PendingOAuthFlow, + opts: { returnTo?: string } = {}, + ): Promise { + const expiresAt = Date.now() + PENDING_FLOW_TTL_MS; + const returnTo = safeReturnTo(opts.returnTo); + if (!this.prisma) { + this.memoryFlows.set(state, { flow: data, returnTo, expiresAt }); + return; + } + await this.prisma.connectorOAuthAttempt + .deleteMany({ where: { expiresAt: { lt: new Date() } } }) + .catch(() => undefined); + await this.prisma.connectorOAuthAttempt.create({ + data: { + stateHash: this.stateHash(state), + userId: data.userId, + connectorId: data.connectorId, + payload: encrypt(JSON.stringify(data), this.encryptionKey(), 'connector-oauth'), + returnTo: returnTo ?? null, + expiresAt: new Date(expiresAt), + }, + }); } - getPendingFlow(state: string): PendingOAuthFlow | undefined { - const flow = this.pendingFlows.get(state); - if (!flow) return undefined; + /** The pending flow for `state`, without consuming it. */ + async getPendingFlow(state: string): Promise { + if (!state) return undefined; + if (!this.prisma) { + const hit = this.memoryFlows.get(state); + if (!hit || hit.expiresAt < Date.now()) return undefined; + return { flow: hit.flow, returnTo: hit.returnTo }; + } + const row = await this.prisma.connectorOAuthAttempt.findUnique({ + where: { stateHash: this.stateHash(state) }, + }); + if (!row || row.expiresAt.getTime() < Date.now()) return undefined; + return this.toRecord(row); + } - // Check expiry - if (Date.now() - flow.createdAt > 10 * 60 * 1000) { - this.pendingFlows.delete(state); - return undefined; + /** The pending flow for `state`, deleted in the same step: usable once. */ + async takePendingFlow(state: string): Promise { + if (!state) return undefined; + if (!this.prisma) { + const hit = this.memoryFlows.get(state); + this.memoryFlows.delete(state); + if (!hit || hit.expiresAt < Date.now()) return undefined; + return { flow: hit.flow, returnTo: hit.returnTo }; } + const row = await this.prisma.connectorOAuthAttempt + .delete({ where: { stateHash: this.stateHash(state) } }) + .catch(() => null); + if (!row || row.expiresAt.getTime() < Date.now()) return undefined; + return this.toRecord(row); + } - return flow; + async deletePendingFlow(state: string): Promise { + await this.takePendingFlow(state); } - deletePendingFlow(state: string): void { - this.pendingFlows.delete(state); + private toRecord(row: { payload: string; returnTo: string | null }): PendingFlowRecord | undefined { + try { + const flow = JSON.parse( + decrypt(row.payload, this.encryptionKey(), 'connector-oauth'), + ) as PendingOAuthFlow; + return { flow, returnTo: safeReturnTo(row.returnTo ?? undefined) }; + } catch (err: any) { + this.logger.warn(`Unreadable pending OAuth flow: ${err?.message || err}`); + return undefined; + } } } + +/** How long a started authorization waits for the user to come back. */ +export const PENDING_FLOW_TTL_MS = 15 * 60 * 1000; + +export interface PendingFlowRecord { + flow: PendingOAuthFlow; + /** Where the dashboard should land afterwards (internal path), if set. */ + returnTo?: string; +} + +/** + * An internal dashboard path, or undefined. Anything else (an absolute URL, + * `//host`, a backslash trick) would be an open redirect. + */ +export function safeReturnTo(value: string | undefined | null): string | undefined { + if (!value || typeof value !== 'string') return undefined; + if (value.length > 500) return undefined; + if (!value.startsWith('/') || value.startsWith('//') || value.includes('\\')) return undefined; + if ([...value].some((ch) => ch.charCodeAt(0) < 0x20)) return undefined; + return value; +} diff --git a/packages/frontend/src/app/connectors/new/page.tsx b/packages/frontend/src/app/connectors/new/page.tsx index 9ee07187..8a6d7ac2 100644 --- a/packages/frontend/src/app/connectors/new/page.tsx +++ b/packages/frontend/src/app/connectors/new/page.tsx @@ -1,6 +1,6 @@ 'use client'; -import { useState } from 'react'; +import { useEffect, useState } from 'react'; import Link from 'next/link'; import { useRouter } from 'next/navigation'; import { useAuth } from '@/lib/auth-context'; @@ -45,6 +45,15 @@ const TYPE_ICONS: Record = { export default function NewConnectorPage() { const { token } = useAuth(); + // What to register in the provider's app, as the server will send it. + const [oauthRedirectUri, setOauthRedirectUri] = useState(null); + useEffect(() => { + if (!token) return; + connectors + .oauthRedirectUri(token) + .then((r) => setOauthRedirectUri(r.redirectUri)) + .catch(() => setOauthRedirectUri(null)); + }, [token]); const router = useRouter(); const [selectedType, setSelectedType] = useState(null); // OData: SAP Gateway mode (catalog, sap-client) and its settings. @@ -569,7 +578,7 @@ export default function NewConnectorPage() {

After creating the connector, you will be redirected to authorize via OAuth2. Tokens will be stored securely.

- Set the Redirect / Callback URI in your OAuth provider to: {typeof window !== 'undefined' ? (window.location.hostname === 'localhost' ? window.location.origin.replace(':3000', ':4000') : window.location.origin) : 'http://localhost:4000'}/api/mcp-oauth/callback + Set the Redirect / Callback URI in your OAuth provider to: {oauthRedirectUri ?? '…'}

diff --git a/packages/frontend/src/app/connectors/oauth/complete/layout.tsx b/packages/frontend/src/app/connectors/oauth/complete/layout.tsx new file mode 100644 index 00000000..fad7c68f --- /dev/null +++ b/packages/frontend/src/app/connectors/oauth/complete/layout.tsx @@ -0,0 +1,9 @@ +import type { Metadata } from 'next'; + +// The URL carries an authorization code until the page strips it: never send +// it on as a Referer. +export const metadata: Metadata = { referrer: 'no-referrer' }; + +export default function Layout({ children }: { children: React.ReactNode }) { + return children; +} diff --git a/packages/frontend/src/app/connectors/oauth/complete/page.tsx b/packages/frontend/src/app/connectors/oauth/complete/page.tsx new file mode 100644 index 00000000..d0e55405 --- /dev/null +++ b/packages/frontend/src/app/connectors/oauth/complete/page.tsx @@ -0,0 +1,82 @@ +'use client'; + +import { Suspense, useEffect, useRef, useState } from 'react'; +import { useRouter, useSearchParams } from 'next/navigation'; +import Link from 'next/link'; +import { useAuth } from '@/lib/auth-context'; +import { connectors } from '@/lib/api'; +import { Card } from '@/components/ui/card'; +import { buttonVariants } from '@/components/ui/button'; +import { cn } from '@/lib/utils'; + +/** + * Where "Authorize with Provider" lands. The provider sends the browser to the + * backend callback, which only checks the state and forwards the code here; + * this page exchanges it in an authenticated request, so the server can check + * that the person completing the authorization is the one who started it. + */ +function CompleteContent() { + const params = useSearchParams(); + const router = useRouter(); + const { token, isLoading } = useAuth(); + const [error, setError] = useState(params.get('error')); + const started = useRef(false); + + const state = params.get('state'); + const code = params.get('code'); + const connectorId = params.get('connectorId'); + + useEffect(() => { + if (error || !state || !code || started.current || isLoading) return; + if (!token) { + // Sign in first, then come back here with the same code. + const here = `${window.location.pathname}${window.location.search}`; + router.replace(`/login?redirect=${encodeURIComponent(here)}`); + return; + } + started.current = true; + // The code is single-use and bound to a verifier the server keeps, but it + // still has no business staying in the address bar or the history. + window.history.replaceState({}, '', '/connectors/oauth/complete'); + connectors + .oauthComplete(state, code, token) + .then((out) => { + const fallback = `/connectors/${out.connectorId}?oauth=success&tools=${out.toolsImported}`; + router.replace(out.returnTo || fallback); + }) + .catch((err: Error) => setError(err.message || 'The authorization could not be completed.')); + }, [error, state, code, token, isLoading, router]); + + if (!error && !state) { + return

Nothing to complete here.

; + } + + if (error) { + return ( +
+

Authorization not completed

+

{error}

+ + {connectorId ? 'Back to the connector' : 'Back to connectors'} + +
+ ); + } + + return

Completing the authorization…

; +} + +export default function OAuthCompletePage() { + return ( +
+ + Loading…

}> + +
+
+
+ ); +} diff --git a/packages/frontend/src/lib/api.ts b/packages/frontend/src/lib/api.ts index af692d0e..46d498d8 100644 --- a/packages/frontend/src/lib/api.ts +++ b/packages/frontend/src/lib/api.ts @@ -425,11 +425,19 @@ export const connectors = { request<{ message: string; created: number; skipped: number; tools: number; errors?: string[] }>('/api/connectors/import-all', { method: 'POST', body: data, token }), healthCheck: (token: string) => request<{ total: number; healthy: number; unhealthy: number; connectors: any[] }>('/api/connectors/health-check', { token }), - oauthAuthorize: (id: string, token: string) => + oauthAuthorize: (id: string, token: string, returnTo?: string) => request<{ authorizationUrl?: string; error?: string }>( `/api/connectors/${id}/oauth/authorize`, - { method: 'POST', token }, + { method: 'POST', token, body: returnTo ? { returnTo } : undefined }, + ), + /** Second half of "Authorize with Provider": exchange the code the provider sent back. */ + oauthComplete: (state: string, code: string, token: string) => + request<{ connectorId: string; toolsImported: number; returnTo?: string }>( + '/api/mcp-oauth/complete', + { method: 'POST', token, body: { state, code } }, ), + oauthRedirectUri: (token: string) => + request<{ redirectUri: string }>('/api/connectors/oauth/redirect-uri', { token }), discoverTools: (id: string, token: string) => request<{ message: string; tools: any[]; skipped?: string[]; error?: string }>( `/api/connectors/${id}/discover-tools`, diff --git a/packages/frontend/tests/e2e/connector-oauth-complete.spec.ts b/packages/frontend/tests/e2e/connector-oauth-complete.spec.ts new file mode 100644 index 00000000..a59b3405 --- /dev/null +++ b/packages/frontend/tests/e2e/connector-oauth-complete.spec.ts @@ -0,0 +1,89 @@ +import { expect, test, type Page } from '@playwright/test'; + +/** + * "Authorize with Provider" completes in the dashboard: the backend callback + * only forwards code + state here, and this page exchanges them in an + * authenticated request, so the server can refuse a user other than the one + * who started the authorization. + */ + +const USER = { + id: 'u1', + email: 'test@example.com', + name: 'Test User', + role: 'ADMIN', + organizationId: 'o1', + emailVerified: true, +}; + +async function mockApi(page: Page, complete: { status: number; body: unknown }) { + const posts: unknown[] = []; + await page.route(/\/api\//, async (route) => { + const req = route.request(); + const url = req.url(); + const json = (body: unknown, status = 200) => + route.fulfill({ status, contentType: 'application/json', body: JSON.stringify(body) }); + if (url.includes('/api/mcp-oauth/complete') && req.method() === 'POST') { + posts.push({ body: req.postDataJSON(), auth: req.headers()['authorization'] }); + return json(complete.body, complete.status); + } + if (url.includes('/api/users/me/onboarding-state')) return json({ onboardingCompletedAt: '2026-01-01T00:00:00Z' }); + if (url.includes('/api/users/me')) return json(USER); + if (url.includes('/api/organizations/current')) return json({ id: 'o1', name: 'Acme', createdAt: '2026-01-01' }); + if (url.includes('/api/organizations/mine')) return json([{ id: 'o1', name: 'Acme', role: 'ADMIN', joinedAt: '2026-01-01' }]); + if (url.includes('/api/license/status')) return json({ plan: 'community', status: 'active' }); + if (url.includes('/api/connectors/c1')) return json({ id: 'c1', name: 'Etsy', type: 'REST', authType: 'OAUTH2', tools: [], isActive: true }); + if (url.includes('/api/connectors')) return json([]); + return json({}); + }); + return posts; +} + +async function signIn(page: Page) { + await page.context().addCookies([{ name: 'amcp_token', value: 'test-token', url: 'http://localhost:3100' }]); + await page.addInitScript((user) => { + localStorage.setItem('amcp_token', 'test-token'); + localStorage.setItem('amcp_user', JSON.stringify(user)); + }, USER); +} + +test('exchanges the code as the signed-in user and lands where the flow asked', async ({ page }) => { + await signIn(page); + const posts = await mockApi(page, { + status: 200, + body: { connectorId: 'c1', toolsImported: 0, returnTo: '/connectors/c1?oauth=success&tools=0' }, + }); + await page.goto('/connectors/oauth/complete?state=st-1&code=co-1'); + await expect(page).toHaveURL(/\/connectors\/c1/); + expect(posts).toEqual([{ body: { state: 'st-1', code: 'co-1' }, auth: 'Bearer test-token' }]); +}); + +test('shows the server refusal instead of a success', async ({ page }) => { + await signIn(page); + await mockApi(page, { + status: 403, + body: { statusCode: 403, message: 'This authorization was started by another account.' }, + }); + await page.goto('/connectors/oauth/complete?state=st-1&code=co-1'); + await expect(page.locator('p[role="alert"]')).toContainText('started by another account'); + // The code does not stay in the address bar. + await expect(page).toHaveURL(/\/connectors\/oauth\/complete$/); +}); + +test('explains an error the provider returned, with a way back', async ({ page }) => { + await signIn(page); + const posts = await mockApi(page, { status: 200, body: {} }); + await page.goto('/connectors/oauth/complete?error=The+authorization+was+cancelled+at+the+provider.&connectorId=c1'); + await expect(page.locator('p[role="alert"]')).toContainText('cancelled at the provider'); + await expect(page.getByRole('link', { name: 'Back to the connector' })).toHaveAttribute('href', '/connectors/c1'); + expect(posts).toHaveLength(0); +}); + +test('sends a signed-out visitor to sign in first, keeping the code', async ({ page }) => { + const posts = await mockApi(page, { status: 200, body: {} }); + await page.goto('/connectors/oauth/complete?state=st-1&code=co-1'); + await expect(page).toHaveURL(/\/login\?redirect=/); + const redirect = new URL(page.url()).searchParams.get('redirect'); + expect(redirect).toBe('/connectors/oauth/complete?state=st-1&code=co-1'); + expect(posts).toHaveLength(0); +});