diff --git a/apps/studio/lib/ai/model.test.ts b/apps/studio/lib/ai/model.test.ts index 7be70a0962a..2d7d2e27658 100644 --- a/apps/studio/lib/ai/model.test.ts +++ b/apps/studio/lib/ai/model.test.ts @@ -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 () => { diff --git a/apps/studio/lib/ai/model.ts b/apps/studio/lib/ai/model.ts index 6903a40d60e..bddb5d87fd2 100644 --- a/apps/studio/lib/ai/model.ts +++ b/apps/studio/lib/ai/model.ts @@ -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 { +export async function getModel(routingKey?: string, isLimited?: boolean): Promise { const hasAwsCredentials = await checkAwsCredentials() const hasOpenAIKey = !!process.env.OPENAI_API_KEY @@ -46,7 +47,8 @@ export async function getModel(routingKey?: string): Promise { // 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 { if (hasOpenAIKey) { return { - model: openai(modelsByProvider.openai), + model: openai(OPENAI_MODEL), } } diff --git a/apps/studio/lib/ai/tool-filter.test.ts b/apps/studio/lib/ai/tool-filter.test.ts index 3f6ecd916a6..03b12fe66f3 100644 --- a/apps/studio/lib/ai/tool-filter.test.ts +++ b/apps/studio/lib/ai/tool-filter.test.ts @@ -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') - }) -}) diff --git a/apps/studio/lib/ai/tool-filter.ts b/apps/studio/lib/ai/tool-filter.ts index 52c7929e971..e16e52b27f3 100644 --- a/apps/studio/lib/ai/tool-filter.ts +++ b/apps/studio/lib/ai/tool-filter.ts @@ -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 = { 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 -} diff --git a/apps/studio/pages/api/ai/sql/generate-v4.ts b/apps/studio/pages/api/ai/sql/generate-v4.ts index 163a145b18b..58d8422da26 100644 --- a/apps/studio/pages/api/ai/sql/generate-v4.ts +++ b/apps/studio/pages/api/ai/sql/generate-v4.ts @@ -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 = {