diff --git a/apps/sim/lib/core/security/input-validation.test.ts b/apps/sim/lib/core/security/input-validation.test.ts index a9ab3785e85..aa7967b7a0f 100644 --- a/apps/sim/lib/core/security/input-validation.test.ts +++ b/apps/sim/lib/core/security/input-validation.test.ts @@ -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'], diff --git a/apps/sim/lib/core/security/input-validation.ts b/apps/sim/lib/core/security/input-validation.ts index 4a1e460d19f..044e899f455 100644 --- a/apps/sim/lib/core/security/input-validation.ts +++ b/apps/sim/lib/core/security/input-validation.ts @@ -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) @@ -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 } } /** @@ -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', '.azuredatabricks.net', '.databricks.azure.us', '.databricks.azure.cn', @@ -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) diff --git a/apps/sim/tools/databricks/cancel_run.ts b/apps/sim/tools/databricks/cancel_run.ts index 0b5ffa3f38a..2875664ffd4 100644 --- a/apps/sim/tools/databricks/cancel_run.ts +++ b/apps/sim/tools/databricks/cancel_run.ts @@ -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 = { @@ -33,13 +34,7 @@ export const cancelRunTool: ToolConfig { - 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', diff --git a/apps/sim/tools/databricks/databricks.test.ts b/apps/sim/tools/databricks/databricks.test.ts new file mode 100644 index 00000000000..deab0b6e617 --- /dev/null +++ b/apps/sim/tools/databricks/databricks.test.ts @@ -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 + +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` + ) + } + ) +}) diff --git a/apps/sim/tools/databricks/execute_sql.ts b/apps/sim/tools/databricks/execute_sql.ts index 8564c2b3416..ae2d66e8898 100644 --- a/apps/sim/tools/databricks/execute_sql.ts +++ b/apps/sim/tools/databricks/execute_sql.ts @@ -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 = @@ -65,13 +66,7 @@ export const executeSqlTool: ToolConfig { - 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', diff --git a/apps/sim/tools/databricks/get_cluster.ts b/apps/sim/tools/databricks/get_cluster.ts index 150b77f72ae..ac1b9b34dc6 100644 --- a/apps/sim/tools/databricks/get_cluster.ts +++ b/apps/sim/tools/databricks/get_cluster.ts @@ -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 = @@ -35,11 +36,7 @@ export const getClusterTool: ToolConfig { - 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() }, diff --git a/apps/sim/tools/databricks/get_job.ts b/apps/sim/tools/databricks/get_job.ts index e96cb48712e..37ca54b6386 100644 --- a/apps/sim/tools/databricks/get_job.ts +++ b/apps/sim/tools/databricks/get_job.ts @@ -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 = { @@ -30,11 +31,7 @@ export const getJobTool: ToolConfig { - 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() }, diff --git a/apps/sim/tools/databricks/get_run.ts b/apps/sim/tools/databricks/get_run.ts index f53190ffb4d..74b316f9edf 100644 --- a/apps/sim/tools/databricks/get_run.ts +++ b/apps/sim/tools/databricks/get_run.ts @@ -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 = { @@ -42,11 +43,7 @@ export const getRunTool: ToolConfig { - 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') diff --git a/apps/sim/tools/databricks/get_run_output.ts b/apps/sim/tools/databricks/get_run_output.ts index 1fefd3da54e..824e09a9feb 100644 --- a/apps/sim/tools/databricks/get_run_output.ts +++ b/apps/sim/tools/databricks/get_run_output.ts @@ -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< @@ -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', diff --git a/apps/sim/tools/databricks/get_statement.ts b/apps/sim/tools/databricks/get_statement.ts index ea32d4487a0..24c49b31fc2 100644 --- a/apps/sim/tools/databricks/get_statement.ts +++ b/apps/sim/tools/databricks/get_statement.ts @@ -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< @@ -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', diff --git a/apps/sim/tools/databricks/list_clusters.ts b/apps/sim/tools/databricks/list_clusters.ts index cc2e31491b3..2b9c6bde4e2 100644 --- a/apps/sim/tools/databricks/list_clusters.ts +++ b/apps/sim/tools/databricks/list_clusters.ts @@ -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 = { @@ -24,13 +25,7 @@ export const listClustersTool: ToolConfig { - 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', diff --git a/apps/sim/tools/databricks/list_jobs.ts b/apps/sim/tools/databricks/list_jobs.ts index 194493c7ed7..0ed73c5e0f5 100644 --- a/apps/sim/tools/databricks/list_jobs.ts +++ b/apps/sim/tools/databricks/list_jobs.ts @@ -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 = { @@ -48,11 +49,7 @@ export const listJobsTool: ToolConfig { - 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) diff --git a/apps/sim/tools/databricks/list_runs.ts b/apps/sim/tools/databricks/list_runs.ts index 69dc333f023..d29a35f5ca1 100644 --- a/apps/sim/tools/databricks/list_runs.ts +++ b/apps/sim/tools/databricks/list_runs.ts @@ -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 = { @@ -73,11 +74,7 @@ export const listRunsTool: ToolConfig { - 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') diff --git a/apps/sim/tools/databricks/list_warehouses.ts b/apps/sim/tools/databricks/list_warehouses.ts index dcb716360ef..5688578c8c8 100644 --- a/apps/sim/tools/databricks/list_warehouses.ts +++ b/apps/sim/tools/databricks/list_warehouses.ts @@ -2,6 +2,7 @@ import type { DatabricksBaseParams, DatabricksListWarehousesResponse, } from '@/tools/databricks/types' +import { databricksUrl } from '@/tools/databricks/utils' import type { ToolConfig } from '@/tools/types' export const listWarehousesTool: ToolConfig< @@ -30,13 +31,7 @@ export const listWarehousesTool: ToolConfig< }, request: { - url: (params) => { - const host = params.host - .trim() - .replace(/^https?:\/\//, '') - .replace(/\/$/, '') - return `https://${host}/api/2.0/sql/warehouses` - }, + url: (params) => databricksUrl(params.host, '/api/2.0/sql/warehouses'), method: 'GET', headers: (params) => ({ Accept: 'application/json', diff --git a/apps/sim/tools/databricks/run_job.ts b/apps/sim/tools/databricks/run_job.ts index a7d93a7db11..f21b94ed6c5 100644 --- a/apps/sim/tools/databricks/run_job.ts +++ b/apps/sim/tools/databricks/run_job.ts @@ -1,5 +1,6 @@ import { getErrorMessage } from '@sim/utils/errors' import type { DatabricksRunJobParams, DatabricksRunJobResponse } from '@/tools/databricks/types' +import { databricksUrl } from '@/tools/databricks/utils' import type { ToolConfig } from '@/tools/types' export const runJobTool: ToolConfig = { @@ -49,13 +50,7 @@ export const runJobTool: ToolConfig { - const host = params.host - .trim() - .replace(/^https?:\/\//, '') - .replace(/\/$/, '') - return `https://${host}/api/2.1/jobs/run-now` - }, + url: (params) => databricksUrl(params.host, '/api/2.1/jobs/run-now'), method: 'POST', headers: (params) => ({ 'Content-Type': 'application/json', diff --git a/apps/sim/tools/databricks/utils.ts b/apps/sim/tools/databricks/utils.ts index c99815027ed..95ae001dc07 100644 --- a/apps/sim/tools/databricks/utils.ts +++ b/apps/sim/tools/databricks/utils.ts @@ -1,3 +1,4 @@ +import { validateDatabricksWorkspaceHost } from '@/lib/core/security/input-validation' import type { DatabricksGenieAgentItem, DatabricksGenieMessage, @@ -66,13 +67,17 @@ export const GENIE_MESSAGE_PARAMS = { }, } as const satisfies ToolConfig['params'] -/** Builds an absolute Databricks REST URL from a workspace host that may carry a scheme or slash. */ +/** + * Builds an absolute Databricks REST URL, refusing any host outside the Databricks workspace + * domains so the access token is only ever sent to Databricks. The host may carry a scheme or + * trailing slash; an `http://` scheme is upgraded to HTTPS as before, not rejected. + */ export function databricksUrl(host: string, path: string): string { - const normalizedHost = host - .trim() - .replace(/^https?:\/\//, '') - .replace(/\/$/, '') - return `https://${normalizedHost}${path}` + const result = validateDatabricksWorkspaceHost(host.trim().replace(/^http:\/\//i, ''), 'host') + if (!result.isValid || !result.sanitized) { + throw new Error(result.error || 'Invalid Databricks workspace host') + } + return `${result.sanitized}${path}` } /** Path of a Genie space: `/api/2.0/genie/spaces/{space_id}`. */