Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
69 changes: 48 additions & 21 deletions apps/sim/lib/credential-groups/managed-mcp-service.ts
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@ import {
type ManagedMcpConnectorId,
requireManagedMcpConnectorUrl,
} from '@/lib/credential-groups/managed-mcp-connectors'
import type { DbOrTx } from '@/lib/db/types'
import type { DbOrTx, DbTransaction } from '@/lib/db/types'
import {
McpDnsResolutionError,
McpDomainNotAllowedError,
Expand Down Expand Up @@ -130,6 +130,27 @@ async function validateServerUrl(url: string): Promise<void> {
}
}

/** A connector input whose URL passed the MCP domain and SSRF checks. */
export interface ValidatedManagedMcpConnectorInput {
input: CreateManagedMcpConnectorInput
url: string
}

/**
* Resolves and checks a connector's URL. The SSRF check resolves DNS, so a caller that joins its
* own transaction runs this before opening it rather than while holding that transaction's locks.
*/
export async function validateManagedMcpConnectorInput(
input: CreateManagedMcpConnectorInput
): Promise<ValidatedManagedMcpConnectorInput> {
const url = resolveManagedMcpConnectorUrl(
input.connectorId,
input.connectorId === 'databricks' ? input.url : undefined
)
await validateServerUrl(url)
return { input, url }
}

function resolveManagedMcpConnectorUrl(
connectorId: ManagedMcpConnectorId,
rawUrl?: string
Expand Down Expand Up @@ -176,37 +197,43 @@ async function retireManagedMcpCredentials(
return retired.map((row) => row.id)
}

interface ManagedMcpConnectorTarget {
workspaceId?: string
organizationId?: string
credentialGroupId: string
userId: string
}

export function createManagedMcpConnector(
params: ManagedMcpConnectorTarget & { input: CreateManagedMcpConnectorInput }
): Promise<ManagedMcpConnectorMutationResult>
/** Joins the caller's transaction with an input validated before that transaction opened. */
export function createManagedMcpConnector(
params: ManagedMcpConnectorTarget & { validated: ValidatedManagedMcpConnectorInput },
executor: DbTransaction
): Promise<ManagedMcpConnectorMutationResult>
export async function createManagedMcpConnector(
params: {
workspaceId?: string
organizationId?: string
credentialGroupId: string
userId: string
input: CreateManagedMcpConnectorInput
},
executor?: DbOrTx
params: ManagedMcpConnectorTarget &
({ input: CreateManagedMcpConnectorInput } | { validated: ValidatedManagedMcpConnectorInput }),
executor?: DbTransaction
): Promise<ManagedMcpConnectorMutationResult> {
const { input, url } =
'validated' in params ? params.validated : await validateManagedMcpConnectorInput(params.input)
const scope = resourceScopeFromOwner(params)
const connector = getManagedMcpConnector(params.input.connectorId)
const url = resolveManagedMcpConnectorUrl(
connector.id,
params.input.connectorId === 'databricks' ? params.input.url : undefined
)
await validateServerUrl(url)
const connector = getManagedMcpConnector(input.connectorId)
const serverId = generateMcpServerId(
scope.kind === 'workspace' ? scope.workspaceId : resourceScopeKey(scope),
url
)
const oauthClientId =
params.input.connectorId === 'databricks' ? params.input.oauthClientId.trim() : null
const oauthClientId = input.connectorId === 'databricks' ? input.oauthClientId.trim() : null
const oauthClientSecret =
params.input.connectorId === 'databricks' && params.input.oauthClientSecret
? (await encryptSecret(params.input.oauthClientSecret)).encrypted
input.connectorId === 'databricks' && input.oauthClientSecret
? (await encryptSecret(input.oauthClientSecret)).encrypted
: null
const name = params.input.connectorId === 'databricks' ? params.input.name.trim() : connector.name
const name = input.connectorId === 'databricks' ? input.name.trim() : connector.name
if (!name)
throw new ManagedMcpConnectorError('Managed MCP connector name is required', 'validation')
if (params.input.connectorId === 'databricks' && !oauthClientId) {
if (input.connectorId === 'databricks' && !oauthClientId) {
throw new ManagedMcpConnectorError('Databricks OAuth Client ID is required', 'validation')
}

Expand Down
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
import { db } from '@sim/db'
import { db, runOutsideTransactionContext } from '@sim/db'
import {
credentialGroup,
mcpServers,
Expand All @@ -16,6 +16,7 @@ import { toRecord } from '@sim/utils/object'
import { eq, inArray, sql } from 'drizzle-orm'
import { afterAll, afterEach, beforeAll, beforeEach, describe, expect, it, vi } from 'vitest'
import { createOrganizationAccountsGroup } from '@/lib/credential-groups/workspace-accounts'
import { tryAcquireAdvisoryXactLock } from '@/lib/db/advisory-locks'
import { approveSearchIntegration } from '@/lib/knowledge/application/search-integrations'
import { defaultLiveSearchPolicy } from '@/lib/sim-search/live/policy-schema'

Expand Down Expand Up @@ -151,6 +152,51 @@ describe('atomic organization live Search MCP setup', () => {
}
)

it('resolves the sign-in server before the approval takes the accounts lock', async () => {
const lockHeldDuringLookup: boolean[] = []
vi.mocked(dns.resolveHostAddresses).mockImplementationOnce(async () => {
/** A separate connection: the lookup must not run inside the approval's transaction. */
const acquired = await runOutsideTransactionContext(() =>
db.transaction((tx) =>
tryAcquireAdvisoryXactLock(
tx,
'search_accounts',
`search-accounts:organization:${ids.organization}`
)
)
)
lockHeldDuringLookup.push(!acquired)
return { addresses: ['93.184.216.34'], preferred: '93.184.216.34' }
})
await approve('fireflies')
expect(lockHeldDuringLookup).toEqual([false])
expect((await snapshot()).servers).toHaveLength(1)
})

it('approves an already-configured provider while DNS is unavailable', async () => {
const first = await approve('fireflies')
const before = (await snapshot()).servers
const lookup = vi.mocked(dns.resolveHostAddresses)
const fixture = lookup.getMockImplementation()
lookup.mockRejectedValue(new Error('DNS unavailable'))
try {
const again = await approve('fireflies')
expect(again.memberAccounts?.groupId).toBe(first.memberAccounts?.groupId)
} finally {
if (fixture) lookup.mockImplementation(fixture)
}
expect((await snapshot()).servers).toEqual(before)
})

it('approves a provider a concurrent approval configured while this lookup failed', async () => {
vi.mocked(dns.resolveHostAddresses).mockImplementationOnce(async () => {
await approve('fireflies')
throw new Error('DNS unavailable')
})
await approve('fireflies')
expect((await snapshot()).servers).toHaveLength(1)
})

it('serializes concurrent approvals into one group and one server per provider', async () => {
const providers = ['fireflies', 'granola', 'notion']
const results = await Promise.all(
Expand Down
12 changes: 9 additions & 3 deletions apps/sim/lib/knowledge/application/search-integrations.ts
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,10 @@ import { refuseCapability } from '@/lib/permission-groups/capabilities'
import { isOrganizationCapabilityWithheld } from '@/lib/permission-groups/capability-assertions'
import { SEARCH_SOURCE_TYPES } from '@/lib/sim-search/connectors'
import { NativeSearchError } from '@/lib/sim-search/live/http'
import { addOrganizationSearchMcpProvider } from '@/lib/sim-search/live/member-setup'
import {
addOrganizationSearchMcpProvider,
prepareSearchMcpProvider,
} from '@/lib/sim-search/live/member-setup'
import {
defaultLiveSearchPolicy,
LIVE_SEARCH_SERVICE_PROVIDERS,
Expand Down Expand Up @@ -150,6 +153,9 @@ export const approveSearchIntegration = defineAuthorizedKnowledgeUseCase({
setWhere: sql`${organizationSearchIntegration.approved} IS DISTINCT FROM ${input.approved}`,
})
.returning({ connectorType: organizationSearchIntegration.connectorType })
const mcpSetup = mcpProvider
? await prepareSearchMcpProvider(context.organizationId, mcpProvider)
: null
let memberAccounts: { groupId: string; changed: boolean } | undefined
const changed =
policy || memberProvider || mcpProvider
Expand All @@ -165,11 +171,11 @@ export const approveSearchIntegration = defineAuthorizedKnowledgeUseCase({
throw new OrchestrationError('validation', error.message)
throw error
})
if (mcpProvider)
if (mcpSetup)
memberAccounts = await addOrganizationSearchMcpProvider(
context.organizationId!,
requirePrincipalSubjectUserId(principal),
mcpProvider,
mcpSetup,
tx
)
if (policy)
Expand Down
94 changes: 75 additions & 19 deletions apps/sim/lib/sim-search/live/member-setup.ts
Original file line number Diff line number Diff line change
@@ -1,20 +1,84 @@
import { mcpServers } from '@sim/db/schema'
import { db } from '@sim/db'
import { credentialGroup, mcpServers } from '@sim/db/schema'
import { and, eq, isNull } from 'drizzle-orm'
import { OrchestrationError } from '@/lib/core/orchestration/types'
import { resourceScopeCondition } from '@/lib/core/resource-scope.server'
import { requireManagedMcpConnectorUrl } from '@/lib/credential-groups/managed-mcp-connectors'
import {
createManagedMcpConnector,
ManagedMcpConnectorError,
type ValidatedManagedMcpConnectorInput,
validateManagedMcpConnectorInput,
} from '@/lib/credential-groups/managed-mcp-service'
import { ensureWorkspaceAccountsGroup } from '@/lib/credential-groups/service'
import type { DbTransaction } from '@/lib/db/types'
import type { ManagedSearchMcpProvider } from '@/lib/sim-search/live/managed-mcp-config'

/**
* A search provider ready for source approval. `validated` is set only when approval will create
* the provider's server.
*/
export interface SearchMcpProviderSetup {
provider: ManagedSearchMcpProvider
validated: ValidatedManagedMcpConnectorInput | null
}

function toSetupError(error: unknown): never {
if (error instanceof ManagedMcpConnectorError)
throw new OrchestrationError(
error.code === 'bad_gateway' ? 'internal' : error.code,
error.message
)
throw error
}

async function hasProviderServer(
organizationId: string,
provider: ManagedSearchMcpProvider
): Promise<boolean> {
const [existing] = await db
.select({ id: mcpServers.id })
.from(mcpServers)
.innerJoin(credentialGroup, eq(credentialGroup.id, mcpServers.credentialGroupId))
.where(
and(
resourceScopeCondition(credentialGroup, { kind: 'organization', organizationId }),
eq(mcpServers.organizationId, organizationId),
eq(mcpServers.managedConnectorId, provider),
isNull(mcpServers.deletedAt)
)
)
.limit(1)
return existing !== undefined
}

/**
* Runs before source approval opens its transaction. A provider whose server does not exist yet
* has that server checked here, because the check resolves DNS and must not run while the
* transaction holds the accounts lock; an already-configured provider needs no check. A failed
* check looks again first, since a concurrent approval may have created the server meanwhile.
*/
export async function prepareSearchMcpProvider(
organizationId: string,
provider: ManagedSearchMcpProvider
): Promise<SearchMcpProviderSetup> {
if (await hasProviderServer(organizationId, provider)) return { provider, validated: null }
try {
return {
provider,
validated: await validateManagedMcpConnectorInput({ connectorId: provider }),
}
} catch (error) {
if (await hasProviderServer(organizationId, provider)) return { provider, validated: null }
toSetupError(error)
Comment thread
waleedlatif1 marked this conversation as resolved.
}
}

/** Joins source approval's transaction, serializing concurrent setup through the accounts lock. */
export async function addOrganizationSearchMcpProvider(
organizationId: string,
userId: string,
provider: ManagedSearchMcpProvider,
{ provider, validated }: SearchMcpProviderSetup,
executor: DbTransaction
): Promise<{ groupId: string; changed: boolean }> {
const group = await ensureWorkspaceAccountsGroup(
Expand Down Expand Up @@ -58,23 +122,15 @@ export async function addOrganizationSearchMcpProvider(
)
return { groupId: group.id, changed: group.created }
}
try {
await createManagedMcpConnector(
{
organizationId,
credentialGroupId: group.id,
userId,
input: { connectorId: provider },
},
executor
/** The server was removed after preparation found it; a retry prepares it again. */
if (!validated)
throw new OrchestrationError(
'conflict',
'Connected accounts changed while adding this source. Try again.'
)
} catch (error) {
if (error instanceof ManagedMcpConnectorError)
throw new OrchestrationError(
error.code === 'bad_gateway' ? 'internal' : error.code,
error.message
)
throw error
}
await createManagedMcpConnector(
{ organizationId, credentialGroupId: group.id, userId, validated },
executor
).catch(toSetupError)
return { groupId: group.id, changed: true }
}
Loading