Skip to content

Commit 8f3eb3a

Browse files
authored
fix(search): validate a new provider server before the approval transaction (#8366)
* fix(search): validate a new provider server before the approval transaction * fix(search): recheck for the provider server after a failed pre-approval lookup
1 parent dbeb732 commit 8f3eb3a

4 files changed

Lines changed: 179 additions & 44 deletions

File tree

‎apps/sim/lib/credential-groups/managed-mcp-service.ts‎

Lines changed: 48 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -20,7 +20,7 @@ import {
2020
type ManagedMcpConnectorId,
2121
requireManagedMcpConnectorUrl,
2222
} from '@/lib/credential-groups/managed-mcp-connectors'
23-
import type { DbOrTx } from '@/lib/db/types'
23+
import type { DbOrTx, DbTransaction } from '@/lib/db/types'
2424
import {
2525
McpDnsResolutionError,
2626
McpDomainNotAllowedError,
@@ -130,6 +130,27 @@ async function validateServerUrl(url: string): Promise<void> {
130130
}
131131
}
132132

133+
/** A connector input whose URL passed the MCP domain and SSRF checks. */
134+
export interface ValidatedManagedMcpConnectorInput {
135+
input: CreateManagedMcpConnectorInput
136+
url: string
137+
}
138+
139+
/**
140+
* Resolves and checks a connector's URL. The SSRF check resolves DNS, so a caller that joins its
141+
* own transaction runs this before opening it rather than while holding that transaction's locks.
142+
*/
143+
export async function validateManagedMcpConnectorInput(
144+
input: CreateManagedMcpConnectorInput
145+
): Promise<ValidatedManagedMcpConnectorInput> {
146+
const url = resolveManagedMcpConnectorUrl(
147+
input.connectorId,
148+
input.connectorId === 'databricks' ? input.url : undefined
149+
)
150+
await validateServerUrl(url)
151+
return { input, url }
152+
}
153+
133154
function resolveManagedMcpConnectorUrl(
134155
connectorId: ManagedMcpConnectorId,
135156
rawUrl?: string
@@ -176,37 +197,43 @@ async function retireManagedMcpCredentials(
176197
return retired.map((row) => row.id)
177198
}
178199

200+
interface ManagedMcpConnectorTarget {
201+
workspaceId?: string
202+
organizationId?: string
203+
credentialGroupId: string
204+
userId: string
205+
}
206+
207+
export function createManagedMcpConnector(
208+
params: ManagedMcpConnectorTarget & { input: CreateManagedMcpConnectorInput }
209+
): Promise<ManagedMcpConnectorMutationResult>
210+
/** Joins the caller's transaction with an input validated before that transaction opened. */
211+
export function createManagedMcpConnector(
212+
params: ManagedMcpConnectorTarget & { validated: ValidatedManagedMcpConnectorInput },
213+
executor: DbTransaction
214+
): Promise<ManagedMcpConnectorMutationResult>
179215
export async function createManagedMcpConnector(
180-
params: {
181-
workspaceId?: string
182-
organizationId?: string
183-
credentialGroupId: string
184-
userId: string
185-
input: CreateManagedMcpConnectorInput
186-
},
187-
executor?: DbOrTx
216+
params: ManagedMcpConnectorTarget &
217+
({ input: CreateManagedMcpConnectorInput } | { validated: ValidatedManagedMcpConnectorInput }),
218+
executor?: DbTransaction
188219
): Promise<ManagedMcpConnectorMutationResult> {
220+
const { input, url } =
221+
'validated' in params ? params.validated : await validateManagedMcpConnectorInput(params.input)
189222
const scope = resourceScopeFromOwner(params)
190-
const connector = getManagedMcpConnector(params.input.connectorId)
191-
const url = resolveManagedMcpConnectorUrl(
192-
connector.id,
193-
params.input.connectorId === 'databricks' ? params.input.url : undefined
194-
)
195-
await validateServerUrl(url)
223+
const connector = getManagedMcpConnector(input.connectorId)
196224
const serverId = generateMcpServerId(
197225
scope.kind === 'workspace' ? scope.workspaceId : resourceScopeKey(scope),
198226
url
199227
)
200-
const oauthClientId =
201-
params.input.connectorId === 'databricks' ? params.input.oauthClientId.trim() : null
228+
const oauthClientId = input.connectorId === 'databricks' ? input.oauthClientId.trim() : null
202229
const oauthClientSecret =
203-
params.input.connectorId === 'databricks' && params.input.oauthClientSecret
204-
? (await encryptSecret(params.input.oauthClientSecret)).encrypted
230+
input.connectorId === 'databricks' && input.oauthClientSecret
231+
? (await encryptSecret(input.oauthClientSecret)).encrypted
205232
: null
206-
const name = params.input.connectorId === 'databricks' ? params.input.name.trim() : connector.name
233+
const name = input.connectorId === 'databricks' ? input.name.trim() : connector.name
207234
if (!name)
208235
throw new ManagedMcpConnectorError('Managed MCP connector name is required', 'validation')
209-
if (params.input.connectorId === 'databricks' && !oauthClientId) {
236+
if (input.connectorId === 'databricks' && !oauthClientId) {
210237
throw new ManagedMcpConnectorError('Databricks OAuth Client ID is required', 'validation')
211238
}
212239

‎apps/sim/lib/knowledge/__integration__/search-mcp-setup.integration.ts‎

Lines changed: 47 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
import { db } from '@sim/db'
1+
import { db, runOutsideTransactionContext } from '@sim/db'
22
import {
33
credentialGroup,
44
mcpServers,
@@ -16,6 +16,7 @@ import { toRecord } from '@sim/utils/object'
1616
import { eq, inArray, sql } from 'drizzle-orm'
1717
import { afterAll, afterEach, beforeAll, beforeEach, describe, expect, it, vi } from 'vitest'
1818
import { createOrganizationAccountsGroup } from '@/lib/credential-groups/workspace-accounts'
19+
import { tryAcquireAdvisoryXactLock } from '@/lib/db/advisory-locks'
1920
import { approveSearchIntegration } from '@/lib/knowledge/application/search-integrations'
2021
import { defaultLiveSearchPolicy } from '@/lib/sim-search/live/policy-schema'
2122

@@ -151,6 +152,51 @@ describe('atomic organization live Search MCP setup', () => {
151152
}
152153
)
153154

155+
it('resolves the sign-in server before the approval takes the accounts lock', async () => {
156+
const lockHeldDuringLookup: boolean[] = []
157+
vi.mocked(dns.resolveHostAddresses).mockImplementationOnce(async () => {
158+
/** A separate connection: the lookup must not run inside the approval's transaction. */
159+
const acquired = await runOutsideTransactionContext(() =>
160+
db.transaction((tx) =>
161+
tryAcquireAdvisoryXactLock(
162+
tx,
163+
'search_accounts',
164+
`search-accounts:organization:${ids.organization}`
165+
)
166+
)
167+
)
168+
lockHeldDuringLookup.push(!acquired)
169+
return { addresses: ['93.184.216.34'], preferred: '93.184.216.34' }
170+
})
171+
await approve('fireflies')
172+
expect(lockHeldDuringLookup).toEqual([false])
173+
expect((await snapshot()).servers).toHaveLength(1)
174+
})
175+
176+
it('approves an already-configured provider while DNS is unavailable', async () => {
177+
const first = await approve('fireflies')
178+
const before = (await snapshot()).servers
179+
const lookup = vi.mocked(dns.resolveHostAddresses)
180+
const fixture = lookup.getMockImplementation()
181+
lookup.mockRejectedValue(new Error('DNS unavailable'))
182+
try {
183+
const again = await approve('fireflies')
184+
expect(again.memberAccounts?.groupId).toBe(first.memberAccounts?.groupId)
185+
} finally {
186+
if (fixture) lookup.mockImplementation(fixture)
187+
}
188+
expect((await snapshot()).servers).toEqual(before)
189+
})
190+
191+
it('approves a provider a concurrent approval configured while this lookup failed', async () => {
192+
vi.mocked(dns.resolveHostAddresses).mockImplementationOnce(async () => {
193+
await approve('fireflies')
194+
throw new Error('DNS unavailable')
195+
})
196+
await approve('fireflies')
197+
expect((await snapshot()).servers).toHaveLength(1)
198+
})
199+
154200
it('serializes concurrent approvals into one group and one server per provider', async () => {
155201
const providers = ['fireflies', 'granola', 'notion']
156202
const results = await Promise.all(

‎apps/sim/lib/knowledge/application/search-integrations.ts‎

Lines changed: 9 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,10 @@ import { refuseCapability } from '@/lib/permission-groups/capabilities'
1616
import { isOrganizationCapabilityWithheld } from '@/lib/permission-groups/capability-assertions'
1717
import { SEARCH_SOURCE_TYPES } from '@/lib/sim-search/connectors'
1818
import { NativeSearchError } from '@/lib/sim-search/live/http'
19-
import { addOrganizationSearchMcpProvider } from '@/lib/sim-search/live/member-setup'
19+
import {
20+
addOrganizationSearchMcpProvider,
21+
prepareSearchMcpProvider,
22+
} from '@/lib/sim-search/live/member-setup'
2023
import {
2124
defaultLiveSearchPolicy,
2225
LIVE_SEARCH_SERVICE_PROVIDERS,
@@ -150,6 +153,9 @@ export const approveSearchIntegration = defineAuthorizedKnowledgeUseCase({
150153
setWhere: sql`${organizationSearchIntegration.approved} IS DISTINCT FROM ${input.approved}`,
151154
})
152155
.returning({ connectorType: organizationSearchIntegration.connectorType })
156+
const mcpSetup = mcpProvider
157+
? await prepareSearchMcpProvider(context.organizationId, mcpProvider)
158+
: null
153159
let memberAccounts: { groupId: string; changed: boolean } | undefined
154160
const changed =
155161
policy || memberProvider || mcpProvider
@@ -165,11 +171,11 @@ export const approveSearchIntegration = defineAuthorizedKnowledgeUseCase({
165171
throw new OrchestrationError('validation', error.message)
166172
throw error
167173
})
168-
if (mcpProvider)
174+
if (mcpSetup)
169175
memberAccounts = await addOrganizationSearchMcpProvider(
170176
context.organizationId!,
171177
requirePrincipalSubjectUserId(principal),
172-
mcpProvider,
178+
mcpSetup,
173179
tx
174180
)
175181
if (policy)

‎apps/sim/lib/sim-search/live/member-setup.ts‎

Lines changed: 75 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -1,20 +1,84 @@
1-
import { mcpServers } from '@sim/db/schema'
1+
import { db } from '@sim/db'
2+
import { credentialGroup, mcpServers } from '@sim/db/schema'
23
import { and, eq, isNull } from 'drizzle-orm'
34
import { OrchestrationError } from '@/lib/core/orchestration/types'
5+
import { resourceScopeCondition } from '@/lib/core/resource-scope.server'
46
import { requireManagedMcpConnectorUrl } from '@/lib/credential-groups/managed-mcp-connectors'
57
import {
68
createManagedMcpConnector,
79
ManagedMcpConnectorError,
10+
type ValidatedManagedMcpConnectorInput,
11+
validateManagedMcpConnectorInput,
812
} from '@/lib/credential-groups/managed-mcp-service'
913
import { ensureWorkspaceAccountsGroup } from '@/lib/credential-groups/service'
1014
import type { DbTransaction } from '@/lib/db/types'
1115
import type { ManagedSearchMcpProvider } from '@/lib/sim-search/live/managed-mcp-config'
1216

17+
/**
18+
* A search provider ready for source approval. `validated` is set only when approval will create
19+
* the provider's server.
20+
*/
21+
export interface SearchMcpProviderSetup {
22+
provider: ManagedSearchMcpProvider
23+
validated: ValidatedManagedMcpConnectorInput | null
24+
}
25+
26+
function toSetupError(error: unknown): never {
27+
if (error instanceof ManagedMcpConnectorError)
28+
throw new OrchestrationError(
29+
error.code === 'bad_gateway' ? 'internal' : error.code,
30+
error.message
31+
)
32+
throw error
33+
}
34+
35+
async function hasProviderServer(
36+
organizationId: string,
37+
provider: ManagedSearchMcpProvider
38+
): Promise<boolean> {
39+
const [existing] = await db
40+
.select({ id: mcpServers.id })
41+
.from(mcpServers)
42+
.innerJoin(credentialGroup, eq(credentialGroup.id, mcpServers.credentialGroupId))
43+
.where(
44+
and(
45+
resourceScopeCondition(credentialGroup, { kind: 'organization', organizationId }),
46+
eq(mcpServers.organizationId, organizationId),
47+
eq(mcpServers.managedConnectorId, provider),
48+
isNull(mcpServers.deletedAt)
49+
)
50+
)
51+
.limit(1)
52+
return existing !== undefined
53+
}
54+
55+
/**
56+
* Runs before source approval opens its transaction. A provider whose server does not exist yet
57+
* has that server checked here, because the check resolves DNS and must not run while the
58+
* transaction holds the accounts lock; an already-configured provider needs no check. A failed
59+
* check looks again first, since a concurrent approval may have created the server meanwhile.
60+
*/
61+
export async function prepareSearchMcpProvider(
62+
organizationId: string,
63+
provider: ManagedSearchMcpProvider
64+
): Promise<SearchMcpProviderSetup> {
65+
if (await hasProviderServer(organizationId, provider)) return { provider, validated: null }
66+
try {
67+
return {
68+
provider,
69+
validated: await validateManagedMcpConnectorInput({ connectorId: provider }),
70+
}
71+
} catch (error) {
72+
if (await hasProviderServer(organizationId, provider)) return { provider, validated: null }
73+
toSetupError(error)
74+
}
75+
}
76+
1377
/** Joins source approval's transaction, serializing concurrent setup through the accounts lock. */
1478
export async function addOrganizationSearchMcpProvider(
1579
organizationId: string,
1680
userId: string,
17-
provider: ManagedSearchMcpProvider,
81+
{ provider, validated }: SearchMcpProviderSetup,
1882
executor: DbTransaction
1983
): Promise<{ groupId: string; changed: boolean }> {
2084
const group = await ensureWorkspaceAccountsGroup(
@@ -58,23 +122,15 @@ export async function addOrganizationSearchMcpProvider(
58122
)
59123
return { groupId: group.id, changed: group.created }
60124
}
61-
try {
62-
await createManagedMcpConnector(
63-
{
64-
organizationId,
65-
credentialGroupId: group.id,
66-
userId,
67-
input: { connectorId: provider },
68-
},
69-
executor
125+
/** The server was removed after preparation found it; a retry prepares it again. */
126+
if (!validated)
127+
throw new OrchestrationError(
128+
'conflict',
129+
'Connected accounts changed while adding this source. Try again.'
70130
)
71-
} catch (error) {
72-
if (error instanceof ManagedMcpConnectorError)
73-
throw new OrchestrationError(
74-
error.code === 'bad_gateway' ? 'internal' : error.code,
75-
error.message
76-
)
77-
throw error
78-
}
131+
await createManagedMcpConnector(
132+
{ organizationId, credentialGroupId: group.id, userId, validated },
133+
executor
134+
).catch(toSetupError)
79135
return { groupId: group.id, changed: true }
80136
}

0 commit comments

Comments
 (0)