diff --git a/apps/studio/lib/ai/ai-details.test.ts b/apps/studio/lib/ai/ai-details.test.ts new file mode 100644 index 00000000000..b1a2b8f0375 --- /dev/null +++ b/apps/studio/lib/ai/ai-details.test.ts @@ -0,0 +1,239 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest' + +import { getOrgAIDetails, getProjectAIDetails } from './ai-details' + +vi.mock('data/organizations/organizations-query', () => ({ + getOrganizations: vi.fn(), +})) + +vi.mock('data/projects/project-detail-query', () => ({ + getProjectDetail: vi.fn(), +})) + +vi.mock('data/subscriptions/org-subscription-query', () => ({ + getOrgSubscription: vi.fn(), +})) + +vi.mock('data/config/project-settings-v2-query', () => ({ + getProjectSettings: vi.fn(), +})) + +vi.mock('hooks/misc/useOrgOptedIntoAi', () => ({ + getAiOptInLevel: vi.fn(), +})) + +vi.mock('components/interfaces/Billing/Subscription/Subscription.utils', () => ({ + subscriptionHasHipaaAddon: vi.fn(), +})) + +vi.mock('data/entitlements/entitlements-query', () => ({ + checkEntitlement: vi.fn(), +})) + +vi.mock('data/fetchers', () => ({ + get: vi.fn(), +})) + +const AUTH = 'Bearer token' +const HEADERS = { 'Content-Type': 'application/json', Authorization: AUTH } + +describe('getOrgAIDetails', () => { + let mockGetOrganizations: ReturnType + let mockGetOrgSubscription: ReturnType + let mockGetAiOptInLevel: ReturnType + let mockSubscriptionHasHipaaAddon: ReturnType + let mockCheckEntitlement: ReturnType + let mockGet: ReturnType + + beforeEach(async () => { + vi.clearAllMocks() + + const orgsQuery = await import('data/organizations/organizations-query') + const subscriptionQuery = await import('data/subscriptions/org-subscription-query') + const aiHook = await import('hooks/misc/useOrgOptedIntoAi') + const subscriptionUtils = + await import('components/interfaces/Billing/Subscription/Subscription.utils') + const entitlementsQuery = await import('data/entitlements/entitlements-query') + const fetchers = await import('data/fetchers') + + mockGetOrganizations = vi.mocked(orgsQuery.getOrganizations) + mockGetOrgSubscription = vi.mocked(subscriptionQuery.getOrgSubscription) + mockGetAiOptInLevel = vi.mocked(aiHook.getAiOptInLevel) + mockSubscriptionHasHipaaAddon = vi.mocked(subscriptionUtils.subscriptionHasHipaaAddon) + mockCheckEntitlement = vi.mocked(entitlementsQuery.checkEntitlement) + mockGet = vi.mocked(fetchers.get) + + mockGetOrgSubscription.mockResolvedValue({ addons: [] }) + mockSubscriptionHasHipaaAddon.mockReturnValue(false) + mockCheckEntitlement.mockResolvedValue({ hasAccess: false }) + mockGet.mockResolvedValue({ data: { signed: false } }) + }) + + it('returns org-level fields', async () => { + mockGetOrganizations.mockResolvedValue([ + { id: 1, slug: 'test-org', plan: { id: 'pro' }, opt_in_tags: [] }, + ]) + mockGetAiOptInLevel.mockReturnValue('schema') + + const result = await getOrgAIDetails({ orgSlug: 'test-org', authorization: AUTH }) + + expect(result).toEqual({ + aiOptInLevel: 'schema', + hasAccessToAdvanceModel: false, + hasHipaaAddon: false, + isDpaSigned: false, + orgId: 1, + planId: 'pro', + }) + }) + + it('returns hasAccessToAdvanceModel true when entitlement grants access', async () => { + mockGetOrganizations.mockResolvedValue([ + { id: 1, slug: 'test-org', plan: { id: 'pro' }, opt_in_tags: [] }, + ]) + mockGetAiOptInLevel.mockReturnValue('schema') + mockCheckEntitlement.mockResolvedValue({ hasAccess: true }) + + const result = await getOrgAIDetails({ orgSlug: 'test-org', authorization: AUTH }) + + expect(result.hasAccessToAdvanceModel).toBe(true) + }) + + it('returns hasHipaaAddon from subscription', async () => { + mockGetOrganizations.mockResolvedValue([ + { id: 1, slug: 'test-org', plan: { id: 'enterprise' }, opt_in_tags: [] }, + ]) + mockGetAiOptInLevel.mockReturnValue('schema') + mockSubscriptionHasHipaaAddon.mockReturnValue(true) + + const result = await getOrgAIDetails({ orgSlug: 'test-org', authorization: AUTH }) + + expect(result.hasHipaaAddon).toBe(true) + }) + + it('returns isDpaSigned true when endpoint returns signed: true', async () => { + mockGetOrganizations.mockResolvedValue([ + { id: 1, slug: 'test-org', plan: { id: 'pro' }, opt_in_tags: [] }, + ]) + mockGetAiOptInLevel.mockReturnValue('schema') + mockGet.mockResolvedValue({ data: { signed: true } }) + + const result = await getOrgAIDetails({ orgSlug: 'test-org', authorization: AUTH }) + + expect(result.isDpaSigned).toBe(true) + }) + + it('returns isDpaSigned undefined when endpoint fails', async () => { + mockGetOrganizations.mockResolvedValue([ + { id: 1, slug: 'test-org', plan: { id: 'pro' }, opt_in_tags: [] }, + ]) + mockGetAiOptInLevel.mockReturnValue('schema') + mockGet.mockResolvedValue({ data: null, error: { message: 'Unauthorized' } }) + + const result = await getOrgAIDetails({ orgSlug: 'test-org', authorization: AUTH }) + + expect(result.isDpaSigned).toBeUndefined() + }) + + it('calls getAiOptInLevel with the matched org opt_in_tags', async () => { + const opt_in_tags = ['AI_SQL_GENERATOR_OPT_IN'] + mockGetOrganizations.mockResolvedValue([ + { id: 1, slug: 'test-org', plan: { id: 'pro' }, opt_in_tags }, + ]) + mockGetAiOptInLevel.mockReturnValue('schema') + + await getOrgAIDetails({ orgSlug: 'test-org', authorization: AUTH }) + + expect(mockGetAiOptInLevel).toHaveBeenCalledWith(opt_in_tags) + }) + + it('forwards authorization headers to all fetches', async () => { + mockGetOrganizations.mockResolvedValue([ + { id: 1, slug: 'test-org', plan: { id: 'pro' }, opt_in_tags: [] }, + ]) + mockGetAiOptInLevel.mockReturnValue('schema') + + await getOrgAIDetails({ orgSlug: 'test-org', authorization: AUTH }) + + expect(mockGetOrganizations).toHaveBeenCalledWith({ headers: HEADERS }) + expect(mockGetOrgSubscription).toHaveBeenCalledWith({ orgSlug: 'test-org' }, undefined, HEADERS) + }) + + it('finds the correct org when multiple orgs are returned', async () => { + mockGetOrganizations.mockResolvedValue([ + { id: 1, slug: 'org-1', plan: { id: 'free' }, opt_in_tags: [] }, + { id: 2, slug: 'test-org', plan: { id: 'pro' }, opt_in_tags: [] }, + ]) + mockGetAiOptInLevel.mockReturnValue('schema') + + const result = await getOrgAIDetails({ orgSlug: 'test-org', authorization: AUTH }) + + expect(result.orgId).toBe(2) + expect(result.planId).toBe('pro') + }) +}) + +describe('getProjectAIDetails', () => { + let mockGetProjectDetail: ReturnType + let mockGetProjectSettings: ReturnType + + beforeEach(async () => { + vi.clearAllMocks() + + const projectQuery = await import('data/projects/project-detail-query') + const settingsQuery = await import('data/config/project-settings-v2-query') + + mockGetProjectDetail = vi.mocked(projectQuery.getProjectDetail) + mockGetProjectSettings = vi.mocked(settingsQuery.getProjectSettings) + }) + + it('returns region and isSensitive', async () => { + mockGetProjectDetail.mockResolvedValue({ ref: 'test-project', region: 'us-east-1' }) + mockGetProjectSettings.mockResolvedValue({ is_sensitive: false }) + + const result = await getProjectAIDetails({ projectRef: 'test-project', authorization: AUTH }) + + expect(result).toEqual({ region: 'us-east-1', isSensitive: false }) + }) + + it('returns isSensitive true when project is marked sensitive', async () => { + mockGetProjectDetail.mockResolvedValue({ ref: 'test-project', region: 'us-east-1' }) + mockGetProjectSettings.mockResolvedValue({ is_sensitive: true }) + + const result = await getProjectAIDetails({ projectRef: 'test-project', authorization: AUTH }) + + expect(result.isSensitive).toBe(true) + }) + + it('returns isSensitive undefined when project settings are unavailable', async () => { + mockGetProjectDetail.mockResolvedValue({ ref: 'test-project', region: 'us-east-1' }) + mockGetProjectSettings.mockResolvedValue(undefined) + + const result = await getProjectAIDetails({ projectRef: 'test-project', authorization: AUTH }) + + expect(result.isSensitive).toBeUndefined() + }) + + it('returns region undefined when project detail is unavailable', async () => { + mockGetProjectDetail.mockResolvedValue(undefined) + mockGetProjectSettings.mockResolvedValue({ is_sensitive: false }) + + const result = await getProjectAIDetails({ projectRef: 'test-project', authorization: AUTH }) + + expect(result.region).toBeUndefined() + }) + + it('forwards authorization headers to all fetches', async () => { + mockGetProjectDetail.mockResolvedValue({ ref: 'test-project', region: 'us-east-1' }) + mockGetProjectSettings.mockResolvedValue({ is_sensitive: false }) + + await getProjectAIDetails({ projectRef: 'test-project', authorization: AUTH }) + + expect(mockGetProjectDetail).toHaveBeenCalledWith({ ref: 'test-project' }, undefined, HEADERS) + expect(mockGetProjectSettings).toHaveBeenCalledWith( + { projectRef: 'test-project' }, + undefined, + HEADERS + ) + }) +}) diff --git a/apps/studio/lib/ai/ai-details.ts b/apps/studio/lib/ai/ai-details.ts new file mode 100644 index 00000000000..e467ef07d57 --- /dev/null +++ b/apps/studio/lib/ai/ai-details.ts @@ -0,0 +1,65 @@ +import { subscriptionHasHipaaAddon } from 'components/interfaces/Billing/Subscription/Subscription.utils' +import { getProjectSettings } from 'data/config/project-settings-v2-query' +import { checkEntitlement } from 'data/entitlements/entitlements-query' +import { get } from 'data/fetchers' +import { getOrganizations } from 'data/organizations/organizations-query' +import { getProjectDetail } from 'data/projects/project-detail-query' +import { getOrgSubscription } from 'data/subscriptions/org-subscription-query' +import { getAiOptInLevel } from 'hooks/misc/useOrgOptedIntoAi' + +export const getOrgAIDetails = async ({ + orgSlug, + authorization, +}: { + orgSlug: string + authorization: string +}) => { + const headers = { + 'Content-Type': 'application/json', + ...(authorization && { Authorization: authorization }), + } + + const [organizations, subscription, advanceModelAccess, dpaSignedStatus] = await Promise.all([ + getOrganizations({ headers }), + getOrgSubscription({ orgSlug }, undefined, headers), + checkEntitlement(orgSlug, 'assistant.advance_model', undefined, headers), + get('/platform/organizations/{slug}/documents/dpa-signed', { + params: { path: { slug: orgSlug } }, + headers, + }), + ]) + + const selectedOrg = organizations.find((org) => org.slug === orgSlug) + + return { + aiOptInLevel: getAiOptInLevel(selectedOrg?.opt_in_tags), + hasAccessToAdvanceModel: advanceModelAccess.hasAccess, + hasHipaaAddon: subscriptionHasHipaaAddon(subscription), + isDpaSigned: dpaSignedStatus.data?.signed, + orgId: selectedOrg?.id, + planId: selectedOrg?.plan.id, + } +} + +export const getProjectAIDetails = async ({ + projectRef, + authorization, +}: { + projectRef: string + authorization: string +}) => { + const headers = { + 'Content-Type': 'application/json', + ...(authorization && { Authorization: authorization }), + } + + const [selectedProject, projectSettings] = await Promise.all([ + getProjectDetail({ ref: projectRef }, undefined, headers), + getProjectSettings({ projectRef }, undefined, headers), + ]) + + return { + region: selectedProject?.region, + isSensitive: projectSettings?.is_sensitive, + } +} diff --git a/apps/studio/lib/ai/braintrust-logger.test.ts b/apps/studio/lib/ai/braintrust-logger.test.ts new file mode 100644 index 00000000000..5343f8c2111 --- /dev/null +++ b/apps/studio/lib/ai/braintrust-logger.test.ts @@ -0,0 +1,79 @@ +import { describe, expect, it } from 'vitest' + +import { isTracingAllowed } from './braintrust-logger' + +const baseAllowed = { + orgHasHipaaAddon: false, + projectIsSensitive: false, + orgIsDpaSigned: false, + projectRegion: 'us-east-1', +} + +describe('isTracingAllowed', () => { + it('allows tracing when all flags are explicitly off/non-EU', () => { + expect(isTracingAllowed(baseAllowed)).toBe(true) + }) + + it('disallows tracing when HIPAA addon is active and project is sensitive', () => { + expect( + isTracingAllowed({ ...baseAllowed, orgHasHipaaAddon: true, projectIsSensitive: true }) + ).toBe(false) + }) + + it('allows tracing when HIPAA addon is active but project is not sensitive', () => { + expect( + isTracingAllowed({ ...baseAllowed, orgHasHipaaAddon: true, projectIsSensitive: false }) + ).toBe(true) + }) + + it('allows tracing when project is sensitive but no HIPAA addon', () => { + expect( + isTracingAllowed({ ...baseAllowed, orgHasHipaaAddon: false, projectIsSensitive: true }) + ).toBe(true) + }) + + it('disallows tracing when DPA is signed', () => { + expect(isTracingAllowed({ ...baseAllowed, orgIsDpaSigned: true })).toBe(false) + }) + + it('disallows tracing for EU regions', () => { + expect(isTracingAllowed({ ...baseAllowed, projectRegion: 'eu-west-1' })).toBe(false) + expect(isTracingAllowed({ ...baseAllowed, projectRegion: 'eu-central-1' })).toBe(false) + }) + + it('allows tracing for non-EU regions', () => { + expect(isTracingAllowed({ ...baseAllowed, projectRegion: 'ap-southeast-1' })).toBe(true) + }) + + it('allows tracing when HIPAA addon is false and is_sensitive is null (DB default)', () => { + expect(isTracingAllowed({ ...baseAllowed, projectIsSensitive: null })).toBe(true) + }) + + it('disallows tracing when HIPAA addon is unknown and is_sensitive is null', () => { + expect( + isTracingAllowed({ ...baseAllowed, orgHasHipaaAddon: undefined, projectIsSensitive: null }) + ).toBe(false) + }) + + it('disallows tracing when flags are undefined (unknown = restricted)', () => { + expect( + isTracingAllowed({ + orgHasHipaaAddon: undefined, + projectIsSensitive: undefined, + orgIsDpaSigned: undefined, + projectRegion: undefined, + }) + ).toBe(false) + expect(isTracingAllowed({ ...baseAllowed, orgIsDpaSigned: undefined })).toBe(false) + expect(isTracingAllowed({ ...baseAllowed, projectRegion: undefined })).toBe(false) + expect(isTracingAllowed({ ...baseAllowed, orgHasHipaaAddon: undefined })).toBe(false) + // projectIsSensitive unknown only matters when orgHasHipaaAddon is also unknown + expect( + isTracingAllowed({ + ...baseAllowed, + orgHasHipaaAddon: undefined, + projectIsSensitive: undefined, + }) + ).toBe(false) + }) +}) diff --git a/apps/studio/lib/ai/braintrust-logger.ts b/apps/studio/lib/ai/braintrust-logger.ts index cae983f9b34..a98e157369f 100644 --- a/apps/studio/lib/ai/braintrust-logger.ts +++ b/apps/studio/lib/ai/braintrust-logger.ts @@ -15,3 +15,31 @@ if (IS_TRACING_ENABLED) { projectId: BRAINTRUST_PROJECT_ID, }) } + +// Checks compliance flags for tracing and returns true only when all checks pass. +// Defaults to disabling tracing when states are unknown. +export function isTracingAllowed({ + orgHasHipaaAddon, + projectIsSensitive, + orgIsDpaSigned, + projectRegion, +}: { + orgHasHipaaAddon: boolean | undefined + projectIsSensitive: boolean | null | undefined + orgIsDpaSigned: boolean | undefined + projectRegion: string | undefined +}) { + // Disable tracing for orgs with a signed (or unknown) DPA status + if (orgIsDpaSigned !== false) return false + + // Disable tracing for EU (or unknown) regions + if (projectRegion === undefined || projectRegion.startsWith('eu-')) return false + + // Disable tracing for orgs with an unknown HIPAA addon state + if (orgHasHipaaAddon === undefined) return false + + // Disable tracing for projects within a HIPAA-enabled org that are sensitive (or unknown sensitivity) + if (orgHasHipaaAddon && projectIsSensitive !== false) return false + + return true +} diff --git a/apps/studio/lib/ai/generate-assistant-response.ts b/apps/studio/lib/ai/generate-assistant-response.ts index 4229fdccd5c..558f67ac5b4 100644 --- a/apps/studio/lib/ai/generate-assistant-response.ts +++ b/apps/studio/lib/ai/generate-assistant-response.ts @@ -28,7 +28,7 @@ export async function generateAssistantResponse({ projectRef, chatId, chatName, - isHipaaEnabled, + allowTracing, userId, orgId, planId, @@ -46,7 +46,7 @@ export async function generateAssistantResponse({ projectRef?: string chatId?: string chatName?: string - isHipaaEnabled?: boolean + allowTracing?: boolean userId?: string orgId?: number planId?: string @@ -56,7 +56,7 @@ export async function generateAssistantResponse({ abortSignal?: AbortSignal onSpanCreated?: (spanId: string) => void }) { - const shouldTrace = IS_TRACING_ENABLED && !isHipaaEnabled + const shouldTrace = allowTracing ?? IS_TRACING_ENABLED const run = async (span?: Span) => { // Only returns last 7 messages @@ -87,7 +87,9 @@ export async function generateAssistantResponse({ const schemasString = aiOptInLevel !== 'disabled' && getSchemas - ? await traced(async () => getSchemas(), { name: 'getSchemas', type: 'function' }) + ? shouldTrace + ? await traced(async () => getSchemas(), { name: 'getSchemas', type: 'function' }) + : await getSchemas() : "You don't have access to any schemas." // Important: do not use dynamic content in the system prompt or Bedrock will not cache it diff --git a/apps/studio/lib/ai/org-ai-details.test.ts b/apps/studio/lib/ai/org-ai-details.test.ts deleted file mode 100644 index ff0ebce2f20..00000000000 --- a/apps/studio/lib/ai/org-ai-details.test.ts +++ /dev/null @@ -1,390 +0,0 @@ -import { beforeEach, describe, expect, it, vi } from 'vitest' - -import { getOrgAIDetails } from './org-ai-details' - -vi.mock('data/organizations/organizations-query', () => ({ - getOrganizations: vi.fn(), -})) - -vi.mock('data/projects/project-detail-query', () => ({ - getProjectDetail: vi.fn(), -})) - -vi.mock('data/subscriptions/org-subscription-query', () => ({ - getOrgSubscription: vi.fn(), -})) - -vi.mock('data/config/project-settings-v2-query', () => ({ - getProjectSettings: vi.fn(), -})) - -vi.mock('hooks/misc/useOrgOptedIntoAi', () => ({ - getAiOptInLevel: vi.fn(), -})) - -vi.mock('components/interfaces/Billing/Subscription/Subscription.utils', () => ({ - subscriptionHasHipaaAddon: vi.fn(), -})) - -vi.mock('data/entitlements/entitlements-query', () => ({ - checkEntitlement: vi.fn(), -})) - -describe('ai/org-ai-details', () => { - let mockGetOrganizations: ReturnType - let mockGetProjectDetail: ReturnType - let mockGetOrgSubscription: ReturnType - let mockGetProjectSettings: ReturnType - let mockGetAiOptInLevel: ReturnType - let mockSubscriptionHasHipaaAddon: ReturnType - let mockCheckEntitlement: ReturnType - - beforeEach(async () => { - vi.clearAllMocks() - - const orgsQuery = await import('data/organizations/organizations-query') - const projectQuery = await import('data/projects/project-detail-query') - const subscriptionQuery = await import('data/subscriptions/org-subscription-query') - const settingsQuery = await import('data/config/project-settings-v2-query') - const aiHook = await import('hooks/misc/useOrgOptedIntoAi') - const subscriptionUtils = - await import('components/interfaces/Billing/Subscription/Subscription.utils') - const entitlementsQuery = await import('data/entitlements/entitlements-query') - - mockGetOrganizations = vi.mocked(orgsQuery.getOrganizations) - mockGetProjectDetail = vi.mocked(projectQuery.getProjectDetail) - mockGetOrgSubscription = vi.mocked(subscriptionQuery.getOrgSubscription) - mockGetProjectSettings = vi.mocked(settingsQuery.getProjectSettings) - mockGetAiOptInLevel = vi.mocked(aiHook.getAiOptInLevel) - mockSubscriptionHasHipaaAddon = vi.mocked(subscriptionUtils.subscriptionHasHipaaAddon) - mockCheckEntitlement = vi.mocked(entitlementsQuery.checkEntitlement) - - // Default mocks for subscription/settings (no HIPAA) - mockGetOrgSubscription.mockResolvedValue({ addons: [] }) - mockGetProjectSettings.mockResolvedValue({ is_sensitive: false }) - mockSubscriptionHasHipaaAddon.mockReturnValue(false) - mockCheckEntitlement.mockResolvedValue({ hasAccess: false }) - }) - - describe('getOrgAIDetails', () => { - it('should fetch organizations and project details', async () => { - const mockOrg = { - id: 1, - slug: 'test-org', - plan: { id: 'pro' }, - opt_in_tags: [], - } - const mockProject = { - id: 1, - organization_id: 1, - ref: 'test-project', - } - - mockGetOrganizations.mockResolvedValue([mockOrg]) - mockGetProjectDetail.mockResolvedValue(mockProject) - mockGetAiOptInLevel.mockReturnValue('full') - - await getOrgAIDetails({ - orgSlug: 'test-org', - authorization: 'Bearer token', - projectRef: 'test-project', - }) - - expect(mockGetOrganizations).toHaveBeenCalledWith({ - headers: { - 'Content-Type': 'application/json', - Authorization: 'Bearer token', - }, - }) - expect(mockGetProjectDetail).toHaveBeenCalledWith({ ref: 'test-project' }, undefined, { - 'Content-Type': 'application/json', - Authorization: 'Bearer token', - }) - }) - - it('should return AI opt-in level and assistant advance-model flag', async () => { - const mockOrg = { - id: 1, - slug: 'test-org', - plan: { id: 'free' }, - opt_in_tags: ['AI_SQL_GENERATOR_OPT_IN'], - } - const mockProject = { - organization_id: 1, - ref: 'test-project', - } - - mockGetOrganizations.mockResolvedValue([mockOrg]) - mockGetProjectDetail.mockResolvedValue(mockProject) - mockGetAiOptInLevel.mockReturnValue('schema_only') - - const result = await getOrgAIDetails({ - orgSlug: 'test-org', - authorization: 'Bearer token', - projectRef: 'test-project', - }) - - expect(result).toEqual({ - aiOptInLevel: 'schema_only', - hasAccessToAdvanceModel: false, - isHipaaEnabled: false, - orgId: 1, - planId: 'free', - }) - }) - - it('should set hasAccessToAdvanceModel when entitlement grants access', async () => { - const mockOrg = { - id: 1, - slug: 'test-org', - plan: { id: 'pro' }, - opt_in_tags: [], - } - const mockProject = { - organization_id: 1, - } - - mockGetOrganizations.mockResolvedValue([mockOrg]) - mockGetProjectDetail.mockResolvedValue(mockProject) - mockGetAiOptInLevel.mockReturnValue('full') - mockCheckEntitlement.mockResolvedValue({ hasAccess: true }) - - const result = await getOrgAIDetails({ - orgSlug: 'test-org', - authorization: 'Bearer token', - projectRef: 'test-project', - }) - - expect(result.hasAccessToAdvanceModel).toBe(true) - }) - - it('should throw error when project and org do not match', async () => { - const mockOrg = { - id: 1, - slug: 'test-org', - plan: { id: 'pro' }, - } - const mockProject = { - organization_id: 2, // Different org ID - } - - mockGetOrganizations.mockResolvedValue([mockOrg]) - mockGetProjectDetail.mockResolvedValue(mockProject) - - await expect( - getOrgAIDetails({ - orgSlug: 'test-org', - authorization: 'Bearer token', - projectRef: 'test-project', - }) - ).rejects.toThrow('Project and organization do not match') - }) - - it('should handle org not found', async () => { - const mockProject = { - organization_id: 1, - } - - mockGetOrganizations.mockResolvedValue([]) // No orgs - mockGetProjectDetail.mockResolvedValue(mockProject) - - await expect( - getOrgAIDetails({ - orgSlug: 'non-existent-org', - authorization: 'Bearer token', - projectRef: 'test-project', - }) - ).rejects.toThrow('Project and organization do not match') - }) - - it('should call getAiOptInLevel with org opt_in_tags', async () => { - const mockOptInTags = ['AI_SQL_GENERATOR_OPT_IN', 'AI_DATA_GENERATOR_OPT_IN'] - const mockOrg = { - id: 1, - slug: 'test-org', - plan: { id: 'pro' }, - opt_in_tags: mockOptInTags, - } - const mockProject = { - organization_id: 1, - } - - mockGetOrganizations.mockResolvedValue([mockOrg]) - mockGetProjectDetail.mockResolvedValue(mockProject) - mockGetAiOptInLevel.mockReturnValue('full') - - await getOrgAIDetails({ - orgSlug: 'test-org', - authorization: 'Bearer token', - projectRef: 'test-project', - }) - - expect(mockGetAiOptInLevel).toHaveBeenCalledWith(mockOptInTags) - }) - - it('should include authorization header when provided', async () => { - const mockOrg = { - id: 1, - slug: 'test-org', - plan: { id: 'pro' }, - } - const mockProject = { - organization_id: 1, - } - - mockGetOrganizations.mockResolvedValue([mockOrg]) - mockGetProjectDetail.mockResolvedValue(mockProject) - mockGetAiOptInLevel.mockReturnValue('full') - - await getOrgAIDetails({ - orgSlug: 'test-org', - authorization: 'Bearer custom-token', - projectRef: 'test-project', - }) - - expect(mockGetOrganizations).toHaveBeenCalledWith({ - headers: expect.objectContaining({ - Authorization: 'Bearer custom-token', - }), - }) - expect(mockGetProjectDetail).toHaveBeenCalledWith( - expect.anything(), - undefined, - expect.objectContaining({ - Authorization: 'Bearer custom-token', - }) - ) - }) - - it('should fetch multiple organizations and find correct one', async () => { - const mockOrgs = [ - { id: 1, slug: 'org-1', plan: { id: 'free' } }, - { id: 2, slug: 'test-org', plan: { id: 'pro' } }, - { id: 3, slug: 'org-3', plan: { id: 'team' } }, - ] - const mockProject = { - organization_id: 2, - } - - mockGetOrganizations.mockResolvedValue(mockOrgs) - mockGetProjectDetail.mockResolvedValue(mockProject) - mockGetAiOptInLevel.mockReturnValue('full') - mockCheckEntitlement.mockResolvedValue({ hasAccess: true }) - - const result = await getOrgAIDetails({ - orgSlug: 'test-org', - authorization: 'Bearer token', - projectRef: 'test-project', - }) - - expect(result.hasAccessToAdvanceModel).toBe(true) - }) - - it('should return isHipaaEnabled true when subscription has HIPAA addon and project is sensitive', async () => { - const mockOrg = { - id: 1, - slug: 'test-org', - plan: { id: 'enterprise' }, - opt_in_tags: [], - } - const mockProject = { organization_id: 1 } - - mockGetOrganizations.mockResolvedValue([mockOrg]) - mockGetProjectDetail.mockResolvedValue(mockProject) - mockGetAiOptInLevel.mockReturnValue('schema') - mockSubscriptionHasHipaaAddon.mockReturnValue(true) - mockGetProjectSettings.mockResolvedValue({ is_sensitive: true }) - - const result = await getOrgAIDetails({ - orgSlug: 'test-org', - authorization: 'Bearer token', - projectRef: 'test-project', - }) - - expect(result.isHipaaEnabled).toBe(true) - }) - - it('should return isHipaaEnabled false when subscription has HIPAA addon but project is not sensitive', async () => { - const mockOrg = { - id: 1, - slug: 'test-org', - plan: { id: 'enterprise' }, - opt_in_tags: [], - } - const mockProject = { organization_id: 1 } - - mockGetOrganizations.mockResolvedValue([mockOrg]) - mockGetProjectDetail.mockResolvedValue(mockProject) - mockGetAiOptInLevel.mockReturnValue('schema') - mockSubscriptionHasHipaaAddon.mockReturnValue(true) - mockGetProjectSettings.mockResolvedValue({ is_sensitive: false }) - - const result = await getOrgAIDetails({ - orgSlug: 'test-org', - authorization: 'Bearer token', - projectRef: 'test-project', - }) - - expect(result.isHipaaEnabled).toBe(false) - }) - - it('should return isHipaaEnabled false when project is sensitive but no HIPAA addon', async () => { - const mockOrg = { - id: 1, - slug: 'test-org', - plan: { id: 'pro' }, - opt_in_tags: [], - } - const mockProject = { organization_id: 1 } - - mockGetOrganizations.mockResolvedValue([mockOrg]) - mockGetProjectDetail.mockResolvedValue(mockProject) - mockGetAiOptInLevel.mockReturnValue('schema') - mockSubscriptionHasHipaaAddon.mockReturnValue(false) - mockGetProjectSettings.mockResolvedValue({ is_sensitive: true }) - - const result = await getOrgAIDetails({ - orgSlug: 'test-org', - authorization: 'Bearer token', - projectRef: 'test-project', - }) - - expect(result.isHipaaEnabled).toBe(false) - }) - - it('should fetch subscription and project settings with authorization headers', async () => { - const mockOrg = { - id: 1, - slug: 'test-org', - plan: { id: 'pro' }, - } - const mockProject = { organization_id: 1 } - - mockGetOrganizations.mockResolvedValue([mockOrg]) - mockGetProjectDetail.mockResolvedValue(mockProject) - mockGetAiOptInLevel.mockReturnValue('schema') - - await getOrgAIDetails({ - orgSlug: 'test-org', - authorization: 'Bearer token', - projectRef: 'test-project', - }) - - const expectedHeaders = { - 'Content-Type': 'application/json', - Authorization: 'Bearer token', - } - - expect(mockGetOrgSubscription).toHaveBeenCalledWith( - { orgSlug: 'test-org' }, - undefined, - expectedHeaders - ) - expect(mockGetProjectSettings).toHaveBeenCalledWith( - { projectRef: 'test-project' }, - undefined, - expectedHeaders - ) - }) - }) -}) diff --git a/apps/studio/lib/ai/org-ai-details.ts b/apps/studio/lib/ai/org-ai-details.ts deleted file mode 100644 index 68e70af900c..00000000000 --- a/apps/studio/lib/ai/org-ai-details.ts +++ /dev/null @@ -1,50 +0,0 @@ -import { subscriptionHasHipaaAddon } from 'components/interfaces/Billing/Subscription/Subscription.utils' -import { getProjectSettings } from 'data/config/project-settings-v2-query' -import { getOrganizations } from 'data/organizations/organizations-query' -import { getProjectDetail } from 'data/projects/project-detail-query' -import { getOrgSubscription } from 'data/subscriptions/org-subscription-query' -import { getAiOptInLevel } from 'hooks/misc/useOrgOptedIntoAi' -import { checkEntitlement } from 'data/entitlements/entitlements-query' - -export const getOrgAIDetails = async ({ - orgSlug, - authorization, - projectRef, -}: { - orgSlug: string - authorization: string - projectRef: string -}) => { - const headers = { - 'Content-Type': 'application/json', - ...(authorization && { Authorization: authorization }), - } - - const [organizations, selectedProject, subscription, projectSettings, advanceModelAccess] = - await Promise.all([ - getOrganizations({ headers }), - getProjectDetail({ ref: projectRef }, undefined, headers), - getOrgSubscription({ orgSlug }, undefined, headers), - getProjectSettings({ projectRef }, undefined, headers), - checkEntitlement(orgSlug, 'assistant.advance_model', undefined, headers), - ]) - - const selectedOrg = organizations.find((org) => org.slug === orgSlug) - - // If the project is not in the organization specific by the org slug, return an error - if (selectedProject?.organization_id !== selectedOrg?.id) { - throw new Error('Project and organization do not match') - } - - const aiOptInLevel = getAiOptInLevel(selectedOrg?.opt_in_tags) - const hasAccessToAdvanceModel = advanceModelAccess.hasAccess - const isHipaaEnabled = subscriptionHasHipaaAddon(subscription) && !!projectSettings?.is_sensitive - - return { - aiOptInLevel, - hasAccessToAdvanceModel, - isHipaaEnabled, - orgId: selectedOrg?.id, - planId: selectedOrg?.plan.id, - } -} diff --git a/apps/studio/lib/api/generate-v4.test.ts b/apps/studio/lib/api/generate-v4.test.ts index ffb69fe453d..1fb2aaacd7b 100644 --- a/apps/studio/lib/api/generate-v4.test.ts +++ b/apps/studio/lib/api/generate-v4.test.ts @@ -1,9 +1,10 @@ +import { UIMessage } from 'ai' +import { sanitizeMessagePart } from 'lib/ai/tools/tool-sanitizer' import { expect, test, vi } from 'vitest' + // End of third-party imports import generateV4 from '../../pages/api/ai/sql/generate-v4' -import { sanitizeMessagePart } from 'lib/ai/tools/tool-sanitizer' -import { UIMessage } from 'ai' vi.mock('lib/ai/tools/tool-sanitizer', () => ({ sanitizeMessagePart: vi.fn((part) => part), @@ -44,10 +45,15 @@ test('generateV4 calls the tool sanitizer', async () => { setHeader: vi.fn(() => mockRes), } - vi.mock('lib/ai/org-ai-details', () => ({ + vi.mock('lib/ai/ai-details', () => ({ getOrgAIDetails: vi.fn().mockResolvedValue({ aiOptInLevel: 'schema_and_log_and_data', hasAccessToAdvanceModel: true, + isDpaSigned: false, + }), + getProjectAIDetails: vi.fn().mockResolvedValue({ + region: 'us-east-1', + isSensitive: false, }), })) diff --git a/apps/studio/lib/api/rate.test.ts b/apps/studio/lib/api/rate.test.ts index 002a46f89fc..197da0c7b4e 100644 --- a/apps/studio/lib/api/rate.test.ts +++ b/apps/studio/lib/api/rate.test.ts @@ -1,4 +1,5 @@ import { expect, test, vi } from 'vitest' + // End of third-party imports import rate from '../../pages/api/ai/feedback/rate' @@ -42,10 +43,15 @@ test('rate calls the tool sanitizer', async () => { setHeader: vi.fn(() => mockRes), } - vi.mock('lib/ai/org-ai-details', () => ({ + vi.mock('lib/ai/ai-details', () => ({ getOrgAIDetails: vi.fn().mockResolvedValue({ aiOptInLevel: 'schema_and_log_and_data', hasAccessToAdvanceModel: true, + isDpaSigned: false, + }), + getProjectAIDetails: vi.fn().mockResolvedValue({ + region: 'us-east-1', + isSensitive: false, }), })) diff --git a/apps/studio/pages/api/ai/code/complete.ts b/apps/studio/pages/api/ai/code/complete.ts index 3af08806aab..b8611bcf0cf 100644 --- a/apps/studio/pages/api/ai/code/complete.ts +++ b/apps/studio/pages/api/ai/code/complete.ts @@ -4,9 +4,9 @@ import { IS_PLATFORM } from 'common' import { source } from 'common-tags' import { executeSql } from 'data/sql/execute-sql-query' import { AiOptInLevel } from 'hooks/misc/useOrgOptedIntoAi' +import { getOrgAIDetails } from 'lib/ai/ai-details' import { getModel } from 'lib/ai/model' import { DEFAULT_COMPLETION_MODEL } from 'lib/ai/model.utils' -import { getOrgAIDetails } from 'lib/ai/org-ai-details' import { EDGE_FUNCTION_PROMPT, GENERAL_PROMPT, @@ -58,7 +58,6 @@ async function handler(req: NextApiRequest, res: NextApiResponse) { const { aiOptInLevel: orgAIOptInLevel } = await getOrgAIDetails({ orgSlug, authorization, - projectRef, }) aiOptInLevel = orgAIOptInLevel diff --git a/apps/studio/pages/api/ai/feedback/rate.ts b/apps/studio/pages/api/ai/feedback/rate.ts index bd52bca4d62..e05fcb8920c 100644 --- a/apps/studio/pages/api/ai/feedback/rate.ts +++ b/apps/studio/pages/api/ai/feedback/rate.ts @@ -3,10 +3,10 @@ import { currentLogger } from 'braintrust' import { IS_PLATFORM } from 'common' import { rateMessageResponseSchema } from 'components/ui/AIAssistantPanel/Message.utils' import type { AiOptInLevel } from 'hooks/misc/useOrgOptedIntoAi' -import { IS_TRACING_ENABLED } from 'lib/ai/braintrust-logger' +import { getOrgAIDetails, getProjectAIDetails } from 'lib/ai/ai-details' +import { IS_TRACING_ENABLED, isTracingAllowed } from 'lib/ai/braintrust-logger' import { getModel } from 'lib/ai/model' import { DEFAULT_COMPLETION_MODEL } from 'lib/ai/model.utils' -import { getOrgAIDetails } from 'lib/ai/org-ai-details' import { sanitizeMessagePart } from 'lib/ai/tools/tool-sanitizer' import apiWrapper from 'lib/api/apiWrapper' import { NextApiRequest, NextApiResponse } from 'next' @@ -54,7 +54,10 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) { const { rating, messages: rawMessages, projectRef, orgSlug, reason, spanId } = data let aiOptInLevel: AiOptInLevel = 'disabled' - let isHipaaEnabled = false + let orgHasHipaaAddon: boolean | undefined + let projectIsSensitive: boolean | undefined + let orgIsDpaSigned: boolean | undefined + let projectRegion: string | undefined if (!IS_PLATFORM) { aiOptInLevel = 'schema' @@ -62,16 +65,16 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) { if (IS_PLATFORM && orgSlug && authorization && projectRef) { try { - // Get organizations and compute opt in level server-side - const { aiOptInLevel: orgAIOptInLevel, isHipaaEnabled: orgIsHipaaEnabled } = - await getOrgAIDetails({ - orgSlug, - authorization, - projectRef, - }) + const [orgDetails, projectDetails] = await Promise.all([ + getOrgAIDetails({ orgSlug, authorization }), + getProjectAIDetails({ projectRef, authorization }), + ]) - aiOptInLevel = orgAIOptInLevel - isHipaaEnabled = orgIsHipaaEnabled + aiOptInLevel = orgDetails.aiOptInLevel + orgHasHipaaAddon = orgDetails.hasHipaaAddon + orgIsDpaSigned = orgDetails.isDpaSigned + projectIsSensitive = projectDetails.isSensitive + projectRegion = projectDetails.region } catch (error) { return res.status(400).json({ error: 'There was an error fetching your organization details', @@ -132,7 +135,11 @@ Instructions: }) // Log feedback to Braintrust if tracing is enabled and span ID is available - if (IS_TRACING_ENABLED && !isHipaaEnabled && spanId) { + if ( + IS_TRACING_ENABLED && + isTracingAllowed({ orgHasHipaaAddon, projectIsSensitive, orgIsDpaSigned, projectRegion }) && + spanId + ) { try { const logger = currentLogger() logger?.logFeedback({ diff --git a/apps/studio/pages/api/ai/sql/generate-v4.ts b/apps/studio/pages/api/ai/sql/generate-v4.ts index 48f85f42e0a..01ca53503cc 100644 --- a/apps/studio/pages/api/ai/sql/generate-v4.ts +++ b/apps/studio/pages/api/ai/sql/generate-v4.ts @@ -4,6 +4,8 @@ import { safeValidateUIMessages } from 'ai' import { IS_PLATFORM } from 'common' import { executeSql } from 'data/sql/execute-sql-query' import type { AiOptInLevel } from 'hooks/misc/useOrgOptedIntoAi' +import { getOrgAIDetails, getProjectAIDetails } from 'lib/ai/ai-details' +import { isTracingAllowed } from 'lib/ai/braintrust-logger' import { generateAssistantResponse } from 'lib/ai/generate-assistant-response' import { getModel } from 'lib/ai/model' import { @@ -14,7 +16,6 @@ import { isKnownAssistantModelId, type AssistantModelId, } from 'lib/ai/model.utils' -import { getOrgAIDetails } from 'lib/ai/org-ai-details' import { getTools } from 'lib/ai/tools' import apiWrapper from 'lib/api/apiWrapper' import { executeQuery } from 'lib/api/self-hosted/query' @@ -107,7 +108,10 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw let aiOptInLevel: AiOptInLevel = 'disabled' let hasAccessToAdvanceModel = false - let isHipaaEnabled = false + let orgHasHipaaAddon: boolean | undefined + let projectIsSensitive: boolean | undefined + let orgIsDpaSigned: boolean | undefined + let projectRegion: string | undefined let orgId: number | undefined let planId: string | undefined @@ -118,24 +122,19 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw if (IS_PLATFORM && orgSlug && authorization && projectRef) { try { - // Get organizations and compute opt in level server-side - const { - aiOptInLevel: orgAIOptInLevel, - hasAccessToAdvanceModel: orgHasAccessToAdvanceModel, - isHipaaEnabled: orgIsHipaaEnabled, - orgId: fetchedOrgId, - planId: fetchedPlanId, - } = await getOrgAIDetails({ - orgSlug, - authorization, - projectRef, - }) + const [orgDetails, projectDetails] = await Promise.all([ + getOrgAIDetails({ orgSlug, authorization }), + getProjectAIDetails({ projectRef, authorization }), + ]) - aiOptInLevel = orgAIOptInLevel - hasAccessToAdvanceModel = orgHasAccessToAdvanceModel - isHipaaEnabled = orgIsHipaaEnabled - orgId = fetchedOrgId - planId = fetchedPlanId + aiOptInLevel = orgDetails.aiOptInLevel + hasAccessToAdvanceModel = orgDetails.hasAccessToAdvanceModel + orgHasHipaaAddon = orgDetails.hasHipaaAddon + orgIsDpaSigned = orgDetails.isDpaSigned + orgId = orgDetails.orgId + planId = orgDetails.planId + projectIsSensitive = projectDetails.isSensitive + projectRegion = projectDetails.region } catch (error) { return res.status(400).json({ error: 'There was an error fetching your organization details', @@ -210,7 +209,12 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw projectRef, chatId, chatName, - isHipaaEnabled, + allowTracing: isTracingAllowed({ + orgHasHipaaAddon, + projectIsSensitive, + orgIsDpaSigned, + projectRegion, + }), userId, orgId, planId, diff --git a/apps/studio/pages/api/ai/sql/policy.ts b/apps/studio/pages/api/ai/sql/policy.ts index 379674bbe41..b99e8776d3d 100644 --- a/apps/studio/pages/api/ai/sql/policy.ts +++ b/apps/studio/pages/api/ai/sql/policy.ts @@ -2,9 +2,9 @@ import { generateText, Output, stepCountIs } from 'ai' import { IS_PLATFORM } from 'common' import { source } from 'common-tags' import type { AiOptInLevel } from 'hooks/misc/useOrgOptedIntoAi' +import { getOrgAIDetails } from 'lib/ai/ai-details' import { getModel } from 'lib/ai/model' import { DEFAULT_COMPLETION_MODEL } from 'lib/ai/model.utils' -import { getOrgAIDetails } from 'lib/ai/org-ai-details' import { RLS_PROMPT } from 'lib/ai/prompts' import { getTools } from 'lib/ai/tools' import apiWrapper from 'lib/api/apiWrapper' @@ -79,7 +79,6 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) { const { aiOptInLevel: orgAIOptInLevel } = await getOrgAIDetails({ orgSlug, authorization, - projectRef, }) aiOptInLevel = orgAIOptInLevel diff --git a/packages/api-types/types/api.d.ts b/packages/api-types/types/api.d.ts index a358c04b180..a8221f6af38 100644 --- a/packages/api-types/types/api.d.ts +++ b/packages/api-types/types/api.d.ts @@ -2359,7 +2359,6 @@ export interface components { security_refresh_token_reuse_interval: number | null security_sb_forwarded_for_enabled: boolean | null security_update_password_require_reauthentication: boolean | null - security_update_password_require_current_password: boolean | null sessions_inactivity_timeout: number | null sessions_single_per_user: boolean | null sessions_tags: string | null @@ -4277,7 +4276,6 @@ export interface components { security_refresh_token_reuse_interval?: number | null security_sb_forwarded_for_enabled?: boolean | null security_update_password_require_reauthentication?: boolean | null - security_update_password_require_current_password?: boolean | null sessions_inactivity_timeout?: number | null sessions_single_per_user?: boolean | null sessions_tags?: string | null diff --git a/packages/api-types/types/platform.d.ts b/packages/api-types/types/platform.d.ts index 22dfa8ca4d6..11657238b19 100644 --- a/packages/api-types/types/platform.d.ts +++ b/packages/api-types/types/platform.d.ts @@ -1854,6 +1854,23 @@ export interface paths { patch?: never trace?: never } + '/platform/organizations/preview-creation': { + parameters: { + query?: never + header?: never + path?: never + cookie?: never + } + get?: never + put?: never + /** Preview tax breakdown for organization creation */ + post: operations['OrganizationsController_previewOrganizationCreation'] + delete?: never + options?: never + head?: never + patch?: never + trace?: never + } '/platform/pg-meta/{ref}/column-privileges': { parameters: { query?: never @@ -7019,7 +7036,6 @@ export interface components { SECURITY_SB_FORWARDED_FOR_ENABLED: boolean SECURITY_UPDATE_PASSWORD_REQUIRE_CURRENT_PASSWORD: boolean SECURITY_UPDATE_PASSWORD_REQUIRE_REAUTHENTICATION: boolean - SECURITY_UPDATE_PASSWORD_REQUIRE_CURRENT_PASSWORD: boolean SESSIONS_INACTIVITY_TIMEOUT: number SESSIONS_SINGLE_PER_USER: boolean SESSIONS_TAGS: string @@ -8201,6 +8217,37 @@ export interface components { name: string schema: string } + PreviewOrganizationCreationBody: { + address?: { + city?: string | null + country: string + line1: string + line2?: string | null + postal_code?: string | null + state?: string | null + } + tax_id?: { + country?: string + type: string + value: string + } + /** @enum {string} */ + tier: 'tier_free' | 'tier_pro' | 'tier_payg' | 'tier_team' + } + PreviewOrganizationCreationResponse: { + currency: string + plan_price: number + tax: { + currency: string + tax_amount: number + tax_rate_percentage: number + total_amount_excluding_tax: number + total_amount_including_tax: number + } | null + /** @enum {string} */ + tax_status: 'calculated' | 'not_applicable' | 'failed' + total: number + } PreviewProjectTransferResponse: { errors: { key: string @@ -10325,7 +10372,6 @@ export interface components { SECURITY_SB_FORWARDED_FOR_ENABLED?: boolean | null SECURITY_UPDATE_PASSWORD_REQUIRE_CURRENT_PASSWORD?: boolean | null SECURITY_UPDATE_PASSWORD_REQUIRE_REAUTHENTICATION?: boolean | null - SECURITY_UPDATE_PASSWORD_REQUIRE_CURRENT_PASSWORD?: boolean | null SESSIONS_INACTIVITY_TIMEOUT?: number | null SESSIONS_SINGLE_PER_USER?: boolean | null SESSIONS_TAGS?: string | null @@ -17249,6 +17295,29 @@ export interface operations { } } } + OrganizationsController_previewOrganizationCreation: { + parameters: { + query?: never + header?: never + path?: never + cookie?: never + } + requestBody: { + content: { + 'application/json': components['schemas']['PreviewOrganizationCreationBody'] + } + } + responses: { + 200: { + headers: { + [name: string]: unknown + } + content: { + 'application/json': components['schemas']['PreviewOrganizationCreationResponse'] + } + } + } + } ColumnPrivilegesController_getColumnPrivileges: { parameters: { query?: never