fix(mcp): recover cleanly from OAuth failures (#5595)

* fix(mcp): improve OAuth failure recovery

* fix(mcp): address OAuth recovery review findings

* fix(mcp): prefer static bearer auth in connection tests

* fix(mcp): preserve discovery failure state

Treat static-header 401s as credential failures, keep OAuth failures pending, and prevent failed refreshes or reflected upstream errors from masquerading as connected state.
This commit is contained in:
Justin Blumencranz
2026-07-17 16:38:21 -07:00
committed by GitHub
parent 7e975e7dae
commit c9baae7f2e
16 changed files with 1447 additions and 194 deletions
@@ -0,0 +1,167 @@
/**
* @vitest-environment node
*/
import type { NextRequest } from 'next/server'
import { beforeEach, describe, expect, it, vi } from 'vitest'
const { mockClearCache, mockDiscoverServerTools, mockSelect, mockUpdateSet } = vi.hoisted(() => ({
mockClearCache: vi.fn(),
mockDiscoverServerTools: vi.fn(),
mockSelect: vi.fn(),
mockUpdateSet: vi.fn(),
}))
vi.mock('@sim/db', () => ({
db: {
select: mockSelect,
update: vi.fn().mockReturnValue({ set: mockUpdateSet }),
},
}))
vi.mock('@/lib/core/utils/with-route-handler', () => ({
withRouteHandler: (handler: unknown) => handler,
}))
vi.mock('@/lib/mcp/middleware', () => ({
withMcpAuth:
() =>
(
handler: (
request: NextRequest,
context: { userId: string; workspaceId: string; requestId: string },
routeContext: { params: Promise<{ id: string }> }
) => Promise<Response>
) =>
(request: NextRequest, routeContext: { params: Promise<{ id: string }> }) =>
handler(
request,
{ userId: 'user-1', workspaceId: 'workspace-1', requestId: 'request-1' },
routeContext
),
}))
vi.mock('@/lib/mcp/service', () => ({
mcpService: {
clearCache: mockClearCache,
discoverServerTools: mockDiscoverServerTools,
},
}))
import { POST } from '@/app/api/mcp/servers/[id]/refresh/route'
const initialServer = {
id: 'server-1',
workspaceId: 'workspace-1',
name: 'OAuth Server',
url: 'https://example.com/mcp',
connectionStatus: 'connected',
lastError: null,
lastConnected: new Date('2026-01-01T00:00:00.000Z'),
toolCount: 4,
statusConfig: { consecutiveFailures: 0, lastSuccessfulDiscovery: null },
}
const persistedServer = {
...initialServer,
connectionStatus: 'disconnected',
lastError: null,
toolCount: 0,
}
function selectRows(rows: unknown[]) {
return {
from: vi.fn().mockReturnValue({
where: vi.fn().mockReturnValue({
limit: vi.fn().mockResolvedValue(rows),
}),
}),
}
}
describe('MCP server refresh route', () => {
beforeEach(() => {
vi.clearAllMocks()
mockSelect.mockReturnValueOnce(selectRows([initialServer]))
mockUpdateSet.mockReturnValue({
where: vi.fn().mockReturnValue({ returning: vi.fn().mockResolvedValue([persistedServer]) }),
})
})
it('preserves the service-persisted OAuth pending status', async () => {
mockDiscoverServerTools.mockRejectedValueOnce(new Error('OAuth authorization required'))
const request = new Request('http://localhost/api/mcp/servers/server-1/refresh', {
method: 'POST',
}) as NextRequest
const response = await POST(request, { params: Promise.resolve({ id: 'server-1' }) })
const body = await response.json()
expect(body.data).toEqual(
expect.objectContaining({
status: 'disconnected',
error: null,
})
)
expect(mockUpdateSet).not.toHaveBeenCalledWith(
expect.objectContaining({ connectionStatus: expect.anything() })
)
})
it('reports the discovery failure when status persistence leaves a stale connected row', async () => {
const reflectedSecret = 'Bearer reflected-static-token'
mockDiscoverServerTools.mockRejectedValueOnce(
new Error(`Upstream reflected ${reflectedSecret}`)
)
mockUpdateSet.mockReturnValueOnce({
where: vi.fn().mockReturnValue({
returning: vi.fn().mockResolvedValue([initialServer]),
}),
})
const request = new Request('http://localhost/api/mcp/servers/server-1/refresh', {
method: 'POST',
}) as NextRequest
const response = await POST(request, { params: Promise.resolve({ id: 'server-1' }) })
const body = await response.json()
expect(body.data).toEqual(
expect.objectContaining({
status: 'disconnected',
error: 'Internal server error',
workflowsUpdated: 0,
})
)
expect(JSON.stringify(body)).not.toContain(reflectedSecret)
expect(mockClearCache).not.toHaveBeenCalled()
})
it('preserves a connected status from a newer successful discovery', async () => {
mockDiscoverServerTools.mockRejectedValueOnce(new Error('Connection failed'))
const newerSuccessfulServer = {
...initialServer,
lastConnected: new Date(Date.now() + 60_000),
toolCount: 7,
}
mockUpdateSet.mockReturnValueOnce({
where: vi.fn().mockReturnValue({
returning: vi.fn().mockResolvedValue([newerSuccessfulServer]),
}),
})
const request = new Request('http://localhost/api/mcp/servers/server-1/refresh', {
method: 'POST',
}) as NextRequest
const response = await POST(request, { params: Promise.resolve({ id: 'server-1' }) })
const body = await response.json()
expect(body.data).toEqual(
expect.objectContaining({
status: 'connected',
error: null,
toolCount: 7,
workflowsUpdated: 0,
})
)
expect(mockClearCache).toHaveBeenCalledWith('workspace-1')
})
})
@@ -2,6 +2,7 @@ import { db } from '@sim/db'
import { mcpServers, workflow, workflowBlocks } from '@sim/db/schema'
import { createLogger } from '@sim/logger'
import { toError } from '@sim/utils/errors'
import { truncate } from '@sim/utils/string'
import { and, eq, inArray, isNull } from 'drizzle-orm'
import type { NextRequest } from 'next/server'
import { mcpServerIdParamsSchema } from '@/lib/api/contracts/mcp'
@@ -9,8 +10,9 @@ import { validationErrorResponse } from '@/lib/api/server'
import { withRouteHandler } from '@/lib/core/utils/with-route-handler'
import { withMcpAuth } from '@/lib/mcp/middleware'
import { mcpService } from '@/lib/mcp/service'
import type { McpServerStatusConfig, McpTool, McpToolSchema } from '@/lib/mcp/types'
import type { McpTool, McpToolSchema } from '@/lib/mcp/types'
import {
categorizeError,
createMcpErrorResponse,
createMcpSuccessResponse,
MCP_TOOL_CORE_PARAMS,
@@ -184,17 +186,10 @@ export const POST = withRouteHandler(
)
}
let connectionStatus: 'connected' | 'disconnected' | 'error' = 'error'
let toolCount = 0
let lastError: string | null = null
let syncResult: SyncResult = { updatedCount: 0, updatedWorkflowIds: [] }
let discoveredTools: McpTool[] = []
const currentStatusConfig: McpServerStatusConfig =
(server.statusConfig as McpServerStatusConfig | null) ?? {
consecutiveFailures: 0,
lastSuccessfulDiscovery: null,
}
let discoveryError: string | null = null
const discoveryStartedAt = new Date()
try {
discoveredTools = await mcpService.discoverServerTools(
@@ -203,10 +198,17 @@ export const POST = withRouteHandler(
workspaceId,
true
)
connectionStatus = 'connected'
toolCount = discoveredTools.length
logger.info(`[${requestId}] Discovered ${toolCount} tools from server ${serverId}`)
logger.info(
`[${requestId}] Discovered ${discoveredTools.length} tools from server ${serverId}`
)
} catch (error) {
discoveryError = truncate(categorizeError(error).message, 200, '')
logger.warn(`[${requestId}] Failed to connect to server ${serverId}`, {
error: discoveryError,
})
}
if (discoveryError === null) {
syncResult = await syncToolSchemasToWorkflows(
workspaceId,
serverId,
@@ -214,37 +216,44 @@ export const POST = withRouteHandler(
requestId,
{ url: server.url ?? undefined, name: server.name ?? undefined }
)
} catch (error) {
connectionStatus = 'error'
lastError =
error instanceof Error
? error.message.split('\n')[0].slice(0, 200)
: 'Connection failed'
logger.warn(`[${requestId}] Failed to connect to server ${serverId}:`, error)
}
const now = new Date()
const newStatusConfig =
connectionStatus === 'connected'
? { consecutiveFailures: 0, lastSuccessfulDiscovery: now.toISOString() }
: {
consecutiveFailures: currentStatusConfig.consecutiveFailures + 1,
lastSuccessfulDiscovery: currentStatusConfig.lastSuccessfulDiscovery,
}
const [refreshedServer] = await db
.update(mcpServers)
.set({
lastToolsRefresh: now,
connectionStatus,
lastError,
lastConnected: connectionStatus === 'connected' ? now : server.lastConnected,
toolCount,
statusConfig: newStatusConfig,
updatedAt: now,
})
.where(eq(mcpServers.id, serverId))
.returning()
.where(
and(
eq(mcpServers.id, serverId),
eq(mcpServers.workspaceId, workspaceId),
isNull(mcpServers.deletedAt)
)
)
.returning({
connectionStatus: mcpServers.connectionStatus,
lastConnected: mcpServers.lastConnected,
lastError: mcpServers.lastError,
toolCount: mcpServers.toolCount,
})
let connectionStatus = refreshedServer?.connectionStatus ?? 'error'
let lastError = refreshedServer ? refreshedServer.lastError : discoveryError
const toolCount = refreshedServer?.toolCount ?? discoveredTools.length
if (discoveryError !== null && connectionStatus === 'connected') {
const newerSuccessWonRace =
refreshedServer?.lastConnected != null &&
refreshedServer.lastConnected > discoveryStartedAt
if (!newerSuccessWonRace) {
connectionStatus = 'disconnected'
lastError = discoveryError
}
}
if (connectionStatus === 'connected') {
await mcpService.clearCache(workspaceId)
@@ -0,0 +1,228 @@
/**
* @vitest-environment node
*/
import { createMockRequest, loggerMock } from '@sim/testing'
import type { NextRequest } from 'next/server'
import { beforeEach, describe, expect, it, vi } from 'vitest'
const {
mockClientOptions,
mockConnect,
mockDetectMcpAuthType,
mockDisconnect,
mockListTools,
mockResolveMcpConfigEnvVars,
mockValidateMcpServerSsrf,
MockMcpSsrfError,
} = vi.hoisted(() => ({
mockClientOptions: vi.fn(),
mockConnect: vi.fn(),
mockDetectMcpAuthType: vi.fn(),
mockDisconnect: vi.fn(),
mockListTools: vi.fn(),
mockResolveMcpConfigEnvVars: vi.fn(),
mockValidateMcpServerSsrf: vi.fn(),
MockMcpSsrfError: class extends Error {},
}))
vi.mock('@/lib/core/utils/with-route-handler', () => ({
withRouteHandler: (handler: unknown) => handler,
}))
vi.mock('@/lib/mcp/client', () => ({
McpClient: class {
constructor(options: unknown) {
mockClientOptions(options)
}
static getVersionInfo() {
return { preferred: '2025-06-18', supported: ['2025-06-18'] }
}
connect = mockConnect
disconnect = mockDisconnect
listTools = mockListTools
getNegotiatedVersion() {
return '2025-06-18'
}
},
}))
vi.mock('@/lib/mcp/domain-check', () => ({
McpDnsResolutionError: class extends Error {},
McpDomainNotAllowedError: class extends Error {},
McpSsrfError: MockMcpSsrfError,
validateMcpDomain: vi.fn(),
validateMcpServerSsrf: mockValidateMcpServerSsrf,
}))
vi.mock('@/lib/mcp/middleware', () => ({
mcpBodyReadErrorResponse: vi.fn(() => null),
readMcpJsonBodyWithLimit: (request: NextRequest) => request.json(),
withMcpAuth:
() =>
(
handler: (
request: NextRequest,
context: { userId: string; workspaceId: string; requestId: string }
) => Promise<Response>
) =>
(request: NextRequest) =>
handler(request, {
userId: 'user-1',
workspaceId: 'workspace-1',
requestId: 'request-1',
}),
}))
vi.mock('@/lib/mcp/oauth', () => ({
detectMcpAuthType: mockDetectMcpAuthType,
}))
vi.mock('@/lib/mcp/resolve-config', () => ({
resolveMcpConfigEnvVars: mockResolveMcpConfigEnvVars,
}))
import { POST } from '@/app/api/mcp/servers/test-connection/route'
const mockLogger = vi.mocked(loggerMock.createLogger).mock.results.at(-1)?.value
function createTestRequest(headers: Record<string, string> = {}) {
return createMockRequest(
'POST',
{
name: 'Dual Auth Server',
transport: 'streamable-http',
url: 'https://example.com/mcp',
headers,
timeout: 10000,
},
{},
'http://localhost/api/mcp/servers/test-connection'
)
}
describe('MCP server test-connection route', () => {
beforeEach(() => {
vi.clearAllMocks()
mockDetectMcpAuthType.mockResolvedValue('oauth')
mockValidateMcpServerSsrf.mockResolvedValue('203.0.113.10')
mockResolveMcpConfigEnvVars.mockImplementation(async (config: unknown) => ({
config,
missingVars: [],
}))
mockConnect.mockResolvedValue(undefined)
mockListTools.mockResolvedValue([])
mockDisconnect.mockResolvedValue(undefined)
})
it('tests configured bearer headers before treating OAuth discovery as mandatory', async () => {
const response = await POST(createTestRequest({ Authorization: 'Bearer static-api-token' }))
const body = await response.json()
expect(response.status).toBe(200)
expect(body.data).toEqual(
expect.objectContaining({ success: true, authType: 'headers', toolCount: 0 })
)
expect(mockDetectMcpAuthType).not.toHaveBeenCalled()
expect(mockClientOptions).toHaveBeenCalledWith(
expect.objectContaining({
config: expect.objectContaining({
headers: { Authorization: 'Bearer static-api-token' },
}),
})
)
expect(mockConnect).toHaveBeenCalledTimes(1)
})
it('returns a header-auth failure when the configured token is rejected', async () => {
mockConnect.mockRejectedValueOnce(new Error('HTTP 401: Unauthorized'))
const response = await POST(createTestRequest({ Authorization: 'Bearer invalid-static-token' }))
const body = await response.json()
expect(response.status).toBe(400)
expect(body.data).toEqual(
expect.objectContaining({
success: false,
authType: 'headers',
error: 'HTTP 401: Unauthorized',
})
)
expect(mockDetectMcpAuthType).not.toHaveBeenCalled()
expect(mockConnect).toHaveBeenCalledTimes(1)
expect(mockDisconnect).toHaveBeenCalledTimes(1)
})
it('does not expose configured credentials echoed by an upstream error', async () => {
const token = 'opaque-static-token'
mockConnect.mockRejectedValueOnce(new Error(`Upstream rejected ${token}`))
const response = await POST(createTestRequest({ Authorization: `Bearer ${token}` }))
const body = await response.json()
expect(response.status).toBe(400)
expect(body.data).toEqual(
expect.objectContaining({ success: false, authType: 'headers', error: 'Connection failed' })
)
expect(mockLogger).toBeDefined()
expect(JSON.stringify(mockLogger?.warn.mock.calls)).not.toContain(token)
})
it('preserves OAuth discovery when no static headers are configured', async () => {
const response = await POST(createTestRequest())
const body = await response.json()
expect(response.status).toBe(200)
expect(body.data).toEqual(
expect.objectContaining({ success: false, authRequired: true, authType: 'oauth' })
)
expect(mockDetectMcpAuthType).toHaveBeenCalledWith('https://example.com/mcp', '203.0.113.10')
expect(mockClientOptions).not.toHaveBeenCalled()
})
it('preserves OAuth discovery when only supplemental headers are configured', async () => {
const response = await POST(createTestRequest({ 'X-Sim-Via': 'workflow' }))
const body = await response.json()
expect(response.status).toBe(200)
expect(body.data).toEqual(
expect.objectContaining({ success: false, authRequired: true, authType: 'oauth' })
)
expect(mockDetectMcpAuthType).toHaveBeenCalledWith('https://example.com/mcp', '203.0.113.10')
expect(mockClientOptions).not.toHaveBeenCalled()
})
it('blocks an env-resolved private URL before forwarding configured credentials', async () => {
const token = 'private-static-token'
mockResolveMcpConfigEnvVars.mockResolvedValueOnce({
config: {
id: 'test-request-1',
name: 'Dual Auth Server',
transport: 'streamable-http',
url: 'http://127.0.0.1/mcp',
headers: { Authorization: `Bearer ${token}` },
timeout: 10000,
retries: 1,
enabled: true,
},
missingVars: [],
})
mockValidateMcpServerSsrf.mockImplementation(async (url: string) => {
if (url === 'http://127.0.0.1/mcp') {
throw new MockMcpSsrfError('Private network targets are not allowed')
}
return '203.0.113.10'
})
const response = await POST(createTestRequest({ Authorization: `Bearer ${token}` }))
const responseText = await response.text()
expect(response.status).toBe(403)
expect(responseText).not.toContain(token)
expect(mockValidateMcpServerSsrf).toHaveBeenNthCalledWith(2, 'http://127.0.0.1/mcp')
expect(mockClientOptions).not.toHaveBeenCalled()
expect(mockDetectMcpAuthType).not.toHaveBeenCalled()
})
})
@@ -1,6 +1,5 @@
import { createLogger } from '@sim/logger'
import { toError } from '@sim/utils/errors'
import { truncate } from '@sim/utils/string'
import type { NextRequest } from 'next/server'
import { mcpServerTestBodySchema } from '@/lib/api/contracts/mcp'
import { withRouteHandler } from '@/lib/core/utils/with-route-handler'
@@ -50,17 +49,35 @@ interface TestConnectionResult {
}
/**
* Extracts a user-friendly error message from connection errors.
* Keeps diagnostic info (timeout, DNS, HTTP status) but strips
* verbose internals (Zod details, full response bodies, stack traces).
* Maps connection failures to allowlisted messages. Upstream response bodies
* may echo configured credentials, so arbitrary error text must not reach API
* responses or logs.
*/
function sanitizeConnectionError(error: unknown): string {
if (!(error instanceof Error)) {
return 'Unknown connection error'
}
const firstLine = error.message.split('\n')[0]
return truncate(firstLine, 200)
const message = error.message.toLowerCase()
if (message.includes('timeout') || message.includes('timed out')) {
return 'Connection timed out'
}
if (message.includes('401') || message.includes('unauthorized')) {
return 'HTTP 401: Unauthorized'
}
if (message.includes('403') || message.includes('forbidden')) {
return 'HTTP 403: Forbidden'
}
if (message.includes('enotfound') || message.includes('could not resolve')) {
return 'MCP server hostname could not be resolved'
}
if (message.includes('econnrefused') || message.includes('connection refused')) {
return 'Connection refused'
}
if (message.includes('certificate') || message.includes('tls') || message.includes('ssl')) {
return 'TLS connection failed'
}
return 'Connection failed'
}
/**
@@ -172,8 +189,14 @@ export const POST = withRouteHandler(
const result: TestConnectionResult = { success: false }
// Skip unauth connect when the server returns an RFC 9728 OAuth challenge.
if (testConfig.url) {
/** An explicit static Bearer token takes precedence over optional OAuth discovery. */
const hasStaticBearerToken = Object.entries(testConfig.headers ?? {}).some(
([name, value]) =>
name.toLowerCase() === 'authorization' && /^Bearer\s+\S+/i.test(value.trim())
)
if (hasStaticBearerToken) {
result.authType = 'headers'
} else if (testConfig.url) {
const detectedAuthType = await detectMcpAuthType(testConfig.url, resolvedIP)
if (detectedAuthType === 'oauth') {
result.authRequired = true
@@ -199,8 +222,8 @@ export const POST = withRouteHandler(
const tools = await client.listTools()
result.toolCount = tools.length
result.success = true
} catch (toolError) {
logger.warn(`[${requestId}] Connection established but could not list tools:`, toolError)
} catch {
logger.warn(`[${requestId}] Connection established but could not list tools`)
result.success = false
result.error = 'Connection established but could not list tools'
result.warnings = result.warnings || []
@@ -224,16 +247,15 @@ export const POST = withRouteHandler(
capabilities: result.supportedCapabilities,
})
} catch (error) {
logger.warn(`[${requestId}] MCP server test failed:`, error)
result.success = false
result.error = sanitizeConnectionError(error)
logger.warn(`[${requestId}] MCP server test failed`, { error: result.error })
} finally {
if (client) {
try {
await client.disconnect()
} catch (disconnectError) {
logger.debug(`[${requestId}] Test client disconnect error (expected):`, disconnectError)
} catch {
logger.debug(`[${requestId}] Test client disconnect error (expected)`)
}
}
}
@@ -23,6 +23,8 @@ import {
mcpServerIdParam,
mcpServerIdUrlKeys,
} from '@/app/workspace/[workspaceId]/settings/[section]/search-params'
import { getRefreshActionState } from '@/app/workspace/[workspaceId]/settings/components/mcp/refresh-action-state'
import { getServerToolsLabel } from '@/app/workspace/[workspaceId]/settings/components/mcp/server-tools-label'
import { RowActionsMenu } from '@/app/workspace/[workspaceId]/settings/components/row-actions-menu'
import { SettingsEmptyState } from '@/app/workspace/[workspaceId]/settings/components/settings-empty-state'
import { SettingsPanel } from '@/app/workspace/[workspaceId]/settings/components/settings-panel'
@@ -61,16 +63,6 @@ function formatTransportLabel(transport: string): string {
.join('-')
}
function formatToolsLabel(tools: McpTool[], connectionStatus?: string): string {
if (connectionStatus === 'error') {
return 'Unable to connect'
}
const count = tools.length
const plural = count !== 1 ? 's' : ''
const names = count > 0 ? `: ${tools.map((t) => t.name).join(', ')}` : ''
return `${count} tool${plural}${names}`
}
interface ServerListItemProps {
canManage: boolean
server: McpServer
@@ -93,8 +85,9 @@ function ServerListItem({
onViewDetails,
}: ServerListItemProps) {
const transportLabel = formatTransportLabel(server.transport || 'http')
const toolsLabel = formatToolsLabel(tools, server.connectionStatus)
const isError = server.connectionStatus === 'error'
const toolsLabel = getServerToolsLabel(tools, server.connectionStatus, server.lastError)
const hasConnectionIssue =
server.connectionStatus === 'error' || server.connectionStatus === 'disconnected'
return (
<div className='flex items-center justify-between gap-3'>
@@ -108,7 +101,7 @@ function ServerListItem({
<p
className={cn(
'truncate text-sm',
isError ? 'text-[var(--text-error)]' : 'text-[var(--text-muted)]'
hasConnectionIssue ? 'text-[var(--text-error)]' : 'text-[var(--text-muted)]'
)}
>
{isRefreshing
@@ -312,19 +305,11 @@ export function MCP() {
}
useEffect(() => {
if (!refreshServerMutation.isSuccess) return
if (!refreshServerMutation.isSuccess && !refreshServerMutation.isError) return
const timeout = window.setTimeout(() => refreshServerMutation.reset(), 3000)
return () => window.clearTimeout(timeout)
// eslint-disable-next-line react-hooks/exhaustive-deps -- mutation object is unstable; isSuccess flag is the trigger
}, [refreshServerMutation.isSuccess])
const refreshingServerId = refreshServerMutation.isPending
? refreshServerMutation.variables?.serverId
: null
const refreshedServerId = refreshServerMutation.isSuccess
? refreshServerMutation.variables?.serverId
: null
const refreshedWorkflowsUpdated = refreshServerMutation.data?.workflowsUpdated
// eslint-disable-next-line react-hooks/exhaustive-deps -- mutation object is unstable; status flags are the triggers
}, [refreshServerMutation.isSuccess, refreshServerMutation.isError])
const editingServer = editingServerId
? (servers.find((s) => s.id === editingServerId) as McpServer | undefined)
@@ -389,15 +374,12 @@ export function MCP() {
if (selectedServer) {
const { server, tools } = selectedServer
const transportLabel = formatTransportLabel(server.transport || 'http')
const refreshLabel =
refreshingServerId === server.id
? 'Refreshing...'
: refreshedServerId === server.id
? refreshedWorkflowsUpdated
? `Synced (${refreshedWorkflowsUpdated} workflow${refreshedWorkflowsUpdated === 1 ? '' : 's'})`
: 'Refreshed'
: 'Refresh tools'
const isCurrentRefresh = refreshServerMutation.variables?.serverId === server.id
const refreshAction = getRefreshActionState({
mutationStatus: isCurrentRefresh ? refreshServerMutation.status : 'idle',
connectionStatus: isCurrentRefresh ? refreshServerMutation.data?.status : undefined,
workflowsUpdated: isCurrentRefresh ? refreshServerMutation.data?.workflowsUpdated : undefined,
})
return (
<SettingsPanel
@@ -407,9 +389,10 @@ export function MCP() {
canEdit
? [
{
text: refreshLabel,
text: refreshAction.text,
textTone: refreshAction.textTone,
onSelect: () => handleRefreshServer(server.id),
disabled: refreshingServerId === server.id || refreshedServerId === server.id,
disabled: refreshAction.disabled,
},
{
text: 'Edit',
@@ -438,11 +421,11 @@ export function MCP() {
</div>
)}
{server.connectionStatus === 'error' && (
{server.connectionStatus !== 'connected' && (
<div className='flex flex-col gap-2'>
<span className='text-[var(--text-muted)] text-caption'>Status</span>
<p className='text-[var(--text-error)] text-sm'>
{server.lastError || 'Unable to connect'}
{getServerToolsLabel([], server.connectionStatus, server.lastError)}
</p>
</div>
)}
@@ -664,7 +647,10 @@ export function MCP() {
tools={tools}
isDeleting={deletingServers.has(server.id)}
isLoadingTools={isLoadingTools}
isRefreshing={refreshingServerId === server.id}
isRefreshing={
refreshServerMutation.isPending &&
refreshServerMutation.variables?.serverId === server.id
}
onRemove={() => handleRemoveServer(server.id)}
onViewDetails={() => handleViewDetails(server.id)}
/>
@@ -0,0 +1,47 @@
import { describe, expect, it } from 'vitest'
import { getRefreshActionState } from '@/app/workspace/[workspaceId]/settings/components/mcp/refresh-action-state'
describe('getRefreshActionState', () => {
it.each(['error', 'disconnected'] as const)(
'shows a retryable red-text Failed state when refresh returns %s',
(status) => {
expect(
getRefreshActionState({
mutationStatus: 'success',
connectionStatus: status,
workflowsUpdated: 0,
})
).toEqual({
text: 'Failed',
textTone: 'error',
disabled: false,
})
}
)
it('shows Failed when the refresh request itself rejects', () => {
expect(
getRefreshActionState({
mutationStatus: 'error',
})
).toEqual({
text: 'Failed',
textTone: 'error',
disabled: false,
})
})
it('preserves successful refresh feedback', () => {
expect(
getRefreshActionState({
mutationStatus: 'success',
connectionStatus: 'connected',
workflowsUpdated: 2,
})
).toEqual({
text: 'Synced (2 workflows)',
textTone: undefined,
disabled: true,
})
})
})
@@ -0,0 +1,37 @@
import type { MutationStatus } from '@tanstack/react-query'
import type { RefreshMcpServerResult } from '@/lib/api/contracts/mcp'
import type { SettingsAction } from '@/app/workspace/[workspaceId]/settings/components/settings-header/settings-header'
interface RefreshActionStateInput {
mutationStatus: MutationStatus
connectionStatus?: RefreshMcpServerResult['status']
workflowsUpdated?: number
}
type RefreshActionState = Pick<SettingsAction, 'text' | 'textTone' | 'disabled'>
export function getRefreshActionState({
mutationStatus,
connectionStatus,
workflowsUpdated,
}: RefreshActionStateInput): RefreshActionState {
if (mutationStatus === 'pending') {
return { text: 'Refreshing...', textTone: undefined, disabled: true }
}
if (
mutationStatus === 'error' ||
(mutationStatus === 'success' && connectionStatus !== 'connected')
) {
return { text: 'Failed', textTone: 'error', disabled: false }
}
if (mutationStatus === 'success') {
const text = workflowsUpdated
? `Synced (${workflowsUpdated} workflow${workflowsUpdated === 1 ? '' : 's'})`
: 'Refreshed'
return { text, textTone: undefined, disabled: true }
}
return { text: 'Refresh tools', textTone: undefined, disabled: false }
}
@@ -0,0 +1,26 @@
import { describe, expect, it } from 'vitest'
import { getServerToolsLabel } from '@/app/workspace/[workspaceId]/settings/components/mcp/server-tools-label'
describe('getServerToolsLabel', () => {
it('shows the persisted server error for errored connections', () => {
expect(getServerToolsLabel([], 'error', 'MCP error -32001: Request timed out')).toBe(
'MCP error -32001: Request timed out'
)
})
it('falls back when an errored connection has no persisted error', () => {
expect(getServerToolsLabel([], 'error', null)).toBe('Unable to connect')
})
it('shows a disconnected state when OAuth was not completed', () => {
expect(getServerToolsLabel([], 'disconnected', null)).toBe('Not Connected')
})
it('shows the persisted error for disconnected connections', () => {
expect(getServerToolsLabel([], 'disconnected', 'Request timed out')).toBe('Request timed out')
})
it('continues showing discovered tools for healthy connections', () => {
expect(getServerToolsLabel([{ name: 'search' }], 'connected', null)).toBe('1 tool: search')
})
})
@@ -0,0 +1,24 @@
import type { McpServer } from '@/lib/api/contracts/mcp'
interface NamedTool {
name: string
}
export function getServerToolsLabel(
tools: NamedTool[],
connectionStatus?: McpServer['connectionStatus'],
lastError?: McpServer['lastError']
): string {
if (connectionStatus === 'error') {
return lastError?.trim() || 'Unable to connect'
}
if (connectionStatus === 'disconnected') {
return lastError?.trim() || 'Not Connected'
}
const count = tools.length
const plural = count !== 1 ? 's' : ''
const names = count > 0 ? `: ${tools.map((tool) => tool.name).join(', ')}` : ''
return `${count} tool${plural}${names}`
}
@@ -19,6 +19,7 @@ const useIsomorphicLayoutEffect = typeof window === 'undefined' ? useEffect : us
export interface SettingsAction {
text: string
textTone?: 'error'
icon?: ComponentType<{ className?: string }>
variant?: 'primary' | 'destructive'
active?: boolean
@@ -69,6 +70,7 @@ function computeSignature(config: SettingsHeaderConfig): string {
back: config.back ? [config.back.text, config.back.icon ? 1 : 0] : null,
actions: config.actions?.map((action) => [
action.text,
action.textTone ?? '',
action.variant ?? '',
action.active ?? false,
action.disabled ?? false,
@@ -155,7 +157,11 @@ export function SettingsHeaderShell({ children }: { children: ReactNode }) {
}
disabled={action.disabled}
>
{action.text}
{action.textTone === 'error' ? (
<span className='text-[var(--text-error)]'>{action.text}</span>
) : (
action.text
)}
</Chip>
)
return action.tooltip ? (
+212
View File
@@ -0,0 +1,212 @@
/**
* @vitest-environment jsdom
*/
import { act, type ReactNode } from 'react'
import { sleep } from '@sim/utils/helpers'
import { QueryClient, QueryClientProvider } from '@tanstack/react-query'
import { createRoot, type Root } from 'react-dom/client'
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
const { mockRequestJson } = vi.hoisted(() => ({
mockRequestJson: vi.fn(),
}))
vi.mock('@/lib/api/client/request', () => ({
requestJson: mockRequestJson,
}))
import {
discoverMcpToolsContract,
listMcpServersContract,
type McpServer,
} from '@/lib/api/contracts/mcp'
import { useForceRefreshMcpTools, useMcpServers, useMcpToolsQuery } from '@/hooks/queries/mcp'
const WORKSPACE_ID = 'workspace-1'
function server(id: string, overrides: Partial<McpServer> = {}): McpServer {
return {
id,
workspaceId: WORKSPACE_ID,
name: id,
transport: 'streamable-http',
url: `https://${id}.example.com/mcp`,
enabled: true,
connectionStatus: 'connected',
createdAt: '2026-01-01T00:00:00.000Z',
updatedAt: '2026-01-01T00:00:00.000Z',
...overrides,
}
}
function renderHookWithClient<T>(useHook: () => T): {
getResult: () => T
unmount: () => void
} {
;(globalThis as { IS_REACT_ACT_ENVIRONMENT?: boolean }).IS_REACT_ACT_ENVIRONMENT = true
const queryClient = new QueryClient({
defaultOptions: { queries: { retry: false } },
})
const container = document.createElement('div')
const root: Root = createRoot(container)
let result: T | undefined
function Probe() {
result = useHook()
return null
}
function Wrapper({ children }: { children: ReactNode }) {
return <QueryClientProvider client={queryClient}>{children}</QueryClientProvider>
}
act(() => {
root.render(
<Wrapper>
<Probe />
</Wrapper>
)
})
return {
getResult: () => {
if (result === undefined) throw new Error('Hook result is not ready')
return result
},
unmount: () => act(() => root.unmount()),
}
}
async function flush() {
await act(async () => {
for (let i = 0; i < 5; i++) {
await Promise.resolve()
await sleep(1)
}
})
}
function mockServers(servers: McpServer[]) {
mockRequestJson.mockImplementation(async (contract) => {
if (contract === listMcpServersContract) {
return { success: true, data: { servers } }
}
if (contract === discoverMcpToolsContract) {
return { success: true, data: { tools: [], totalCount: 0, byServer: {} } }
}
throw new Error('Unexpected MCP request')
})
}
describe('useMcpToolsQuery', () => {
beforeEach(() => {
vi.clearAllMocks()
})
afterEach(() => {
vi.restoreAllMocks()
})
it('does not auto-discover disconnected or errored OAuth servers', async () => {
mockServers([
server('oauth-disconnected', { authType: 'oauth', connectionStatus: 'disconnected' }),
server('oauth-error', { authType: 'oauth', connectionStatus: 'error' }),
])
const { unmount } = renderHookWithClient(() => useMcpToolsQuery(WORKSPACE_ID))
await flush()
expect(mockRequestJson).toHaveBeenCalledTimes(1)
expect(mockRequestJson).toHaveBeenCalledWith(
listMcpServersContract,
expect.objectContaining({ query: { workspaceId: WORKSPACE_ID } })
)
unmount()
})
it('continues discovering connected OAuth and disconnected non-OAuth servers', async () => {
mockServers([
server('oauth-connected', { authType: 'oauth', connectionStatus: 'connected' }),
server('headers-disconnected', { authType: 'headers', connectionStatus: 'disconnected' }),
])
const { unmount } = renderHookWithClient(() => useMcpToolsQuery(WORKSPACE_ID))
await flush()
expect(mockRequestJson).toHaveBeenCalledTimes(3)
expect(mockRequestJson).toHaveBeenCalledWith(
discoverMcpToolsContract,
expect.objectContaining({
query: { workspaceId: WORKSPACE_ID, serverId: 'oauth-connected' },
})
)
expect(mockRequestJson).toHaveBeenCalledWith(
discoverMcpToolsContract,
expect.objectContaining({
query: { workspaceId: WORKSPACE_ID, serverId: 'headers-disconnected' },
})
)
unmount()
})
it('refreshes the server list after a connected OAuth discovery fails', async () => {
let serverListRequests = 0
mockRequestJson.mockImplementation(async (contract) => {
if (contract === listMcpServersContract) {
serverListRequests++
const connectionStatus = serverListRequests === 1 ? 'connected' : 'disconnected'
return {
success: true,
data: {
servers: [server('oauth-server', { authType: 'oauth', connectionStatus })],
},
}
}
if (contract === discoverMcpToolsContract) {
throw new Error('OAuth authorization required')
}
throw new Error('Unexpected MCP request')
})
const { unmount } = renderHookWithClient(() => useMcpToolsQuery(WORKSPACE_ID))
await flush()
expect(serverListRequests).toBe(2)
expect(
mockRequestJson.mock.calls.filter(([contract]) => contract === discoverMcpToolsContract)
).toHaveLength(1)
unmount()
})
it('does not force-refresh disconnected OAuth servers', async () => {
mockServers([
server('oauth-disconnected', { authType: 'oauth', connectionStatus: 'disconnected' }),
server('headers-connected', { authType: 'headers', connectionStatus: 'connected' }),
])
const { getResult, unmount } = renderHookWithClient(() => ({
servers: useMcpServers(WORKSPACE_ID),
refresh: useForceRefreshMcpTools(),
}))
await flush()
await act(async () => {
await getResult().refresh.mutateAsync(WORKSPACE_ID)
})
const discoveryCalls = mockRequestJson.mock.calls.filter(
([contract]) => contract === discoverMcpToolsContract
)
expect(discoveryCalls).toHaveLength(1)
expect(discoveryCalls[0]?.[1]).toEqual(
expect.objectContaining({
query: { workspaceId: WORKSPACE_ID, refresh: true, serverId: 'headers-connected' },
})
)
unmount()
})
})
+28 -6
View File
@@ -125,20 +125,31 @@ async function fetchMcpTools(
}
}
function isServerEligibleForDiscovery(server: McpServer, workspaceId: string): boolean {
return (
server.enabled &&
server.workspaceId === workspaceId &&
(server.authType !== 'oauth' || server.connectionStatus === 'connected')
)
}
/**
* Workspace aggregate derived from N parallel per-server queries via
* `useQueries`. One slow server cannot block the others.
*/
export function useMcpToolsQuery(workspaceId: string) {
const queryClient = useQueryClient()
const { data: servers, isLoading: serversLoading } = useMcpServers(workspaceId)
// Skip disabled rows (would 404 → negative-cache) and rows from a previous
// workspace (keepPreviousData on useMcpServers).
/**
* Skip disabled rows, rows retained from a previous workspace, and OAuth rows
* that require explicit authorization before discovery can succeed.
*/
const serverIds = useMemo(
() =>
servers
? servers
.filter((s) => s.enabled && s.workspaceId === workspaceId)
.filter((server) => isServerEligibleForDiscovery(server, workspaceId))
.map((s) => s.id)
.sort()
: [],
@@ -148,8 +159,17 @@ export function useMcpToolsQuery(workspaceId: string) {
const results = useQueries({
queries: serverIds.map((serverId) => ({
queryKey: mcpKeys.serverToolsList(workspaceId, serverId),
queryFn: ({ signal }: { signal?: AbortSignal }) =>
fetchMcpTools(workspaceId, false, signal, serverId),
queryFn: async ({ signal }: { signal?: AbortSignal }) => {
try {
return await fetchMcpTools(workspaceId, false, signal, serverId)
} catch (error) {
await queryClient.invalidateQueries(
{ queryKey: mcpKeys.serversList(workspaceId) },
{ cancelRefetch: false }
)
throw error
}
},
enabled: !!workspaceId,
retry: false,
staleTime: MCP_SERVER_TOOLS_STALE_TIME,
@@ -203,7 +223,9 @@ export function useForceRefreshMcpTools() {
mutationFn: async (workspaceId: string) => {
const allServers =
queryClient.getQueryData<McpServer[]>(mcpKeys.serversList(workspaceId)) ?? []
const servers = allServers.filter((s) => s.enabled && s.workspaceId === workspaceId)
const servers = allServers.filter((server) =>
isServerEligibleForDiscovery(server, workspaceId)
)
const results = await Promise.allSettled(
servers.map(async (server) => {
const tools = await fetchMcpTools(workspaceId, true, undefined, server.id)
+109 -1
View File
@@ -3,6 +3,20 @@
*/
import { beforeEach, describe, expect, it, vi } from 'vitest'
const { mockLogger, mockSdkConnect } = vi.hoisted(() => ({
mockLogger: {
debug: vi.fn(),
error: vi.fn(),
info: vi.fn(),
warn: vi.fn(),
},
mockSdkConnect: vi.fn().mockResolvedValue(undefined),
}))
vi.mock('@sim/logger', () => ({
createLogger: () => mockLogger,
}))
/**
* Capture the notification handler registered via `client.setNotificationHandler()`.
* This lets us simulate the MCP SDK delivering a `tools/list_changed` notification.
@@ -14,7 +28,7 @@ vi.mock('@modelcontextprotocol/sdk/client/index.js', () => ({
class {
constructor() {
Object.assign(this, {
connect: vi.fn().mockResolvedValue(undefined),
connect: mockSdkConnect,
close: vi.fn().mockResolvedValue(undefined),
getServerVersion: vi.fn().mockReturnValue('2025-06-18'),
getServerCapabilities: vi.fn().mockReturnValue({ tools: { listChanged: true } }),
@@ -65,6 +79,7 @@ describe('McpClient notification handler', () => {
beforeEach(() => {
capturedNotificationHandler = null
vi.clearAllMocks()
mockSdkConnect.mockResolvedValue(undefined)
})
it('fires onToolsChanged when a notification arrives while connected', async () => {
@@ -115,6 +130,99 @@ describe('McpClient notification handler', () => {
expect(capturedNotificationHandler).toBeNull()
})
it('uses the server connection timeout for the initialize request', async () => {
const client = new McpClient({
config: { ...createConfig(), timeout: 12_345 },
securityPolicy: { requireConsent: false, auditLevel: 'basic' },
})
await client.connect()
expect(mockSdkConnect).toHaveBeenCalledWith(expect.anything(), { timeout: 12_345 })
})
it('normalizes invalid connection timeouts before calling the SDK', async () => {
const client = new McpClient({
config: { ...createConfig(), timeout: -1 },
securityPolicy: { requireConsent: false, auditLevel: 'basic' },
})
await client.connect()
expect(mockSdkConnect).toHaveBeenCalledWith(expect.anything(), { timeout: 30_000 })
})
it('logs connection diagnostics without header values', async () => {
const client = new McpClient({
config: {
...createConfig(),
authType: 'headers',
headers: { Authorization: 'Bearer do-not-log', 'X-API-Key': 'also-secret' },
timeout: 12_345,
},
securityPolicy: { requireConsent: false, auditLevel: 'basic' },
})
await client.connect()
expect(mockLogger.info).toHaveBeenCalledWith(
expect.stringContaining('Successfully connected'),
expect.objectContaining({
authType: 'headers',
headerNames: ['Authorization', 'X-API-Key'],
hasUnresolvedEnvRefs: false,
phase: 'initialize',
outcome: 'connected',
timeoutMs: 12_345,
})
)
expect(JSON.stringify(mockLogger.info.mock.calls)).not.toContain('do-not-log')
expect(JSON.stringify(mockLogger.info.mock.calls)).not.toContain('also-secret')
})
it('classifies initialize timeouts in connection diagnostics', async () => {
mockSdkConnect.mockRejectedValueOnce(new Error('MCP error -32001: Request timed out'))
const client = new McpClient({
config: {
...createConfig(),
headers: { Authorization: 'Bearer do-not-log' },
},
securityPolicy: { requireConsent: false, auditLevel: 'basic' },
})
await expect(client.connect()).rejects.toThrow('Request timed out')
expect(mockLogger.error).toHaveBeenCalledWith(
expect.stringContaining('Failed to connect'),
expect.objectContaining({
phase: 'initialize',
outcome: 'timeout',
timeoutMs: 30_000,
error: expect.objectContaining({
name: 'Error',
}),
})
)
expect(JSON.stringify(mockLogger.error.mock.calls)).not.toContain('do-not-log')
})
it('does not log opaque credentials echoed by MCP errors', async () => {
const secret = 'opaque-credential-without-a-known-prefix'
mockSdkConnect.mockRejectedValueOnce(new Error(`Upstream rejected ${secret}`))
const client = new McpClient({
config: {
...createConfig(),
authType: 'headers',
headers: { 'X-Custom-Credential': secret },
},
securityPolicy: { requireConsent: false, auditLevel: 'basic' },
})
await expect(client.connect()).rejects.toThrow('Upstream rejected')
expect(JSON.stringify(mockLogger.error.mock.calls)).not.toContain(secret)
})
it('passes configured headers for OAuth transports as well as header auth transports', () => {
const authProvider = {} as unknown as NonNullable<McpClientOptions['authProvider']>
new McpClient({
+69 -6
View File
@@ -9,9 +9,10 @@ import {
ToolListChangedNotificationSchema,
} from '@modelcontextprotocol/sdk/types.js'
import { createLogger } from '@sim/logger'
import { getErrorMessage } from '@sim/utils/errors'
import { describeError, getErrorMessage } from '@sim/utils/errors'
import { getMaxExecutionTimeout } from '@/lib/core/execution-limits'
import { createPinnedFetch } from '@/lib/core/security/input-validation.server'
import { sanitizeForLogging } from '@/lib/core/security/redaction'
import { McpOauthRedirectRequired } from '@/lib/mcp/oauth'
import {
type McpClientOptions,
@@ -29,9 +30,30 @@ import {
type McpVersionInfo,
} from '@/lib/mcp/types'
import { MCP_CLIENT_CONSTANTS } from '@/lib/mcp/utils'
import { createEnvVarPattern } from '@/executor/utils/reference-validation'
const logger = createLogger('McpClient')
type ConnectionOutcome =
| 'started'
| 'connected'
| 'authorization_required'
| 'timeout'
| 'unauthorized'
| 'cancelled'
| 'error'
function classifyConnectionOutcome(error: unknown): ConnectionOutcome {
if (error instanceof McpOauthRedirectRequired || error instanceof UnauthorizedError) {
return 'authorization_required'
}
const message = getErrorMessage(error, '').toLowerCase()
if (message.includes('connection attempt cancelled')) return 'cancelled'
if (message.includes('timeout') || message.includes('timed out')) return 'timeout'
if (message.includes('401') || message.includes('unauthorized')) return 'unauthorized'
return 'error'
}
interface McpClientConnectOptions {
isCancelled?: () => boolean
}
@@ -90,10 +112,34 @@ export class McpClient {
* for `notifications/tools/list_changed` after connecting.
*/
async connect(options: McpClientConnectOptions = {}): Promise<void> {
logger.info(`Connecting to MCP server: ${this.config.name} (${this.config.transport})`)
const startedAt = Date.now()
const configuredTimeout = this.config.timeout
const timeoutMs =
configuredTimeout !== undefined && Number.isFinite(configuredTimeout) && configuredTimeout > 0
? Math.min(Math.floor(configuredTimeout), getMaxExecutionTimeout())
: MCP_CLIENT_CONSTANTS.CLIENT_TIMEOUT
const headerNames = Object.keys(this.config.headers ?? {}).sort()
const hasUnresolvedEnvRefs = [
this.config.url ?? '',
...Object.values(this.config.headers ?? {}),
].some((value) => createEnvVarPattern().test(value))
const diagnostics = {
serverId: this.config.id,
authType: this.config.authType ?? (headerNames.length > 0 ? 'headers' : 'none'),
headerNames,
hasUnresolvedEnvRefs,
phase: 'initialize',
timeoutMs,
}
logger.info(`Connecting to MCP server: ${this.config.name} (${this.config.transport})`, {
...diagnostics,
outcome: 'started' satisfies ConnectionOutcome,
})
try {
await this.client.connect(this.transport)
await this.client.connect(this.transport, {
timeout: timeoutMs,
})
if (options.isCancelled?.()) {
await this.client.close().catch((error) => {
logger.warn(`Error closing cancelled connection to ${this.config.name}:`, error)
@@ -116,17 +162,34 @@ export class McpClient {
const serverVersion = this.client.getServerVersion()
logger.info(`Successfully connected to MCP server: ${this.config.name}`, {
...diagnostics,
durationMs: Date.now() - startedAt,
outcome: 'connected' satisfies ConnectionOutcome,
protocolVersion: serverVersion,
})
} catch (error) {
this.isConnected = false
if (error instanceof McpOauthRedirectRequired || error instanceof UnauthorizedError) {
const errorMessage = getErrorMessage(error, 'Unknown error')
const describedError = describeError(error)
const outcome = classifyConnectionOutcome(error)
logger.error(`Failed to connect to MCP server ${this.config.name}`, {
...diagnostics,
durationMs: Date.now() - startedAt,
error: {
name: sanitizeForLogging(describedError.name, 100),
code: describedError.code ? sanitizeForLogging(describedError.code, 100) : undefined,
errno: describedError.errno ? sanitizeForLogging(describedError.errno, 100) : undefined,
syscall: describedError.syscall
? sanitizeForLogging(describedError.syscall, 100)
: undefined,
},
outcome,
})
if (outcome === 'authorization_required') {
this.connectionStatus.lastError = undefined
throw error
}
const errorMessage = getErrorMessage(error, 'Unknown error')
this.connectionStatus.lastError = errorMessage
logger.error(`Failed to connect to MCP server ${this.config.name}:`, error)
throw new McpConnectionError(errorMessage, this.config.name)
}
}
+220 -10
View File
@@ -1,6 +1,9 @@
/**
* @vitest-environment node
*/
import { UnauthorizedError } from '@modelcontextprotocol/sdk/client/auth.js'
import { loggerMock } from '@sim/testing'
import { beforeEach, describe, expect, it, vi } from 'vitest'
const {
@@ -14,16 +17,19 @@ const {
mockValidateSsrf,
mockIsDomainAllowed,
mockCacheAdapter,
mockUpdateSet,
mockUpdateReturning,
} = vi.hoisted(() => {
const mockListTools = vi.fn()
const mockConnect = vi.fn()
const mockDisconnect = vi.fn()
const mockUpdateReturning = vi.fn().mockResolvedValue([{ id: 'server-1' }])
// In-memory cache adapter so the service never touches the real Redis the
// local .env points at (unreachable in CI/sandbox → hangs). Honors TTL via
// an expiry timestamp so negative-cache assertions behave like production.
const cacheStore = new Map<string, { tools: unknown[]; expiry: number }>()
const mockCacheAdapter = {
get: async (key: string) => {
get: vi.fn(async (key: string) => {
const entry = cacheStore.get(key)
if (!entry) return null
if (entry.expiry <= Date.now()) {
@@ -31,16 +37,16 @@ const {
return null
}
return entry
},
set: async (key: string, tools: unknown[], ttlMs: number) => {
}),
set: vi.fn(async (key: string, tools: unknown[], ttlMs: number) => {
cacheStore.set(key, { tools, expiry: Date.now() + ttlMs })
},
delete: async (key: string) => {
}),
delete: vi.fn(async (key: string) => {
cacheStore.delete(key)
},
clear: async () => {
}),
clear: vi.fn(async () => {
cacheStore.clear()
},
}),
dispose: () => {},
}
return {
@@ -67,6 +73,10 @@ const {
mockValidateDomain: vi.fn(),
mockValidateSsrf: vi.fn(),
mockIsDomainAllowed: vi.fn(() => true),
mockUpdateReturning,
mockUpdateSet: vi.fn().mockReturnValue({
where: vi.fn().mockReturnValue({ returning: mockUpdateReturning }),
}),
}
})
@@ -80,13 +90,12 @@ vi.mock('@sim/db', () => {
})
return thenable
}
const setter = vi.fn().mockReturnValue({ where: vi.fn().mockResolvedValue(undefined) })
return {
db: {
select: vi.fn().mockReturnValue({
from: vi.fn().mockReturnValue({ where }),
}),
update: vi.fn().mockReturnValue({ set: setter }),
update: vi.fn().mockReturnValue({ set: mockUpdateSet }),
insert: vi.fn(),
delete: vi.fn(),
},
@@ -126,6 +135,8 @@ vi.mock('@/lib/mcp/storage', () => ({
import { mcpService } from '@/lib/mcp/service'
import { McpOauthAuthorizationRequiredError } from '@/lib/mcp/types'
const mockLogger = vi.mocked(loggerMock.createLogger).mock.results.at(-1)?.value
const WORKSPACE_ID = 'workspace-test'
const USER_ID = 'user-test'
@@ -173,6 +184,8 @@ describe('McpService.discoverTools per-server caching', () => {
)
mockConnect.mockResolvedValue(undefined)
mockDisconnect.mockResolvedValue(undefined)
mockUpdateReturning.mockReset()
mockUpdateReturning.mockResolvedValue([{ id: 'server-1' }])
// The McpService singleton holds cache state across imports.
await mcpService.clearCache()
})
@@ -321,6 +334,70 @@ describe('McpService.discoverTools per-server caching', () => {
expect(mockListTools).not.toHaveBeenCalled()
})
it('persists and negative-caches UnauthorizedError for a headers-auth server', async () => {
const reflectedCredential = 'Bearer static-secret-for-bulk-discovery'
mockGetWorkspaceServersRows.mockResolvedValue([
dbRow('mcp-a', 'A', {
statusConfig: { consecutiveFailures: 0, lastSuccessfulDiscovery: null },
}),
])
mockListTools.mockRejectedValueOnce(
new UnauthorizedError(`Rejected Authorization: ${reflectedCredential}`)
)
const first = await mcpService.discoverTools(USER_ID, WORKSPACE_ID)
expect(first).toEqual([])
await vi.waitFor(() => {
expect(mockUpdateSet).toHaveBeenCalledWith(
expect.objectContaining({
connectionStatus: 'disconnected',
lastError: 'Authentication failed',
statusConfig: { consecutiveFailures: 1, lastSuccessfulDiscovery: null },
})
)
expect(mockCacheAdapter.set).toHaveBeenCalledWith(
`workspace:${WORKSPACE_ID}:server:mcp-a:failure`,
[],
expect.any(Number)
)
})
expect(JSON.stringify(mockUpdateSet.mock.calls)).not.toContain(reflectedCredential)
expect(JSON.stringify(mockCacheAdapter.set.mock.calls)).not.toContain(reflectedCredential)
expect(JSON.stringify(mockLogger?.warn.mock.calls)).not.toContain(reflectedCredential)
mockListTools.mockClear()
const second = await mcpService.discoverTools(USER_ID, WORKSPACE_ID)
expect(second).toEqual([])
expect(mockListTools).not.toHaveBeenCalled()
})
it('keeps UnauthorizedError soft-pending for an OAuth server', async () => {
mockGetWorkspaceServersRows.mockResolvedValue([dbRow('mcp-a', 'A', { authType: 'oauth' })])
mockResolveEnvVars.mockRejectedValue(new UnauthorizedError('OAuth token rejected'))
const first = await mcpService.discoverTools(USER_ID, WORKSPACE_ID)
expect(first).toEqual([])
await vi.waitFor(() => {
expect(mockUpdateSet).toHaveBeenCalledWith(
expect.objectContaining({
connectionStatus: 'disconnected',
lastError: null,
})
)
})
expect(mockCacheAdapter.set).not.toHaveBeenCalledWith(
`workspace:${WORKSPACE_ID}:server:mcp-a:failure`,
[],
expect.any(Number)
)
mockResolveEnvVars.mockClear()
await mcpService.discoverTools(USER_ID, WORKSPACE_ID)
expect(mockResolveEnvVars).toHaveBeenCalledTimes(1)
})
it('successful discoverServerTools clears the negative cache', async () => {
mockGetWorkspaceServersRows.mockResolvedValue([dbRow('mcp-a', 'A')])
mockListTools.mockRejectedValueOnce(new Error('Request timed out'))
@@ -366,4 +443,137 @@ describe('McpService.discoverTools per-server caching', () => {
expect(after.map((t) => t.name)).toEqual(['a1'])
expect(mockListTools).toHaveBeenCalledTimes(1)
})
it('persists a per-server discovery failure before rethrowing it', async () => {
mockGetWorkspaceServersRows.mockResolvedValue([
dbRow('mcp-a', 'A', {
statusConfig: { consecutiveFailures: 0, lastSuccessfulDiscovery: null },
}),
])
mockListTools.mockRejectedValueOnce(new Error('Request timed out'))
await expect(mcpService.discoverServerTools(USER_ID, 'mcp-a', WORKSPACE_ID)).rejects.toThrow(
'Request timed out'
)
expect(mockUpdateSet).toHaveBeenCalledWith(
expect.objectContaining({
connectionStatus: 'disconnected',
lastError: 'Request timed out',
statusConfig: { consecutiveFailures: 1, lastSuccessfulDiscovery: null },
})
)
})
it('persists and negative-caches per-server UnauthorizedError for headers auth', async () => {
const reflectedCredential = 'Bearer static-secret-for-server-discovery'
mockGetWorkspaceServersRows.mockResolvedValue([
dbRow('mcp-a', 'A', {
statusConfig: { consecutiveFailures: 0, lastSuccessfulDiscovery: null },
}),
])
mockListTools.mockRejectedValueOnce(
new UnauthorizedError(`Rejected Authorization: ${reflectedCredential}`)
)
await expect(mcpService.discoverServerTools(USER_ID, 'mcp-a', WORKSPACE_ID)).rejects.toThrow(
reflectedCredential
)
expect(mockUpdateSet).toHaveBeenCalledWith(
expect.objectContaining({
connectionStatus: 'disconnected',
lastError: 'Authentication failed',
statusConfig: { consecutiveFailures: 1, lastSuccessfulDiscovery: null },
})
)
expect(JSON.stringify(mockUpdateSet.mock.calls)).not.toContain(reflectedCredential)
expect(JSON.stringify(mockCacheAdapter.set.mock.calls)).not.toContain(reflectedCredential)
expect(JSON.stringify(mockLogger?.warn.mock.calls)).not.toContain(reflectedCredential)
mockListTools.mockClear()
await expect(mcpService.discoverServerTools(USER_ID, 'mcp-a', WORKSPACE_ID)).rejects.toThrow(
'cooldown'
)
expect(mockListTools).not.toHaveBeenCalled()
})
it('keeps per-server UnauthorizedError soft-pending for OAuth auth', async () => {
mockGetWorkspaceServersRows.mockResolvedValue([dbRow('mcp-a', 'A', { authType: 'oauth' })])
mockResolveEnvVars.mockRejectedValue(new UnauthorizedError('OAuth token rejected'))
await expect(mcpService.discoverServerTools(USER_ID, 'mcp-a', WORKSPACE_ID)).rejects.toThrow(
'OAuth token rejected'
)
expect(mockUpdateSet).toHaveBeenCalledWith(
expect.objectContaining({
connectionStatus: 'disconnected',
lastError: null,
})
)
expect(mockCacheAdapter.set).not.toHaveBeenCalledWith(
`workspace:${WORKSPACE_ID}:server:mcp-a:failure`,
[],
expect.any(Number)
)
mockResolveEnvVars.mockClear()
await expect(mcpService.discoverServerTools(USER_ID, 'mcp-a', WORKSPACE_ID)).rejects.toThrow(
'OAuth token rejected'
)
expect(mockResolveEnvVars).toHaveBeenCalledTimes(1)
})
it('promotes the persisted server status to error on the third consecutive failure', async () => {
mockGetWorkspaceServersRows.mockResolvedValue([
dbRow('mcp-a', 'A', {
statusConfig: { consecutiveFailures: 2, lastSuccessfulDiscovery: null },
}),
])
mockListTools.mockRejectedValueOnce(new Error('Connection refused'))
await expect(mcpService.discoverServerTools(USER_ID, 'mcp-a', WORKSPACE_ID)).rejects.toThrow(
'Connection refused'
)
expect(mockUpdateSet).toHaveBeenCalledWith(
expect.objectContaining({
connectionStatus: 'error',
statusConfig: { consecutiveFailures: 3, lastSuccessfulDiscovery: null },
})
)
})
it('persists OAuth-required discovery as disconnected without a failure error', async () => {
mockGetWorkspaceServersRows.mockResolvedValue([dbRow('mcp-a', 'A')])
mockListTools.mockRejectedValueOnce(new McpOauthAuthorizationRequiredError('mcp-a', 'A'))
await expect(mcpService.discoverServerTools(USER_ID, 'mcp-a', WORKSPACE_ID)).rejects.toThrow(
'OAuth authorization required'
)
expect(mockUpdateSet).toHaveBeenCalledWith(
expect.objectContaining({
connectionStatus: 'disconnected',
lastError: null,
})
)
})
it('does not negative-cache a failure older than a successful discovery', async () => {
mockGetWorkspaceServersRows.mockResolvedValue([dbRow('mcp-a', 'A')])
mockListTools.mockRejectedValueOnce(new Error('Older request failed'))
mockUpdateReturning.mockResolvedValueOnce([])
await expect(mcpService.discoverServerTools(USER_ID, 'mcp-a', WORKSPACE_ID)).rejects.toThrow(
'Older request failed'
)
mockListTools.mockResolvedValueOnce([tool('a1', 'mcp-a')])
const tools = await mcpService.discoverServerTools(USER_ID, 'mcp-a', WORKSPACE_ID)
expect(tools.map((tool) => tool.name)).toEqual(['a1'])
expect(mockListTools).toHaveBeenCalledTimes(2)
})
})
+170 -84
View File
@@ -5,7 +5,7 @@ import { mcpServers } from '@sim/db/schema'
import { createLogger } from '@sim/logger'
import { getErrorMessage } from '@sim/utils/errors'
import { sleep } from '@sim/utils/helpers'
import { and, eq, isNull } from 'drizzle-orm'
import { and, eq, isNull, lte, or } from 'drizzle-orm'
import { isTest } from '@/lib/core/config/env-flags'
import { generateRequestId } from '@/lib/core/utils/request'
import { McpClient } from '@/lib/mcp/client'
@@ -66,6 +66,28 @@ type DiscoveryOutcome =
// exemption survives the getErrorMessage call.
| { kind: 'error'; message: string; originalError: unknown }
type ServerStatusUpdate =
| { outcome: 'connected'; toolCount: number }
| { outcome: 'failed'; error: string; discoveryStartedAt?: Date }
function isOauthAuthorizationError(error: unknown, authType: McpServerConfig['authType']): boolean {
return (
error instanceof McpOauthAuthorizationRequiredError ||
(authType === 'oauth' && error instanceof UnauthorizedError)
)
}
function getDiscoveryFailureMessage(
error: unknown,
authType: McpServerConfig['authType'],
fallback: string
): string {
if (authType !== 'oauth' && error instanceof UnauthorizedError) {
return 'Authentication failed'
}
return getErrorMessage(error, fallback)
}
class McpService {
private cacheAdapter: McpCacheStorageAdapter
private readonly cacheTimeout = MCP_CONSTANTS.CACHE_TIMEOUT
@@ -311,11 +333,30 @@ class McpService {
private async updateServerStatus(
serverId: string,
workspaceId: string,
success: boolean,
error?: string,
toolCount?: number
): Promise<void> {
update: ServerStatusUpdate
): Promise<boolean> {
try {
const now = new Date()
if (update.outcome === 'connected') {
await db
.update(mcpServers)
.set({
connectionStatus: 'connected',
lastConnected: now,
lastError: null,
toolCount: update.toolCount,
lastToolsRefresh: now,
statusConfig: {
consecutiveFailures: 0,
lastSuccessfulDiscovery: now.toISOString(),
},
updatedAt: now,
})
.where(eq(mcpServers.id, serverId))
return true
}
const [currentServer] = await db
.select({ statusConfig: mcpServers.statusConfig })
.from(mcpServers)
@@ -337,49 +378,42 @@ class McpService {
lastSuccessfulDiscovery: storedConfig?.lastSuccessfulDiscovery ?? null,
}
const now = new Date()
const newFailures = currentConfig.consecutiveFailures + 1
const isErrorState = newFailures >= MCP_CONSTANTS.MAX_CONSECUTIVE_FAILURES
if (success) {
await db
.update(mcpServers)
.set({
connectionStatus: 'connected',
lastConnected: now,
lastError: null,
toolCount: toolCount ?? 0,
lastToolsRefresh: now,
statusConfig: {
consecutiveFailures: 0,
lastSuccessfulDiscovery: now.toISOString(),
},
updatedAt: now,
})
.where(eq(mcpServers.id, serverId))
} else {
const newFailures = currentConfig.consecutiveFailures + 1
const isErrorState = newFailures >= MCP_CONSTANTS.MAX_CONSECUTIVE_FAILURES
await db
.update(mcpServers)
.set({
connectionStatus: isErrorState ? 'error' : 'disconnected',
lastError: error || 'Unknown error',
statusConfig: {
consecutiveFailures: newFailures,
lastSuccessfulDiscovery: currentConfig.lastSuccessfulDiscovery,
},
updatedAt: now,
})
.where(eq(mcpServers.id, serverId))
if (isErrorState) {
logger.warn(
`Server ${serverId} marked as error after ${newFailures} consecutive failures`
const updatedServers = await db
.update(mcpServers)
.set({
connectionStatus: isErrorState ? 'error' : 'disconnected',
lastError: update.error || 'Unknown error',
statusConfig: {
consecutiveFailures: newFailures,
lastSuccessfulDiscovery: currentConfig.lastSuccessfulDiscovery,
},
updatedAt: now,
})
.where(
and(
eq(mcpServers.id, serverId),
eq(mcpServers.workspaceId, workspaceId),
isNull(mcpServers.deletedAt),
update.discoveryStartedAt
? or(
isNull(mcpServers.lastConnected),
lte(mcpServers.lastConnected, update.discoveryStartedAt)
)
: undefined
)
}
)
.returning({ id: mcpServers.id })
if (isErrorState && updatedServers.length > 0) {
logger.warn(`Server ${serverId} marked as error after ${newFailures} consecutive failures`)
}
return updatedServers.length > 0
} catch (err) {
logger.error(`Failed to update server status for ${serverId}:`, err)
return false
}
}
@@ -390,9 +424,10 @@ class McpService {
private async markServerUnhealthy(
workspaceId: string,
serverId: string,
error: unknown
error: unknown,
authType: McpServerConfig['authType']
): Promise<void> {
if (error instanceof McpOauthAuthorizationRequiredError || error instanceof UnauthorizedError) {
if (isOauthAuthorizationError(error, authType)) {
return
}
try {
@@ -406,6 +441,40 @@ class McpService {
}
}
private async markServerOauthPending(
serverId: string,
workspaceId: string,
discoveryStartedAt?: Date
): Promise<boolean> {
try {
const updatedServers = await db
.update(mcpServers)
.set({
connectionStatus: 'disconnected',
lastError: null,
updatedAt: new Date(),
})
.where(
and(
eq(mcpServers.id, serverId),
eq(mcpServers.workspaceId, workspaceId),
isNull(mcpServers.deletedAt),
discoveryStartedAt
? or(
isNull(mcpServers.lastConnected),
lte(mcpServers.lastConnected, discoveryStartedAt)
)
: undefined
)
)
.returning({ id: mcpServers.id })
return updatedServers.length > 0
} catch (error) {
logger.warn(`Failed to mark OAuth server ${serverId} disconnected:`, error)
return false
}
}
private async isServerUnhealthy(workspaceId: string, serverId: string): Promise<boolean> {
try {
const entry = await this.cacheAdapter.get(failureCacheKey(workspaceId, serverId))
@@ -429,6 +498,7 @@ class McpService {
forceRefresh = false
): Promise<McpTool[]> {
const requestId = generateRequestId()
const discoveryStartedAt = new Date()
try {
logger.info(`[${requestId}] Discovering MCP tools for workspace ${workspaceId}`)
@@ -479,15 +549,12 @@ class McpService {
await client.disconnect()
}
} catch (error) {
if (
error instanceof McpOauthAuthorizationRequiredError ||
error instanceof UnauthorizedError
) {
if (isOauthAuthorizationError(error, config.authType)) {
return { kind: 'oauth-pending' }
}
return {
kind: 'error',
message: getErrorMessage(error, 'Unknown error'),
message: getDiscoveryFailureMessage(error, config.authType, 'Unknown error'),
originalError: error,
}
}
@@ -516,7 +583,10 @@ class McpService {
fetchedCount++
allTools.push(...outcome.tools)
deferredSideEffects.push(
this.updateServerStatus(server.id, workspaceId, true, undefined, outcome.tools.length)
this.updateServerStatus(server.id, workspaceId, {
outcome: 'connected',
toolCount: outcome.tools.length,
})
)
cacheWrites.push(
this.cacheAdapter
@@ -536,18 +606,9 @@ class McpService {
// Mark disconnected so the UI surfaces the re-auth button.
logger.info(`[${requestId}] Skipping server ${server.name}: OAuth authorization pending`)
deferredSideEffects.push(
db
.update(mcpServers)
.set({
connectionStatus: 'disconnected',
lastError: null,
updatedAt: new Date(),
})
.where(eq(mcpServers.id, server.id))
.then(() => undefined)
.catch((err) => {
logger.warn(`[${requestId}] Failed to mark server ${server.id} disconnected:`, err)
})
this.markServerOauthPending(server.id, workspaceId, discoveryStartedAt).then(
() => undefined
)
)
return
}
@@ -561,13 +622,26 @@ class McpService {
`[${requestId}] Failed to discover tools from server ${server.name}: ${outcome.message}`
)
deferredSideEffects.push(
this.updateServerStatus(server.id, workspaceId, false, outcome.message),
this.markServerUnhealthy(workspaceId, server.id, outcome.originalError),
this.cacheAdapter
.delete(serverCacheKey(workspaceId, server.id))
.catch((err) =>
logger.warn(`[${requestId}] Cache delete failed for ${server.name}:`, err)
)
this.updateServerStatus(server.id, workspaceId, {
outcome: 'failed',
error: outcome.message,
discoveryStartedAt,
}).then(async (statusApplied) => {
if (!statusApplied) return
await Promise.allSettled([
this.markServerUnhealthy(
workspaceId,
server.id,
outcome.originalError,
server.authType
),
this.cacheAdapter
.delete(serverCacheKey(workspaceId, server.id))
.catch((err) =>
logger.warn(`[${requestId}] Cache delete failed for ${server.name}:`, err)
),
])
})
)
})
@@ -636,6 +710,7 @@ class McpService {
forceRefresh: boolean
): Promise<McpTool[]> {
const requestId = generateRequestId()
const discoveryStartedAt = new Date()
const maxRetries = 2
if (!forceRefresh) {
@@ -658,6 +733,7 @@ class McpService {
}
for (let attempt = 0; attempt < maxRetries; attempt++) {
let authType: McpServerConfig['authType']
try {
logger.info(
`[${requestId}] Discovering tools from server ${serverId} for user ${userId}${attempt > 0 ? ` (attempt ${attempt + 1})` : ''}`
@@ -667,6 +743,7 @@ class McpService {
if (!config) {
throw new Error(`Server ${serverId} not found or not accessible`)
}
authType = config.authType
const { config: resolvedConfig, resolvedIP } = await this.resolveConfigEnvVars(
config,
@@ -685,7 +762,10 @@ class McpService {
logger.warn(`[${requestId}] Cache write failed for ${config.name}:`, err)
),
this.clearServerFailure(workspaceId, serverId),
this.updateServerStatus(serverId, workspaceId, true, undefined, tools.length),
this.updateServerStatus(serverId, workspaceId, {
outcome: 'connected',
toolCount: tools.length,
}),
])
return tools
} finally {
@@ -701,14 +781,23 @@ class McpService {
continue
}
// Drop positive cache so a follow-up doesn't return stale tools.
await Promise.allSettled([
this.cacheAdapter
.delete(serverCacheKey(workspaceId, serverId))
.catch((err) =>
logger.warn(`[${requestId}] Cache delete failed for ${serverId}:`, err)
),
this.markServerUnhealthy(workspaceId, serverId, error),
])
const statusApplied = isOauthAuthorizationError(error, authType)
? await this.markServerOauthPending(serverId, workspaceId, discoveryStartedAt)
: await this.updateServerStatus(serverId, workspaceId, {
outcome: 'failed',
error: getDiscoveryFailureMessage(error, authType, 'Connection failed'),
discoveryStartedAt,
})
if (statusApplied) {
await Promise.allSettled([
this.cacheAdapter
.delete(serverCacheKey(workspaceId, serverId))
.catch((err) =>
logger.warn(`[${requestId}] Cache delete failed for ${serverId}:`, err)
),
this.markServerUnhealthy(workspaceId, serverId, error, authType),
])
}
throw error
}
}
@@ -747,10 +836,7 @@ class McpService {
error: undefined,
})
} catch (error) {
if (
error instanceof McpOauthAuthorizationRequiredError ||
error instanceof UnauthorizedError
) {
if (isOauthAuthorizationError(error, config.authType)) {
summaries.push({
id: config.id,
name: config.name,
@@ -771,7 +857,7 @@ class McpService {
status: 'error',
toolCount: 0,
lastSeen: undefined,
error: getErrorMessage(error, 'Connection failed'),
error: getDiscoveryFailureMessage(error, config.authType, 'Connection failed'),
})
}
}