chore: update to ai sdk v5 beta (#1074)
This commit is contained in:
parent
7d8e71383f
commit
4c281fe09d
54 changed files with 1372 additions and 1060 deletions
112
app/(chat)/api/chat/[id]/stream/route.ts
Normal file
112
app/(chat)/api/chat/[id]/stream/route.ts
Normal file
|
|
@ -0,0 +1,112 @@
|
|||
import { auth } from '@/app/(auth)/auth';
|
||||
import {
|
||||
getChatById,
|
||||
getMessagesByChatId,
|
||||
getStreamIdsByChatId,
|
||||
} from '@/lib/db/queries';
|
||||
import type { Chat } from '@/lib/db/schema';
|
||||
import { ChatSDKError } from '@/lib/errors';
|
||||
import type { ChatMessage } from '@/lib/types';
|
||||
import { createUIMessageStream, JsonToSseTransformStream } from 'ai';
|
||||
import { getStreamContext } from '../../route';
|
||||
import { differenceInSeconds } from 'date-fns';
|
||||
|
||||
export async function GET(
|
||||
_: Request,
|
||||
{ params }: { params: Promise<{ id: string }> },
|
||||
) {
|
||||
const { id: chatId } = await params;
|
||||
|
||||
const streamContext = getStreamContext();
|
||||
const resumeRequestedAt = new Date();
|
||||
|
||||
if (!streamContext) {
|
||||
return new Response(null, { status: 204 });
|
||||
}
|
||||
|
||||
if (!chatId) {
|
||||
return new ChatSDKError('bad_request:api').toResponse();
|
||||
}
|
||||
|
||||
const session = await auth();
|
||||
|
||||
if (!session?.user) {
|
||||
return new ChatSDKError('unauthorized:chat').toResponse();
|
||||
}
|
||||
|
||||
let chat: Chat;
|
||||
|
||||
try {
|
||||
chat = await getChatById({ id: chatId });
|
||||
} catch {
|
||||
return new ChatSDKError('not_found:chat').toResponse();
|
||||
}
|
||||
|
||||
if (!chat) {
|
||||
return new ChatSDKError('not_found:chat').toResponse();
|
||||
}
|
||||
|
||||
if (chat.visibility === 'private' && chat.userId !== session.user.id) {
|
||||
return new ChatSDKError('forbidden:chat').toResponse();
|
||||
}
|
||||
|
||||
const streamIds = await getStreamIdsByChatId({ chatId });
|
||||
|
||||
if (!streamIds.length) {
|
||||
return new ChatSDKError('not_found:stream').toResponse();
|
||||
}
|
||||
|
||||
const recentStreamId = streamIds.at(-1);
|
||||
|
||||
if (!recentStreamId) {
|
||||
return new ChatSDKError('not_found:stream').toResponse();
|
||||
}
|
||||
|
||||
const emptyDataStream = createUIMessageStream<ChatMessage>({
|
||||
execute: () => {},
|
||||
});
|
||||
|
||||
const stream = await streamContext.resumableStream(recentStreamId, () =>
|
||||
emptyDataStream.pipeThrough(new JsonToSseTransformStream()),
|
||||
);
|
||||
|
||||
/*
|
||||
* 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 = createUIMessageStream<ChatMessage>({
|
||||
execute: ({ writer }) => {
|
||||
writer.write({
|
||||
type: 'data-appendMessage',
|
||||
data: JSON.stringify(mostRecentMessage),
|
||||
transient: true,
|
||||
});
|
||||
},
|
||||
});
|
||||
|
||||
return new Response(
|
||||
restoredStream.pipeThrough(new JsonToSseTransformStream()),
|
||||
{ status: 200 },
|
||||
);
|
||||
}
|
||||
|
||||
return new Response(stream, { status: 200 });
|
||||
}
|
||||
|
|
@ -1,8 +1,9 @@
|
|||
import {
|
||||
appendClientMessage,
|
||||
appendResponseMessages,
|
||||
createDataStream,
|
||||
convertToModelMessages,
|
||||
createUIMessageStream,
|
||||
JsonToSseTransformStream,
|
||||
smoothStream,
|
||||
stepCountIs,
|
||||
streamText,
|
||||
} from 'ai';
|
||||
import { auth, type UserType } from '@/app/(auth)/auth';
|
||||
|
|
@ -13,11 +14,10 @@ import {
|
|||
getChatById,
|
||||
getMessageCountByUserId,
|
||||
getMessagesByChatId,
|
||||
getStreamIdsByChatId,
|
||||
saveChat,
|
||||
saveMessages,
|
||||
} from '@/lib/db/queries';
|
||||
import { generateUUID, getTrailingMessageId } from '@/lib/utils';
|
||||
import { convertToUIMessages, generateUUID } from '@/lib/utils';
|
||||
import { generateTitleFromUserMessage } from '../../actions';
|
||||
import { createDocument } from '@/lib/ai/tools/create-document';
|
||||
import { updateDocument } from '@/lib/ai/tools/update-document';
|
||||
|
|
@ -33,15 +33,16 @@ import {
|
|||
type ResumableStreamContext,
|
||||
} from 'resumable-stream';
|
||||
import { after } from 'next/server';
|
||||
import type { Chat } from '@/lib/db/schema';
|
||||
import { differenceInSeconds } from 'date-fns';
|
||||
import { ChatSDKError } from '@/lib/errors';
|
||||
import type { ChatMessage } from '@/lib/types';
|
||||
import type { ChatModel } from '@/lib/ai/models';
|
||||
import type { VisibilityType } from '@/components/visibility-selector';
|
||||
|
||||
export const maxDuration = 60;
|
||||
|
||||
let globalStreamContext: ResumableStreamContext | null = null;
|
||||
|
||||
function getStreamContext() {
|
||||
export function getStreamContext() {
|
||||
if (!globalStreamContext) {
|
||||
try {
|
||||
globalStreamContext = createResumableStreamContext({
|
||||
|
|
@ -72,8 +73,17 @@ export async function POST(request: Request) {
|
|||
}
|
||||
|
||||
try {
|
||||
const { id, message, selectedChatModel, selectedVisibilityType } =
|
||||
requestBody;
|
||||
const {
|
||||
id,
|
||||
message,
|
||||
selectedChatModel,
|
||||
selectedVisibilityType,
|
||||
}: {
|
||||
id: string;
|
||||
message: ChatMessage;
|
||||
selectedChatModel: ChatModel['id'];
|
||||
selectedVisibilityType: VisibilityType;
|
||||
} = requestBody;
|
||||
|
||||
const session = await auth();
|
||||
|
||||
|
|
@ -111,13 +121,8 @@ export async function POST(request: Request) {
|
|||
}
|
||||
}
|
||||
|
||||
const previousMessages = await getMessagesByChatId({ id });
|
||||
|
||||
const messages = appendClientMessage({
|
||||
// @ts-expect-error: todo add type conversion from DBMessage[] to UIMessage[]
|
||||
messages: previousMessages,
|
||||
message,
|
||||
});
|
||||
const messagesFromDb = await getMessagesByChatId({ id });
|
||||
const uiMessages = [message, ...convertToUIMessages(messagesFromDb)];
|
||||
|
||||
const { longitude, latitude, city, country } = geolocation(request);
|
||||
|
||||
|
|
@ -135,7 +140,7 @@ export async function POST(request: Request) {
|
|||
id: message.id,
|
||||
role: 'user',
|
||||
parts: message.parts,
|
||||
attachments: message.experimental_attachments ?? [],
|
||||
attachments: [],
|
||||
createdAt: new Date(),
|
||||
},
|
||||
],
|
||||
|
|
@ -144,13 +149,13 @@ export async function POST(request: Request) {
|
|||
const streamId = generateUUID();
|
||||
await createStreamId({ streamId, chatId: id });
|
||||
|
||||
const stream = createDataStream({
|
||||
execute: (dataStream) => {
|
||||
const stream = createUIMessageStream({
|
||||
execute: ({ writer: dataStream }) => {
|
||||
const result = streamText({
|
||||
model: myProvider.languageModel(selectedChatModel),
|
||||
system: systemPrompt({ selectedChatModel, requestHints }),
|
||||
messages,
|
||||
maxSteps: 5,
|
||||
messages: convertToModelMessages(uiMessages),
|
||||
stopWhen: stepCountIs(5),
|
||||
experimental_activeTools:
|
||||
selectedChatModel === 'chat-model-reasoning'
|
||||
? []
|
||||
|
|
@ -161,7 +166,6 @@ export async function POST(request: Request) {
|
|||
'requestSuggestions',
|
||||
],
|
||||
experimental_transform: smoothStream({ chunking: 'word' }),
|
||||
experimental_generateMessageId: generateUUID,
|
||||
tools: {
|
||||
getWeather,
|
||||
createDocument: createDocument({ session, dataStream }),
|
||||
|
|
@ -171,42 +175,6 @@ export async function POST(request: Request) {
|
|||
dataStream,
|
||||
}),
|
||||
},
|
||||
onFinish: async ({ response }) => {
|
||||
if (session.user?.id) {
|
||||
try {
|
||||
const assistantId = getTrailingMessageId({
|
||||
messages: response.messages.filter(
|
||||
(message) => message.role === 'assistant',
|
||||
),
|
||||
});
|
||||
|
||||
if (!assistantId) {
|
||||
throw new Error('No assistant message found!');
|
||||
}
|
||||
|
||||
const [, assistantMessage] = appendResponseMessages({
|
||||
messages: [message],
|
||||
responseMessages: response.messages,
|
||||
});
|
||||
|
||||
await saveMessages({
|
||||
messages: [
|
||||
{
|
||||
id: assistantId,
|
||||
chatId: id,
|
||||
role: assistantMessage.role,
|
||||
parts: assistantMessage.parts,
|
||||
attachments:
|
||||
assistantMessage.experimental_attachments ?? [],
|
||||
createdAt: new Date(),
|
||||
},
|
||||
],
|
||||
});
|
||||
} catch (_) {
|
||||
console.error('Failed to save chat');
|
||||
}
|
||||
}
|
||||
},
|
||||
experimental_telemetry: {
|
||||
isEnabled: isProductionEnvironment,
|
||||
functionId: 'stream-text',
|
||||
|
|
@ -215,11 +183,27 @@ export async function POST(request: Request) {
|
|||
|
||||
result.consumeStream();
|
||||
|
||||
result.mergeIntoDataStream(dataStream, {
|
||||
sendReasoning: true,
|
||||
dataStream.merge(
|
||||
result.toUIMessageStream({
|
||||
sendReasoning: true,
|
||||
}),
|
||||
);
|
||||
},
|
||||
generateId: generateUUID,
|
||||
onFinish: async ({ messages }) => {
|
||||
await saveMessages({
|
||||
messages: messages.map((message) => ({
|
||||
id: message.id,
|
||||
role: message.role,
|
||||
parts: message.parts,
|
||||
createdAt: new Date(),
|
||||
attachments: [],
|
||||
chatId: id,
|
||||
})),
|
||||
});
|
||||
},
|
||||
onError: () => {
|
||||
onError: (error) => {
|
||||
console.log(error);
|
||||
return 'Oops, an error occurred!';
|
||||
},
|
||||
});
|
||||
|
|
@ -228,7 +212,9 @@ export async function POST(request: Request) {
|
|||
|
||||
if (streamContext) {
|
||||
return new Response(
|
||||
await streamContext.resumableStream(streamId, () => stream),
|
||||
await streamContext.resumableStream(streamId, () =>
|
||||
stream.pipeThrough(new JsonToSseTransformStream()),
|
||||
),
|
||||
);
|
||||
} else {
|
||||
return new Response(stream);
|
||||
|
|
@ -240,101 +226,6 @@ 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 });
|
||||
}
|
||||
|
||||
const { searchParams } = new URL(request.url);
|
||||
const chatId = searchParams.get('chatId');
|
||||
|
||||
if (!chatId) {
|
||||
return new ChatSDKError('bad_request:api').toResponse();
|
||||
}
|
||||
|
||||
const session = await auth();
|
||||
|
||||
if (!session?.user) {
|
||||
return new ChatSDKError('unauthorized:chat').toResponse();
|
||||
}
|
||||
|
||||
let chat: Chat;
|
||||
|
||||
try {
|
||||
chat = await getChatById({ id: chatId });
|
||||
} catch {
|
||||
return new ChatSDKError('not_found:chat').toResponse();
|
||||
}
|
||||
|
||||
if (!chat) {
|
||||
return new ChatSDKError('not_found:chat').toResponse();
|
||||
}
|
||||
|
||||
if (chat.visibility === 'private' && chat.userId !== session.user.id) {
|
||||
return new ChatSDKError('forbidden:chat').toResponse();
|
||||
}
|
||||
|
||||
const streamIds = await getStreamIdsByChatId({ chatId });
|
||||
|
||||
if (!streamIds.length) {
|
||||
return new ChatSDKError('not_found:stream').toResponse();
|
||||
}
|
||||
|
||||
const recentStreamId = streamIds.at(-1);
|
||||
|
||||
if (!recentStreamId) {
|
||||
return new ChatSDKError('not_found:stream').toResponse();
|
||||
}
|
||||
|
||||
const emptyDataStream = createDataStream({
|
||||
execute: () => {},
|
||||
});
|
||||
|
||||
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) {
|
||||
const { searchParams } = new URL(request.url);
|
||||
const id = searchParams.get('id');
|
||||
|
|
|
|||
|
|
@ -1,27 +1,25 @@
|
|||
import { z } from 'zod';
|
||||
|
||||
const textPartSchema = z.object({
|
||||
text: z.string().min(1).max(2000),
|
||||
type: z.enum(['text']),
|
||||
text: z.string().min(1).max(2000),
|
||||
});
|
||||
|
||||
const filePartSchema = z.object({
|
||||
type: z.enum(['file']),
|
||||
mediaType: z.enum(['image/jpeg', 'image/png']),
|
||||
name: z.string().min(1).max(100),
|
||||
url: z.string().url(),
|
||||
});
|
||||
|
||||
const partSchema = z.union([textPartSchema, filePartSchema]);
|
||||
|
||||
export const postRequestBodySchema = z.object({
|
||||
id: z.string().uuid(),
|
||||
message: z.object({
|
||||
id: z.string().uuid(),
|
||||
createdAt: z.coerce.date(),
|
||||
role: z.enum(['user']),
|
||||
content: z.string().min(1).max(2000),
|
||||
parts: z.array(textPartSchema),
|
||||
experimental_attachments: z
|
||||
.array(
|
||||
z.object({
|
||||
url: z.string().url(),
|
||||
name: z.string().min(1).max(2000),
|
||||
contentType: z.enum(['image/png', 'image/jpg', 'image/jpeg']),
|
||||
}),
|
||||
)
|
||||
.optional(),
|
||||
parts: z.array(partSchema),
|
||||
}),
|
||||
selectedChatModel: z.enum(['chat-model', 'chat-model-reasoning']),
|
||||
selectedVisibilityType: z.enum(['public', 'private']),
|
||||
|
|
|
|||
|
|
@ -6,8 +6,7 @@ import { Chat } from '@/components/chat';
|
|||
import { getChatById, getMessagesByChatId } from '@/lib/db/queries';
|
||||
import { DataStreamHandler } from '@/components/data-stream-handler';
|
||||
import { DEFAULT_CHAT_MODEL } from '@/lib/ai/models';
|
||||
import type { DBMessage } from '@/lib/db/schema';
|
||||
import type { Attachment, UIMessage } from 'ai';
|
||||
import { convertToUIMessages } from '@/lib/utils';
|
||||
|
||||
export default async function Page(props: { params: Promise<{ id: string }> }) {
|
||||
const params = await props.params;
|
||||
|
|
@ -38,18 +37,7 @@ export default async function Page(props: { params: Promise<{ id: string }> }) {
|
|||
id,
|
||||
});
|
||||
|
||||
function convertToUIMessages(messages: Array<DBMessage>): Array<UIMessage> {
|
||||
return messages.map((message) => ({
|
||||
id: message.id,
|
||||
parts: message.parts as UIMessage['parts'],
|
||||
role: message.role as UIMessage['role'],
|
||||
// Note: content will soon be deprecated in @ai-sdk/react
|
||||
content: '',
|
||||
createdAt: message.createdAt,
|
||||
experimental_attachments:
|
||||
(message.attachments as Array<Attachment>) ?? [],
|
||||
}));
|
||||
}
|
||||
const uiMessages = convertToUIMessages(messagesFromDb);
|
||||
|
||||
const cookieStore = await cookies();
|
||||
const chatModelFromCookie = cookieStore.get('chat-model');
|
||||
|
|
@ -59,14 +47,14 @@ export default async function Page(props: { params: Promise<{ id: string }> }) {
|
|||
<>
|
||||
<Chat
|
||||
id={chat.id}
|
||||
initialMessages={convertToUIMessages(messagesFromDb)}
|
||||
initialMessages={uiMessages}
|
||||
initialChatModel={DEFAULT_CHAT_MODEL}
|
||||
initialVisibilityType={chat.visibility}
|
||||
isReadonly={session?.user?.id !== chat.userId}
|
||||
session={session}
|
||||
autoResume={true}
|
||||
/>
|
||||
<DataStreamHandler id={id} />
|
||||
<DataStreamHandler />
|
||||
</>
|
||||
);
|
||||
}
|
||||
|
|
@ -75,14 +63,14 @@ export default async function Page(props: { params: Promise<{ id: string }> }) {
|
|||
<>
|
||||
<Chat
|
||||
id={chat.id}
|
||||
initialMessages={convertToUIMessages(messagesFromDb)}
|
||||
initialMessages={uiMessages}
|
||||
initialChatModel={chatModelFromCookie.value}
|
||||
initialVisibilityType={chat.visibility}
|
||||
isReadonly={session?.user?.id !== chat.userId}
|
||||
session={session}
|
||||
autoResume={true}
|
||||
/>
|
||||
<DataStreamHandler id={id} />
|
||||
<DataStreamHandler />
|
||||
</>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import { AppSidebar } from '@/components/app-sidebar';
|
|||
import { SidebarInset, SidebarProvider } from '@/components/ui/sidebar';
|
||||
import { auth } from '../(auth)/auth';
|
||||
import Script from 'next/script';
|
||||
import { DataStreamProvider } from '@/components/data-stream-provider';
|
||||
|
||||
export const experimental_ppr = true;
|
||||
|
||||
|
|
@ -21,10 +22,12 @@ export default async function Layout({
|
|||
src="https://cdn.jsdelivr.net/pyodide/v0.23.4/full/pyodide.js"
|
||||
strategy="beforeInteractive"
|
||||
/>
|
||||
<SidebarProvider defaultOpen={!isCollapsed}>
|
||||
<AppSidebar user={session?.user} />
|
||||
<SidebarInset>{children}</SidebarInset>
|
||||
</SidebarProvider>
|
||||
<DataStreamProvider>
|
||||
<SidebarProvider defaultOpen={!isCollapsed}>
|
||||
<AppSidebar user={session?.user} />
|
||||
<SidebarInset>{children}</SidebarInset>
|
||||
</SidebarProvider>
|
||||
</DataStreamProvider>
|
||||
</>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -32,7 +32,7 @@ export default async function Page() {
|
|||
session={session}
|
||||
autoResume={false}
|
||||
/>
|
||||
<DataStreamHandler id={id} />
|
||||
<DataStreamHandler />
|
||||
</>
|
||||
);
|
||||
}
|
||||
|
|
@ -49,7 +49,7 @@ export default async function Page() {
|
|||
session={session}
|
||||
autoResume={false}
|
||||
/>
|
||||
<DataStreamHandler id={id} />
|
||||
<DataStreamHandler />
|
||||
</>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue