diff --git a/apps/sim/lib/mcp/service.ts b/apps/sim/lib/mcp/service.ts index 7d93b58112b..1826df7318f 100644 --- a/apps/sim/lib/mcp/service.ts +++ b/apps/sim/lib/mcp/service.ts @@ -492,6 +492,18 @@ class McpService { return client } + /** An operation owns this unpooled client and must disconnect it in its finalizer. */ + async openManagedMcpSession( + serverId: string, + scope: ResourceScope, + auth: McpOauthCredentials, + signal: AbortSignal + ): Promise> { + const config = await this.getServerConfig(serverId, scope) + if (!config) throw new Error('Managed MCP server is unavailable') + return this.createManagedOauthClient(config, auth, signal) + } + async discoverManagedMcpTools( serverId: string, scope: string | ResourceScope, diff --git a/apps/sim/lib/sim-search/live/account-session.ts b/apps/sim/lib/sim-search/live/account-session.ts index cde7178ed15..b230a4aca56 100644 --- a/apps/sim/lib/sim-search/live/account-session.ts +++ b/apps/sim/lib/sim-search/live/account-session.ts @@ -3,7 +3,7 @@ import type { ResourceOwner } from '@/lib/core/resource-scope' import { resourceScopeFromOwner } from '@/lib/core/resource-scope' import type { PinnedConnectionPool } from '@/lib/core/security/input-validation.server' import type { ResolvedLiveAccount } from '@/lib/sim-search/live/accounts' -import { createCodaMcpClient, readCodaMcp, searchCodaMcp } from '@/lib/sim-search/live/coda-mcp' +import { readCodaMcp, searchCodaMcp } from '@/lib/sim-search/live/coda-mcp' import { readFirefliesMcp, searchFirefliesMcp } from '@/lib/sim-search/live/fireflies-mcp' import { createAdminGitLabSession } from '@/lib/sim-search/live/gitlab-admin' import { readGranolaMcp, searchGranolaMcp } from '@/lib/sim-search/live/granola-mcp' @@ -39,6 +39,8 @@ export interface LiveAccountSession { policy: LiveSearchPolicy /** Service verification covered a bounded subset of the source's configured users. */ servicePartial: boolean + /** Releases the operation-owned provider transport after all reads and checks settle. */ + close(): Promise search(input: NativeSearchInput): Promise /** True only when the document is inside the source boundary and the member may read it. */ verify(document: Reference): Promise @@ -104,10 +106,9 @@ export async function openLiveAccountSession( 'unavailable', 'This provider does not support managed MCP Search.' ) - return provider === 'coda' - ? createCodaMcpClient(owner, userId, account.id, signal, input.searches) - : createManagedSearchMcpClient(owner, userId, account.id, provider, signal, input.searches) + return createManagedSearchMcpClient(owner, userId, account.id, provider, signal, input.searches) } + const memberPolicy = livePolicyFor(input.policies, provider) const mcp = await openMcp() const searchMcp = (search: NativeSearchInput) => { if (!mcp) throw new NativeSearchError('unavailable', 'Managed MCP connection unavailable.') @@ -176,9 +177,15 @@ export async function openLiveAccountSession( verify: (document: Reference) => verifyPolicy(document, document.accessMetadata), } } - const boundary = await sourceBoundary(livePolicyFor(input.policies, provider)) + const boundary = await sourceBoundary(memberPolicy).catch(async (error: unknown) => { + await mcp?.close() + throw error + }) return { + async close() { + await mcp?.close() + }, policy: boundary.policy, servicePartial: boundary.partial, async search(search) { diff --git a/apps/sim/lib/sim-search/live/application.test.ts b/apps/sim/lib/sim-search/live/application.test.ts index 7af995cd48b..92573d1f542 100644 --- a/apps/sim/lib/sim-search/live/application.test.ts +++ b/apps/sim/lib/sim-search/live/application.test.ts @@ -44,10 +44,9 @@ vi.mock('@/lib/sim-search/live/policy-store', () => ({ livePolicyFor: vi.fn(() => defaultLiveSearchPolicy()), })) vi.mock('@/lib/sim-search/live/managed-mcp', () => ({ - createManagedSearchMcpClient: async () => ({ call: mocks.mcpCall }), + createManagedSearchMcpClient: async () => ({ call: mocks.mcpCall, close: async () => {} }), })) vi.mock('@/lib/sim-search/live/coda-mcp', () => ({ - createCodaMcpClient: vi.fn(), searchCodaMcp: vi.fn(), readCodaMcp: vi.fn(), })) diff --git a/apps/sim/lib/sim-search/live/application.ts b/apps/sim/lib/sim-search/live/application.ts index 08304f4748b..86d3239bc79 100644 --- a/apps/sim/lib/sim-search/live/application.ts +++ b/apps/sim/lib/sim-search/live/application.ts @@ -607,12 +607,13 @@ export const searchLiveKnowledge = defineAuthorizedKnowledgeUseCase({ results: [], } } + let session: LiveAccountSession | undefined try { signal.throwIfAborted() const resolved = await measureSearchStage('live.resolve', () => resolveListedLiveAccount(input, userId, account) ) - const session = await measureSearchStage('live.session', () => + session = await measureSearchStage('live.session', () => openLiveAccountSession({ owner: input, userId, @@ -623,9 +624,10 @@ export const searchLiveKnowledge = defineAuthorizedKnowledgeUseCase({ searches: natives.length, }) ) + const currentSession = session return await Promise.all( natives.map((target) => - searchQuery(account, resolved, session, target.native, statusFor(target)).catch( + searchQuery(account, resolved, currentSession, target.native, statusFor(target)).catch( (error) => failed(error, target) ) ) @@ -634,6 +636,7 @@ export const searchLiveKnowledge = defineAuthorizedKnowledgeUseCase({ return natives.map((target) => failed(error, target)) } finally { settled.abort() + await session?.close() } } let searched: SearchedQuery[] @@ -760,8 +763,9 @@ export const readLiveDocument = defineAuthorizedKnowledgeUseCase({ : AbortSignal.timeout(15_000) const pool = createPinnedConnectionPool() let document: NativeDocument + let session: LiveAccountSession | undefined try { - const session = await openLiveAccountSession({ + session = await openLiveAccountSession({ owner: input, userId, resolved, @@ -774,7 +778,10 @@ export const readLiveDocument = defineAuthorizedKnowledgeUseCase({ 'not_found', 'Document is outside your organization’s search scope' ) - document = await measureSearchStage('live.read', () => session.read(reference, input.filters)) + const currentSession = session + document = await measureSearchStage('live.read', () => + currentSession.read(reference, input.filters) + ) /** Readers degrade section failures to warnings, so the signal decides cancellation. */ signal.throwIfAborted() const current = await session.verifyCurrent(document) @@ -786,7 +793,11 @@ export const readLiveDocument = defineAuthorizedKnowledgeUseCase({ 'Document is outside your organization’s search scope' ) } finally { - pool.destroy() + try { + await session?.close() + } finally { + pool.destroy() + } } if (!matchesLiveFilters(document, input.documentId, reference.provider, input.filters)) throw new OrchestrationError('not_found', 'Document is outside the selected search filters') diff --git a/apps/sim/lib/sim-search/live/coda-mcp.ts b/apps/sim/lib/sim-search/live/coda-mcp.ts index b39e2d36740..4ed21f3b094 100644 --- a/apps/sim/lib/sim-search/live/coda-mcp.ts +++ b/apps/sim/lib/sim-search/live/coda-mcp.ts @@ -1,28 +1,14 @@ import { createLogger } from '@sim/logger' -import type { ResourceOwner } from '@/lib/core/resource-scope' import { parseCodaResourceUri } from '@/lib/sim-search/live/coda-uri' import { hasDateBounds, nativeText } from '@/lib/sim-search/live/dates' import { array, NativeSearchError, object, string } from '@/lib/sim-search/live/http' -import { - createManagedSearchMcpClient, - type ManagedSearchMcpClient, -} from '@/lib/sim-search/live/managed-mcp' +import type { ManagedSearchMcpClient } from '@/lib/sim-search/live/managed-mcp' import type { NativeDocument, NativePage, NativeSearchInput } from '@/lib/sim-search/live/types' const logger = createLogger('CodaMcpSearch') export interface CodaMcpClient extends ManagedSearchMcpClient {} -export function createCodaMcpClient( - owner: ResourceOwner, - userId: string, - credentialId: string, - signal: AbortSignal, - searches = 1 -): Promise { - return createManagedSearchMcpClient(owner, userId, credentialId, 'coda', signal, searches) -} - function requireCodaUri(uri: string): string { const parsed = parseCodaResourceUri(uri) if (!parsed) diff --git a/apps/sim/lib/sim-search/live/managed-mcp.integration.ts b/apps/sim/lib/sim-search/live/managed-mcp.integration.ts new file mode 100644 index 00000000000..3beb8aa1e8f --- /dev/null +++ b/apps/sim/lib/sim-search/live/managed-mcp.integration.ts @@ -0,0 +1,677 @@ +import { mkdir, writeFile } from 'node:fs/promises' +import { createServer } from 'node:http' +import { dirname } from 'node:path' +import { Server } from '@modelcontextprotocol/sdk/server/index.js' +import { StreamableHTTPServerTransport } from '@modelcontextprotocol/sdk/server/streamableHttp.js' +import { CallToolRequestSchema, ListToolsRequestSchema } from '@modelcontextprotocol/sdk/types.js' +import { db } from '@sim/db' +import { + credential, + credentialGroupEnrollment, + mcpServers, + member, + organization, + organizationSearchIntegration, + user, +} from '@sim/db/schema' +import { readTestRedisUrl } from '@sim/db/testing/test-infrastructure' +import { createSessionPrincipal } from '@sim/testing/factories/principal.factory' +import { createDeferred } from '@sim/testing/helpers/deferred' +import { sleep } from '@sim/utils/helpers' +import { generateId } from '@sim/utils/id' +import { eq, inArray } from 'drizzle-orm' +import { afterAll, afterEach, beforeAll, describe, expect, it, vi } from 'vitest' +import { env } from '@/lib/core/config/env' +import { runWithOutboundOrganization } from '@/lib/core/network/context.server' +import { createManagedMcpConnector } from '@/lib/credential-groups/managed-mcp-service' +import { createViewerCredentialGroupEnrollment } from '@/lib/credential-groups/self-enrollment' +import { ensureWorkspaceAccountsGroup } from '@/lib/credential-groups/service' +import { encryptManagedMcpTokens } from '@/lib/credentials/managed-mcp' +import * as pinnedFetch from '@/lib/mcp/pinned-fetch' +import { readLiveDocument, searchLiveKnowledge } from '@/lib/sim-search/live/application' +import { createManagedSearchMcpClient } from '@/lib/sim-search/live/managed-mcp' +import { ResolvedSecretTraceRegistry } from '@/executor/utils/resolved-secret-trace-registry' + +const RESOURCE = 'https://mcp.lucid.app/mcp/readonly' +const DOCUMENT = '00000000-0000-4000-8000-000000000001' +const SECOND_DOCUMENT = '00000000-0000-4000-8000-000000000002' +const TITLE = 'Synthetic topology' +const editUrl = `https://lucid.app/lucidchart/${DOCUMENT}/edit` +const actors = [0, 1].map(() => ({ + userId: generateId(), + organizationId: generateId(), + credentialId: `mcp-cg-${generateId()}`, + token: generateId(), +})) +actors.push({ + userId: generateId(), + organizationId: actors[0].organizationId, + credentialId: `mcp-cg-${generateId()}`, + token: generateId(), +}) +const sessions = new Map() +const serverIds = new Map() +const protocols: Server[] = [] +const openTransports = new Set<{ close(): Promise }>() +const events: { method: string; actor: number; at: number }[] = [] +const measurements: { name: string; durationMs: number; initializations: number }[] = [] +let origin = '' +let setupDelay = 0 +let onTool: ((name: string, args: Record) => Promise) | undefined +let failDiscovery = false +let onInitialize: (() => Promise) | undefined +let onDiscovery: (() => Promise) | undefined + +function payload(value: unknown) { + return { content: [{ type: 'text' as const, text: JSON.stringify(value) }] } +} + +const providerServer = createServer(async (request, response) => { + try { + const token = request.headers.authorization?.replace(/^Bearer /, '') + const actor = actors.findIndex((candidate) => candidate.token === token) + if (actor < 0) { + response.writeHead(401).end() + return + } + const id = request.headers['mcp-session-id'] + let session = typeof id === 'string' ? sessions.get(id) : undefined + if (session && session.actor !== actor) { + response.writeHead(403).end() + return + } + if (!session) { + if (id || request.method !== 'POST') { + response.writeHead(404).end() + return + } + events.push({ method: 'initialize', actor, at: performance.now() }) + await onInitialize?.() + await sleep(setupDelay) + const protocol = new Server( + { name: 'synthetic-lucid', version: '1' }, + { capabilities: { tools: {} } } + ) + protocol.setRequestHandler(ListToolsRequestSchema, async ({ params }) => { + events.push({ method: 'tools/list', actor, at: performance.now() }) + await onDiscovery?.() + if (failDiscovery && params?.cursor) throw new Error('Synthetic discovery failure') + return { + ...(failDiscovery ? { nextCursor: 'second-page' } : {}), + tools: ['search', 'fetch', 'lucid_get_document_metadata', 'write_document'].map( + (name) => ({ + name, + inputSchema: { + type: 'object' as const, + properties: { + query: { type: 'string' }, + product: { type: 'array', items: { type: 'string' } }, + last_modified_after: { type: 'string' }, + }, + additionalProperties: name !== 'search', + }, + }) + ), + } + }) + protocol.setRequestHandler(CallToolRequestSchema, async ({ params }) => { + events.push({ method: params.name, actor, at: performance.now() }) + await onTool?.(params.name, params.arguments ?? {}) + const documentId = + params.arguments?.query === 'second topology' || + params.arguments?.document_id === SECOND_DOCUMENT + ? SECOND_DOCUMENT + : DOCUMENT + const documentTitle = documentId === SECOND_DOCUMENT ? 'Second topology' : TITLE + const documentUrl = `https://lucid.app/lucidchart/${documentId}/edit` + if (params.name === 'search') + return payload({ results: [{ id: documentId, title: documentTitle, url: documentUrl }] }) + if (params.name === 'lucid_get_document_metadata') + return payload({ + documentId, + title: documentTitle, + product: 'lucidchart', + viewUrl: documentUrl, + version: 7, + pageCount: 1, + lastModified: '2026-09-01T12:00:00Z', + }) + const manifest = { + document_id: DOCUMENT, + title: TITLE, + edit_url: editUrl, + metadata: { page_count: 1, page_region_counts: [1] }, + } + if (params.arguments?.metadata_only) return payload(manifest) + return payload({ + ...manifest, + metadata: { ...manifest.metadata, page_index: 1 }, + page_id: 'page-0', + page_index: 1, + text: JSON.stringify({ + pages: [ + { + pageId: 'page-0', + pageTitle: 'Topology', + pageIndex: 0, + totalChunks: 1, + requestedChunks: [ + { + chunkIndex: 0, + data: { nodes: [{ label: `Private diagram ${actor}` }], edges: [] }, + }, + ], + }, + ], + }), + }) + }) + const transport = new StreamableHTTPServerTransport({ + sessionIdGenerator: generateId, + enableJsonResponse: true, + onsessioninitialized: (sessionId) => { + sessions.set(sessionId, { transport, actor }) + }, + }) + protocols.push(protocol) + await protocol.connect(transport) + session = { transport, actor } + } + await session.transport.handleRequest(request, response) + } catch { + if (!response.headersSent) response.writeHead(500) + response.end() + } +}) + +beforeAll(async () => { + Object.assign(env, { REDIS_URL: readTestRedisUrl(), EGRESS_ALLOWED_HOSTS: '127.0.0.1' }) + await new Promise((resolve) => providerServer.listen(0, '127.0.0.1', resolve)) + const address = providerServer.address() + if (!address || typeof address === 'string') throw new Error('Fixture failed to bind') + origin = `http://127.0.0.1:${address.port}/mcp` + const guarded = pinnedFetch.createGuardedMcpFetch + vi.spyOn(pinnedFetch, 'createGuardedMcpFetch').mockImplementation((url) => { + if (url !== RESOURCE) throw new Error('Unexpected MCP fixture origin') + const transport = guarded(origin) + openTransports.add(transport) + return { + fetch: (input, init) => { + if (new URL(input instanceof Request ? input.url : input).href !== RESOURCE) + throw new Error('Unexpected MCP fixture destination') + return transport.fetch(origin, init) + }, + close: async () => { + await transport.close() + openTransports.delete(transport) + }, + } + }) + for (const actor of actors) { + await db.insert(user).values({ + id: actor.userId, + name: 'Session fixture', + email: `${actor.userId}@fixture.test`, + emailVerified: true, + createdAt: new Date(), + updatedAt: new Date(), + }) + await db + .insert(organization) + .values({ id: actor.organizationId, name: 'Session fixture', slug: actor.organizationId }) + .onConflictDoNothing() + await db.insert(member).values({ + id: generateId(), + organizationId: actor.organizationId, + userId: actor.userId, + role: 'owner', + }) + const group = await ensureWorkspaceAccountsGroup( + { kind: 'organization', organizationId: actor.organizationId }, + actor.userId + ) + let serverId = serverIds.get(actor.organizationId) + if (!serverId) { + const { mcpServer } = await db.transaction((tx) => + createManagedMcpConnector( + { + organizationId: actor.organizationId, + credentialGroupId: group.id, + userId: actor.userId, + validated: { input: { connectorId: 'lucid' }, url: RESOURCE }, + }, + tx + ) + ) + serverId = mcpServer.id + serverIds.set(actor.organizationId, serverId) + } + const { enrollment } = await createViewerCredentialGroupEnrollment({ + organizationId: actor.organizationId, + credentialGroupId: group.id, + userId: actor.userId, + }) + await db + .update(credentialGroupEnrollment) + .set({ status: 'completed' }) + .where(eq(credentialGroupEnrollment.id, enrollment.id)) + await db.insert(credential).values({ + id: actor.credentialId, + organizationId: actor.organizationId, + type: 'managed_mcp', + displayName: 'Session fixture', + grantedAt: new Date(), + credentialGroupEnrollmentId: enrollment.id, + mcpServerId: serverId, + mcpOauthConfigVersion: ( + await db + .select({ version: mcpServers.oauthConfigVersion }) + .from(mcpServers) + .where(eq(mcpServers.id, serverId)) + )[0].version, + managedOauthStatus: 'active', + mcpTools: [], + encryptedOauthTokenSet: await encryptManagedMcpTokens({ + access_token: actor.token, + token_type: 'Bearer', + }), + }) + await db + .insert(organizationSearchIntegration) + .values({ organizationId: actor.organizationId, connectorType: 'lucid', approved: true }) + .onConflictDoNothing() + } +}) + +afterEach(async () => { + for (const transport of openTransports) await transport.close() + openTransports.clear() +}) + +afterAll(async () => { + if (process.env.MANAGED_MCP_SESSION_REPORT_PATH) { + await mkdir(dirname(process.env.MANAGED_MCP_SESSION_REPORT_PATH), { recursive: true }) + await writeFile( + process.env.MANAGED_MCP_SESSION_REPORT_PATH, + JSON.stringify({ measurements, events, openTransports: openTransports.size }, null, 2) + ) + } + await Promise.all(protocols.map((protocol) => protocol.close())) + await new Promise((resolve, reject) => { + providerServer.close((error) => (error ? reject(error) : resolve())) + providerServer.closeAllConnections() + }) + await db.delete(organization).where( + inArray( + organization.id, + actors.map((actor) => actor.organizationId) + ) + ) + await db.delete(user).where( + inArray( + user.id, + actors.map((actor) => actor.userId) + ) + ) +}) + +function search(index = 0, signal?: AbortSignal) { + const actor = actors[index] + return searchLiveKnowledge.execute({ + principal: createSessionPrincipal({ userId: actor.userId, sessionId: generateId() }), + input: { + organizationId: actor.organizationId, + query: 'topology', + topK: 10, + filters: { source: 'lucid' }, + signal, + }, + }) +} +function read(documentId: string, index = 0, signal?: AbortSignal) { + const actor = actors[index] + return readLiveDocument.execute({ + principal: createSessionPrincipal({ userId: actor.userId, sessionId: generateId() }), + input: { + organizationId: actor.organizationId, + documentId, + limit: 8, + resultSecretRegistry: new ResolvedSecretTraceRegistry([]), + signal, + }, + }) +} + +/** Actual SDK HTTP + current database grants, including final document verification. */ +describe('managed Search operation sessions', () => { + it('reads a complete diagram with one handshake per authorized operation and disposes both transports', async () => { + setupDelay = 150 + try { + const start = performance.now() + const offset = events.length + const found = await search() + expect(found.results, JSON.stringify(found.live)).toHaveLength(1) + const document = await read(found.results[0].documentId) + expect(JSON.stringify(document)).toContain('Private diagram 0') + const initializations = events + .slice(offset) + .filter((event) => event.method === 'initialize').length + measurements.push({ + name: 'search and complete read', + durationMs: performance.now() - start, + initializations, + }) + expect(openTransports.size).toBe(0) + expect(initializations).toBe(2) + } finally { + setupDelay = 0 + } + }) + it('multiplexes native queries on one connection without mixing their results', async () => { + const actor = actors[0] + const bothEntered = createDeferred() + const entered: string[] = [] + const finished: string[] = [] + const offset = events.length + let timer: ReturnType | undefined + onTool = async (name, args) => { + if (name !== 'search') return + const query = String(args.query) + entered.push(query) + if (entered.length === 1) timer = setTimeout(() => bothEntered.resolve(), 2_000) + if (entered.length === 2) bothEntered.resolve() + await bothEntered.promise + expect(entered).toHaveLength(2) + if (query === 'topology') await sleep(25) + finished.push(query) + } + try { + const found = await searchLiveKnowledge.execute({ + principal: createSessionPrincipal({ userId: actor.userId, sessionId: generateId() }), + input: { + organizationId: actor.organizationId, + query: 'topology', + topK: 10, + filters: { source: 'lucid' }, + nativeQueries: [ + { provider: 'lucid', query: 'topology', kind: 'lucidchart' }, + { provider: 'lucid', query: 'second topology', kind: 'lucidchart' }, + ], + }, + }) + expect(finished).toEqual(['second topology', 'topology']) + expect(found.results.map((result) => result.documentName).sort()).toEqual([ + 'Second topology', + TITLE, + ]) + expect(found.retrieval.status).toBe('complete') + expect( + found.live?.accounts.map(({ queryIndex, status }) => ({ queryIndex, status })) + ).toEqual([ + { queryIndex: 0, status: 'ok' }, + { queryIndex: 1, status: 'ok' }, + ]) + expect(events.slice(offset).filter((event) => event.method === 'initialize')).toHaveLength(1) + expect(openTransports.size).toBe(0) + } finally { + clearTimeout(timer) + bothEntered.resolve() + onTool = undefined + } + }) + it('isolates simultaneous users and rejects another organization’s signed document reference', async () => { + const results = await Promise.all([search(0), search(1)]) + const documents = await Promise.all( + results.map((result, index) => read(result.results[0].documentId, index)) + ) + expect(JSON.stringify(documents[0])).toContain('Private diagram 0') + expect(JSON.stringify(documents[0])).not.toContain('Private diagram 1') + expect(JSON.stringify(documents[1])).toContain('Private diagram 1') + const before = events.length + await expect(read(results[0].results[0].documentId, 1)).rejects.toMatchObject({ + code: 'not_found', + }) + expect(events.length).toBe(before) + expect(openTransports.size).toBe(0) + }) + it('withholds a document revoked during content fetch and closes the transport', async () => { + const found = await search() + expect(found.results).toHaveLength(1) + onTool = async (name, args) => { + if (name === 'fetch' && !args.metadata_only) + await db + .update(credential) + .set({ revokedAt: new Date() }) + .where(eq(credential.id, actors[0].credentialId)) + } + try { + await expect(read(found.results[0].documentId)).rejects.toMatchObject({ status: 'reconnect' }) + expect(openTransports.size).toBe(0) + } finally { + onTool = undefined + await db + .update(credential) + .set({ revokedAt: null }) + .where(eq(credential.id, actors[0].credentialId)) + } + }) + it('closes a connection when complete tool discovery fails', async () => { + failDiscovery = true + const offset = events.length + try { + const found = await search() + expect(found.results).toHaveLength(0) + expect(found.live?.accounts[0].status).toBe('unavailable') + expect(events.slice(offset).map((event) => event.method)).toEqual([ + 'initialize', + 'tools/list', + 'tools/list', + ]) + expect(openTransports.size).toBe(0) + } finally { + failDiscovery = false + } + }) + it.each(['initialize', 'discovery', 'content'] as const)( + 'releases transport on cancellation during %s', + async (phase) => { + const found = await search() + expect(found.results).toHaveLength(1) + const entered = createDeferred() + const release = createDeferred() + const block = async () => { + entered.resolve() + await release.promise + } + const controller = new AbortController() + if (phase === 'initialize') onInitialize = block + else if (phase === 'discovery') onDiscovery = block + else + onTool = async (name, args) => { + if (name === 'fetch' && !args.metadata_only) await block() + } + const pending = read(found.results[0].documentId, 0, controller.signal) + const rejected = expect(pending).rejects.toThrow() + try { + await entered.promise + controller.abort(new Error('Fixture caller cancelled')) + await rejected + expect(openTransports.size).toBe(0) + } finally { + onInitialize = undefined + onDiscovery = undefined + onTool = undefined + release.resolve() + await pending.catch(() => {}) + } + const next = await search() + expect(next.results).toHaveLength(1) + expect(openTransports.size).toBe(0) + } + ) + + it.each(['schema', 'budget'] as const)( + 'enforces the %s boundary before sending a tool request', + async (boundary) => { + const actor = actors[0] + await runWithOutboundOrganization(actor.organizationId, async () => { + const client = await createManagedSearchMcpClient( + { organizationId: actor.organizationId }, + actor.userId, + actor.credentialId, + 'lucid', + AbortSignal.timeout(10_000) + ) + try { + if (boundary === 'budget') { + for (let i = 0; i < 12; i++) await client.call('search', { query: 'topology' }) + } + const before = events.length + if (boundary === 'schema') { + await expect(client.call('write_document', {})).rejects.toThrow('read-only') + await expect(client.call('search', { query: 42 })).rejects.toThrow('schema') + } else + await expect(client.call('search', { query: 'topology' })).rejects.toThrow( + 'request limit' + ) + expect(events.length).toBe(before) + } finally { + await client.close() + } + }) + expect(openTransports.size).toBe(0) + } + ) + + it('does not retain transport when a persisted policy is invalid', async () => { + await db + .update(organization) + .set({ metadata: { liveSearchPolicies: { lucid: { mode: 'invalid-policy' } } } }) + .where(eq(organization.id, actors[0].organizationId)) + try { + const found = await search() + expect(found.results).toHaveLength(0) + expect(found.live?.accounts[0].status).toBe('unavailable') + expect(openTransports.size).toBe(0) + } finally { + await db + .update(organization) + .set({ metadata: null }) + .where(eq(organization.id, actors[0].organizationId)) + } + }) + + it.each(['before', 'during'] as const)( + 'rejects revocation %s a call within an open operation', + async (phase) => { + const actor = actors[0] + await runWithOutboundOrganization(actor.organizationId, async () => { + const client = await createManagedSearchMcpClient( + { organizationId: actor.organizationId }, + actor.userId, + actor.credentialId, + 'lucid', + AbortSignal.timeout(10_000) + ) + const revoke = async () => { + await db + .update(credential) + .set({ revokedAt: new Date() }) + .where(eq(credential.id, actor.credentialId)) + } + try { + await client.call('search', { query: 'topology' }) + if (phase === 'before') await revoke() + else onTool = revoke + const before = events.length + await expect(client.call('search', { query: 'topology' })).rejects.toMatchObject({ + status: 'reconnect', + }) + if (phase === 'before') expect(events.length).toBe(before) + } finally { + onTool = undefined + await client.close() + await db + .update(credential) + .set({ revokedAt: null }) + .where(eq(credential.id, actor.credentialId)) + } + }) + expect(openTransports.size).toBe(0) + } + ) + + it('keeps two personal grants on the same server in separate operation sessions', async () => { + const [first, second] = await Promise.all([search(0), search(2)]) + const documents = await Promise.all([ + read(first.results[0].documentId, 0), + read(second.results[0].documentId, 2), + ]) + expect(JSON.stringify(documents[0])).toContain('Private diagram 0') + expect(JSON.stringify(documents[0])).not.toContain('Private diagram 2') + expect(JSON.stringify(documents[1])).toContain('Private diagram 2') + expect(openTransports.size).toBe(0) + }) + + it('reloads a rotated persisted token after a challenge without sharing the operation session', async () => { + const actor = actors[0] + await runWithOutboundOrganization(actor.organizationId, async () => { + const client = await createManagedSearchMcpClient( + { organizationId: actor.organizationId }, + actor.userId, + actor.credentialId, + 'lucid', + AbortSignal.timeout(10_000) + ) + try { + await client.call('search', { query: 'topology' }) + actor.token = generateId() + await db + .update(credential) + .set({ + encryptedOauthTokenSet: await encryptManagedMcpTokens({ + access_token: actor.token, + token_type: 'Bearer', + }), + }) + .where(eq(credential.id, actor.credentialId)) + const before = events.length + expect(await client.call('search', { query: 'topology' })).toMatchObject({ + results: [{ id: DOCUMENT }], + }) + expect(events.slice(before).map((event) => event.method)).toEqual(['search']) + } finally { + await client.close() + } + }) + expect(openTransports.size).toBe(0) + }) + + it('rejects replacement of the grant inside an operation before sending another request', async () => { + const actor = actors[0] + await runWithOutboundOrganization(actor.organizationId, async () => { + const client = await createManagedSearchMcpClient( + { organizationId: actor.organizationId }, + actor.userId, + actor.credentialId, + 'lucid', + AbortSignal.timeout(10_000) + ) + try { + await client.call('search', { query: 'topology' }) + await db + .update(credential) + .set({ grantedAt: new Date() }) + .where(eq(credential.id, actor.credentialId)) + const before = events.length + await expect(client.call('search', { query: 'topology' })).rejects.toThrow( + 'connection changed' + ) + expect(events.length).toBe(before) + } finally { + await client.close() + } + }) + expect(openTransports.size).toBe(0) + }) +}) diff --git a/apps/sim/lib/sim-search/live/managed-mcp.test.ts b/apps/sim/lib/sim-search/live/managed-mcp.test.ts index a72d87d2c6a..2d0a2dc4d59 100644 --- a/apps/sim/lib/sim-search/live/managed-mcp.test.ts +++ b/apps/sim/lib/sim-search/live/managed-mcp.test.ts @@ -1,94 +1,9 @@ -import { mcpServiceMock, mcpServiceMockFns } from '@sim/testing/mocks/mcp-service.mock' import { describe, expect, it, vi } from 'vitest' - -const mocks = vi.hoisted(() => ({ runtime: vi.fn(), auth: vi.fn() })) -vi.mock('@/lib/sim-search/live/mcp-accounts', () => ({ loadOwnManagedMcpRuntime: mocks.runtime })) -vi.mock('@/lib/mcp/service', () => mcpServiceMock) -vi.mock('@/lib/mcp/application/managed-auth-provider', () => ({ - createManagedMcpAuthProvider: mocks.auth, -})) - import { NativeSearchError } from '@/lib/sim-search/live/http' -import { createManagedSearchMcpClient } from '@/lib/sim-search/live/managed-mcp' import { managedMcpPayload } from '@/lib/sim-search/live/managed-mcp-payload' -/** Failure modes: a write tool escapes the allowlist; replaced grants stay usable; payloads exhaust memory; schema drift changes tool meaning. */ +/** Failure modes: payloads exhaust memory or reflect sensitive provider errors. */ describe('managed search MCP read boundary', () => { - it('rejects write tools, changed grants, and invalid wire arguments before provider execution', async () => { - const runtime = { - mcpServerId: 'server', - credentialId: 'mine', - scope: { kind: 'organization', organizationId: 'org' }, - oauthConfigVersion: 1, - grantedAt: new Date(0), - } - mocks.runtime.mockResolvedValue(runtime) - mcpServiceMockFns.mockDiscoverManagedMcpTools.mockResolvedValue([ - { - name: 'fireflies_get_transcripts', - inputSchema: { - type: 'object', - properties: { keyword: { type: 'string' } }, - additionalProperties: false, - }, - }, - { name: 'fireflies_share_meeting', inputSchema: { type: 'object' } }, - ]) - mcpServiceMockFns.mockExecuteManagedMcpTool.mockImplementation(async () => { - throw new Error('Must not execute') - }) - const client = await createManagedSearchMcpClient( - { organizationId: 'org' }, - 'person', - 'mine', - 'fireflies', - new AbortController().signal - ) - await expect(client.call('fireflies_share_meeting', {})).rejects.toThrow('read-only') - await expect(client.call('fireflies_get_transcripts', { keyword: 42 })).rejects.toMatchObject({ - status: 'unavailable', - message: - 'Fireflies rejected these search arguments. Its current tool schema is incompatible with this query.', - }) - mocks.runtime.mockResolvedValue({ ...runtime, grantedAt: new Date(1) }) - await expect(client.call('fireflies_get_transcripts', { keyword: 'term' })).rejects.toThrow( - 'connection changed' - ) - }) - - it('withholds a response when its member grant is revoked during the provider request', async () => { - let revoked = false - mocks.runtime.mockImplementation(async () => { - if (revoked) throw new NativeSearchError('reconnect', 'Member grant was revoked') - return { - mcpServerId: 'server', - credentialId: 'mine', - scope: { kind: 'organization', organizationId: 'org' }, - oauthConfigVersion: 1, - grantedAt: new Date(0), - } - }) - mcpServiceMockFns.mockDiscoverManagedMcpTools.mockResolvedValue([ - { name: 'fireflies_get_transcripts', inputSchema: { type: 'object' } }, - ]) - mcpServiceMockFns.mockExecuteManagedMcpTool.mockImplementation(async () => { - revoked = true - return { - structuredContent: { transcripts: [{ id: 'secret-meeting', title: 'Revoked content' }] }, - } - }) - const client = await createManagedSearchMcpClient( - { organizationId: 'org' }, - 'person', - 'mine', - 'fireflies', - new AbortController().signal - ) - await expect( - client.call('fireflies_get_transcripts', { keyword: 'term', scope: 'all' }) - ).rejects.toThrow('revoked') - }) - it('rejects oversized and failed tool payloads without exposing provider errors', () => { expect(() => managedMcpPayload( diff --git a/apps/sim/lib/sim-search/live/managed-mcp.ts b/apps/sim/lib/sim-search/live/managed-mcp.ts index feedc1e89bf..ee56fa6000e 100644 --- a/apps/sim/lib/sim-search/live/managed-mcp.ts +++ b/apps/sim/lib/sim-search/live/managed-mcp.ts @@ -27,7 +27,7 @@ export async function createManagedSearchMcpClient( provider: ManagedSearchMcpProvider, signal: AbortSignal, searches = 1 -): Promise { +): Promise }> { signal.throwIfAborted() const label = MANAGED_MCP_CONNECTORS[provider].name const initial = await loadOwnManagedMcpRuntime(owner, userId, credentialId, provider) @@ -43,13 +43,18 @@ export async function createManagedSearchMcpClient( return current } const loadProvider = async () => createManagedMcpAuthProvider(await loadCurrent()) - const tools = await mcpService.discoverManagedMcpTools( + const session = await mcpService.openManagedMcpSession( initial.mcpServerId, initial.scope, { credentialId, loadProvider }, - signal, - { requireComplete: true } + signal ) + const tools = await session + .listTools(signal, { requireComplete: true }) + .catch(async (error: unknown) => { + await session.disconnect() + throw error + }) const allowed: readonly string[] = MANAGED_SEARCH_MCP_READ_TOOLS[provider] const byName = new Map( tools.filter((tool) => allowed.includes(tool.name)).map((tool) => [tool.name, tool]) @@ -57,6 +62,7 @@ export async function createManagedSearchMcpClient( const budget = 12 * Math.min(4, Math.max(1, searches)) let requests = 0 return { + close: () => session.disconnect(), hasTool: (name) => byName.has(name), hasArgument(name, path) { let schema: Record = toRecord(byName.get(name)?.inputSchema) @@ -88,15 +94,10 @@ export async function createManagedSearchMcpClient( `${label} rejected these search arguments. Its current tool schema is incompatible with this query.` ) await loadCurrent() - const result = await mcpService.executeManagedMcpTool({ - connectionId: credentialId, - serverId: initial.mcpServerId, - scope: initial.scope, - toolCall: { name, arguments: args }, - loadAuthProvider: loadProvider, - signal, - timeoutMs: 10_000, - }) + const result = await session.callTool( + { name, arguments: args }, + { signal, timeoutMs: 10_000 } + ) await loadCurrent() return managedMcpPayload(result, label) }, diff --git a/apps/sim/lib/sim-search/live/providers.ts b/apps/sim/lib/sim-search/live/providers.ts index 891900575dc..177e8613ee5 100644 --- a/apps/sim/lib/sim-search/live/providers.ts +++ b/apps/sim/lib/sim-search/live/providers.ts @@ -64,7 +64,7 @@ export const LIVE_SEARCH_PROVIDERS = { transport: 'managed_mcp', guide: { syntax: - 'Nonempty document-title keywords, at most 400 characters. Results are relevance-ranked, not guaranteed literal title matches. The provider returns at most 200 relevance-ranked candidates; Sim verifies metadata for at most 10. Search has no continuation.', + 'Nonempty document-title keywords, at most 400 characters. Results are relevance-ranked, not guaranteed literal title matches. The provider returns at most 200 relevance-ranked candidates; Sim verifies metadata for at most 10. Search has no continuation and cannot enumerate the account. For a browse-all request without a title or topic, ask for one instead of guessing keywords.', scope: 'kind lucidchart or lucidspark selects a product; omit to search both. To search shape text within a known document, set project to its UUID or Lucid URL and use one literal substring of at most 200 characters. Dates use modification time; sorting and end dates apply only to retrieved candidates, not the entire account.', example: 'deployment architecture', @@ -264,7 +264,7 @@ export const LIVE_SEARCH_PROVIDERS = { transport: 'managed_mcp', guide: { syntax: - 'Natural-language or plain keyword content search through Notion MCP. Search terms are required even with dates or sorting. Availability depends on the connected account and plan; results are restricted to Notion pages, excluding connected apps.', + 'Natural-language or plain keyword content search through Notion MCP. Search terms are required even with dates or sorting; this search cannot enumerate the account. For a browse-all request without a title or topic, ask for one instead of guessing keywords. Availability depends on the connected account and plan; results are restricted to Notion pages, excluding connected apps.', scope: 'project optionally takes a known Notion page URL when the advertised tool supports page scoping. Dates use explicit last-edited timestamps; results without those timestamps cannot satisfy date filters. Read a result for page content.', example: 'deployment rollback checklist', diff --git a/apps/sim/scripts/test-search-hubspot-live.ts b/apps/sim/scripts/test-search-hubspot-live.ts index 68a023b9ab5..e76c49c12f0 100644 --- a/apps/sim/scripts/test-search-hubspot-live.ts +++ b/apps/sim/scripts/test-search-hubspot-live.ts @@ -80,64 +80,71 @@ const ids = (documents: { id: string }[]) => for (const test of config.queries) { await check(test.name, async () => { const provider = await client() - const input = { - query: test.query, - native: { provider: 'hubspot' as const, query: test.query, kind: test.kind }, - filters: test.filters, - scopes: [], - limit: 25, - } - const result = await searchHubSpotMcp(provider, input) - assert( - !result.partial && !result.nextCursor && !result.hasMore, - 'Fixture must fit one complete page' - ) - assert.deepEqual( - ids(result.documents), - [...test.expectedIds].sort(), - 'Exact CRM record identity differs' - ) - for (const document of result.documents) { - assert(document.url && document.content, 'Result must include source evidence') - if (test.filters?.startDate) - assert( - Date.parse(document.modifiedAt!) >= Date.parse(test.filters.startDate), - 'Record precedes lower bound' - ) - if (test.filters?.endDate) - assert( - Date.parse(document.modifiedAt!) < Date.parse(test.filters.endDate), - 'Record exceeds exclusive upper bound' - ) - } - if (test.filters?.sortBy) { - const dates = result.documents.map((document) => Date.parse(document.modifiedAt!)) - assert(dates.every(Number.isFinite), 'Sorted records must include valid modification dates') - assert.deepEqual( - dates, - [...dates].sort((a, b) => (test.filters?.sortBy === 'oldest' ? a - b : b - a)) + try { + const input = { + query: test.query, + native: { provider: 'hubspot' as const, query: test.query, kind: test.kind }, + filters: test.filters, + scopes: [], + limit: 25, + } + const result = await searchHubSpotMcp(provider, input) + assert( + !result.partial && !result.nextCursor && !result.hasMore, + 'Fixture must fit one complete page' ) - } - if (result.documents.length) { - const read = await readHubSpotMcp(provider, result.documents[0].id) - assert.equal(read.id, result.documents[0].id) - assert.equal(new URL(read.url).pathname, new URL(result.documents[0].url).pathname) - assert(read.content.includes('hs_object_id:'), 'Read must include returned record properties') - } - if (result.documents.length > 1) { - const first = await searchHubSpotMcp(provider, { ...input, limit: 1 }) - assert(first.nextCursor, 'First page must advertise continuation') - const second = await searchHubSpotMcp(provider, { - ...input, - limit: 1, - native: { ...input.native, cursor: first.nextCursor }, - }) assert.deepEqual( - ids([...first.documents, ...second.documents]), - ids(result.documents.slice(0, 2)) + ids(result.documents), + [...test.expectedIds].sort(), + 'Exact CRM record identity differs' ) - if (result.documents.length === 2) - assert.equal(second.nextCursor, undefined, 'Final page must terminate') + for (const document of result.documents) { + assert(document.url && document.content, 'Result must include source evidence') + if (test.filters?.startDate) + assert( + Date.parse(document.modifiedAt!) >= Date.parse(test.filters.startDate), + 'Record precedes lower bound' + ) + if (test.filters?.endDate) + assert( + Date.parse(document.modifiedAt!) < Date.parse(test.filters.endDate), + 'Record exceeds exclusive upper bound' + ) + } + if (test.filters?.sortBy) { + const dates = result.documents.map((document) => Date.parse(document.modifiedAt!)) + assert(dates.every(Number.isFinite), 'Sorted records must include valid modification dates') + assert.deepEqual( + dates, + [...dates].sort((a, b) => (test.filters?.sortBy === 'oldest' ? a - b : b - a)) + ) + } + if (result.documents.length) { + const read = await readHubSpotMcp(provider, result.documents[0].id) + assert.equal(read.id, result.documents[0].id) + assert.equal(new URL(read.url).pathname, new URL(result.documents[0].url).pathname) + assert( + read.content.includes('hs_object_id:'), + 'Read must include returned record properties' + ) + } + if (result.documents.length > 1) { + const first = await searchHubSpotMcp(provider, { ...input, limit: 1 }) + assert(first.nextCursor, 'First page must advertise continuation') + const second = await searchHubSpotMcp(provider, { + ...input, + limit: 1, + native: { ...input.native, cursor: first.nextCursor }, + }) + assert.deepEqual( + ids([...first.documents, ...second.documents]), + ids(result.documents.slice(0, 2)) + ) + if (result.documents.length === 2) + assert.equal(second.nextCursor, undefined, 'Final page must terminate') + } + } finally { + await provider.close() } }) }