diff --git a/apps/studio/data/config/project-settings-v2-query.ts b/apps/studio/data/config/project-settings-v2-query.ts index c957b930eaa..abec2e80236 100644 --- a/apps/studio/data/config/project-settings-v2-query.ts +++ b/apps/studio/data/config/project-settings-v2-query.ts @@ -19,13 +19,15 @@ export type ProjectSettings = components['schemas']['ProjectSettingsResponse'] & export async function getProjectSettings( { projectRef }: ProjectSettingsVariables, - signal?: AbortSignal + signal?: AbortSignal, + headers?: Record ) { if (!projectRef) throw new Error('projectRef is required') const { data, error } = await get('/platform/projects/{ref}/settings', { params: { path: { ref: projectRef } }, signal, + headers, }) if (error) handleError(error) diff --git a/apps/studio/data/subscriptions/org-subscription-query.ts b/apps/studio/data/subscriptions/org-subscription-query.ts index 4cc1d9292b5..82ba5b9979d 100644 --- a/apps/studio/data/subscriptions/org-subscription-query.ts +++ b/apps/studio/data/subscriptions/org-subscription-query.ts @@ -13,13 +13,15 @@ export type OrgSubscriptionVariables = { export async function getOrgSubscription( { orgSlug }: OrgSubscriptionVariables, - signal?: AbortSignal + signal?: AbortSignal, + headers?: Record ) { if (!orgSlug) throw new Error('orgSlug is required') const { error, data } = await get('/platform/organizations/{slug}/billing/subscription', { params: { path: { slug: orgSlug } }, signal, + headers, }) if (error) handleError(error) diff --git a/apps/studio/lib/ai/generate-assistant-response.ts b/apps/studio/lib/ai/generate-assistant-response.ts index 66ef0922aec..d0bc75ba733 100644 --- a/apps/studio/lib/ai/generate-assistant-response.ts +++ b/apps/studio/lib/ai/generate-assistant-response.ts @@ -45,6 +45,8 @@ export async function generateAssistantResponse({ getSchemas?: () => Promise projectRef?: string chatName?: string + // TODO(mattrossman): use for excluding HIPAA projects from assistant tracing + isHipaaEnabled?: boolean promptProviderOptions?: Record providerOptions?: Record abortSignal?: AbortSignal diff --git a/apps/studio/lib/ai/org-ai-details.test.ts b/apps/studio/lib/ai/org-ai-details.test.ts index 2a913fe95c0..e0926a2b4d4 100644 --- a/apps/studio/lib/ai/org-ai-details.test.ts +++ b/apps/studio/lib/ai/org-ai-details.test.ts @@ -10,25 +10,53 @@ 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(), +})) + describe('ai/org-ai-details', () => { let mockGetOrganizations: ReturnType let mockGetProjectDetail: ReturnType + let mockGetOrgSubscription: ReturnType + let mockGetProjectSettings: ReturnType let mockGetAiOptInLevel: ReturnType + let mockSubscriptionHasHipaaAddon: 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' + ) 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) + + // Default mocks for subscription/settings (no HIPAA) + mockGetOrgSubscription.mockResolvedValue({ addons: [] }) + mockGetProjectSettings.mockResolvedValue({ is_sensitive: false }) + mockSubscriptionHasHipaaAddon.mockReturnValue(false) }) describe('getOrgAIDetails', () => { @@ -92,6 +120,7 @@ describe('ai/org-ai-details', () => { expect(result).toEqual({ aiOptInLevel: 'schema_only', isLimited: true, + isHipaaEnabled: false, }) }) @@ -239,5 +268,112 @@ describe('ai/org-ai-details', () => { expect(result.isLimited).toBe(false) // Pro plan }) + + 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 index 7f9cf99576e..6457d390914 100644 --- a/apps/studio/lib/ai/org-ai-details.ts +++ b/apps/studio/lib/ai/org-ai-details.ts @@ -1,5 +1,8 @@ +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' export const getOrgAIDetails = async ({ @@ -16,9 +19,11 @@ export const getOrgAIDetails = async ({ ...(authorization && { Authorization: authorization }), } - const [organizations, selectedProject] = await Promise.all([ + const [organizations, selectedProject, subscription, projectSettings] = await Promise.all([ getOrganizations({ headers }), getProjectDetail({ ref: projectRef }, undefined, headers), + getOrgSubscription({ orgSlug }, undefined, headers), + getProjectSettings({ projectRef }, undefined, headers), ]) const selectedOrg = organizations.find((org) => org.slug === orgSlug) @@ -30,9 +35,11 @@ export const getOrgAIDetails = async ({ const aiOptInLevel = getAiOptInLevel(selectedOrg?.opt_in_tags) const isLimited = selectedOrg?.plan.id === 'free' + const isHipaaEnabled = subscriptionHasHipaaAddon(subscription) && !!projectSettings?.is_sensitive return { aiOptInLevel, isLimited, + isHipaaEnabled, } } diff --git a/apps/studio/pages/api/ai/sql/generate-v4.ts b/apps/studio/pages/api/ai/sql/generate-v4.ts index 92de85ef6d6..354a30f4b93 100644 --- a/apps/studio/pages/api/ai/sql/generate-v4.ts +++ b/apps/studio/pages/api/ai/sql/generate-v4.ts @@ -89,6 +89,7 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse) { let aiOptInLevel: AiOptInLevel = 'disabled' let isLimited = false + let isHipaaEnabled = false if (!IS_PLATFORM) { aiOptInLevel = 'schema' @@ -97,7 +98,11 @@ 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, isLimited: orgAILimited } = await getOrgAIDetails({ + const { + aiOptInLevel: orgAIOptInLevel, + isLimited: orgAILimited, + isHipaaEnabled: orgIsHipaaEnabled, + } = await getOrgAIDetails({ orgSlug, authorization, projectRef, @@ -105,6 +110,7 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse) { aiOptInLevel = orgAIOptInLevel isLimited = orgAILimited + isHipaaEnabled = orgIsHipaaEnabled } catch (error) { return res.status(400).json({ error: 'There was an error fetching your organization details', @@ -174,6 +180,7 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse) { getSchemas: aiOptInLevel !== 'disabled' ? getSchemas : undefined, projectRef, chatName, + isHipaaEnabled, promptProviderOptions, providerOptions, abortSignal: abortController.signal,