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
136 changes: 114 additions & 22 deletions apps/sim/app/api/mcp/oauth/callback/route.test.ts
Original file line number Diff line number Diff line change
@@ -1,14 +1,20 @@
import { runInNewContext } from 'node:vm'
import { BroadcastChannel } from 'node:worker_threads'
import {
authMockFns,
dbChainMockFns,
mcpOauthMock,
mcpOauthMockFns,
resetDbChainMock,
} from '@sim/testing'
import { flushMicrotasks } from '@sim/testing/helpers/async'
import { emcnMock } from '@sim/testing/mocks/emcn.mock'
import { getMockLogger } from '@sim/testing/mocks/logger.mock'
import { mcpServiceMock, mcpServiceMockFns } from '@sim/testing/mocks/mcp-service.mock'
import { NextRequest } from 'next/server'
import { beforeEach, describe, expect, it, vi } from 'vitest'
import { CredentialGroupOAuthStateVersionError } from '@/lib/credential-groups/oauth-attempt-version'
import { connectCredentialGroupInPopup } from '@/lib/credential-groups/oauth-popup'

const {
mockAuthenticateEnrollment,
Expand All @@ -24,6 +30,7 @@ const {

vi.mock('@/lib/mcp/oauth', () => mcpOauthMock)
vi.mock('@/lib/mcp/service', () => mcpServiceMock)
vi.mock('@sim/emcn', () => emcnMock)
vi.mock('@/lib/credential-groups/application/enrollment-auth', () => ({
credentialGroupOAuthAttemptPrincipal: mockAuthenticateEnrollment,
}))
Expand Down Expand Up @@ -88,30 +95,41 @@ describe('MCP OAuth callback route', () => {
mockEnforceCallbackRateLimit.mockResolvedValue(null)
})

it.each([undefined, 'denied'])(
'finishes a direct connection without the invitation form: %s',
async (error) => {
const completionId = '00000000-0000-4000-8000-000000000002'
mockConsumeManagedAttempt.mockResolvedValueOnce({
state: 'mcp_cg_direct',
organizationId: 'organization-1',
invitationToken: 'invitation-token',
mcpServerId: 'server-1',
completionId,
returnTo: 'integrations',
})
const response = await GET(
new NextRequest(
`http://localhost:3000/api/mcp/oauth/callback?state=mcp_cg_direct&${error ? 'error=denied' : 'code=code-1'}`
)
it.each([
[undefined, null],
['access_denied', 'denied'],
['invalid_request', 'failed'],
['invalid_scope', 'failed'],
['server_error', 'provider_unavailable'],
['temporarily_unavailable', 'provider_unavailable'],
['unexpected-secret-value', 'failed'],
])('finishes a direct connection without the invitation form: %s', async (error, expected) => {
const completionId = '00000000-0000-4000-8000-000000000002'
mockConsumeManagedAttempt.mockResolvedValueOnce({
state: 'mcp_cg_direct',
organizationId: 'organization-1',
invitationToken: 'invitation-token',
mcpServerId: 'server-1',
completionId,
returnTo: 'integrations',
})
const response = await GET(
new NextRequest(
`http://localhost:3000/api/mcp/oauth/callback?state=mcp_cg_direct&${error ? `error=${error}&error_description=private-provider-detail` : 'code=code-1'}`
)
const destination = new URL(response.headers.get('location')!, 'http://localhost:3000')
expect(destination.pathname).toBe('/credential-groups/complete')
expect(destination.searchParams.get('completionId')).toBe(completionId)
expect(destination.searchParams.get('organizationId')).toBe('organization-1')
expect(destination.searchParams.get('oauth')).toBe(error ?? null)
)
const destination = new URL(response.headers.get('location')!, 'http://localhost:3000')
expect(destination.pathname).toBe('/credential-groups/complete')
expect(destination.searchParams.get('completionId')).toBe(completionId)
expect(destination.searchParams.get('organizationId')).toBe('organization-1')
expect(destination.searchParams.get('oauth')).toBe(expected)
if (error) {
const warnings = JSON.stringify(getMockLogger('McpOauthCallbackAPI').warn.mock.calls)
expect(warnings).toContain(error === 'unexpected-secret-value' ? 'unknown' : error)
expect(warnings).not.toContain('unexpected-secret-value')
expect(warnings).not.toContain('private-provider-detail')
}
)
})

it('performs the token exchange through the SSRF-guarded mcpAuthGuarded wrapper', async () => {
const request = new NextRequest(
Expand Down Expand Up @@ -185,4 +203,78 @@ describe('MCP OAuth callback route', () => {
expect(mockAuthenticateEnrollment).not.toHaveBeenCalled()
expect(mockCompleteManagedMcpOAuth).not.toHaveBeenCalled()
})

it.each(['missing', 'version'])(
'delivers a %s managed-state failure to only the initiating popup',
async (kind) => {
vi.useFakeTimers({ toFake: ['setTimeout', 'clearTimeout'] })
vi.stubGlobal('BroadcastChannel', BroadcastChannel)
vi.stubGlobal('window', {
open: () => ({ closed: false, opener: null, location: { replace() {} }, close() {} }),
})
const controller = new AbortController()
const completionId = '00000000-0000-4000-8000-000000000010'
const secondId = '00000000-0000-4000-8000-000000000011'
const state = 'mcp_cg_first'
let first: string | undefined
let second: string | undefined
const start = (nonce: string) => async () => ({
invitationLink: 'http://localhost:3000/credential-groups/enroll/fixture',
authorizationUrl: `https://oauth.example.com/authorize?state=${nonce}`,
})
const outcomes = [
connectCredentialGroupInPopup(completionId, start(state), controller.signal).then(
() => {
first = 'connected'
},
(error: Error) => {
first = error.message
}
),
connectCredentialGroupInPopup(secondId, start('mcp_cg_second'), controller.signal).then(
() => {
second = 'connected'
},
(error: Error) => {
second = error.message
}
),
]
const observer = new BroadcastChannel('mcp-oauth')
const success = new BroadcastChannel(`sim:credential-group-oauth:${secondId}`)
try {
await flushMicrotasks()
if (kind === 'missing') mockConsumeManagedAttempt.mockResolvedValueOnce(null)
else
mockConsumeManagedAttempt.mockRejectedValueOnce(
new CredentialGroupOAuthStateVersionError()
)
const response = await GET(
new NextRequest(`http://localhost:3000/api/mcp/oauth/callback?state=${state}&code=unused`)
)
const body = await response.text()
const script = body.match(/<script>([\s\S]*?)<\/script>/)?.[1]
if (!script) throw new Error('Callback did not provide its popup handoff')
const delivered = new Promise<void>((resolve) => {
observer.onmessage = () => resolve()
})
runInNewContext(script, { BroadcastChannel, setTimeout() {} })
await delivered
await vi.waitFor(() =>
expect(first).toBe('This connection attempt expired. Try connecting your account again.')
)
expect(second).toBeUndefined()
expect(mockCompleteManagedMcpOAuth).not.toHaveBeenCalled()
success.postMessage('connected')
await outcomes[1]
expect(second).toBe('connected')
} finally {
controller.abort()
observer.close()
success.close()
await Promise.all(outcomes)
vi.useRealTimers()
}
}
)
})
33 changes: 31 additions & 2 deletions apps/sim/app/api/mcp/oauth/callback/route.ts
Original file line number Diff line number Diff line change
Expand Up @@ -100,7 +100,13 @@ async function completeManagedMcpCallback(params: {
throw error
}
if (!attempt) {
return htmlClose('Invalid or expired authorization state.', false, 'invalid_state')
return htmlClose(
'Invalid or expired authorization state.',
false,
'invalid_state',
undefined,
params.state
)
}
const failureRedirect = (oauth: CredentialGroupOAuthFailure) =>
attempt.completionId
Expand All @@ -110,7 +116,30 @@ async function completeManagedMcpCallback(params: {
attempt.returnTo === 'integrations' ? attempt.organizationId : undefined
)
: createCredentialGroupEnrollmentRedirect(attempt.invitationToken, { oauth })
if (params.error) return failureRedirect('denied')
if (params.error) {
const errorCode = [
'invalid_request',
'unauthorized_client',
'access_denied',
'unsupported_response_type',
'invalid_scope',
'server_error',
'temporarily_unavailable',
].includes(params.error)
? params.error
: 'unknown'
logger.warn('Managed MCP authorization returned a provider error', {
phase: 'provider_authorization',
errorCode,
})
return failureRedirect(
errorCode === 'access_denied'
? 'denied'
: errorCode === 'server_error' || errorCode === 'temporarily_unavailable'
? 'provider_unavailable'
: 'failed'
)
}
if (!params.code) {
return failureRedirect('failed')
}
Expand Down
21 changes: 21 additions & 0 deletions apps/sim/lib/credential-groups/oauth-popup.ts
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import { toast } from '@sim/emcn'
import { toError } from '@sim/utils/errors'
import { isRecordLike } from '@sim/utils/object'
import type { OrganizationAccountConnectionResponse } from '@/lib/api/contracts/organization-accounts'
import {
CREDENTIAL_GROUP_OAUTH_FAILURE_MESSAGES,
Expand Down Expand Up @@ -28,13 +29,16 @@ export async function connectCredentialGroupInPopup(
return new Promise((resolve, reject) => {
const controller = new AbortController()
const channel = new BroadcastChannel(credentialGroupOAuthCompletionChannel(completionId))
const mcpChannel = new BroadcastChannel('mcp-oauth')
let mcpState: string | undefined
let settled = false
const finish = (error?: Error) => {
if (settled) return
settled = true
controller.abort()
clearTimeout(timer)
channel.close()
mcpChannel.close()
signal.removeEventListener('abort', abort)
toast.dismiss(notice)
if (error) reject(error)
Expand All @@ -59,6 +63,19 @@ export async function connectCredentialGroupInPopup(
else if (isCredentialGroupOAuthFailure(data))
finish(new Error(CREDENTIAL_GROUP_OAUTH_FAILURE_MESSAGES[data]))
}
/** Invalid state has no stored completion ID; only accept a failure for this attempt's nonce. */
mcpChannel.onmessage = ({ data }: MessageEvent<unknown>) => {
if (
mcpState &&
isRecordLike(data) &&
data.type === 'mcp-oauth' &&
data.state === mcpState &&
data.ok === false &&
data.reason === 'invalid_state'
) {
finish(new Error(CREDENTIAL_GROUP_OAUTH_FAILURE_MESSAGES.expired))
}
}
signal.addEventListener('abort', abort, { once: true })
if (signal.aborted) {
abort()
Expand All @@ -71,6 +88,10 @@ export async function connectCredentialGroupInPopup(
abort()
return
}
const state = result.authorizationUrl
? new URL(result.authorizationUrl).searchParams.get('state')
: null
if (state?.startsWith('mcp_cg_')) mcpState = state
popup.location.replace(result.authorizationUrl ?? result.invitationLink)
})
.catch((error) => {
Expand Down
Loading