mirror of
https://github.com/simstudioai/sim.git
synced 2026-09-24 15:45:35 +08:00
fix(mcp): bound OAuth callback steps and log each phase to pinpoint hangs (#5807)
* fix(mcp): bound OAuth callback steps and log each phase to pinpoint hangs The MCP OAuth callback could hang indefinitely after burning state, with no error, no timeout, and no log — the request stalls somewhere between the burn and token persistence with no observable I/O. Wrap each awaited callback step (loadPreregisteredClient, mcpAuthGuarded, clearVerifier, discoverServerTools) with a per-step timeout so a stalled operation surfaces as a labeled error instead of hanging forever, and log start/done for every step, every guarded OAuth fetch phase (validate/request/read-body), and the token-exchange DB writes. The last start without a matching done names the exact stall point. * fix(mcp): widen callback step bounds and drop resolved IP from fetch logs - Raise mcpAuthGuarded step bound 60s->120s and discoverServerTools 30s->60s so a legitimately slow multi-leg exchange (each internal fetch is already 30s-bounded) can't be falsely reported as failed. - Log only a pinned boolean, not the resolved IP, so self-hosted private endpoint addresses don't reach log sinks. * fix(mcp): tolerate non-thenable returns in callback timedStep helper Promise.resolve(fn()) so a synchronously-returning (e.g. mocked) step can't throw on .catch — restores the callback route unit tests.
This commit is contained in:
@@ -25,6 +25,46 @@ const logger = createLogger('McpOauthCallbackAPI')
|
||||
|
||||
export const dynamic = 'force-dynamic'
|
||||
|
||||
class OauthCallbackStepTimeout extends Error {
|
||||
constructor(step: string, ms: number) {
|
||||
super(`MCP OAuth callback step "${step}" did not settle within ${ms}ms`)
|
||||
this.name = 'OauthCallbackStepTimeout'
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Times and bounds one awaited step of the callback so a stalled operation
|
||||
* surfaces as a labeled, logged error instead of hanging the request forever.
|
||||
* The losing promise is not cancelled (a wedged DB/socket op can't be), so it
|
||||
* settles in the background with its rejection swallowed; the point is that the
|
||||
* request stops waiting on it and the logs name the exact step that stalled.
|
||||
*/
|
||||
async function timedStep<T>(step: string, ms: number, fn: () => Promise<T>): Promise<T> {
|
||||
const start = Date.now()
|
||||
logger.info(`OAuth callback step start: ${step}`)
|
||||
const work = Promise.resolve(fn())
|
||||
work.catch(() => {})
|
||||
let timer: ReturnType<typeof setTimeout> | undefined
|
||||
try {
|
||||
const value = await Promise.race([
|
||||
work,
|
||||
new Promise<never>((_, reject) => {
|
||||
timer = setTimeout(() => reject(new OauthCallbackStepTimeout(step, ms)), ms)
|
||||
timer.unref?.()
|
||||
}),
|
||||
])
|
||||
logger.info(`OAuth callback step done: ${step} (${Date.now() - start}ms)`)
|
||||
return value
|
||||
} catch (error) {
|
||||
logger.error(`OAuth callback step failed: ${step} (${Date.now() - start}ms)`, {
|
||||
error: toError(error).message,
|
||||
})
|
||||
throw error
|
||||
} finally {
|
||||
clearTimeout(timer)
|
||||
}
|
||||
}
|
||||
|
||||
function escapeHtml(value: string): string {
|
||||
return value
|
||||
.replace(/&/g, '&')
|
||||
@@ -145,8 +185,9 @@ export const GET = withRouteHandler(async (request: NextRequest) => {
|
||||
serverId
|
||||
)
|
||||
}
|
||||
const serverUrl = server.url
|
||||
try {
|
||||
assertSafeOauthServerUrl(server.url)
|
||||
assertSafeOauthServerUrl(serverUrl)
|
||||
} catch {
|
||||
return respond(
|
||||
'MCP OAuth requires https (or http://localhost for development).',
|
||||
@@ -157,16 +198,22 @@ export const GET = withRouteHandler(async (request: NextRequest) => {
|
||||
}
|
||||
|
||||
// Burn state before token exchange so a replayed callback cannot reuse it.
|
||||
await clearState(row.id, 'callback:burn-before-exchange')
|
||||
await timedStep('clearState(burn)', 10_000, () =>
|
||||
clearState(row.id, 'callback:burn-before-exchange')
|
||||
)
|
||||
|
||||
const preregistered = await loadPreregisteredClient(server.id)
|
||||
const preregistered = await timedStep('loadPreregisteredClient', 15_000, () =>
|
||||
loadPreregisteredClient(server.id)
|
||||
)
|
||||
const provider = new SimMcpOauthProvider({ row, preregistered })
|
||||
let result: Awaited<ReturnType<typeof mcpAuthGuarded>>
|
||||
try {
|
||||
result = await mcpAuthGuarded(provider, {
|
||||
serverUrl: server.url,
|
||||
authorizationCode: code,
|
||||
})
|
||||
result = await timedStep('mcpAuthGuarded', 120_000, () =>
|
||||
mcpAuthGuarded(provider, {
|
||||
serverUrl,
|
||||
authorizationCode: code,
|
||||
})
|
||||
)
|
||||
} catch (e) {
|
||||
logger.error('Token exchange failed during MCP OAuth callback', e)
|
||||
return respond(
|
||||
@@ -176,7 +223,11 @@ export const GET = withRouteHandler(async (request: NextRequest) => {
|
||||
server.id
|
||||
)
|
||||
} finally {
|
||||
await clearVerifier(row.id)
|
||||
await timedStep('clearVerifier', 10_000, () => clearVerifier(row.id)).catch((e) =>
|
||||
logger.error('Failed to clear PKCE verifier after MCP OAuth callback', {
|
||||
error: toError(e).message,
|
||||
})
|
||||
)
|
||||
}
|
||||
|
||||
if (result !== 'AUTHORIZED') {
|
||||
@@ -185,7 +236,9 @@ export const GET = withRouteHandler(async (request: NextRequest) => {
|
||||
|
||||
try {
|
||||
// forceRefresh: skip any stale cache from before re-auth.
|
||||
await mcpService.discoverServerTools(session.user.id, server.id, server.workspaceId, true)
|
||||
await timedStep('discoverServerTools', 60_000, () =>
|
||||
mcpService.discoverServerTools(session.user.id, server.id, server.workspaceId, true)
|
||||
)
|
||||
} catch (e) {
|
||||
logger.warn('Post-auth tools refresh failed', toError(e).message)
|
||||
}
|
||||
|
||||
@@ -170,18 +170,22 @@ export async function saveClientInformation(
|
||||
info: OAuthClientInformationMixed
|
||||
): Promise<void> {
|
||||
const encrypted = await encryptClientInformation(info)
|
||||
logger.info('Persisting MCP OAuth client information', { rowId })
|
||||
await db
|
||||
.update(mcpServerOauth)
|
||||
.set({ clientInformation: encrypted, updatedAt: new Date() })
|
||||
.where(eq(mcpServerOauth.id, rowId))
|
||||
logger.info('Persisted MCP OAuth client information', { rowId })
|
||||
}
|
||||
|
||||
export async function saveTokens(rowId: string, tokens: OAuthTokens): Promise<void> {
|
||||
const encrypted = await encryptTokens(tokens)
|
||||
logger.info('Persisting MCP OAuth tokens', { rowId })
|
||||
await db
|
||||
.update(mcpServerOauth)
|
||||
.set({ tokens: encrypted, lastRefreshedAt: new Date(), updatedAt: new Date() })
|
||||
.where(eq(mcpServerOauth.id, rowId))
|
||||
logger.info('Persisted MCP OAuth tokens', { rowId })
|
||||
}
|
||||
|
||||
export async function saveCodeVerifier(rowId: string, verifier: string): Promise<void> {
|
||||
|
||||
@@ -1,9 +1,12 @@
|
||||
import type { FetchLike } from '@modelcontextprotocol/sdk/shared/transport.js'
|
||||
import { createLogger } from '@sim/logger'
|
||||
import type { Agent } from 'undici'
|
||||
import { createPinnedFetchWithDispatcher } from '@/lib/core/security/input-validation.server'
|
||||
import { validateMcpServerSsrf } from '@/lib/mcp/domain-check'
|
||||
import { McpError } from '@/lib/mcp/types'
|
||||
|
||||
const logger = createLogger('McpOauthFetch')
|
||||
|
||||
/** Pinned fetch for the live MCP transport, plus a handle to release its sockets. */
|
||||
export interface PinnedMcpFetch {
|
||||
fetch: typeof fetch
|
||||
@@ -151,13 +154,17 @@ function releaseStreamOnSettle(
|
||||
export function createSsrfGuardedMcpFetch(timeoutMs: number = OAUTH_FETCH_TIMEOUT_MS): FetchLike {
|
||||
return (async (url, init) => {
|
||||
const target = typeof url === 'string' ? url : url.href
|
||||
const host = URL.canParse(target) ? new URL(target).host : target
|
||||
const startedAt = Date.now()
|
||||
const timeoutSignal = AbortSignal.timeout(timeoutMs)
|
||||
// Bound every phase — validation, request, body read — by the deadline + caller signal.
|
||||
const signal = init?.signal ? AbortSignal.any([init.signal, timeoutSignal]) : timeoutSignal
|
||||
// Per-request Agent must be torn down (finally): a one-shot leg never reuses its socket.
|
||||
let dispatcher: Agent | undefined
|
||||
try {
|
||||
logger.info('OAuth guarded fetch: validating', { host })
|
||||
const resolvedIP = await withDeadline(validateMcpServerSsrf(target), signal)
|
||||
logger.info('OAuth guarded fetch: requesting', { host, pinned: Boolean(resolvedIP) })
|
||||
let response: Response
|
||||
if (resolvedIP) {
|
||||
const pinned = createPinnedFetchWithDispatcher(resolvedIP, {
|
||||
@@ -175,11 +182,23 @@ export function createSsrfGuardedMcpFetch(timeoutMs: number = OAUTH_FETCH_TIMEOU
|
||||
if (contentType.includes('text/event-stream')) {
|
||||
const streamed = releaseStreamOnSettle(response, dispatcher, signal)
|
||||
dispatcher = undefined // teardown ownership moved to releaseStreamOnSettle
|
||||
logger.info('OAuth guarded fetch: streaming response', {
|
||||
host,
|
||||
status: response.status,
|
||||
ms: Date.now() - startedAt,
|
||||
})
|
||||
return streamed
|
||||
}
|
||||
return await bufferUnderDeadline(response, signal)
|
||||
logger.info('OAuth guarded fetch: reading body', { host, status: response.status })
|
||||
const buffered = await bufferUnderDeadline(response, signal)
|
||||
logger.info('OAuth guarded fetch: done', {
|
||||
host,
|
||||
status: response.status,
|
||||
ms: Date.now() - startedAt,
|
||||
})
|
||||
return buffered
|
||||
} catch (error) {
|
||||
const host = URL.canParse(target) ? new URL(target).host : target
|
||||
logger.warn('OAuth guarded fetch: failed', { host, ms: Date.now() - startedAt })
|
||||
// Relabel only our own deadline — by reason identity, not signal state (which may
|
||||
// abort independently just after the deadline).
|
||||
if (timeoutSignal.aborted && error === timeoutSignal.reason) {
|
||||
|
||||
Reference in New Issue
Block a user