From 75af1320f4c5390252dd3fedaffedc2876db556f Mon Sep 17 00:00:00 2001 From: Jeremy Date: Wed, 7 May 2025 16:02:53 -0700 Subject: [PATCH] fix: race condition after resumable stream ends (#986) --- app/(chat)/api/chat/route.ts | 46 +++++++++++++++++++++++++++++++---- components/chat.tsx | 19 ++++++++------- hooks/use-auto-resume.ts | 47 ++++++++++++++++++++++++++++++++++++ lib/db/queries.ts | 2 +- lib/types.ts | 1 + package.json | 2 +- tests/routes/chat.test.ts | 46 +++++++++++++++++++++++++++++++++-- 7 files changed, 145 insertions(+), 18 deletions(-) create mode 100644 hooks/use-auto-resume.ts create mode 100644 lib/types.ts diff --git a/app/(chat)/api/chat/route.ts b/app/(chat)/api/chat/route.ts index 4fe241b..705ef2d 100644 --- a/app/(chat)/api/chat/route.ts +++ b/app/(chat)/api/chat/route.ts @@ -34,6 +34,7 @@ import { } from 'resumable-stream'; import { after } from 'next/server'; import type { Chat } from '@/lib/db/schema'; +import { differenceInSeconds } from 'date-fns'; export const maxDuration = 60; @@ -245,6 +246,7 @@ export async function POST(request: Request) { export async function GET(request: Request) { const streamContext = getStreamContext(); + const resumeRequestedAt = new Date(); if (!streamContext) { return new Response(null, { status: 204 }); @@ -295,12 +297,46 @@ export async function GET(request: Request) { execute: () => {}, }); - return new Response( - await streamContext.resumableStream(recentStreamId, () => emptyDataStream), - { - status: 200, - }, + const stream = await streamContext.resumableStream( + recentStreamId, + () => emptyDataStream, ); + + /* + * For when the generation is streaming during SSR + * but the resumable stream has concluded at this point. + */ + if (!stream) { + const messages = await getMessagesByChatId({ id: chatId }); + const mostRecentMessage = messages.at(-1); + + if (!mostRecentMessage) { + return new Response(emptyDataStream, { status: 200 }); + } + + if (mostRecentMessage.role !== 'assistant') { + return new Response(emptyDataStream, { status: 200 }); + } + + const messageCreatedAt = new Date(mostRecentMessage.createdAt); + + if (differenceInSeconds(resumeRequestedAt, messageCreatedAt) > 15) { + return new Response(emptyDataStream, { status: 200 }); + } + + const restoredStream = createDataStream({ + execute: (buffer) => { + buffer.writeData({ + type: 'append-message', + message: JSON.stringify(mostRecentMessage), + }); + }, + }); + + return new Response(restoredStream, { status: 200 }); + } + + return new Response(stream, { status: 200 }); } export async function DELETE(request: Request) { diff --git a/components/chat.tsx b/components/chat.tsx index b2c4e51..f4ad074 100644 --- a/components/chat.tsx +++ b/components/chat.tsx @@ -18,6 +18,7 @@ import { toast } from './toast'; import type { Session } from 'next-auth'; import { useSearchParams } from 'next/navigation'; import { useChatVisibility } from '@/hooks/use-chat-visibility'; +import { useAutoResume } from '@/hooks/use-auto-resume'; export function Chat({ id, @@ -54,6 +55,7 @@ export function Chat({ stop, reload, experimental_resume, + data, } = useChat({ id, initialMessages, @@ -77,15 +79,6 @@ export function Chat({ }, }); - useEffect(() => { - if (autoResume) { - experimental_resume(); - } - - // note: this hook has no dependencies since it only needs to run once - // eslint-disable-next-line react-hooks/exhaustive-deps - }, []); - const searchParams = useSearchParams(); const query = searchParams.get('query'); @@ -111,6 +104,14 @@ export function Chat({ const [attachments, setAttachments] = useState>([]); const isArtifactVisible = useArtifactSelector((state) => state.isVisible); + useAutoResume({ + autoResume, + initialMessages, + experimental_resume, + data, + setMessages, + }); + return ( <>
diff --git a/hooks/use-auto-resume.ts b/hooks/use-auto-resume.ts new file mode 100644 index 0000000..3b728df --- /dev/null +++ b/hooks/use-auto-resume.ts @@ -0,0 +1,47 @@ +'use client'; + +import { useEffect } from 'react'; +import type { UIMessage } from 'ai'; +import type { UseChatHelpers } from '@ai-sdk/react'; +import type { DataPart } from '@/lib/types'; + +export interface UseAutoResumeParams { + autoResume: boolean; + initialMessages: UIMessage[]; + experimental_resume: UseChatHelpers['experimental_resume']; + data: UseChatHelpers['data']; + setMessages: UseChatHelpers['setMessages']; +} + +export function useAutoResume({ + autoResume, + initialMessages, + experimental_resume, + data, + setMessages, +}: UseAutoResumeParams) { + useEffect(() => { + if (!autoResume) return; + + const mostRecentMessage = initialMessages.at(-1); + + if (mostRecentMessage?.role === 'user') { + experimental_resume(); + } + + // we intentionally run this once + // eslint-disable-next-line react-hooks/exhaustive-deps + }, []); + + useEffect(() => { + if (!data) return; + if (data.length === 0) return; + + const dataPart = data[0] as DataPart; + + if (dataPart.type === 'append-message') { + const message = JSON.parse(dataPart.message) as UIMessage; + setMessages([...initialMessages, message]); + } + }, [data, initialMessages, setMessages]); +} diff --git a/lib/db/queries.ts b/lib/db/queries.ts index e87e2f5..e1215db 100644 --- a/lib/db/queries.ts +++ b/lib/db/queries.ts @@ -500,7 +500,7 @@ export async function getStreamIdsByChatId({ chatId }: { chatId: string }) { .select({ id: stream.id }) .from(stream) .where(eq(stream.chatId, chatId)) - .orderBy(desc(stream.createdAt)) + .orderBy(asc(stream.createdAt)) .execute(); return streamIds.map(({ id }) => id); diff --git a/lib/types.ts b/lib/types.ts new file mode 100644 index 0000000..7376116 --- /dev/null +++ b/lib/types.ts @@ -0,0 +1 @@ +export type DataPart = { type: 'append-message'; message: string }; diff --git a/package.json b/package.json index 5bd9757..cba75ec 100644 --- a/package.json +++ b/package.json @@ -1,6 +1,6 @@ { "name": "ai-chatbot", - "version": "3.0.19", + "version": "3.0.20", "private": true, "scripts": { "dev": "next dev --turbo", diff --git a/tests/routes/chat.test.ts b/tests/routes/chat.test.ts index 16099a7..3d744da 100644 --- a/tests/routes/chat.test.ts +++ b/tests/routes/chat.test.ts @@ -144,7 +144,7 @@ test.describe ); }); - test('Ada cannot resume chat generation that has ended', async ({ + test('Ada can resume chat generation that has ended during request', async ({ adaContext, }) => { const chatId = generateUUID(); @@ -191,6 +191,48 @@ test.describe secondResponse.text(), ]); + expect(secondResponseContent).toContain('append-message'); + }); + + test('Ada cannot resume chat generation that has ended', async ({ + adaContext, + }) => { + const chatId = generateUUID(); + + const firstResponse = await adaContext.request.post('/api/chat', { + data: { + id: chatId, + message: { + id: generateUUID(), + role: 'user', + content: 'Help me write an essay about Silcon Valley', + parts: [ + { + type: 'text', + text: 'Help me write an essay about Silicon Valley', + }, + ], + createdAt: new Date().toISOString(), + }, + selectedChatModel: 'chat-model', + selectedVisibilityType: 'private', + }, + }); + + const firstStatusCode = firstResponse.status(); + expect(firstStatusCode).toBe(200); + + await firstResponse.text(); + await new Promise((resolve) => setTimeout(resolve, 15 * 1000)); + await new Promise((resolve) => setTimeout(resolve, 15000)); + const secondResponse = await adaContext.request.get( + `/api/chat?chatId=${chatId}`, + ); + + const secondStatusCode = secondResponse.status(); + expect(secondStatusCode).toBe(200); + + const secondResponseContent = await secondResponse.text(); expect(secondResponseContent).toEqual(''); }); @@ -266,7 +308,7 @@ test.describe }, }); - await new Promise((resolve) => setTimeout(resolve, 1000)); + await new Promise((resolve) => setTimeout(resolve, 10 * 1000)); const secondRequest = babbageContext.request.get( `/api/chat?chatId=${chatId}`,