fix: race condition after resumable stream ends (#986)
This commit is contained in:
parent
566b01f367
commit
75af1320f4
7 changed files with 145 additions and 18 deletions
|
|
@ -34,6 +34,7 @@ import {
|
||||||
} from 'resumable-stream';
|
} from 'resumable-stream';
|
||||||
import { after } from 'next/server';
|
import { after } from 'next/server';
|
||||||
import type { Chat } from '@/lib/db/schema';
|
import type { Chat } from '@/lib/db/schema';
|
||||||
|
import { differenceInSeconds } from 'date-fns';
|
||||||
|
|
||||||
export const maxDuration = 60;
|
export const maxDuration = 60;
|
||||||
|
|
||||||
|
|
@ -245,6 +246,7 @@ export async function POST(request: Request) {
|
||||||
|
|
||||||
export async function GET(request: Request) {
|
export async function GET(request: Request) {
|
||||||
const streamContext = getStreamContext();
|
const streamContext = getStreamContext();
|
||||||
|
const resumeRequestedAt = new Date();
|
||||||
|
|
||||||
if (!streamContext) {
|
if (!streamContext) {
|
||||||
return new Response(null, { status: 204 });
|
return new Response(null, { status: 204 });
|
||||||
|
|
@ -295,12 +297,46 @@ export async function GET(request: Request) {
|
||||||
execute: () => {},
|
execute: () => {},
|
||||||
});
|
});
|
||||||
|
|
||||||
return new Response(
|
const stream = await streamContext.resumableStream(
|
||||||
await streamContext.resumableStream(recentStreamId, () => emptyDataStream),
|
recentStreamId,
|
||||||
{
|
() => emptyDataStream,
|
||||||
status: 200,
|
|
||||||
},
|
|
||||||
);
|
);
|
||||||
|
|
||||||
|
/*
|
||||||
|
* 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) {
|
export async function DELETE(request: Request) {
|
||||||
|
|
|
||||||
|
|
@ -18,6 +18,7 @@ import { toast } from './toast';
|
||||||
import type { Session } from 'next-auth';
|
import type { Session } from 'next-auth';
|
||||||
import { useSearchParams } from 'next/navigation';
|
import { useSearchParams } from 'next/navigation';
|
||||||
import { useChatVisibility } from '@/hooks/use-chat-visibility';
|
import { useChatVisibility } from '@/hooks/use-chat-visibility';
|
||||||
|
import { useAutoResume } from '@/hooks/use-auto-resume';
|
||||||
|
|
||||||
export function Chat({
|
export function Chat({
|
||||||
id,
|
id,
|
||||||
|
|
@ -54,6 +55,7 @@ export function Chat({
|
||||||
stop,
|
stop,
|
||||||
reload,
|
reload,
|
||||||
experimental_resume,
|
experimental_resume,
|
||||||
|
data,
|
||||||
} = useChat({
|
} = useChat({
|
||||||
id,
|
id,
|
||||||
initialMessages,
|
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 searchParams = useSearchParams();
|
||||||
const query = searchParams.get('query');
|
const query = searchParams.get('query');
|
||||||
|
|
||||||
|
|
@ -111,6 +104,14 @@ export function Chat({
|
||||||
const [attachments, setAttachments] = useState<Array<Attachment>>([]);
|
const [attachments, setAttachments] = useState<Array<Attachment>>([]);
|
||||||
const isArtifactVisible = useArtifactSelector((state) => state.isVisible);
|
const isArtifactVisible = useArtifactSelector((state) => state.isVisible);
|
||||||
|
|
||||||
|
useAutoResume({
|
||||||
|
autoResume,
|
||||||
|
initialMessages,
|
||||||
|
experimental_resume,
|
||||||
|
data,
|
||||||
|
setMessages,
|
||||||
|
});
|
||||||
|
|
||||||
return (
|
return (
|
||||||
<>
|
<>
|
||||||
<div className="flex flex-col min-w-0 h-dvh bg-background">
|
<div className="flex flex-col min-w-0 h-dvh bg-background">
|
||||||
|
|
|
||||||
47
hooks/use-auto-resume.ts
Normal file
47
hooks/use-auto-resume.ts
Normal 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]);
|
||||||
|
}
|
||||||
|
|
@ -500,7 +500,7 @@ export async function getStreamIdsByChatId({ chatId }: { chatId: string }) {
|
||||||
.select({ id: stream.id })
|
.select({ id: stream.id })
|
||||||
.from(stream)
|
.from(stream)
|
||||||
.where(eq(stream.chatId, chatId))
|
.where(eq(stream.chatId, chatId))
|
||||||
.orderBy(desc(stream.createdAt))
|
.orderBy(asc(stream.createdAt))
|
||||||
.execute();
|
.execute();
|
||||||
|
|
||||||
return streamIds.map(({ id }) => id);
|
return streamIds.map(({ id }) => id);
|
||||||
|
|
|
||||||
1
lib/types.ts
Normal file
1
lib/types.ts
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
export type DataPart = { type: 'append-message'; message: string };
|
||||||
|
|
@ -1,6 +1,6 @@
|
||||||
{
|
{
|
||||||
"name": "ai-chatbot",
|
"name": "ai-chatbot",
|
||||||
"version": "3.0.19",
|
"version": "3.0.20",
|
||||||
"private": true,
|
"private": true,
|
||||||
"scripts": {
|
"scripts": {
|
||||||
"dev": "next dev --turbo",
|
"dev": "next dev --turbo",
|
||||||
|
|
|
||||||
|
|
@ -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,
|
adaContext,
|
||||||
}) => {
|
}) => {
|
||||||
const chatId = generateUUID();
|
const chatId = generateUUID();
|
||||||
|
|
@ -191,6 +191,48 @@ test.describe
|
||||||
secondResponse.text(),
|
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('');
|
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(
|
const secondRequest = babbageContext.request.get(
|
||||||
`/api/chat?chatId=${chatId}`,
|
`/api/chat?chatId=${chatId}`,
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue