mirror of
https://github.com/supabase/supabase.git
synced 2026-10-05 01:15:03 +03:00
feat(studio): notebook read tools (#48908)
## Summary - Adds `list_notebooks` (cursor-paginated) and `get_notebook` AI tools in `lib/ai/tools/notebook-tools.ts`, modeled directly on `report-tools.ts`: server-side `getContent`/`getNotebook` with the `authorization` header forwarded, zod-validated input. - `get_notebook` resolves every cell and exposes `unchecked_sql` as a plain `sql` field for the agent to read — display only, per the `safe-sql-execution` skill; nothing here executes SQL. - Registers both tools in `lib/ai/tools/index.ts` (same platform branch as reports) and in `lib/ai/tool-filter.ts`'s `toolSetValidationSchema` + `TOOL_CATEGORY_MAP` (`SCHEMA` tier). - Adds an optional `headers` param to `content-infinite-query.ts`'s `getContent`, mirroring the sibling `content-query.ts`, so the cursor-paginated fetch can carry the `Authorization` header from a server context. - New tools are behind the Explorer feature flag. Stacked on #48907 (1.4 — notebook query and mutation hooks), per the Notebooks implementation plan (stack 2.1). Resolves FE-4081 Resolves FE-4080 ## Test plan - [x] `pnpm exec tsc --noEmit` — no new errors - [x] `pnpm exec vitest run lib/ai/tools/notebook-tools.test.ts lib/ai/tools/index.test.ts lib/ai/tools/report-tools.test.ts data/content/notebooks` — 36/36 passing - [x] `pnpm --filter studio run lint` — no new warnings - [x] `pnpm exec prettier --check` on changed files — clean <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **New Features** * Added AI tools to list project notebooks with pagination. * Added AI support for retrieving notebook markdown and resolved SQL cell content. * Notebook tools now respect project and authorization context. * Notebook features are available only when Explorer access is enabled. * Content requests can forward custom request headers. * **Tests** * Added coverage for notebook tools, Explorer access, feature flags, authorization, pagination, and error handling. <!-- end of auto-generated comment: release notes by coderabbit.ai -->
This commit is contained in:
1 parent
241bb11c06
commit
7798e42435
14 files changed
+622
-4
No files matched your search
@@ -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,
|
||||
})
|
||||
|
||||
|
||||
@@ -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<ReturnType<typeof GetServerFlags>>
|
||||
|
||||
const TEST_EMAIL = trustedUserEmail('user@example.com')
|
||||
|
||||
vi.mock('common', () => ({ IS_PLATFORM: true }))
|
||||
vi.mock('@/lib/server/configcat', async (importOriginal) => {
|
||||
const actual = await importOriginal<typeof import('@/lib/server/configcat')>()
|
||||
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)
|
||||
})
|
||||
})
|
||||
@@ -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<boolean> {
|
||||
if (!IS_PLATFORM) return false
|
||||
|
||||
const flags = await getServerFlags(userEmail)
|
||||
return flags.some((flag) => flag.settingKey === 'explorer' && flag.settingValue === true)
|
||||
}
|
||||
@@ -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<string, ToolCategory> = {
|
||||
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,
|
||||
|
||||
@@ -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')
|
||||
})
|
||||
})
|
||||
@@ -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 }) : {}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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<GetUserContentResponse>({
|
||||
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<GetUserContentByIdResponse>({
|
||||
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<GetUserContentByIdResponse>({
|
||||
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<APIErrorBody>({ 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()
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -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 }
|
||||
}
|
||||
}
|
||||
}),
|
||||
}
|
||||
},
|
||||
}),
|
||||
}
|
||||
}
|
||||
@@ -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' })
|
||||
)
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,73 @@
|
||||
import { getClient, PollingMode, User } from '@configcat/sdk/node'
|
||||
|
||||
let serverClient: ReturnType<typeof getClient>
|
||||
|
||||
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<string, string>) {
|
||||
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<string, string>
|
||||
) {
|
||||
const client = getServerClient()
|
||||
|
||||
if (!client) {
|
||||
return []
|
||||
}
|
||||
|
||||
return client.getAllValuesAsync(buildUser(userEmail, customAttributes))
|
||||
}
|
||||
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
|
||||
|
||||
Generated
+8
@@ -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)
|
||||
|
||||
Reference in new issue
Block a user