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
6 changes: 6 additions & 0 deletions apps/sim/lib/core/security/input-validation.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -615,6 +615,12 @@ describe('validateServiceNowInstanceUrl (vendor-hosted allowlist)', () => {
expect(result.sanitized).toBe('https://acme.servicenowservices.com/api/now')
})

it.concurrent('drops a trailing FQDN dot, which TLS hostname verification rejects', () => {
const result = validateServiceNowInstanceUrl('https://acme.service-now.com./api/now')
expect(result.isValid).toBe(true)
expect(result.sanitized).toBe('https://acme.service-now.com/api/now')
})

it.concurrent.each([
['https://support.acme.com', 'vanity CNAME'],
['https://acme.service-now.com.evil.com', 'lookalike suffix'],
Expand Down
21 changes: 16 additions & 5 deletions apps/sim/lib/core/security/input-validation.ts
Original file line number Diff line number Diff line change
Expand Up @@ -1167,6 +1167,9 @@ function validateVendorHostedUrl(
if (!urlResult.isValid) return urlResult

const parsed = new URL(candidate)
// A trailing FQDN dot names the same host, but TLS hostname verification rejects it.
const fullyQualified = parsed.hostname.endsWith('.')
if (fullyQualified) parsed.hostname = parsed.hostname.slice(0, -1)
const hostname = parsed.hostname.toLowerCase()
const allowed = suffixes.some(
(suffix) => (allowBareSuffix && hostname === suffix.slice(1)) || hostname.endsWith(suffix)
Expand All @@ -1185,7 +1188,8 @@ function validateVendorHostedUrl(
}
}

return { isValid: true, sanitized: sanitize === 'origin' ? parsed.origin : candidate }
if (sanitize === 'origin') return { isValid: true, sanitized: parsed.origin }
return { isValid: true, sanitized: fullyQualified ? parsed.href : candidate }
}

/**
Expand Down Expand Up @@ -1267,15 +1271,20 @@ export function validateWorkdayTenantUrl(
}

/**
* Every production Databricks control-plane DNS zone, mirroring `ALL_ENVS` in the
* Databricks SDK (`databricks/sdk/environments.py`). The SDK's `.dev.*`/`.staging.*`
* zones are internal and deliberately omitted; the ones that are subdomains of a
* zone listed here (e.g. `.staging.cloud.databricks.com`) match by suffix anyway.
* Every production Databricks control-plane DNS zone. All but `.cloud.databricks.mil` mirror
* `ALL_ENVS` in the Databricks SDK (`databricks/sdk/environments.py`); the DoD zone comes from the
* Databricks AWS GovCloud docs. The SDK's `.dev.*`/`.staging.*` zones are internal and
* deliberately omitted; the ones that are subdomains of a zone listed here (e.g.
* `.staging.cloud.databricks.com`) match by suffix anyway. `.databricks.com` admits workspace
* custom URLs (`acme.databricks.com`) and subsumes the AWS and GCP zones, which stay listed so
* the rejection message names them.
*/
const DATABRICKS_ALLOWED_HOST_SUFFIXES = [
'.cloud.databricks.com',
'.cloud.databricks.us',
'.cloud.databricks.mil',
'.gcp.databricks.com',
'.databricks.com',
Comment thread
waleedlatif1 marked this conversation as resolved.
'.azuredatabricks.net',
'.databricks.azure.us',
'.databricks.azure.cn',
Expand All @@ -1288,6 +1297,8 @@ const DATABRICKS_ALLOWED_HOST_SUFFIXES = [
* every REST call is made against it. Example valid hosts:
* - dbc-1234abcd-5678.cloud.databricks.com (AWS)
* - dbc-1234abcd-5678.cloud.databricks.us (AWS GovCloud)
* - dbc-1234abcd-5678.cloud.databricks.mil (AWS GovCloud DoD)
* - acme.databricks.com (workspace custom URL)
* - adb-1234567890123456.7.azuredatabricks.net (Azure)
* - adb-1234567890123456.7.databricks.azure.us (Azure US Government)
* - adb-1234567890123456.7.databricks.azure.cn (Azure China)
Expand Down
9 changes: 2 additions & 7 deletions apps/sim/tools/databricks/cancel_run.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ import type {
DatabricksCancelRunParams,
DatabricksCancelRunResponse,
} from '@/tools/databricks/types'
import { databricksUrl } from '@/tools/databricks/utils'
import type { ToolConfig } from '@/tools/types'

export const cancelRunTool: ToolConfig<DatabricksCancelRunParams, DatabricksCancelRunResponse> = {
Expand Down Expand Up @@ -33,13 +34,7 @@ export const cancelRunTool: ToolConfig<DatabricksCancelRunParams, DatabricksCanc
},

request: {
url: (params) => {
const host = params.host
.trim()
.replace(/^https?:\/\//, '')
.replace(/\/$/, '')
return `https://${host}/api/2.1/jobs/runs/cancel`
},
url: (params) => databricksUrl(params.host, '/api/2.1/jobs/runs/cancel'),
method: 'POST',
headers: (params) => ({
'Content-Type': 'application/json',
Expand Down
111 changes: 111 additions & 0 deletions apps/sim/tools/databricks/databricks.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,111 @@
import { inputValidationMock } from '@sim/testing'
import { partialToolRegistry } from '@sim/testing/mocks/tool-registry.mock'
import { getErrorMessage } from '@sim/utils/errors'
import { describe, expect, it, vi } from 'vitest'

vi.mock('@/lib/core/security/input-validation.server', () => inputValidationMock)

import * as databricksTools from '@/tools/databricks'
import { executeTool } from '@/tools/index'
import { tools } from '@/tools/registry'

/** Registers only this service's configs in the global registry mock; the full one is ~6,000 modules. */
Object.assign(tools, partialToolRegistry(databricksTools))

type ToolParams = Record<string, unknown>

function urlBuilder(toolId: string): (params: ToolParams) => string {
const url = tools[toolId].request.url
if (typeof url !== 'function') throw new Error(`${toolId} has a static url`)
return url as (params: ToolParams) => string
}

/** Every identifier any Databricks tool reads while building its URL. */
const REQUEST_PARAMS = {
apiKey: 'dapi-test-token',
spaceId: 'space1',
conversationId: 'conv1',
messageId: 'msg1',
attachmentId: 'att1',
statementId: 'stmt1',
clusterId: 'cluster1',
jobId: 1,
runId: 2,
warehouseId: 'wh1',
content: 'question',
sql: 'SELECT 1',
rating: 'POSITIVE',
}

/** The validator's refusal, naming the `host` param and every allowlisted Databricks domain. */
const HOST_ALLOWLIST_ERROR =
'host must be a Databricks-hosted domain (e.g., *.cloud.databricks.com, *.cloud.databricks.us, *.cloud.databricks.mil, *.gcp.databricks.com, *.databricks.com, *.azuredatabricks.net, *.databricks.azure.us, *.databricks.azure.cn)'

const DATABRICKS_TOOL_IDS = Object.keys(tools).filter((id) => id.startsWith('databricks_'))

describe('databricks workspace host allowlist', () => {
it.each([
'attacker.example.com',
'https://attacker.example.com/',
'dbc-1.cloud.databricks.com.attacker.example.com',
'attacker.example.com/dbc-1.cloud.databricks.com',
'dbc-1.cloud.databricks.com@attacker.example.com',
'databricks.com',
'acme-databricks.com',
])('refuses %s in every tool with the host allowlist error', (host) => {
const notRefusedByAllowlist = DATABRICKS_TOOL_IDS.map((id) => {
try {
return `${id}: built ${urlBuilder(id)({ ...REQUEST_PARAMS, host })}`
} catch (error) {
return `${id}: ${getErrorMessage(error)}`
}
}).filter((outcome) => !outcome.endsWith(`: ${HOST_ALLOWLIST_ERROR}`))
expect(notRefusedByAllowlist).toEqual([])
})

it('fails the tool call for a foreign host with the allowlist error', async () => {
const result = await executeTool('databricks_list_clusters', {
host: 'attacker.example.com',
apiKey: 'dapi-test-token',
})

expect(result.success).toBe(false)
expect(result.error).toContain(HOST_ALLOWLIST_ERROR)
})

it.each([
['dbc-a1b2.cloud.databricks.com', 'https://dbc-a1b2.cloud.databricks.com'],
[' https://dbc-a1b2.cloud.databricks.com/ ', 'https://dbc-a1b2.cloud.databricks.com'],
['http://dbc-a1b2.cloud.databricks.com', 'https://dbc-a1b2.cloud.databricks.com'],
['adb-123.4.azuredatabricks.net', 'https://adb-123.4.azuredatabricks.net'],
['https://123.4.gcp.databricks.com/', 'https://123.4.gcp.databricks.com'],
['dbc-a1b2.cloud.databricks.us', 'https://dbc-a1b2.cloud.databricks.us'],
['adb-123.4.databricks.azure.us', 'https://adb-123.4.databricks.azure.us'],
['adb-123.4.databricks.azure.cn', 'https://adb-123.4.databricks.azure.cn'],
['dbc-a1b2.cloud.databricks.mil', 'https://dbc-a1b2.cloud.databricks.mil'],
['https://acme.databricks.com/', 'https://acme.databricks.com'],
])('builds the same request URLs for workspace host %s', (host, origin) => {
const params = { ...REQUEST_PARAMS, host }
expect(urlBuilder('databricks_list_clusters')(params)).toBe(`${origin}/api/2.0/clusters/list`)
expect(urlBuilder('databricks_execute_sql')(params)).toBe(`${origin}/api/2.0/sql/statements/`)
expect(urlBuilder('databricks_get_job')(params)).toBe(`${origin}/api/2.1/jobs/get?job_id=1`)
expect(urlBuilder('databricks_get_run_output')(params)).toBe(
`${origin}/api/2.1/jobs/runs/get-output?run_id=2`
)
expect(urlBuilder('databricks_genie_get_message')(params)).toBe(
`${origin}/api/2.0/genie/spaces/space1/conversations/conv1/messages/msg1`
)
})

it.each([
['dbc-a1b2.cloud.databricks.com.', 'https://dbc-a1b2.cloud.databricks.com'],
['https://dbc-a1b2.cloud.databricks.com.:8443/', 'https://dbc-a1b2.cloud.databricks.com:8443'],
])(
'drops the trailing FQDN dot of %s, which the old tools kept and Bun TLS rejects',
(host, origin) => {
expect(urlBuilder('databricks_list_clusters')({ ...REQUEST_PARAMS, host })).toBe(
`${origin}/api/2.0/clusters/list`
)
}
)
})
9 changes: 2 additions & 7 deletions apps/sim/tools/databricks/execute_sql.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ import type {
DatabricksExecuteSqlParams,
DatabricksExecuteSqlResponse,
} from '@/tools/databricks/types'
import { databricksUrl } from '@/tools/databricks/utils'
import type { ToolConfig } from '@/tools/types'

export const executeSqlTool: ToolConfig<DatabricksExecuteSqlParams, DatabricksExecuteSqlResponse> =
Expand Down Expand Up @@ -65,13 +66,7 @@ export const executeSqlTool: ToolConfig<DatabricksExecuteSqlParams, DatabricksEx
},

request: {
url: (params) => {
const host = params.host
.trim()
.replace(/^https?:\/\//, '')
.replace(/\/$/, '')
return `https://${host}/api/2.0/sql/statements/`
},
url: (params) => databricksUrl(params.host, '/api/2.0/sql/statements/'),
method: 'POST',
headers: (params) => ({
'Content-Type': 'application/json',
Expand Down
7 changes: 2 additions & 5 deletions apps/sim/tools/databricks/get_cluster.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ import type {
DatabricksGetClusterParams,
DatabricksGetClusterResponse,
} from '@/tools/databricks/types'
import { databricksUrl } from '@/tools/databricks/utils'
import type { ToolConfig } from '@/tools/types'

export const getClusterTool: ToolConfig<DatabricksGetClusterParams, DatabricksGetClusterResponse> =
Expand Down Expand Up @@ -35,11 +36,7 @@ export const getClusterTool: ToolConfig<DatabricksGetClusterParams, DatabricksGe

request: {
url: (params) => {
const host = params.host
.trim()
.replace(/^https?:\/\//, '')
.replace(/\/$/, '')
const url = new URL(`https://${host}/api/2.0/clusters/get`)
const url = new URL(databricksUrl(params.host, '/api/2.0/clusters/get'))
url.searchParams.set('cluster_id', params.clusterId.trim())
return url.toString()
},
Expand Down
7 changes: 2 additions & 5 deletions apps/sim/tools/databricks/get_job.ts
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import type { DatabricksGetJobParams, DatabricksGetJobResponse } from '@/tools/databricks/types'
import { databricksUrl } from '@/tools/databricks/utils'
import type { ToolConfig } from '@/tools/types'

export const getJobTool: ToolConfig<DatabricksGetJobParams, DatabricksGetJobResponse> = {
Expand Down Expand Up @@ -30,11 +31,7 @@ export const getJobTool: ToolConfig<DatabricksGetJobParams, DatabricksGetJobResp

request: {
url: (params) => {
const host = params.host
.trim()
.replace(/^https?:\/\//, '')
.replace(/\/$/, '')
const url = new URL(`https://${host}/api/2.1/jobs/get`)
const url = new URL(databricksUrl(params.host, '/api/2.1/jobs/get'))
url.searchParams.set('job_id', String(params.jobId))
return url.toString()
},
Expand Down
7 changes: 2 additions & 5 deletions apps/sim/tools/databricks/get_run.ts
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import type { DatabricksGetRunParams, DatabricksGetRunResponse } from '@/tools/databricks/types'
import { databricksUrl } from '@/tools/databricks/utils'
import type { ToolConfig } from '@/tools/types'

export const getRunTool: ToolConfig<DatabricksGetRunParams, DatabricksGetRunResponse> = {
Expand Down Expand Up @@ -42,11 +43,7 @@ export const getRunTool: ToolConfig<DatabricksGetRunParams, DatabricksGetRunResp

request: {
url: (params) => {
const host = params.host
.trim()
.replace(/^https?:\/\//, '')
.replace(/\/$/, '')
const url = new URL(`https://${host}/api/2.1/jobs/runs/get`)
const url = new URL(databricksUrl(params.host, '/api/2.1/jobs/runs/get'))
url.searchParams.set('run_id', String(params.runId))
if (params.includeHistory) url.searchParams.set('include_history', 'true')
if (params.includeResolvedValues) url.searchParams.set('include_resolved_values', 'true')
Expand Down
10 changes: 3 additions & 7 deletions apps/sim/tools/databricks/get_run_output.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ import type {
DatabricksGetRunOutputParams,
DatabricksGetRunOutputResponse,
} from '@/tools/databricks/types'
import { databricksUrl } from '@/tools/databricks/utils'
import type { ToolConfig } from '@/tools/types'

export const getRunOutputTool: ToolConfig<
Expand Down Expand Up @@ -36,13 +37,8 @@ export const getRunOutputTool: ToolConfig<
},

request: {
url: (params) => {
const host = params.host
.trim()
.replace(/^https?:\/\//, '')
.replace(/\/$/, '')
return `https://${host}/api/2.1/jobs/runs/get-output?run_id=${params.runId}`
},
url: (params) =>
databricksUrl(params.host, `/api/2.1/jobs/runs/get-output?run_id=${params.runId}`),
method: 'GET',
headers: (params) => ({
Accept: 'application/json',
Expand Down
10 changes: 3 additions & 7 deletions apps/sim/tools/databricks/get_statement.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ import type {
DatabricksExecuteSqlResponse,
DatabricksGetStatementParams,
} from '@/tools/databricks/types'
import { databricksUrl } from '@/tools/databricks/utils'
import type { ToolConfig } from '@/tools/types'

export const getStatementTool: ToolConfig<
Expand Down Expand Up @@ -36,13 +37,8 @@ export const getStatementTool: ToolConfig<
},

request: {
url: (params) => {
const host = params.host
.trim()
.replace(/^https?:\/\//, '')
.replace(/\/$/, '')
return `https://${host}/api/2.0/sql/statements/${params.statementId.trim()}`
},
url: (params) =>
databricksUrl(params.host, `/api/2.0/sql/statements/${params.statementId.trim()}`),
method: 'GET',
headers: (params) => ({
Accept: 'application/json',
Expand Down
9 changes: 2 additions & 7 deletions apps/sim/tools/databricks/list_clusters.ts
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import type { DatabricksBaseParams, DatabricksListClustersResponse } from '@/tools/databricks/types'
import { databricksUrl } from '@/tools/databricks/utils'
import type { ToolConfig } from '@/tools/types'

export const listClustersTool: ToolConfig<DatabricksBaseParams, DatabricksListClustersResponse> = {
Expand All @@ -24,13 +25,7 @@ export const listClustersTool: ToolConfig<DatabricksBaseParams, DatabricksListCl
},

request: {
url: (params) => {
const host = params.host
.trim()
.replace(/^https?:\/\//, '')
.replace(/\/$/, '')
return `https://${host}/api/2.0/clusters/list`
},
url: (params) => databricksUrl(params.host, '/api/2.0/clusters/list'),
method: 'GET',
headers: (params) => ({
Accept: 'application/json',
Expand Down
7 changes: 2 additions & 5 deletions apps/sim/tools/databricks/list_jobs.ts
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import type { DatabricksListJobsParams, DatabricksListJobsResponse } from '@/tools/databricks/types'
import { databricksUrl } from '@/tools/databricks/utils'
import type { ToolConfig } from '@/tools/types'

export const listJobsTool: ToolConfig<DatabricksListJobsParams, DatabricksListJobsResponse> = {
Expand Down Expand Up @@ -48,11 +49,7 @@ export const listJobsTool: ToolConfig<DatabricksListJobsParams, DatabricksListJo

request: {
url: (params) => {
const host = params.host
.trim()
.replace(/^https?:\/\//, '')
.replace(/\/$/, '')
const url = new URL(`https://${host}/api/2.1/jobs/list`)
const url = new URL(databricksUrl(params.host, '/api/2.1/jobs/list'))
if (params.limit) url.searchParams.set('limit', String(params.limit))
if (params.offset) url.searchParams.set('offset', String(params.offset))
if (params.name) url.searchParams.set('name', params.name)
Expand Down
7 changes: 2 additions & 5 deletions apps/sim/tools/databricks/list_runs.ts
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import type { DatabricksListRunsParams, DatabricksListRunsResponse } from '@/tools/databricks/types'
import { databricksUrl } from '@/tools/databricks/utils'
import type { ToolConfig } from '@/tools/types'

export const listRunsTool: ToolConfig<DatabricksListRunsParams, DatabricksListRunsResponse> = {
Expand Down Expand Up @@ -73,11 +74,7 @@ export const listRunsTool: ToolConfig<DatabricksListRunsParams, DatabricksListRu

request: {
url: (params) => {
const host = params.host
.trim()
.replace(/^https?:\/\//, '')
.replace(/\/$/, '')
const url = new URL(`https://${host}/api/2.1/jobs/runs/list`)
const url = new URL(databricksUrl(params.host, '/api/2.1/jobs/runs/list'))
if (params.jobId) url.searchParams.set('job_id', String(params.jobId))
if (params.activeOnly) url.searchParams.set('active_only', 'true')
if (params.completedOnly) url.searchParams.set('completed_only', 'true')
Expand Down
Loading
Loading