chore(studio): extend Assistant response time and handle timeouts (#50892)

## 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


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## 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.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
Saxon FletcherandClaude Opus 5.5 authored and GitHub committed 2026-09-28 12:44:39 +10:00
1 parent 0eb08cb9f0
commit 5e59b6047e
17 files changed
+374 -71

No files matched your search

@@ -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 = ({
<ConversationContent className="w-full py-8 mb-10">
{renderedMessages}
<div className="w-full max-w-3xl mx-auto">
{error && (
{(error || isTimedOut) && (
<AlertError
error={
isContextExceededError
? ASSISTANT_ERRORS['context-exceeded']
: IS_PLATFORM
? ASSISTANT_ERRORS['default']
: error
}
error={displayError}
showErrorPrefix={false}
showInstructions={false}
subject="Sorry, I'm having trouble responding right now."
subject={
isTimedOut
? 'Assistant response timed out'
: "Sorry, I'm having trouble responding right now."
}
additionalActions={
<div className="flex items-center gap-x-2 mr-auto">
{isContextExceededError ? (
@@ -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<typeof MessagePartSwitcher>[0]['part']
function Provider({
children,
isLoading = false,
isLastMessage = true,
}: PropsWithChildren<{ isLoading?: boolean; isLastMessage?: boolean }>) {
return (
<MessageProvider
messageInfo={{ id: 'message-1', isLoading, isLastMessage, state: 'idle' }}
messageActions={{
onDelete: () => {},
onEdit: () => {},
onBranch: () => {},
onCancelEdit: () => {},
}}
>
{children}
</MessageProvider>
)
}
describe('MessagePartSwitcher', () => {
it('keeps consecutive generic tool parts as direct siblings', () => {
const reasoningPart = {
@@ -22,10 +44,10 @@ describe('MessagePartSwitcher', () => {
} satisfies ToolUIPart
const { container } = customRender(
<>
<Provider>
<MessagePartSwitcher part={reasoningPart} />
<MessagePartSwitcher part={toolPart} />
</>
</Provider>
)
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(
<Provider isLoading>
<MessagePartSwitcher part={part} />
</Provider>
)
expect(container.querySelector('.animate-spin')).not.toBeNull()
rerender(
<Provider>
<MessagePartSwitcher part={part} />
</Provider>
)
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(
<Provider isLoading>
<MessagePartSwitcher part={part} />
</Provider>
)
expect(queryByText('Response interrupted')).toBeNull()
rerender(
<Provider>
<MessagePartSwitcher part={part} />
</Provider>
)
expect(getByText('Response interrupted')).toBeInTheDocument()
}
)
it('does not restart an interrupted indicator when another message is streaming', () => {
const { container, getByText } = customRender(
<Provider isLoading isLastMessage={false}>
<MessagePartSwitcher part={{ type: 'reasoning', state: 'streaming', text: '' }} />
</Provider>
)
expect(getByText('Response interrupted')).toBeInTheDocument()
expect(container.querySelector('.animate-spin')).toBeNull()
})
})
@@ -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<VercelMessage['parts']>[number]) =>
export const MessagePartSwitcher = memo(
function MessagePartSwitcher({ part }: { part: NonNullable<VercelMessage['parts']>[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 (
<Tool
icon={<CircleStop strokeWidth={1.5} size={12} className="text-foreground-muted" />}
label="Response interrupted"
>
{part.type === 'reasoning' ? part.text : undefined}
</Tool>
)
}
const content = (() => {
switch (part.type) {
case 'dynamic-tool': {
@@ -95,6 +95,7 @@ function FeedMessage({
id={message.id}
message={message}
isLoading={isLoading}
isLastMessage
isAfterEditedMessage={false}
isBeingEdited={false}
addToolApprovalResponse={addToolApprovalResponse}
@@ -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)
})
})
@@ -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<typeof assistantMessageMetadataSchema>
/** 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.
*
+5
View File
@@ -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.'
@@ -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<string, any>
providerOptions?: Record<string, any>
abortSignal?: AbortSignal
timeout?: TimeoutConfiguration<ToolSet>
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<string, unknown>) => {
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<string, unknown> = {
@@ -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<typeof ai.streamText>[0])
+6 -6
View File
@@ -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
)
: []
@@ -7,6 +7,8 @@ vi.mock('common', () => ({
IS_PLATFORM: true,
}))
const executeOptions = { toolCallId: 'test', messages: [], context: {} }
describe('ai/tools/incident-tools', () => {
let mockFetch: ReturnType<typeof vi.fn>
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)
})
})
})
})
+4 -2
View File
@@ -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) {
+21 -20
View File
@@ -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 }
},
+5 -5
View File
@@ -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)
+2 -2
View File
@@ -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
)
+2 -2
View File
@@ -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
+93 -5
View File
@@ -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<typeof streamText>
)
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<typeof streamText>
)
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))
})
+19 -4
View File
@@ -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) {