diff --git a/apps/studio/data/content/content-infinite-query.ts b/apps/studio/data/content/content-infinite-query.ts index d6d2e9d84d3..706a897eff7 100644 --- a/apps/studio/data/content/content-infinite-query.ts +++ b/apps/studio/data/content/content-infinite-query.ts @@ -17,7 +17,8 @@ interface GetContentVariables { export async function getContent( { projectRef, type, name, limit = 10, sort, cursor }: GetContentVariables, - signal?: AbortSignal + signal?: AbortSignal, + headers?: HeadersInit ) { if (typeof projectRef === 'undefined') { throw new Error('projectRef is required for getContent') @@ -37,6 +38,7 @@ export async function getContent( cursor, }, }, + headers, signal, }) diff --git a/apps/studio/lib/ai/is-explorer-enabled.test.ts b/apps/studio/lib/ai/is-explorer-enabled.test.ts new file mode 100644 index 00000000000..9a376031d80 --- /dev/null +++ b/apps/studio/lib/ai/is-explorer-enabled.test.ts @@ -0,0 +1,57 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest' + +import { isExplorerEnabled } from './is-explorer-enabled' +import { trustedUserEmail, type getServerFlags as GetServerFlags } from '@/lib/server/configcat' + +type Flags = Awaited> + +const TEST_EMAIL = trustedUserEmail('user@example.com') + +vi.mock('common', () => ({ IS_PLATFORM: true })) +vi.mock('@/lib/server/configcat', async (importOriginal) => { + const actual = await importOriginal() + return { ...actual, getServerFlags: vi.fn() } +}) + +describe('isExplorerEnabled', () => { + beforeEach(async () => { + vi.clearAllMocks() + const common = await import('common') + vi.spyOn(common, 'IS_PLATFORM', 'get').mockReturnValue(true) + }) + + it('returns false when self-hosted, without calling getServerFlags', async () => { + const common = await import('common') + vi.spyOn(common, 'IS_PLATFORM', 'get').mockReturnValue(false) + + const { getServerFlags } = await import('@/lib/server/configcat') + const result = await isExplorerEnabled(TEST_EMAIL) + + expect(result).toBe(false) + expect(getServerFlags).not.toHaveBeenCalled() + }) + + it('returns true when the explorer flag resolves true for this user', async () => { + const { getServerFlags } = await import('@/lib/server/configcat') + vi.mocked(getServerFlags).mockResolvedValue([ + { settingKey: 'explorer', settingValue: true }, + { settingKey: 'other_flag', settingValue: false }, + ] satisfies Flags) + + const result = await isExplorerEnabled(TEST_EMAIL) + + expect(result).toBe(true) + expect(getServerFlags).toHaveBeenCalledWith(TEST_EMAIL) + }) + + it('returns false when the explorer flag is absent or false', async () => { + const { getServerFlags } = await import('@/lib/server/configcat') + vi.mocked(getServerFlags).mockResolvedValue([ + { settingKey: 'explorer', settingValue: false }, + ] satisfies Flags) + + const result = await isExplorerEnabled(TEST_EMAIL) + + expect(result).toBe(false) + }) +}) diff --git a/apps/studio/lib/ai/is-explorer-enabled.ts b/apps/studio/lib/ai/is-explorer-enabled.ts new file mode 100644 index 00000000000..cb1acbab3e6 --- /dev/null +++ b/apps/studio/lib/ai/is-explorer-enabled.ts @@ -0,0 +1,19 @@ +import { IS_PLATFORM } from 'common' + +import { getServerFlags, type TrustedUserEmail } from '@/lib/server/configcat' + +// Notebooks haven't shipped yet — they're gated behind the Explorer feature flag on the +// client (`useFlag('explorer')`). The AI tools that read notebook content (list_notebooks, +// get_notebook in lib/ai/tools/notebook-tools.ts) must not be advertised to the assistant until +// this resolves true for the requesting user, otherwise the assistant will believe notebooks +// exist for users who can't yet see them in the UI. +// +// Takes TrustedUserEmail, not a plain string: callers must go through +// lib/server/configcat's trustedUserEmail() to prove this came from a verified source +// (a decoded, signature-checked JWT claim), not a raw request param. +export async function isExplorerEnabled(userEmail?: TrustedUserEmail): Promise { + if (!IS_PLATFORM) return false + + const flags = await getServerFlags(userEmail) + return flags.some((flag) => flag.settingKey === 'explorer' && flag.settingValue === true) +} diff --git a/apps/studio/lib/ai/tool-filter.ts b/apps/studio/lib/ai/tool-filter.ts index 3969897cf71..e74977ab37e 100644 --- a/apps/studio/lib/ai/tool-filter.ts +++ b/apps/studio/lib/ai/tool-filter.ts @@ -37,6 +37,8 @@ export const toolSetValidationSchema = z.record( 'list_policies', 'list_reports', 'get_report', + 'list_notebooks', + 'get_notebook', // Fallback tools for self-hosted 'getSchemaTables', @@ -88,6 +90,8 @@ export const TOOL_CATEGORY_MAP: Record = { list_policies: TOOL_CATEGORIES.SCHEMA, list_reports: TOOL_CATEGORIES.SCHEMA, get_report: TOOL_CATEGORIES.SCHEMA, + list_notebooks: TOOL_CATEGORIES.SCHEMA, + get_notebook: TOOL_CATEGORIES.SCHEMA, getSchemaTables: TOOL_CATEGORIES.SCHEMA, getRlsKnowledge: TOOL_CATEGORIES.SCHEMA, getFunctions: TOOL_CATEGORIES.SCHEMA, diff --git a/apps/studio/lib/ai/tools/index.test.ts b/apps/studio/lib/ai/tools/index.test.ts index fb927e1bb29..983c21d73bf 100644 --- a/apps/studio/lib/ai/tools/index.test.ts +++ b/apps/studio/lib/ai/tools/index.test.ts @@ -83,4 +83,18 @@ describe('ai/tools getTools', () => { expect(tools).toHaveProperty('fallback_tool') expect(tools).not.toHaveProperty('list_tables') }) + + it('excludes notebook tools when isExplorerEnabled is not set', async () => { + const tools = await getTools(BASE_PARAMS) + + expect(tools).not.toHaveProperty('list_notebooks') + expect(tools).not.toHaveProperty('get_notebook') + }) + + it('includes notebook tools only when isExplorerEnabled is true', async () => { + const tools = await getTools({ ...BASE_PARAMS, isExplorerEnabled: true }) + + expect(tools).toHaveProperty('list_notebooks') + expect(tools).toHaveProperty('get_notebook') + }) }) diff --git a/apps/studio/lib/ai/tools/index.ts b/apps/studio/lib/ai/tools/index.ts index 8dad88afdb6..6b487d0c8a7 100644 --- a/apps/studio/lib/ai/tools/index.ts +++ b/apps/studio/lib/ai/tools/index.ts @@ -5,6 +5,7 @@ import { filterToolsByOptInLevel } from '../tool-filter' import { getFallbackTools } from './fallback-tools' import { getIncidentTools } from './incident-tools' import { getMcpTools } from './mcp-tools' +import { getNotebookTools } from './notebook-tools' import { getReportTools } from './report-tools' import { getSchemaTools } from './schema-tools' import { getStudioTools } from './studio-tools' @@ -19,6 +20,7 @@ export const getTools = async ({ accessToken, baseUrl, supportMode, + isExplorerEnabled, signal, }: { projectRef: string @@ -28,6 +30,10 @@ export const getTools = async ({ accessToken?: string baseUrl?: string supportMode?: boolean + // Notebooks haven't shipped yet — they live behind the Explorer feature flag, so the + // assistant must not advertise list_notebooks/get_notebook until the caller confirms the + // flag is on for this user. + isExplorerEnabled?: boolean // Required: tools fetched from the remote MCP server hold an HTTP connection // that is closed when this signal aborts (i.e. when the request ends). signal: AbortSignal @@ -72,6 +78,7 @@ export const getTools = async ({ connectionString, }), ...getReportTools({ projectRef, authorization }), + ...(isExplorerEnabled ? getNotebookTools({ projectRef, authorization }) : {}), ...(baseUrl ? getIncidentTools({ baseUrl }) : {}), } } diff --git a/apps/studio/lib/ai/tools/notebook-tools.test.ts b/apps/studio/lib/ai/tools/notebook-tools.test.ts new file mode 100644 index 00000000000..eadd9f10305 --- /dev/null +++ b/apps/studio/lib/ai/tools/notebook-tools.test.ts @@ -0,0 +1,210 @@ +import { components } from 'api-types' +import { HttpResponse } from 'msw' +import { describe, expect, it } from 'vitest' + +import { getNotebookTools } from './notebook-tools' +import { addAPIMock, type APIErrorBody } from '@/tests/lib/msw' + +type GetUserContentByIdResponse = components['schemas']['GetUserContentByIdResponse'] +type GetUserContentResponse = components['schemas']['GetUserContentResponse'] + +const NOTEBOOK_CONTENT = { + schema_version: 1, + cells: [ + { _tag: 'markdown_cell', id: 'cell-1', text: '# Signup funnel' }, + { + _tag: 'database_cell', + id: 'cell-2', + sql: 'select * from auth.users limit 100', + row_limit: 100, + }, + { + _tag: 'log_cell', + id: 'cell-3', + sql: 'select timestamp, event_message from edge_logs limit 10', + time_range: { _tag: 'relative_time_range', unit: 'hour', amount: 1 }, + }, + ], +} + +describe('ai/tools/notebook-tools', () => { + describe('getNotebookTools', () => { + it('should return list_notebooks and get_notebook tools', () => { + const tools = getNotebookTools() + + expect(Object.keys(tools)).toEqual(['list_notebooks', 'get_notebook']) + }) + + it('should not require approval to read notebooks', () => { + const tools = getNotebookTools() + + expect(tools.list_notebooks.needsApproval).toBeUndefined() + expect(tools.get_notebook.needsApproval).toBeUndefined() + }) + }) + + describe('list_notebooks', () => { + it('should list notebooks with summary fields, forwarding the authorization header and cursor', async () => { + let capturedRequest: Request | undefined + + addAPIMock({ + method: 'get', + path: '/platform/projects/:ref/content', + response: ({ request }) => { + capturedRequest = request + return HttpResponse.json({ + cursor: 'next-page', + data: [ + { + id: 'notebook-1', + name: 'Signup funnel', + description: undefined, + visibility: 'project', + favorite: false, + folder_id: null, + inserted_at: '2026-01-01T00:00:00.000Z', + updated_at: '2026-01-01T00:00:00.000Z', + owner_id: 1, + owner: { id: 1, username: 'test' }, + updated_by: { id: 1, username: 'test' }, + project_id: 1, + type: 'notebook', + content: NOTEBOOK_CONTENT, + }, + ], + } as unknown as GetUserContentResponse) + }, + }) + + const tools = getNotebookTools({ + projectRef: 'test-project', + authorization: 'Bearer token', + }) + if (!tools.list_notebooks.execute) throw new Error('execute is undefined') + + const result = await tools.list_notebooks.execute( + { limit: 20, cursor: 'prev-page' }, + { toolCallId: 'test', messages: [] } + ) + + expect(capturedRequest?.headers.get('authorization')).toBe('Bearer token') + const url = new URL(capturedRequest!.url) + expect(url.pathname).toContain('/projects/test-project/content') + expect(url.searchParams.get('type')).toBe('notebook') + expect(url.searchParams.get('limit')).toBe('20') + expect(url.searchParams.get('cursor')).toBe('prev-page') + + expect(result).toEqual({ + notebooks: [ + { + id: 'notebook-1', + name: 'Signup funnel', + description: undefined, + visibility: 'project', + updated_at: '2026-01-01T00:00:00.000Z', + cell_count: 3, + }, + ], + cursor: 'next-page', + }) + }) + }) + + describe('get_notebook', () => { + it('should resolve markdown text and SQL for every cell', async () => { + addAPIMock({ + method: 'get', + path: '/platform/projects/:ref/content/item/:id', + response: () => + HttpResponse.json({ + id: 'notebook-1', + name: 'Signup funnel', + description: undefined, + visibility: 'project', + favorite: false, + folder_id: null, + inserted_at: '2026-01-01T00:00:00.000Z', + updated_at: '2026-01-01T00:00:00.000Z', + owner_id: 1, + project_id: 1, + type: 'notebook', + content: NOTEBOOK_CONTENT, + } as unknown as GetUserContentByIdResponse), + }) + + const tools = getNotebookTools({ projectRef: 'test-project' }) + if (!tools.get_notebook.execute) throw new Error('execute is undefined') + + const result = await tools.get_notebook.execute( + { id: 'notebook-1' }, + { toolCallId: 'test', messages: [] } + ) + + expect(result).toEqual({ + id: 'notebook-1', + name: 'Signup funnel', + description: undefined, + visibility: 'project', + cells: [ + { _tag: 'markdown_cell', id: 'cell-1', text: '# Signup funnel' }, + { + _tag: 'database_cell', + id: 'cell-2', + row_limit: 100, + sql: 'select * from auth.users limit 100', + }, + { + _tag: 'log_cell', + id: 'cell-3', + time_range: { _tag: 'relative_time_range', unit: 'hour', amount: 1 }, + sql: 'select timestamp, event_message from edge_logs limit 10', + }, + ], + }) + }) + + it('should throw when the content id is not a notebook', async () => { + addAPIMock({ + method: 'get', + path: '/platform/projects/:ref/content/item/:id', + response: () => + HttpResponse.json({ + id: 'snippet-1', + name: 'My query', + description: undefined, + visibility: 'user', + favorite: false, + folder_id: null, + inserted_at: '2026-01-01T00:00:00.000Z', + updated_at: '2026-01-01T00:00:00.000Z', + owner_id: 1, + project_id: 1, + type: 'sql', + content: { content_id: 'snippet-1', sql: 'select 1', schema_version: '1' }, + }), + }) + + const tools = getNotebookTools({ projectRef: 'test-project' }) + if (!tools.get_notebook.execute) throw new Error('execute is undefined') + + await expect( + tools.get_notebook.execute({ id: 'snippet-1' }, { toolCallId: 'test', messages: [] }) + ).rejects.toThrow('is not a notebook') + }) + + it('should throw when the content id does not exist', async () => { + addAPIMock({ + method: 'get', + path: '/platform/projects/:ref/content/item/:id', + response: () => HttpResponse.json({ message: 'Not found' }, { status: 404 }), + }) + + const tools = getNotebookTools({ projectRef: 'test-project' }) + if (!tools.get_notebook.execute) throw new Error('execute is undefined') + + await expect( + tools.get_notebook.execute({ id: 'missing' }, { toolCallId: 'test', messages: [] }) + ).rejects.toThrow() + }) + }) +}) diff --git a/apps/studio/lib/ai/tools/notebook-tools.ts b/apps/studio/lib/ai/tools/notebook-tools.ts new file mode 100644 index 00000000000..1e5b53c62c7 --- /dev/null +++ b/apps/studio/lib/ai/tools/notebook-tools.ts @@ -0,0 +1,85 @@ +import { tool } from 'ai' +import { z } from 'zod' + +import { getContent } from '@/data/content/content-infinite-query' +import { getNotebook } from '@/data/content/notebooks/notebook-query' +import type { Notebooks } from '@/types' + +export type NotebookToolsContext = { + projectRef?: string + authorization?: string +} + +export const getNotebookTools = (ctx: NotebookToolsContext = {}) => { + const { projectRef, authorization } = ctx + const authHeaders = authorization ? { Authorization: authorization } : undefined + + return { + list_notebooks: tool({ + description: 'List the notebooks saved for this project', + inputSchema: z.object({ + cursor: z + .string() + .optional() + .describe('Cursor from a previous call, used to fetch the next page.'), + limit: z + .number() + .int() + .positive() + .max(100) + .default(20) + .describe('Max number of notebooks to return.'), + }), + execute: async ({ cursor, limit }) => { + const { content, cursor: nextCursor } = await getContent( + { projectRef, type: 'notebook', limit, cursor }, + undefined, + authHeaders + ) + + return { + notebooks: content.map((notebook) => ({ + id: notebook.id, + name: notebook.name, + description: notebook.description, + visibility: notebook.visibility, + updated_at: notebook.updated_at, + cell_count: (notebook.content as Notebooks.Content).cells.length, + })), + cursor: nextCursor, + } + }, + }), + get_notebook: tool({ + description: + 'Get a single notebook by id, including the markdown text and resolved SQL of every cell.', + inputSchema: z.object({ + id: z.string().describe('The id of the notebook to fetch.'), + }), + execute: async ({ id }) => { + const notebook = await getNotebook({ projectRef, id }, undefined, authHeaders) + + return { + id: notebook.id, + name: notebook.name, + description: notebook.description, + visibility: notebook.visibility, + cells: notebook.content.cells.map((cell) => { + switch (cell._tag) { + case 'markdown_cell': + return cell + case 'database_cell': { + const { unchecked_sql, ...rest } = cell + return { ...rest, sql: unchecked_sql } + } + case 'log_cell': { + const { unchecked_sql, ...rest } = cell + return { ...rest, sql: unchecked_sql } + } + } + }), + } + }, + }), + } +} diff --git a/apps/studio/lib/server/configcat.test.ts b/apps/studio/lib/server/configcat.test.ts new file mode 100644 index 00000000000..417c6b688c3 --- /dev/null +++ b/apps/studio/lib/server/configcat.test.ts @@ -0,0 +1,127 @@ +import * as configcat from '@configcat/sdk/node' +import type { IConfigCatClient } from '@configcat/sdk/node' +import { beforeEach, describe, expect, it, vi } from 'vitest' + +vi.mock('@configcat/sdk/node', () => ({ + getClient: vi.fn(), + PollingMode: { + LazyLoad: 'LazyLoad', + }, + User: vi.fn(), +})) + +describe('lib/server/configcat getServerFlags', () => { + const mockClient = { + getAllValuesAsync: vi.fn(), + } + + beforeEach(() => { + vi.clearAllMocks() + vi.resetModules() + vi.unstubAllEnvs() + vi.mocked(configcat.getClient).mockReturnValue(mockClient as unknown as IConfigCatClient) + }) + + it('should return empty array and skip getClient when no env vars are present', async () => { + const { getServerFlags, trustedUserEmail } = await import('./configcat') + const result = await getServerFlags(trustedUserEmail('test@example.com')) + + expect(result).toEqual([]) + expect(configcat.getClient).not.toHaveBeenCalled() + }) + + it('should prefer the proxy over the direct SDK key when both are configured', async () => { + vi.stubEnv('NEXT_PUBLIC_CONFIGCAT_SDK_KEY', 'test-sdk-key') + vi.stubEnv('NEXT_PUBLIC_CONFIGCAT_PROXY_URL', 'https://proxy.example.com') + mockClient.getAllValuesAsync.mockResolvedValue([]) + + const { getServerFlags, trustedUserEmail } = await import('./configcat') + await getServerFlags(trustedUserEmail('test@example.com')) + + expect(configcat.getClient).toHaveBeenCalledTimes(1) + expect(configcat.getClient).toHaveBeenCalledWith('configcat-proxy/frontend-v2', 'LazyLoad', { + baseUrl: 'https://proxy.example.com', + }) + }) + + it('should fall back to the direct SDK key when no proxy URL is configured', async () => { + vi.stubEnv('NEXT_PUBLIC_CONFIGCAT_SDK_KEY', 'test-sdk-key') + mockClient.getAllValuesAsync.mockResolvedValue([]) + + const { getServerFlags, trustedUserEmail } = await import('./configcat') + await getServerFlags(trustedUserEmail('test@example.com')) + + expect(configcat.getClient).toHaveBeenCalledWith('test-sdk-key', 'LazyLoad') + }) + + it('should call getAllValuesAsync with a user built from the given email', async () => { + vi.stubEnv('NEXT_PUBLIC_CONFIGCAT_SDK_KEY', 'test-sdk-key') + const mockValues = [{ settingKey: 'explorer', settingValue: true }] + mockClient.getAllValuesAsync.mockResolvedValue(mockValues) + + const { getServerFlags, trustedUserEmail } = await import('./configcat') + const result = await getServerFlags(trustedUserEmail('test@example.com')) + + expect(configcat.User).toHaveBeenCalledWith( + 'test@example.com', + undefined, + undefined, + expect.any(Object) + ) + expect(result).toEqual(mockValues) + }) + + it('reuses the same client across calls instead of creating a new one each time', async () => { + vi.stubEnv('NEXT_PUBLIC_CONFIGCAT_SDK_KEY', 'test-sdk-key') + mockClient.getAllValuesAsync.mockResolvedValue([]) + + const { getServerFlags, trustedUserEmail } = await import('./configcat') + await getServerFlags(trustedUserEmail('a@example.com')) + await getServerFlags(trustedUserEmail('b@example.com')) + + expect(configcat.getClient).toHaveBeenCalledTimes(1) + }) + + describe('is_staff targeting attribute', () => { + beforeEach(() => { + vi.stubEnv('NEXT_PUBLIC_CONFIGCAT_SDK_KEY', 'test-sdk-key') + mockClient.getAllValuesAsync.mockResolvedValue([]) + }) + + it('is true for a real @supabase.com/@supabase.io email', async () => { + const { getServerFlags, trustedUserEmail } = await import('./configcat') + await getServerFlags(trustedUserEmail('person@supabase.io')) + + expect(configcat.User).toHaveBeenCalledWith( + 'person@supabase.io', + undefined, + undefined, + expect.objectContaining({ is_staff: 'true' }) + ) + }) + + it('is false for a domain that merely contains "@supabase." as a substring', async () => { + const { getServerFlags, trustedUserEmail } = await import('./configcat') + await getServerFlags(trustedUserEmail('attacker@supabase.evil.com')) + + expect(configcat.User).toHaveBeenCalledWith( + 'attacker@supabase.evil.com', + undefined, + undefined, + expect.objectContaining({ is_staff: 'false' }) + ) + }) + + it('is false when no email is given', async () => { + const { getServerFlags } = await import('./configcat') + await getServerFlags(undefined) + + expect(configcat.User).toHaveBeenCalledWith( + 'anonymous', + undefined, + undefined, + expect.objectContaining({ is_staff: 'false' }) + ) + }) + }) +}) diff --git a/apps/studio/lib/server/configcat.ts b/apps/studio/lib/server/configcat.ts new file mode 100644 index 00000000000..9de056afd42 --- /dev/null +++ b/apps/studio/lib/server/configcat.ts @@ -0,0 +1,73 @@ +import { getClient, PollingMode, User } from '@configcat/sdk/node' + +let serverClient: ReturnType + +export type TrustedUserEmail = string & { readonly __trustedUserEmailBrand: never } + +// Promotes an email to TrustedUserEmail. Only call this with an email whose origin is already +// verified — e.g. `claims.email` from a JWT that apiAuthenticate has confirmed — never with a +// raw request body/query param. +export function trustedUserEmail(email: string | undefined): TrustedUserEmail | undefined { + return email as TrustedUserEmail | undefined +} + +const STAFF_EMAIL_DOMAINS = ['supabase.com', 'supabase.io'] + +function isStaffEmail(email: string): boolean { + const domain = email.slice(email.lastIndexOf('@') + 1).toLowerCase() + return STAFF_EMAIL_DOMAINS.includes(domain) +} + +function buildUser(userEmail?: TrustedUserEmail, customAttributes?: Record) { + const _customAttributes = { + ...customAttributes, + is_staff: (!!userEmail && isStaffEmail(userEmail)).toString(), + } + + return new User(userEmail ?? 'anonymous', undefined, undefined, _customAttributes) +} + +function getServerClient() { + if (serverClient) return serverClient + + const proxyUrl = process.env.NEXT_PUBLIC_CONFIGCAT_PROXY_URL + const sdkKey = process.env.NEXT_PUBLIC_CONFIGCAT_SDK_KEY + + if (!proxyUrl && !sdkKey) { + console.log('Skipping server ConfigCat set up as env vars are not present') + return undefined + } + + try { + if (proxyUrl) { + serverClient = getClient('configcat-proxy/frontend-v2', PollingMode.LazyLoad, { + baseUrl: proxyUrl, + }) + return serverClient + } + + if (sdkKey) { + serverClient = getClient(sdkKey, PollingMode.LazyLoad) + return serverClient + } + + return undefined + } catch (error) { + const message = error instanceof Error ? error.message : String(error) + console.error(`Failed to get server ConfigCat client: ${message}`) + return undefined + } +} + +export async function getServerFlags( + userEmail?: TrustedUserEmail, + customAttributes?: Record +) { + const client = getServerClient() + + if (!client) { + return [] + } + + return client.getAllValuesAsync(buildUser(userEmail, customAttributes)) +} diff --git a/apps/studio/package.json b/apps/studio/package.json index e3dff9d1c64..2e747959ece 100644 --- a/apps/studio/package.json +++ b/apps/studio/package.json @@ -46,6 +46,7 @@ "@ai-sdk/provider-utils": "^4.0.19", "@ai-sdk/react": "^3.0.118", "@aws-sdk/credential-providers": "^3.1041.0", + "@configcat/sdk": "^1.1.0", "@dagrejs/dagre": "^1.0.4", "@dnd-kit/core": "^6.1.0", "@dnd-kit/modifiers": "^9.0.0", diff --git a/apps/studio/pages/api/ai/sql/generate-v4.ts b/apps/studio/pages/api/ai/sql/generate-v4.ts index 407131986eb..ed76f31dd4f 100644 --- a/apps/studio/pages/api/ai/sql/generate-v4.ts +++ b/apps/studio/pages/api/ai/sql/generate-v4.ts @@ -15,6 +15,7 @@ import { } from '@/lib/ai/assistant-message-metadata' import { isTracingAllowed } from '@/lib/ai/braintrust-logger' import { generateAssistantResponse } from '@/lib/ai/generate-assistant-response' +import { isExplorerEnabled } from '@/lib/ai/is-explorer-enabled' import { getModel } from '@/lib/ai/model' import { DEFAULT_ASSISTANT_ADVANCE_MODEL_ID, @@ -28,6 +29,7 @@ import { getTools } from '@/lib/ai/tools' import { apiWrapper } from '@/lib/api/apiWrapper' import { executeQuery } from '@/lib/api/self-hosted/query' import { getURL } from '@/lib/helpers' +import { trustedUserEmail } from '@/lib/server/configcat' export const maxDuration = 120 @@ -148,6 +150,8 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw } } + const explorerEnabled = await isExplorerEnabled(trustedUserEmail(claims?.email)) + const envThrottled = process.env.IS_THROTTLED !== 'false' let effectiveModel: AssistantModelId = requestedModel ?? DEFAULT_ASSISTANT_ADVANCE_MODEL_ID @@ -184,6 +188,7 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw accessToken, baseUrl: getURL(), supportMode, + isExplorerEnabled: explorerEnabled, signal: abortController.signal, }) diff --git a/apps/studio/pages/api/ai/sql/policy.ts b/apps/studio/pages/api/ai/sql/policy.ts index f394b534b3b..998a0a2b7a4 100644 --- a/apps/studio/pages/api/ai/sql/policy.ts +++ b/apps/studio/pages/api/ai/sql/policy.ts @@ -1,3 +1,4 @@ +import type { JwtPayload } from '@supabase/supabase-js' import { generateText, Output, stepCountIs } from 'ai' import { IS_PLATFORM } from 'common' import { source } from 'common-tags' @@ -6,11 +7,13 @@ import { z } from 'zod' import type { AiOptInLevel } from '@/hooks/misc/useOrgOptedIntoAi' import { getAIDetails } from '@/lib/ai/ai-details' +import { isExplorerEnabled } from '@/lib/ai/is-explorer-enabled' import { getModel } from '@/lib/ai/model' import { DEFAULT_COMPLETION_MODEL } from '@/lib/ai/model.utils' import { RLS_PROMPT } from '@/lib/ai/prompts' import { getTools } from '@/lib/ai/tools' import { apiWrapper } from '@/lib/api/apiWrapper' +import { trustedUserEmail } from '@/lib/server/configcat' const policySchema = z.object({ sql: z.string().describe('The generated Postgres CREATE POLICY statement.'), @@ -40,19 +43,19 @@ const requestBodySchema = z.object({ message: z.string().optional(), }) -async function handler(req: NextApiRequest, res: NextApiResponse) { +async function handler(req: NextApiRequest, res: NextApiResponse, claims?: JwtPayload) { const { method } = req switch (method) { case 'POST': - return handlePost(req, res) + return handlePost(req, res, claims) default: res.setHeader('Allow', ['POST']) res.status(405).json({ data: null, error: { message: `Method ${method} Not Allowed` } }) } } -export async function handlePost(req: NextApiRequest, res: NextApiResponse) { +export async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: JwtPayload) { const authorization = req.headers.authorization const accessToken = authorization?.replace('Bearer ', '') @@ -87,6 +90,8 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) { } } + const explorerEnabled = await isExplorerEnabled(trustedUserEmail(claims?.email)) + try { const { modelParams, error: modelError } = await getModel({ provider: 'openai', @@ -113,6 +118,7 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) { authorization, aiOptInLevel, accessToken, + isExplorerEnabled: explorerEnabled, signal: toolsAbortController.signal, }) diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index 4f6a7828163..368781bcaea 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -906,6 +906,9 @@ importers: '@aws-sdk/credential-providers': specifier: ^3.1041.0 version: 3.1041.0 + '@configcat/sdk': + specifier: ^1.1.0 + version: 1.1.0 '@dagrejs/dagre': specifier: ^1.0.4 version: 1.0.4 @@ -3462,6 +3465,9 @@ packages: resolution: {integrity: sha512-ooWCrlZP11i8GImSjTHYHLkvFDP48nS4+204nGb1RiX/WXYHmJA2III9/e2DWVabCESdW7hBAEzHRqUn9OUVvQ==} engines: {node: '>=0.1.90'} + '@configcat/sdk@1.1.0': + resolution: {integrity: sha512-oSueg7jIk6in+XreRjnV+U6lQfgu5W2w6TqtO1CjEcyR2zRrVC7fLLZoySHnRlPwExXrvbp4olDYWEbfj1zL0A==} + '@contentlayer2/cli@0.4.3': resolution: {integrity: sha512-ZJ+Iiu2rVI50x60XoqnrsO/Q8eqFX5AlP1L0U/3ygaAas3tnOqTzQZ1UsxYQMpJzcLok24ddlhKfQKbCMUJPiQ==} @@ -18979,6 +18985,8 @@ snapshots: '@colors/colors@1.5.0': optional: true + '@configcat/sdk@1.1.0': {} + '@contentlayer2/cli@0.4.3(esbuild@0.28.1)(markdown-wasm@1.2.0)(supports-color@8.1.1)': dependencies: '@contentlayer2/core': 0.4.3(esbuild@0.28.1)(markdown-wasm@1.2.0)(supports-color@8.1.1)