From 5e59b6047e493759f2c7b1e76004e719bc808bda Mon Sep 17 00:00:00 2001 From: Saxon Fletcher Date: Mon, 28 Sep 2026 12:44:39 +1000 Subject: [PATCH] chore(studio): extend Assistant response time and handle timeouts (#50892) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Problem - Assistant responses were capped at 120 seconds and 10 steps, which is too short for longer reasoning or multi-step tool work. - When the hosting platform ended a request at that limit, the connection just dropped. The user got no explanation, and "Thinking…" and tool rows kept spinning. - Studio's own tools ignored the request's abort signal, so a stop, disconnect or deadline couldn't cancel their in-flight requests. - Aborted responses never closed their Braintrust span. Under TanStack Start, the remote MCP client was only released on `res.on('close')`, which the adapter never emits. ## Solution Uses AI SDK options instead of custom stream handling: - `maxDuration` goes to 300s and the step limit to 20. `streamText({ timeout: { totalMs } })` stops the response at 270s, leaving time to finish the stream before the platform cutoff. - `toUIMessageStream({ messageMetadata })` marks an aborted response `timedOut: true`. `Chat` ignores `abort` chunks, so the client reads this flag instead and shows a timeout alert with Retry. The flag is saved with the message, so the alert survives a reload. - `toUIMessageStream({ onEnd })` aborts the request whenever the stream ends, releasing the MCP client on both runtimes. `streamText({ onAbort })` ends the Braintrust span. - Studio tools pass the SDK's `abortSignal` to their fetches. MCP tools already did. - Reasoning and server-tool rows that never finished show "Response interrupted" instead of a spinner or "Ran X ✓". There's no per-tool timeout. Approved SQL and migrations can legitimately run longer, and aborting the HTTP request doesn't stop the query in Postgres. ## Review instructions 1. Run the unit tests: `cd apps/studio && pnpm vitest run lib/api/generate-v4.test.ts lib/ai components/ui/AIAssistantPanel` 2. To see a timeout without waiting 4.5 minutes, temporarily set `ASSISTANT_TIMEOUT_MS` in `apps/studio/lib/ai/assistant-timeout.ts` to `15_000` and run `pnpm dev:studio`. 3. Ask the Assistant something that needs several tool calls or long reasoning, for example "Audit my schema for missing indexes and RLS gaps, then write the fixes." 4. After 15 seconds, check that: - the response stops and a "Assistant response timed out" alert appears with Retry - any in-progress reasoning or tool row shows "Response interrupted" instead of spinning - Retry starts a new response - reloading the page still shows the alert on that chat 5. Stop a response with the Stop button before the deadline. It should stop without the timeout alert. 6. With the default 270s, confirm that a normal response completes as before. ## Checklist Check all before review: - [ ] I have read [CONTRIBUTING.md](https://github.com/supabase/supabase/blob/master/CONTRIBUTING.md) - [ ] If I wrote a new docs topic or edited an existing topic, I used the `/write-the-docs` or `/edit-the-docs` skill, which references [WORD_LIST](https://github.com/supabase/supabase/blob/master/apps/docs/WORD_LIST.md) and the docs [CONTRIBUTING](https://github.com/supabase/supabase/blob/master/apps/docs/CONTRIBUTING.md) guide ## Summary by CodeRabbit * **Improvements** * AI assistant responses can now run for up to five minutes, supporting longer requests. * When a response times out, the assistant displays a message suggesting you retry or ask for a smaller change. * Incomplete responses now show a “Response interrupted” notice, and loading indicators stop when generation ends. --------- Co-authored-by: Claude Opus 5.5 --- .../ui/AIAssistantPanel/AssistantChat.tsx | 27 +++-- .../AIAssistantPanel/Message.Parts.test.tsx | 87 +++++++++++++++- .../ui/AIAssistantPanel/Message.Parts.tsx | 36 ++++++- .../Message.performance.test.tsx | 1 + .../lib/ai/assistant-message-metadata.test.ts | 25 ++++- .../lib/ai/assistant-message-metadata.ts | 9 ++ apps/studio/lib/ai/assistant-timeout.ts | 5 + .../lib/ai/generate-assistant-response.ts | 23 ++++- apps/studio/lib/ai/tools/fallback-tools.ts | 12 +-- .../lib/ai/tools/incident-tools.test.ts | 34 +++++-- apps/studio/lib/ai/tools/incident-tools.ts | 6 +- apps/studio/lib/ai/tools/notebook-tools.ts | 41 ++++---- apps/studio/lib/ai/tools/report-tools.ts | 10 +- apps/studio/lib/ai/tools/schema-tools.ts | 4 +- apps/studio/lib/ai/tools/studio-tools.ts | 4 +- apps/studio/lib/api/generate-v4.test.ts | 98 ++++++++++++++++++- apps/studio/pages/api/ai/sql/generate-v4.ts | 23 ++++- 17 files changed, 374 insertions(+), 71 deletions(-) create mode 100644 apps/studio/lib/ai/assistant-timeout.ts diff --git a/apps/studio/components/ui/AIAssistantPanel/AssistantChat.tsx b/apps/studio/components/ui/AIAssistantPanel/AssistantChat.tsx index 96a7546c00e..5e8b516ce6d 100644 --- a/apps/studio/components/ui/AIAssistantPanel/AssistantChat.tsx +++ b/apps/studio/components/ui/AIAssistantPanel/AssistantChat.tsx @@ -38,7 +38,11 @@ import { useLocalStorageQuery } from '@/hooks/misc/useLocalStorage' import { useOrgAiOptInLevel } from '@/hooks/misc/useOrgOptedIntoAi' import { useSelectedOrganizationQuery } from '@/hooks/misc/useSelectedOrganization' import { useSelectedProjectQuery } from '@/hooks/misc/useSelectedProject' -import type { AssistantMessageMetadata } from '@/lib/ai/assistant-message-metadata' +import { + isTimedOutMessage, + type AssistantMessageMetadata, +} from '@/lib/ai/assistant-message-metadata' +import { ASSISTANT_TIMEOUT_MESSAGE } from '@/lib/ai/assistant-timeout' import { getParallelApprovalIdsToReject } from '@/lib/ai/message-utils' import { IS_PLATFORM } from '@/lib/constants' import { uuidv4 } from '@/lib/helpers' @@ -289,6 +293,11 @@ export const AssistantChat = ({ (error.message?.includes('context_length_exceeded') || error.message?.includes('exceeds the context window')) + const isTimedOut = !error && !isChatLoading && isTimedOutMessage(chatMessages.at(-1)) + let displayError = IS_PLATFORM ? ASSISTANT_ERRORS['default'] : error + if (isContextExceededError) displayError = ASSISTANT_ERRORS['context-exceeded'] + if (isTimedOut) displayError = { message: ASSISTANT_TIMEOUT_MESSAGE } + const editedMessageIndex = editingMessageId ? chatMessages.findIndex((message) => message.id === editingMessageId) : -1 @@ -577,18 +586,16 @@ export const AssistantChat = ({ {renderedMessages}
- {error && ( + {(error || isTimedOut) && ( {isContextExceededError ? ( diff --git a/apps/studio/components/ui/AIAssistantPanel/Message.Parts.test.tsx b/apps/studio/components/ui/AIAssistantPanel/Message.Parts.test.tsx index 8f8042d8377..8470a88e2be 100644 --- a/apps/studio/components/ui/AIAssistantPanel/Message.Parts.test.tsx +++ b/apps/studio/components/ui/AIAssistantPanel/Message.Parts.test.tsx @@ -1,11 +1,33 @@ import type { ToolUIPart } from 'ai' +import { type PropsWithChildren } from 'react' import { describe, expect, it } from 'vitest' +import { MessageProvider } from './Message.Context' import { MessagePartSwitcher } from './Message.Parts' import { customRender } from '@/tests/lib/custom-render' type MessagePart = Parameters[0]['part'] +function Provider({ + children, + isLoading = false, + isLastMessage = true, +}: PropsWithChildren<{ isLoading?: boolean; isLastMessage?: boolean }>) { + return ( + {}, + onEdit: () => {}, + onBranch: () => {}, + onCancelEdit: () => {}, + }} + > + {children} + + ) +} + describe('MessagePartSwitcher', () => { it('keeps consecutive generic tool parts as direct siblings', () => { const reasoningPart = { @@ -22,10 +44,10 @@ describe('MessagePartSwitcher', () => { } satisfies ToolUIPart const { container } = customRender( - <> + - + ) const toolRows = container.querySelectorAll('.tool-item') @@ -33,4 +55,65 @@ describe('MessagePartSwitcher', () => { expect(toolRows[0].nextElementSibling).toBe(toolRows[1]) expect(toolRows[0]).toHaveClass('max-w-3xl') }) + + it.each([ + { type: 'reasoning', state: 'streaming', text: 'Still thinking' }, + { type: 'tool-execute_sql', state: 'input-streaming', toolCallId: 'sql-1' }, + { type: 'tool-create_notebook', state: 'input-streaming', toolCallId: 'notebook-1' }, + { type: 'tool-update_notebook', state: 'input-streaming', toolCallId: 'notebook-2' }, + { type: 'tool-query_logs', state: 'input-available', toolCallId: 'logs-1', input: {} }, + ] satisfies MessagePart[])('stops the $type indicator when the request ends', (part) => { + const { container, getByText, rerender } = customRender( + + + + ) + expect(container.querySelector('.animate-spin')).not.toBeNull() + + rerender( + + + + ) + expect(getByText('Response interrupted')).toBeInTheDocument() + expect(container.querySelector('.animate-spin')).toBeNull() + }) + + it.each([ + { type: 'tool-search_docs', state: 'input-available', toolCallId: 'docs-1', input: {} }, + { + type: 'dynamic-tool', + toolName: 'list_tables', + state: 'input-available', + toolCallId: 'mcp-1', + input: {}, + }, + ] satisfies MessagePart[])( + 'marks a $type call that never returned as interrupted once the request ends', + (part) => { + const { queryByText, getByText, rerender } = customRender( + + + + ) + expect(queryByText('Response interrupted')).toBeNull() + + rerender( + + + + ) + expect(getByText('Response interrupted')).toBeInTheDocument() + } + ) + + it('does not restart an interrupted indicator when another message is streaming', () => { + const { container, getByText } = customRender( + + + + ) + expect(getByText('Response interrupted')).toBeInTheDocument() + expect(container.querySelector('.animate-spin')).toBeNull() + }) }) diff --git a/apps/studio/components/ui/AIAssistantPanel/Message.Parts.tsx b/apps/studio/components/ui/AIAssistantPanel/Message.Parts.tsx index 8eb9895e9e9..50603b2f871 100644 --- a/apps/studio/components/ui/AIAssistantPanel/Message.Parts.tsx +++ b/apps/studio/components/ui/AIAssistantPanel/Message.Parts.tsx @@ -1,6 +1,12 @@ import { UIMessage as VercelMessage } from '@ai-sdk/react' -import { type DynamicToolUIPart, type ReasoningUIPart, type TextUIPart, type ToolUIPart } from 'ai' -import { BrainIcon, CheckIcon, Loader2 } from 'lucide-react' +import { + isToolUIPart, + type DynamicToolUIPart, + type ReasoningUIPart, + type TextUIPart, + type ToolUIPart, +} from 'ai' +import { BrainIcon, CheckIcon, CircleStop, Loader2 } from 'lucide-react' import { memo, type ReactNode } from 'react' import { cn } from 'ui' @@ -336,6 +342,32 @@ const isCompactToolPart = (part: NonNullable[number]) => export const MessagePartSwitcher = memo( function MessagePartSwitcher({ part }: { part: NonNullable[number] }) { + const { isLoading, isLastMessage } = useMessageInfoContext() + const isActiveMessage = isLoading && isLastMessage + // Compact rows and query_logs run on the server, so `input-available` means the tool never + // returned. Other tools wait in that state for the user to act. + const isServerToolAwaitingOutput = + isToolUIPart(part) && + part.state === 'input-available' && + (isCompactToolPart(part) || + part.type === 'tool-query_logs' || + (part.type === 'dynamic-tool' && part.toolName === 'query_logs')) + const isIncompletePart = + (part.type === 'reasoning' && part.state === 'streaming') || + (isToolUIPart(part) && part.state === 'input-streaming') || + isServerToolAwaitingOutput + + if (!isActiveMessage && isIncompletePart) { + return ( + } + label="Response interrupted" + > + {part.type === 'reasoning' ? part.text : undefined} + + ) + } + const content = (() => { switch (part.type) { case 'dynamic-tool': { diff --git a/apps/studio/components/ui/AIAssistantPanel/Message.performance.test.tsx b/apps/studio/components/ui/AIAssistantPanel/Message.performance.test.tsx index bdd8b19a2d2..eb6dac5013a 100644 --- a/apps/studio/components/ui/AIAssistantPanel/Message.performance.test.tsx +++ b/apps/studio/components/ui/AIAssistantPanel/Message.performance.test.tsx @@ -95,6 +95,7 @@ function FeedMessage({ id={message.id} message={message} isLoading={isLoading} + isLastMessage isAfterEditedMessage={false} isBeingEdited={false} addToolApprovalResponse={addToolApprovalResponse} diff --git a/apps/studio/lib/ai/assistant-message-metadata.test.ts b/apps/studio/lib/ai/assistant-message-metadata.test.ts index 2206121e900..f617c3c00c2 100644 --- a/apps/studio/lib/ai/assistant-message-metadata.test.ts +++ b/apps/studio/lib/ai/assistant-message-metadata.test.ts @@ -3,6 +3,7 @@ import { describe, expect, it } from 'vitest' import { assistantMessageMetadataSchema, + isTimedOutMessage, messagesIncludeLogsSnippets, } from '@/lib/ai/assistant-message-metadata' @@ -10,8 +11,8 @@ function userMessage(id: string, text: string, metadata?: unknown): UIMessage { return { id, role: 'user', parts: [{ type: 'text', text }], metadata } as UIMessage } -function assistantMessage(id: string, text: string): UIMessage { - return { id, role: 'assistant', parts: [{ type: 'text', text }] } as UIMessage +function assistantMessage(id: string, text: string, metadata?: unknown): UIMessage { + return { id, role: 'assistant', parts: [{ type: 'text', text }], metadata } as UIMessage } describe('assistantMessageMetadataSchema', () => { @@ -88,3 +89,23 @@ describe('messagesIncludeLogsSnippets', () => { expect(messagesIncludeLogsSnippets([userMessage('1', 'hi', 'not an object')])).toBe(false) }) }) + +describe('isTimedOutMessage', () => { + it('detects an assistant response the server stopped at the deadline', () => { + expect(isTimedOutMessage(assistantMessage('1', 'partial', { timedOut: true }))).toBe(true) + }) + + it('is false for a completed response', () => { + expect(isTimedOutMessage(assistantMessage('1', 'done'))).toBe(false) + expect(isTimedOutMessage(assistantMessage('1', 'done', { timedOut: false }))).toBe(false) + }) + + it('ignores the flag on user messages', () => { + expect(isTimedOutMessage(userMessage('1', 'hi', { timedOut: true }))).toBe(false) + }) + + it('is false rather than throwing on a missing message or malformed metadata', () => { + expect(isTimedOutMessage(undefined)).toBe(false) + expect(isTimedOutMessage(assistantMessage('1', 'partial', { timedOut: 'yes' }))).toBe(false) + }) +}) diff --git a/apps/studio/lib/ai/assistant-message-metadata.ts b/apps/studio/lib/ai/assistant-message-metadata.ts index d30af08387e..21e90065064 100644 --- a/apps/studio/lib/ai/assistant-message-metadata.ts +++ b/apps/studio/lib/ai/assistant-message-metadata.ts @@ -10,11 +10,20 @@ export const assistantMessageMetadataSchema = z * carried by each snippet's own fence in the message text. */ containsLogsSnippets: z.boolean().optional(), + /** Set by the server when this response was cut off at the Assistant's deadline. */ + timedOut: z.boolean().optional(), }) .optional() export type AssistantMessageMetadata = z.infer +/** Whether the server stopped this assistant response at the Assistant's deadline. */ +export function isTimedOutMessage(message: UIMessage | undefined): boolean { + if (message?.role !== 'assistant') return false + const metadata = assistantMessageMetadataSchema.safeParse(message.metadata) + return metadata.success && metadata.data?.timedOut === true +} + /** * Whether any user message in the conversation attached a logs query. * diff --git a/apps/studio/lib/ai/assistant-timeout.ts b/apps/studio/lib/ai/assistant-timeout.ts new file mode 100644 index 00000000000..453c72c699c --- /dev/null +++ b/apps/studio/lib/ai/assistant-timeout.ts @@ -0,0 +1,5 @@ +// Leaves time to finish the stream before the hosting platform's 300-second cutoff, set by +// `maxDuration` in pages/api/ai/sql/generate-v4.ts (Next.js) and vite.config.ts (TanStack Start). +export const ASSISTANT_TIMEOUT_MS = 270_000 +export const ASSISTANT_TIMEOUT_MESSAGE = + 'The Assistant took too long to respond. Retry, or ask for a smaller change.' diff --git a/apps/studio/lib/ai/generate-assistant-response.ts b/apps/studio/lib/ai/generate-assistant-response.ts index 18eac67cfee..94086372e41 100644 --- a/apps/studio/lib/ai/generate-assistant-response.ts +++ b/apps/studio/lib/ai/generate-assistant-response.ts @@ -5,6 +5,7 @@ import { type LanguageModel, type ModelMessage, type SystemModelMessage, + type TimeoutConfiguration, type ToolSet, type UIMessage, } from 'ai' @@ -48,6 +49,7 @@ export async function generateAssistantResponse({ providerOptions, requestedModel, abortSignal, + timeout, onSpanCreated, }: { messages: UIMessage[] @@ -72,6 +74,7 @@ export async function generateAssistantResponse({ systemProviderOptions?: Record providerOptions?: Record abortSignal?: AbortSignal + timeout?: TimeoutConfiguration onSpanCreated?: (spanId: string) => void }) { const shouldTrace = allowTracing ?? IS_TRACING_ENABLED @@ -126,14 +129,24 @@ export async function generateAssistantResponse({ const streamTextFn = shouldTrace ? tracedStreamText : ai.streamText + // onEnd still fires after an abort once a step has finished, so end the span only once. + let isSpanEnded = false + const endSpan = (metadata: Record) => { + if (!span || isSpanEnded) return + isSpanEnded = true + span.log({ metadata }) + span.end() + } + return streamTextFn({ model, instructions: systemMessage, - stopWhen: isStepCount(10), + stopWhen: isStepCount(20), messages: coreMessages, ...(providerOptions && { providerOptions }), tools, ...(abortSignal && { abortSignal }), + ...(timeout && { timeout }), ...(span && { onEnd: ({ steps, finishReason }) => { const metadata: Record = { @@ -147,8 +160,12 @@ export async function generateAssistantResponse({ } } } - span.log({ metadata }) - span.end() + endSpan(metadata) + }, + // The call aborts on either the request signal or `timeout`, so an unaborted + // request signal means the deadline stopped it. + onAbort: () => { + endSpan({ isAborted: true, isTimedOut: !abortSignal?.aborted }) }, }), } satisfies Parameters[0]) diff --git a/apps/studio/lib/ai/tools/fallback-tools.ts b/apps/studio/lib/ai/tools/fallback-tools.ts index b8479f9020e..925d191f2c6 100644 --- a/apps/studio/lib/ai/tools/fallback-tools.ts +++ b/apps/studio/lib/ai/tools/fallback-tools.ts @@ -34,7 +34,7 @@ export const getFallbackTools = ({ inputSchema: z.object({ schemas: z.array(z.string()).describe('The schema names to get the definitions for'), }), - execute: async ({ schemas }) => { + execute: async ({ schemas }, { abortSignal }) => { try { const { result } = includeSchemaMetadata ? await executeSql( @@ -43,7 +43,7 @@ export const getFallbackTools = ({ connectionString, sql: getEntityDefinitionsSql({ schemas }), }, - undefined, + abortSignal, headers, IS_PLATFORM ? undefined : executeQuery ) @@ -84,7 +84,7 @@ export const getFallbackTools = ({ inputSchema: z.object({ schemas: z.array(z.string()).describe('The schema names to get the policies for'), }), - execute: async ({ schemas }) => { + execute: async ({ schemas }, { abortSignal }) => { const data = includeSchemaMetadata ? await getDatabasePolicies( { @@ -92,7 +92,7 @@ export const getFallbackTools = ({ connectionString, schemas, }, - undefined + abortSignal ) : [] @@ -355,7 +355,7 @@ export const getFallbackTools = ({ inputSchema: z.object({ schemas: z.array(z.string()).describe('The schema names to get the functions for'), }), - execute: async ({ schemas }) => { + execute: async ({ schemas }, { abortSignal }) => { try { const data = includeSchemaMetadata ? await getDatabaseFunctions( @@ -363,7 +363,7 @@ export const getFallbackTools = ({ projectRef, connectionString, }, - undefined, + abortSignal, headers ) : [] diff --git a/apps/studio/lib/ai/tools/incident-tools.test.ts b/apps/studio/lib/ai/tools/incident-tools.test.ts index 82c27313fe1..f7b73a8e77f 100644 --- a/apps/studio/lib/ai/tools/incident-tools.test.ts +++ b/apps/studio/lib/ai/tools/incident-tools.test.ts @@ -7,6 +7,8 @@ vi.mock('common', () => ({ IS_PLATFORM: true, })) +const executeOptions = { toolCallId: 'test', messages: [], context: {} } + describe('ai/tools/incident-tools', () => { let mockFetch: ReturnType let mockAbortSignal: AbortSignal @@ -52,7 +54,7 @@ describe('ai/tools/incident-tools', () => { vi.spyOn(common, 'IS_PLATFORM', 'get').mockReturnValue(false) const tools = getIncidentTools({ baseUrl: 'https://supabase.com/dashboard' }) - const result = await (tools.get_active_incidents.execute as any)({}) + const result = await (tools.get_active_incidents.execute as any)({}, executeOptions) expect(result).toEqual({ incidents: [], @@ -95,7 +97,7 @@ describe('ai/tools/incident-tools', () => { }) const tools = getIncidentTools({ baseUrl: 'https://supabase.com/dashboard' }) - const result = await (tools.get_active_incidents.execute as any)({}) + const result = await (tools.get_active_incidents.execute as any)({}, executeOptions) expect(result).toEqual({ incidents: [], @@ -123,7 +125,7 @@ describe('ai/tools/incident-tools', () => { }) const tools = getIncidentTools({ baseUrl: 'https://supabase.com/dashboard' }) - const result = await (tools.get_active_incidents.execute as any)({}) + const result = await (tools.get_active_incidents.execute as any)({}, executeOptions) expect((result as any).incidents).toEqual([ { @@ -162,7 +164,7 @@ describe('ai/tools/incident-tools', () => { }) const tools = getIncidentTools({ baseUrl: 'https://supabase.com/dashboard' }) - const result = await (tools.get_active_incidents.execute as any)({}) + const result = await (tools.get_active_incidents.execute as any)({}, executeOptions) expect((result as any).incidents).toHaveLength(2) expect((result as any).message).toContain('2 active incidents') @@ -175,7 +177,7 @@ describe('ai/tools/incident-tools', () => { mockFetch.mockRejectedValue(new Error('Network error')) const tools = getIncidentTools({ baseUrl: 'https://supabase.com/dashboard' }) - const result = await (tools.get_active_incidents.execute as any)({}) + const result = await (tools.get_active_incidents.execute as any)({}, executeOptions) expect(result).toEqual({ incidents: [], @@ -193,7 +195,7 @@ describe('ai/tools/incident-tools', () => { }) const tools = getIncidentTools({ baseUrl: 'https://supabase.com/dashboard' }) - const result = await (tools.get_active_incidents.execute as any)({}) + const result = await (tools.get_active_incidents.execute as any)({}, executeOptions) expect(result).toEqual({ incidents: [], @@ -220,6 +222,26 @@ describe('ai/tools/incident-tools', () => { const callArgs = mockFetch.mock.calls[0] expect(callArgs[1].signal).toBeInstanceOf(AbortSignal) }) + + it('cancels the request when the Assistant request is aborted', async () => { + mockFetch.mockResolvedValue({ + ok: true, + json: async () => [], + }) + const abortController = new AbortController() + + const tools = getIncidentTools({ baseUrl: 'https://supabase.com/dashboard' }) + if (!tools.get_active_incidents.execute) throw new Error('execute is undefined') + await tools.get_active_incidents.execute( + {}, + { ...executeOptions, abortSignal: abortController.signal } + ) + + const { signal } = mockFetch.mock.calls[0][1] + expect(signal.aborted).toBe(false) + abortController.abort() + expect(signal.aborted).toBe(true) + }) }) }) }) diff --git a/apps/studio/lib/ai/tools/incident-tools.ts b/apps/studio/lib/ai/tools/incident-tools.ts index 2275b8ef5a2..dd4277f54e3 100644 --- a/apps/studio/lib/ai/tools/incident-tools.ts +++ b/apps/studio/lib/ai/tools/incident-tools.ts @@ -15,7 +15,7 @@ export const getIncidentTools = ({ baseUrl }: { baseUrl: string }) => ({ description: 'Check for active incidents. Use this tool when the user reports issues with any Supabase service, including the database, authentication, realtime, storage, and functions. Possible problems include, but are not limited to, connection issues, timeouts, service unavailability, authentication failures, or unexpected errors.', inputSchema: z.object({}), - execute: async () => { + execute: async (_input, { abortSignal }) => { if (!IS_PLATFORM) { return { incidents: [], @@ -25,7 +25,9 @@ export const getIncidentTools = ({ baseUrl }: { baseUrl: string }) => ({ try { const response = await fetch(`${baseUrl}/api/incident-status`, { - signal: AbortSignal.timeout(5_000), + signal: abortSignal + ? AbortSignal.any([abortSignal, AbortSignal.timeout(5_000)]) + : AbortSignal.timeout(5_000), }) if (!response.ok) { diff --git a/apps/studio/lib/ai/tools/notebook-tools.ts b/apps/studio/lib/ai/tools/notebook-tools.ts index 798d61fa721..40ab468767a 100644 --- a/apps/studio/lib/ai/tools/notebook-tools.ts +++ b/apps/studio/lib/ai/tools/notebook-tools.ts @@ -128,8 +128,8 @@ export const getNotebookTools = (ctx: NotebookToolsContext = {}) => { description: 'List the databases available for this project — the primary and any read replicas.', inputSchema: z.object({}), - execute: async () => { - const databases = await getReadReplicas({ projectRef }, undefined, authHeaders) + execute: async (_input, { abortSignal }) => { + const databases = await getReadReplicas({ projectRef }, abortSignal, authHeaders) return { databases: (databases ?? []).map((database) => ({ @@ -162,10 +162,10 @@ export const getNotebookTools = (ctx: NotebookToolsContext = {}) => { 'Field to sort notebooks by. There is no "updated_at" sort — use "inserted_at" for creation order.' ), }), - execute: async ({ cursor, limit, sort_by }) => { + execute: async ({ cursor, limit, sort_by }, { abortSignal }) => { const { content, cursor: nextCursor } = await getContent( { projectRef, type: 'notebook', limit, cursor, sort: sort_by }, - undefined, + abortSignal, authHeaders ) @@ -188,8 +188,8 @@ export const getNotebookTools = (ctx: NotebookToolsContext = {}) => { 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) + execute: async ({ id }, { abortSignal }) => { + const notebook = await getNotebook({ projectRef, id }, abortSignal, authHeaders) // toWireNotebook discards the `unchecked_sql` brand for display purposes only — the // result is returned to the agent, never written back. @@ -215,8 +215,8 @@ export const getNotebookTools = (ctx: NotebookToolsContext = {}) => { ), }), needsApproval: true, - execute: async ({ id, expected_updated_at }) => { - const notebook = await getNotebook({ projectRef, id }, undefined, authHeaders) + execute: async ({ id, expected_updated_at }, { abortSignal }) => { + const notebook = await getNotebook({ projectRef, id }, abortSignal, authHeaders) if (notebook.updated_at !== expected_updated_at) { throw new NotebookToolError( @@ -238,7 +238,7 @@ export const getNotebookTools = (ctx: NotebookToolsContext = {}) => { let replicaLookupFailure: { error: unknown } | undefined if (shouldLookupReplica) { try { - databases = await getReadReplicas({ projectRef }, undefined, authHeaders) + databases = await getReadReplicas({ projectRef }, abortSignal, authHeaders) } catch (error) { replicaLookupFailure = { error } } @@ -260,6 +260,7 @@ export const getNotebookTools = (ctx: NotebookToolsContext = {}) => { sql: acceptUntrustedLogsSql(cell.unchecked_sql), range: resolveLogTimeRange(cell.time_range), endpoint: QUERY_SOURCE_REGISTRY.logs.endpoint, + signal: abortSignal, headers: authHeaders, }) if (result.error) throw result.error @@ -304,7 +305,7 @@ export const getNotebookTools = (ctx: NotebookToolsContext = {}) => { sql: limitedSql.sql, isStatementTimeoutDisabled: true, }, - undefined, + abortSignal, authHeaders ) cells.push({ @@ -351,13 +352,13 @@ export const getNotebookTools = (ctx: NotebookToolsContext = {}) => { ), }), needsApproval: true, - execute: async ({ name, description, content }) => { + execute: async ({ name, description, content }, { abortSignal }) => { if ( content.cells.some( (cell) => cell._tag === 'database_cell' && cell.database_identifier !== undefined ) ) { - const databases = await getReadReplicas({ projectRef }, undefined, authHeaders) + const databases = await getReadReplicas({ projectRef }, abortSignal, authHeaders) assertValidDatabaseIdentifiers( content.cells, new Set((databases ?? []).map((database) => database.identifier)) @@ -385,7 +386,7 @@ export const getNotebookTools = (ctx: NotebookToolsContext = {}) => { description, content: { schema_version: content.schema_version, cells }, }, - undefined, + abortSignal, authHeaders ) @@ -407,7 +408,7 @@ export const getNotebookTools = (ctx: NotebookToolsContext = {}) => { ), }), needsApproval: true, - execute: async ({ id, expected_updated_at, operations }) => { + execute: async ({ id, expected_updated_at, operations }, { abortSignal }) => { const newCells = operations.flatMap((operation) => operation._tag === 'insert_cell' || operation._tag === 'replace_cell' ? [operation.cell] @@ -418,14 +419,14 @@ export const getNotebookTools = (ctx: NotebookToolsContext = {}) => { (cell) => cell._tag === 'database_cell' && cell.database_identifier !== undefined ) ) { - const databases = await getReadReplicas({ projectRef }, undefined, authHeaders) + const databases = await getReadReplicas({ projectRef }, abortSignal, authHeaders) assertValidDatabaseIdentifiers( newCells, new Set((databases ?? []).map((database) => database.identifier)) ) } - const notebook = await getNotebook({ projectRef, id }, undefined, authHeaders) + const notebook = await getNotebook({ projectRef, id }, abortSignal, authHeaders) if (notebook.updated_at !== expected_updated_at) { throw new NotebookToolError( @@ -467,7 +468,7 @@ export const getNotebookTools = (ctx: NotebookToolsContext = {}) => { description: notebook.description ?? undefined, content: { schema_version: result.notebook.schema_version, cells }, }, - undefined, + abortSignal, authHeaders ) @@ -487,10 +488,10 @@ export const getNotebookTools = (ctx: NotebookToolsContext = {}) => { id: z.string().describe('The id of the notebook to delete.'), }), needsApproval: true, - execute: async ({ id }) => { - const notebook = await getNotebook({ projectRef, id }, undefined, authHeaders) + execute: async ({ id }, { abortSignal }) => { + const notebook = await getNotebook({ projectRef, id }, abortSignal, authHeaders) - await deleteContents({ projectRef: projectRef ?? '', ids: [id] }, undefined, authHeaders) + await deleteContents({ projectRef: projectRef ?? '', ids: [id] }, abortSignal, authHeaders) return { id, name: notebook.name } }, diff --git a/apps/studio/lib/ai/tools/report-tools.ts b/apps/studio/lib/ai/tools/report-tools.ts index 414090f2275..8a22c1e6022 100644 --- a/apps/studio/lib/ai/tools/report-tools.ts +++ b/apps/studio/lib/ai/tools/report-tools.ts @@ -26,10 +26,10 @@ export const getReportTools = (ctx: ReportToolsContext = {}) => { .default(20) .describe('Max number of reports to return.'), }), - execute: async ({ limit }) => { + execute: async ({ limit }, { abortSignal }) => { const { content } = await getContent( { projectRef, type: 'report', limit }, - undefined, + abortSignal, authHeaders ) @@ -49,8 +49,8 @@ export const getReportTools = (ctx: ReportToolsContext = {}) => { inputSchema: z.object({ id: z.string().describe('The id of the report to fetch.'), }), - execute: async ({ id }) => { - const report = await getContentById({ projectRef, id }, undefined, authHeaders) + execute: async ({ id }, { abortSignal }) => { + const report = await getContentById({ projectRef, id }, abortSignal, authHeaders) if (report.type !== 'report') { throw new Error(`Content ${id} is not a report (type: ${report.type})`) } @@ -63,7 +63,7 @@ export const getReportTools = (ctx: ReportToolsContext = {}) => { // A SQL-block chart's `id` is the id of its linked `type: 'sql'` content row. const snippet = await getContentById( { projectRef, id: chart.id }, - undefined, + abortSignal, authHeaders ).catch(() => null) diff --git a/apps/studio/lib/ai/tools/schema-tools.ts b/apps/studio/lib/ai/tools/schema-tools.ts index 749e3c3cd6f..e2c48f5fead 100644 --- a/apps/studio/lib/ai/tools/schema-tools.ts +++ b/apps/studio/lib/ai/tools/schema-tools.ts @@ -17,14 +17,14 @@ export const getSchemaTools = ({ inputSchema: z.object({ schemas: z.array(z.string()).describe('The schema names to get the policies for'), }), - execute: async ({ schemas }) => { + execute: async ({ schemas }, { abortSignal }) => { const data = await getDatabasePolicies( { projectRef, connectionString, schemas, }, - undefined, + abortSignal, authorization ? { Authorization: authorization } : undefined ) diff --git a/apps/studio/lib/ai/tools/studio-tools.ts b/apps/studio/lib/ai/tools/studio-tools.ts index 74c59a048fc..1f311cd9f66 100644 --- a/apps/studio/lib/ai/tools/studio-tools.ts +++ b/apps/studio/lib/ai/tools/studio-tools.ts @@ -72,13 +72,13 @@ export const getStudioTools = (ctx: StudioToolsContext = {}) => { 'Asks the user to execute a SQL statement and return the results. Requires user approval before executing.', inputSchema: executeSqlInputSchema, needsApproval: true, - execute: async ({ sql }) => { + execute: async ({ sql }, { abortSignal }) => { // The `needsApproval: true` gate on this tool means the user has // explicitly approved this AI-generated SQL before execute runs — // that approval is the user gesture that promotes untrusted to safe. const { result } = await executeSql( { projectRef, connectionString, sql: acceptUntrustedSql(untrustedSql(sql)) }, - undefined, + abortSignal, authHeaders ) return result diff --git a/apps/studio/lib/api/generate-v4.test.ts b/apps/studio/lib/api/generate-v4.test.ts index 1b82828975d..9d1934837c5 100644 --- a/apps/studio/lib/api/generate-v4.test.ts +++ b/apps/studio/lib/api/generate-v4.test.ts @@ -1,8 +1,9 @@ import { safeSql } from '@supabase/pg-meta' -import { UIMessage } from 'ai' +import { pipeUIMessageStreamToResponse, streamText, UIMessage } from 'ai' import { expect, test, vi } from 'vitest' import generateV4 from '../../pages/api/ai/sql/generate-v4' +import { ASSISTANT_TIMEOUT_MS } from '@/lib/ai/assistant-timeout' import { getTools } from '@/lib/ai/tools' import { sanitizeMessagePart } from '@/lib/ai/tools/tool-sanitizer' @@ -37,13 +38,28 @@ vi.mock('ai', async () => { const actual = await vi.importActual('ai') return { ...actual, - streamText: vi.fn().mockReturnValue({ - pipeUIMessageStreamToResponse: vi.fn(), + streamText: vi.fn().mockImplementation(() => ({ + stream: new ReadableStream({ + start(controller) { + controller.enqueue({ type: 'start' }) + controller.close() + }, + }), + })), + // Consume the response, as the real Node response writer does. + pipeUIMessageStreamToResponse: vi.fn(async ({ stream }) => { + const chunks: unknown[] = [] + const reader = stream.getReader() + while (true) { + const { done, value } = await reader.read() + if (done) return chunks + chunks.push(value) + } }), } }) -test('generateV4 calls the tool sanitizer', async () => { +function createMocks() { const mockReq = { method: 'POST', headers: { @@ -80,7 +96,16 @@ test('generateV4 calls the tool sanitizer', async () => { on: vi.fn(), } - await generateV4(mockReq as any, mockRes as any) + return { mockRes, callGenerateV4: () => generateV4(mockReq as any, mockRes as any) } +} + +test('generateV4 calls the tool sanitizer', async () => { + const { mockRes, callGenerateV4 } = createMocks() + + await callGenerateV4() + expect(pipeUIMessageStreamToResponse).toHaveBeenCalledOnce() + await vi.mocked(pipeUIMessageStreamToResponse).mock.results[0].value + expect(mockRes.status).not.toHaveBeenCalledWith(500) expect(sanitizeMessagePart).toHaveBeenCalled() expect(getTools).toHaveBeenCalledWith( @@ -92,3 +117,66 @@ test('generateV4 calls the tool sanitizer', async () => { // opened in getTools is torn down when the stream finishes or the client drops expect(mockRes.on).toHaveBeenCalledWith('close', expect.any(Function)) }) + +test('generateV4 streams a tool result that continues the previous assistant message', async () => { + vi.mocked(streamText).mockClear() + vi.mocked(pipeUIMessageStreamToResponse).mockClear() + // After an approval, streamText runs the approved tool first and streams its result into the + // assistant message the client already has, so the tool call never appears in this stream. + vi.mocked(streamText).mockImplementationOnce( + () => + ({ + stream: new ReadableStream({ + start(controller) { + controller.enqueue({ type: 'start' }) + controller.enqueue({ + type: 'tool-result', + toolCallId: 'test-tool-call-id', + toolName: 'render_page', + input: {}, + output: { status: 'ready' }, + }) + controller.close() + }, + }), + }) as unknown as ReturnType + ) + const { callGenerateV4 } = createMocks() + + await callGenerateV4() + const chunks = await vi.mocked(pipeUIMessageStreamToResponse).mock.results[0].value + + expect(chunks).toContainEqual( + expect.objectContaining({ type: 'tool-output-available', toolCallId: 'test-tool-call-id' }) + ) +}) + +test('generateV4 flags a response the deadline stopped and releases the request', async () => { + vi.mocked(streamText).mockClear() + vi.mocked(pipeUIMessageStreamToResponse).mockClear() + vi.mocked(streamText).mockImplementationOnce( + () => + ({ + stream: new ReadableStream({ + start(controller) { + controller.enqueue({ type: 'start' }) + controller.enqueue({ type: 'abort', reason: 'signal timed out' }) + controller.close() + }, + }), + }) as unknown as ReturnType + ) + const { callGenerateV4 } = createMocks() + + await callGenerateV4() + const chunks = await vi.mocked(pipeUIMessageStreamToResponse).mock.results[0].value + + expect(chunks).toContainEqual({ type: 'message-metadata', messageMetadata: { timedOut: true } }) + const params = vi.mocked(streamText).mock.calls[0][0] + expect(params.timeout).toEqual({ totalMs: expect.any(Number) }) + const { totalMs } = params.timeout as { totalMs: number } + expect(totalMs).toBeGreaterThan(0) + expect(totalMs).toBeLessThanOrEqual(ASSISTANT_TIMEOUT_MS) + // Ending the stream aborts the request signal, which closes the remote MCP client. + await vi.waitFor(() => expect(params.abortSignal?.aborted).toBe(true)) +}) diff --git a/apps/studio/pages/api/ai/sql/generate-v4.ts b/apps/studio/pages/api/ai/sql/generate-v4.ts index e163fd7f66c..46a8f795406 100644 --- a/apps/studio/pages/api/ai/sql/generate-v4.ts +++ b/apps/studio/pages/api/ai/sql/generate-v4.ts @@ -12,7 +12,9 @@ import { NO_SCHEMA_ACCESS_MESSAGE } from '@/lib/ai/assistant-context' import { assistantMessageMetadataSchema, messagesIncludeLogsSnippets, + type AssistantMessageMetadata, } from '@/lib/ai/assistant-message-metadata' +import { ASSISTANT_TIMEOUT_MS } from '@/lib/ai/assistant-timeout' import { isTracingAllowed } from '@/lib/ai/braintrust-logger' import { generateAssistantResponse } from '@/lib/ai/generate-assistant-response' import { isExplorerEnabled } from '@/lib/ai/is-explorer-enabled' @@ -31,7 +33,7 @@ import { executeQuery } from '@/lib/api/self-hosted/query' import { getURL } from '@/lib/helpers' import { trustedUserEmail } from '@/lib/server/configcat' -export const maxDuration = 120 +export const maxDuration = 300 export const config = { api: { @@ -75,6 +77,7 @@ const requestBodySchema = z.object({ }) async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: JwtPayload) { + const requestStartedAt = Date.now() const authorization = req.headers.authorization const accessToken = authorization?.replace('Bearer ', '') @@ -176,8 +179,8 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw const abortController = new AbortController() req.on('close', () => abortController.abort()) req.on('aborted', () => abortController.abort()) - // Fires when the response finishes streaming or the connection drops, which - // is what tears down the remote MCP connection opened in getTools. + // Fires when the connection drops. Aborting tears down the remote MCP connection opened + // in getTools. The TanStack adapter doesn't emit it, so settling the pipe below also aborts. res.on('close', () => abortController.abort()) const tools = await getTools({ @@ -237,6 +240,7 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw requestedModel, systemProviderOptions, abortSignal: abortController.signal, + timeout: { totalMs: Math.max(0, ASSISTANT_TIMEOUT_MS - (Date.now() - requestStartedAt)) }, onSpanCreated: (spanId) => { res.setHeader('x-braintrust-span-id', spanId) }, @@ -265,13 +269,24 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw return JSON.stringify(error) }, + // The browser never receives an abort caused by its own disconnect, so any abort that + // reaches it is the deadline. Chat ignores abort chunks, so flag the message instead. + messageMetadata: ({ part }): AssistantMessageMetadata => + part.type === 'abort' ? { timedOut: true } : undefined, }) - pipeUIMessageStreamToResponse({ + // Keep this asynchronous so the TanStack adapter can return the streaming + // Response immediately. Handle piping failures after headers have been sent. + // Abort here rather than in toUIMessageStream's onEnd: that callback rebuilds the response + // message, which fails on approval continuations without the client's original messages. + void pipeUIMessageStreamToResponse({ response: res, stream, headers: { 'Content-Encoding': 'none' }, }) + .catch((error) => console.error('Error piping Assistant stream:', error)) + // Runs when the stream finishes, aborts, or is cancelled. + .finally(() => abortController.abort()) } catch (error) { console.error('Error in handlePost:', error) if (error instanceof Error) {