Context reduction (#37032)

* context saving

* fix model name

* add results comment

* fix errors
This commit is contained in:
Saxon Fletcher authored and GitHub committed 2025-07-11 16:03:37 +10:00
1 parent 3cf6a2d17f
commit fd670667ab
5 files changed
+42 -247

No files matched your search

+2 -2
View File
@@ -1,7 +1,7 @@
import { openai } from '@ai-sdk/openai'
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
import * as bedrockModule from './bedrock'
import { getModel, ModelErrorMessage, modelsByProvider } from './model'
import { getModel, ModelErrorMessage } from './model'
vi.mock('@ai-sdk/openai', () => ({
openai: vi.fn(() => 'openai-model'),
@@ -43,7 +43,7 @@ describe('getModel', () => {
const { model } = await getModel('test-key')
expect(model).toEqual('openai-model')
expect(openai).toHaveBeenCalledWith(modelsByProvider.openai)
expect(openai).toHaveBeenCalledWith('gpt-4.1-2025-04-14')
})
it('should return error when neither AWS credentials nor OPENAI_API_KEY is available', async () => {
+13 -11
View File
@@ -7,16 +7,17 @@ import {
selectBedrockRegion,
} from './bedrock'
export const modelsByProvider = {
bedrock: {
us1: 'us.anthropic.claude-3-7-sonnet-20250219-v1:0',
us2: 'us.anthropic.claude-3-7-sonnet-20250219-v1:0',
us3: 'us.anthropic.claude-3-7-sonnet-20250219-v1:0',
eu: 'eu.anthropic.claude-3-7-sonnet-20250219-v1:0',
},
openai: 'gpt-4.1-2025-04-14',
export const regionMap = {
us1: 'us',
us2: 'us',
us3: 'us',
eu: 'eu',
}
const SONNET_MODEL = 'anthropic.claude-3-7-sonnet-20250219-v1:0'
const HAIKU_MODEL = 'anthropic.claude-3-5-haiku-20241022-v1:0'
const OPENAI_MODEL = 'gpt-4.1-2025-04-14'
export type ModelSuccess = {
model: LanguageModel
error?: never
@@ -38,7 +39,7 @@ export const ModelErrorMessage =
* An optional routing key can be provided to distribute requests across
* different Bedrock regions.
*/
export async function getModel(routingKey?: string): Promise<ModelResponse> {
export async function getModel(routingKey?: string, isLimited?: boolean): Promise<ModelResponse> {
const hasAwsCredentials = await checkAwsCredentials()
const hasOpenAIKey = !!process.env.OPENAI_API_KEY
@@ -46,7 +47,8 @@ export async function getModel(routingKey?: string): Promise<ModelResponse> {
// Select the Bedrock region based on the routing key
const bedrockRegion: BedrockRegion = routingKey ? await selectBedrockRegion(routingKey) : 'us1'
const bedrock = bedrockForRegion(bedrockRegion)
const modelName = modelsByProvider.bedrock[bedrockRegion]
const model = isLimited ? HAIKU_MODEL : SONNET_MODEL
const modelName = `${regionMap[bedrockRegion]}.${model}`
return {
model: bedrock(modelName),
@@ -55,7 +57,7 @@ export async function getModel(routingKey?: string): Promise<ModelResponse> {
if (hasOpenAIKey) {
return {
model: openai(modelsByProvider.openai),
model: openai(OPENAI_MODEL),
}
}
+6 -147
View File
@@ -9,7 +9,6 @@ import {
createPrivacyMessageTool,
toolSetValidationSchema,
transformToolResult,
checkNetworkExtensionsAndAdjustOptInLevel,
DatabaseExtension,
} from './tool-filter'
@@ -18,7 +17,6 @@ describe('TOOL_CATEGORY_MAP', () => {
expect(TOOL_CATEGORY_MAP['display_query']).toBe(TOOL_CATEGORIES.UI)
expect(TOOL_CATEGORY_MAP['list_tables']).toBe(TOOL_CATEGORIES.SCHEMA)
expect(TOOL_CATEGORY_MAP['get_logs']).toBe(TOOL_CATEGORIES.LOG)
expect(TOOL_CATEGORY_MAP['execute_sql']).toBe(TOOL_CATEGORIES.DATA)
})
})
@@ -40,8 +38,6 @@ describe('tool allowance by opt-in level', () => {
get_logs: { execute: vitest.fn().mockResolvedValue({ status: 'success' }) },
get_advisors: { execute: vitest.fn().mockResolvedValue({ status: 'success' }) },
get_log_counts: { execute: vitest.fn().mockResolvedValue({ status: 'success' }) },
// Data tools
execute_sql: { execute: vitest.fn().mockResolvedValue({ status: 'success' }) },
} as unknown as ToolSet
const filtered = filterToolsByOptInLevel(mockTools, optInLevel as any)
@@ -104,7 +100,7 @@ describe('tool allowance by opt-in level', () => {
expect(tools).not.toContain('execute_sql')
})
it('should return all tools for schema_and_log_and_data opt-in level', () => {
it('should return all tools for schema_and_log_and_data opt-in level (excluding execute_sql)', () => {
const tools = getAllowedTools('schema_and_log_and_data')
expect(tools).toContain('display_query')
expect(tools).toContain('display_edge_function')
@@ -117,7 +113,7 @@ describe('tool allowance by opt-in level', () => {
expect(tools).toContain('get_logs')
expect(tools).toContain('get_advisors')
expect(tools).toContain('get_log_counts')
expect(tools).toContain('execute_sql')
expect(tools).not.toContain('execute_sql')
})
})
@@ -137,8 +133,6 @@ describe('filterToolsByOptInLevel', () => {
get_logs: { execute: vitest.fn().mockResolvedValue({ status: 'success' }) },
get_advisors: { execute: vitest.fn().mockResolvedValue({ status: 'success' }) },
get_log_counts: { execute: vitest.fn().mockResolvedValue({ status: 'success' }) },
// Data tools
execute_sql: { execute: vitest.fn().mockResolvedValue({ status: 'success' }) },
// Unknown tool - should be filtered out entirely
some_other_tool: { execute: vitest.fn().mockResolvedValue({ status: 'success' }) },
} as unknown as ToolSet
@@ -195,7 +189,6 @@ describe('filterToolsByOptInLevel', () => {
'get_logs',
'get_advisors',
'get_log_counts',
'execute_sql',
])
})
@@ -210,21 +203,16 @@ describe('filterToolsByOptInLevel', () => {
'get_logs',
'get_advisors',
'get_log_counts',
'execute_sql',
])
})
it('should stub log and execute tools for schema opt-in level', async () => {
it('should stub log tools for schema opt-in level', async () => {
const tools = filterToolsByOptInLevel(mockTools, 'schema')
await expectStubsFor(tools, ['get_logs', 'get_advisors', 'get_log_counts', 'execute_sql'])
await expectStubsFor(tools, ['get_logs', 'get_advisors', 'get_log_counts'])
})
it('should stub execute tool for schema_and_log opt-in level', async () => {
const tools = filterToolsByOptInLevel(mockTools, 'schema_and_log')
await expectStubsFor(tools, ['execute_sql'])
})
// No execute_sql tool, so nothing additional to stub for schema_and_log opt-in level
it('should not stub any tools for schema_and_log_and_data opt-in level', async () => {
const tools = filterToolsByOptInLevel(mockTools, 'schema_and_log_and_data')
@@ -340,7 +328,6 @@ describe('toolSetValidationSchema', () => {
list_edge_functions: { parameters: z.object({}), execute: vitest.fn() },
list_branches: { parameters: z.object({}), execute: vitest.fn() },
get_logs: { parameters: z.object({}), execute: vitest.fn() },
execute_sql: { parameters: z.object({}), execute: vitest.fn() },
search_docs: { parameters: z.object({}), execute: vitest.fn() },
get_advisors: { parameters: z.object({}), execute: vitest.fn() },
display_query: { parameters: z.object({}), execute: vitest.fn() },
@@ -354,137 +341,9 @@ describe('toolSetValidationSchema', () => {
// Test with missing tool
const incompleteTools = { ...allExpectedTools }
delete (incompleteTools as any).execute_sql
delete (incompleteTools as any).search_docs
const incompleteValidationResult = toolSetValidationSchema.safeParse(incompleteTools)
expect(incompleteValidationResult.success).toBe(true) // Should still pass as we allow subsets
})
})
describe('checkNetworkExtensionsAndAdjustOptInLevel', () => {
const createMockExtension = (name: string, installedVersion?: string): DatabaseExtension => ({
comment: null,
default_version: '1.0',
installed_version: installedVersion || null,
name,
schema: 'public',
})
it('should return the same opt-in level when no extensions are provided', () => {
const result = checkNetworkExtensionsAndAdjustOptInLevel(null, 'schema_and_log_and_data')
expect(result).toBe('schema_and_log_and_data')
})
it('should return the same opt-in level when extensions array is empty', () => {
const result = checkNetworkExtensionsAndAdjustOptInLevel([], 'schema_and_log_and_data')
expect(result).toBe('schema_and_log_and_data')
})
it('should return the same opt-in level when extensions array is undefined', () => {
const result = checkNetworkExtensionsAndAdjustOptInLevel(undefined, 'schema_and_log_and_data')
expect(result).toBe('schema_and_log_and_data')
})
it('should not downgrade when no dangerous network extensions are installed', () => {
const extensions = [
createMockExtension('postgis', '3.2.0'),
createMockExtension('uuid-ossp', '1.1'),
createMockExtension('pg_trgm', '1.6'),
]
const result = checkNetworkExtensionsAndAdjustOptInLevel(extensions, 'schema_and_log_and_data')
expect(result).toBe('schema_and_log_and_data')
})
it('should downgrade from schema_and_log_and_data to schema_and_log when pg_net is installed', () => {
const extensions = [
createMockExtension('postgis', '3.2.0'),
createMockExtension('pg_net', '0.7.1'), // installed
createMockExtension('uuid-ossp', '1.1'),
]
const result = checkNetworkExtensionsAndAdjustOptInLevel(extensions, 'schema_and_log_and_data')
expect(result).toBe('schema_and_log')
})
it('should downgrade from schema_and_log_and_data to schema_and_log when http extension is installed', () => {
const extensions = [
createMockExtension('postgis', '3.2.0'),
createMockExtension('http', '1.5.0'), // installed
createMockExtension('uuid-ossp', '1.1'),
]
const result = checkNetworkExtensionsAndAdjustOptInLevel(extensions, 'schema_and_log_and_data')
expect(result).toBe('schema_and_log')
})
it('should downgrade when both pg_net and http extensions are installed', () => {
const extensions = [
createMockExtension('pg_net', '0.7.1'), // installed
createMockExtension('http', '1.5.0'), // installed
createMockExtension('uuid-ossp', '1.1'),
]
const result = checkNetworkExtensionsAndAdjustOptInLevel(extensions, 'schema_and_log_and_data')
expect(result).toBe('schema_and_log')
})
it('should not downgrade when dangerous extensions are available but not installed', () => {
const extensions = [
createMockExtension('pg_net'), // not installed (no installed_version)
createMockExtension('http'), // not installed (no installed_version)
createMockExtension('postgis', '3.2.0'), // installed
]
const result = checkNetworkExtensionsAndAdjustOptInLevel(extensions, 'schema_and_log_and_data')
expect(result).toBe('schema_and_log_and_data')
})
it('should not downgrade when opt-in level is already at or below schema_and_log', () => {
const extensions = [
createMockExtension('pg_net', '0.7.1'), // installed
]
// Test with schema_and_log - should remain unchanged
const result1 = checkNetworkExtensionsAndAdjustOptInLevel(extensions, 'schema_and_log')
expect(result1).toBe('schema_and_log')
// Test with schema - should remain unchanged
const result2 = checkNetworkExtensionsAndAdjustOptInLevel(extensions, 'schema')
expect(result2).toBe('schema')
// Test with disabled - should remain unchanged
const result3 = checkNetworkExtensionsAndAdjustOptInLevel(extensions, 'disabled')
expect(result3).toBe('disabled')
})
it('should handle extensions with empty string installed_version as not installed', () => {
const extensions = [
{
comment: null,
default_version: '1.0',
installed_version: '', // empty string should be treated as not installed
name: 'pg_net',
schema: 'public',
},
]
const result = checkNetworkExtensionsAndAdjustOptInLevel(extensions, 'schema_and_log_and_data')
expect(result).toBe('schema_and_log_and_data')
})
it('should handle extensions with null installed_version as not installed', () => {
const extensions = [
{
comment: null,
default_version: '1.0',
installed_version: null, // explicitly null
name: 'pg_net',
schema: 'public',
},
]
const result = checkNetworkExtensionsAndAdjustOptInLevel(extensions, 'schema_and_log_and_data')
expect(result).toBe('schema_and_log_and_data')
})
})
-33
View File
@@ -24,7 +24,6 @@ export const toolSetValidationSchema = z.record(
'list_edge_functions',
'list_branches',
'get_logs',
'execute_sql',
'search_docs',
'get_advisors',
@@ -102,9 +101,6 @@ export const TOOL_CATEGORY_MAP: Record<string, ToolCategory> = {
get_logs: TOOL_CATEGORIES.LOG,
get_advisors: TOOL_CATEGORIES.LOG,
get_log_counts: TOOL_CATEGORIES.LOG,
// Data tools - MCP only
execute_sql: TOOL_CATEGORIES.DATA,
}
/**
@@ -189,32 +185,3 @@ export function filterToolsByOptInLevel(tools: ToolSet, aiOptInLevel: AiOptInLev
})
)
}
/**
* If either `pg_net` or `http` extension is enabled, we cap the opt-in level to a maximum
* of `schema_and_log`, effectively disabling the `execute_sql` tool (which requires data-level access).
*/
export function checkNetworkExtensionsAndAdjustOptInLevel(
dbExtensions: DatabaseExtension[] | null | undefined,
currentOptInLevel: AiOptInLevel
): AiOptInLevel {
if (!dbExtensions || dbExtensions.length === 0) {
return currentOptInLevel
}
// List of dangerous network extensions that can make external connections
const dangerousNetworkExtensions = ['pg_net', 'http']
// Check if any dangerous network extensions are installed (have installed_version)
const hasNetworkExtensions = dbExtensions
.filter((ext) => !!ext.installed_version)
.some((ext) => dangerousNetworkExtensions.includes(ext.name))
// If network extensions are installed and current opt-in level is full data access,
// downgrade to schema_and_log to disable execute_sql tool
if (hasNetworkExtensions && currentOptInLevel === 'schema_and_log_and_data') {
return 'schema_and_log'
}
return currentOptInLevel
}
+21 -54
View File
@@ -1,7 +1,6 @@
import pgMeta from '@supabase/pg-meta'
import { streamText, tool, ToolSet } from 'ai'
import { source } from 'common-tags'
import crypto from 'crypto'
import { NextApiRequest, NextApiResponse } from 'next'
import { z } from 'zod'
@@ -15,15 +14,9 @@ import apiWrapper from 'lib/api/apiWrapper'
import { queryPgMetaSelfHosted } from 'lib/self-hosted'
import { getUnifiedLogsChart } from 'data/logs/unified-logs-chart-query'
import { getUnifiedLogs } from 'data/logs/unified-logs-infinite-query'
import { getDatabaseExtensions } from 'data/database-extensions/database-extensions-query'
import { QuerySearchParamsType } from 'components/interfaces/UnifiedLogs/UnifiedLogs.types'
import { createSupabaseMCPClient } from 'lib/ai/supabase-mcp'
import {
filterToolsByOptInLevel,
toolSetValidationSchema,
transformToolResult,
checkNetworkExtensionsAndAdjustOptInLevel,
} from 'lib/ai/tool-filter'
import { filterToolsByOptInLevel, toolSetValidationSchema } from 'lib/ai/tool-filter'
import { getTools } from './tools'
export const maxDuration = 120
@@ -74,7 +67,19 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse) {
return res.status(400).json({ error: 'Invalid request body', issues: parseError.issues })
}
const { messages, projectRef, connectionString, orgSlug, chatName } = data
const { messages: rawMessages, projectRef, connectionString, orgSlug, chatName } = data
// Server-side safety: limit to last 5 messages and remove `results` property to prevent accidental leakage.
// Results property is used to cache results client-side after queries are run
// Tool results will still be included in history sent to model
const messages = (rawMessages || []).slice(-5).map((msg: any) => {
if (msg && msg.role === 'assistant' && 'results' in msg) {
const cleanedMsg = { ...msg }
delete cleanedMsg.results
return cleanedMsg
}
return msg
})
// Get organizations and compute opt in level server-side
const [organizations, projects] = await Promise.all([
@@ -104,29 +109,9 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse) {
const aiOptInLevel = getAiOptInLevel(selectedOrg?.opt_in_tags)
// Check enabled database extensions to potentially restrict data access
let effectiveAiOptInLevel = aiOptInLevel
const isLimited = messages.length > 4
try {
let headers = new Headers()
if (authorization) headers.set('Authorization', authorization)
const dbExtensions = await getDatabaseExtensions(
{ projectRef, connectionString },
undefined,
headers
)
effectiveAiOptInLevel = checkNetworkExtensionsAndAdjustOptInLevel(
dbExtensions,
effectiveAiOptInLevel
)
} catch (error) {
console.error('Failed to fetch database extensions:', error)
effectiveAiOptInLevel = 'disabled'
}
const { model, error: modelError } = await getModel(projectRef) // use project ref as routing key
const { model, error: modelError } = await getModel(projectRef, isLimited) // use project ref as routing key
if (modelError) {
return res.status(500).json({ error: modelError.message })
@@ -305,8 +290,8 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse) {
limit: z
.number()
.min(1)
.max(100)
.default(20)
.max(20)
.default(10)
.describe('Maximum number of logs to return (1-100, defaults to 20)'),
}),
execute: async (args) => {
@@ -424,7 +409,7 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse) {
const availableMcpTools = await mcpClient.tools()
// Filter tools based on the (potentially modified) AI opt-in level
const allowedMcpTools = filterToolsByOptInLevel(availableMcpTools, effectiveAiOptInLevel)
const allowedMcpTools = filterToolsByOptInLevel(availableMcpTools, aiOptInLevel)
// Validate that only known tools are provided
const { data: validatedTools, error: validationError } =
@@ -438,29 +423,11 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse) {
})
}
// Modify the execute_sql tool to add manualToolCallId (if it exists)
mcpTools = {
...validatedTools,
...(validatedTools.execute_sql && {
execute_sql: transformToolResult(validatedTools.execute_sql, (result) => {
const manualToolCallId = `manual_${crypto.randomUUID()}`
if (typeof result === 'object') {
return { ...result, manualToolCallId }
} else {
console.warn('execute_sql result is not an object, cannot add manualToolCallId')
return {
error: 'Internal error: Unexpected tool result format',
manualToolCallId,
}
}
}),
}),
}
mcpTools = { ...validatedTools }
}
// Filter local tools based on the (potentially modified) AI opt-in level
const filteredLocalTools = filterToolsByOptInLevel(localTools, effectiveAiOptInLevel)
const filteredLocalTools = filterToolsByOptInLevel(localTools, aiOptInLevel)
// Combine MCP tools with filtered local tools
const tools: ToolSet = {