fix: race condition after resumable stream ends (#986)

This commit is contained in:
Jeremy 2025-05-07 16:02:53 -07:00 committed by GitHub
parent 566b01f367
commit 75af1320f4
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 145 additions and 18 deletions

View file

@ -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) {

View file

@ -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<Array<Attachment>>([]);
const isArtifactVisible = useArtifactSelector((state) => state.isVisible);
useAutoResume({
autoResume,
initialMessages,
experimental_resume,
data,
setMessages,
});
return (
<>
<div className="flex flex-col min-w-0 h-dvh bg-background">

47
hooks/use-auto-resume.ts Normal file
View file

@ -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]);
}

View file

@ -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);

1
lib/types.ts Normal file
View file

@ -0,0 +1 @@
export type DataPart = { type: 'append-message'; message: string };

View file

@ -1,6 +1,6 @@
{
"name": "ai-chatbot",
"version": "3.0.19",
"version": "3.0.20",
"private": true,
"scripts": {
"dev": "next dev --turbo",

View file

@ -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}`,