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:
Charis authored and GitHub committed 2026-08-11 08:40:51 -04:00
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)
})
})
+19
View File
@@ -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)
}
+4
View File
@@ -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,
+14
View File
@@ -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')
})
})
+7
View File
@@ -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 }
}
}
}),
}
},
}),
}
}
+127
View File
@@ -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' })
)
})
})
})
+73
View File
@@ -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))
}
+1
View File
@@ -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,
})
+9 -3
View File
@@ -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,
})
+8
View File
@@ -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)