🎄 merry christmas: ai sdk v6 beta + tool approval (#1361)
This commit is contained in:
parent
6e5b883cf2
commit
4d3ba8d9fe
21 changed files with 429 additions and 202 deletions
|
|
@ -13,9 +13,7 @@ import {
|
|||
type ResumableStreamContext,
|
||||
} from "resumable-stream";
|
||||
import { auth, type UserType } from "@/app/(auth)/auth";
|
||||
import type { VisibilityType } from "@/components/visibility-selector";
|
||||
import { entitlementsByUserType } from "@/lib/ai/entitlements";
|
||||
import type { ChatModel } from "@/lib/ai/models";
|
||||
import { type RequestHints, systemPrompt } from "@/lib/ai/prompts";
|
||||
import { getLanguageModel } from "@/lib/ai/providers";
|
||||
import { createDocument } from "@/lib/ai/tools/create-document";
|
||||
|
|
@ -32,6 +30,7 @@ import {
|
|||
saveChat,
|
||||
saveMessages,
|
||||
updateChatTitleById,
|
||||
updateMessage,
|
||||
} from "@/lib/db/queries";
|
||||
import type { DBMessage } from "@/lib/db/schema";
|
||||
import { ChatSDKError } from "@/lib/errors";
|
||||
|
|
@ -75,17 +74,8 @@ export async function POST(request: Request) {
|
|||
}
|
||||
|
||||
try {
|
||||
const {
|
||||
id,
|
||||
message,
|
||||
selectedChatModel,
|
||||
selectedVisibilityType,
|
||||
}: {
|
||||
id: string;
|
||||
message: ChatMessage;
|
||||
selectedChatModel: ChatModel["id"];
|
||||
selectedVisibilityType: VisibilityType;
|
||||
} = requestBody;
|
||||
const { id, message, messages, selectedChatModel, selectedVisibilityType } =
|
||||
requestBody;
|
||||
|
||||
const session = await auth();
|
||||
|
||||
|
|
@ -104,6 +94,9 @@ export async function POST(request: Request) {
|
|||
return new ChatSDKError("rate_limit:chat").toResponse();
|
||||
}
|
||||
|
||||
// Check if this is a tool approval flow (all messages sent)
|
||||
const isToolApprovalFlow = Boolean(messages);
|
||||
|
||||
const chat = await getChatById({ id });
|
||||
let messagesFromDb: DBMessage[] = [];
|
||||
let titlePromise: Promise<string> | null = null;
|
||||
|
|
@ -112,9 +105,11 @@ export async function POST(request: Request) {
|
|||
if (chat.userId !== session.user.id) {
|
||||
return new ChatSDKError("forbidden:chat").toResponse();
|
||||
}
|
||||
// Only fetch messages if chat already exists
|
||||
messagesFromDb = await getMessagesByChatId({ id });
|
||||
} else {
|
||||
// Only fetch messages if chat already exists and not tool approval
|
||||
if (!isToolApprovalFlow) {
|
||||
messagesFromDb = await getMessagesByChatId({ id });
|
||||
}
|
||||
} else if (message?.role === "user") {
|
||||
// Save chat immediately with placeholder title
|
||||
await saveChat({
|
||||
id,
|
||||
|
|
@ -127,7 +122,10 @@ export async function POST(request: Request) {
|
|||
titlePromise = generateTitleFromUserMessage({ message });
|
||||
}
|
||||
|
||||
const uiMessages = [...convertToUIMessages(messagesFromDb), message];
|
||||
// Use all messages for tool approval, otherwise DB messages + new message
|
||||
const uiMessages = isToolApprovalFlow
|
||||
? (messages as ChatMessage[])
|
||||
: [...convertToUIMessages(messagesFromDb), message as ChatMessage];
|
||||
|
||||
const { longitude, latitude, city, country } = geolocation(request);
|
||||
|
||||
|
|
@ -138,24 +136,29 @@ export async function POST(request: Request) {
|
|||
country,
|
||||
};
|
||||
|
||||
await saveMessages({
|
||||
messages: [
|
||||
{
|
||||
chatId: id,
|
||||
id: message.id,
|
||||
role: "user",
|
||||
parts: message.parts,
|
||||
attachments: [],
|
||||
createdAt: new Date(),
|
||||
},
|
||||
],
|
||||
});
|
||||
// Only save user messages to the database (not tool approval responses)
|
||||
if (message?.role === "user") {
|
||||
await saveMessages({
|
||||
messages: [
|
||||
{
|
||||
chatId: id,
|
||||
id: message.id,
|
||||
role: "user",
|
||||
parts: message.parts,
|
||||
attachments: [],
|
||||
createdAt: new Date(),
|
||||
},
|
||||
],
|
||||
});
|
||||
}
|
||||
|
||||
const streamId = generateUUID();
|
||||
await createStreamId({ streamId, chatId: id });
|
||||
|
||||
const stream = createUIMessageStream({
|
||||
execute: ({ writer: dataStream }) => {
|
||||
// Pass original messages for tool approval continuation
|
||||
originalMessages: isToolApprovalFlow ? uiMessages : undefined,
|
||||
execute: async ({ writer: dataStream }) => {
|
||||
// Handle title generation in parallel
|
||||
if (titlePromise) {
|
||||
titlePromise.then((title) => {
|
||||
|
|
@ -171,7 +174,7 @@ export async function POST(request: Request) {
|
|||
const result = streamText({
|
||||
model: getLanguageModel(selectedChatModel),
|
||||
system: systemPrompt({ selectedChatModel, requestHints }),
|
||||
messages: convertToModelMessages(uiMessages),
|
||||
messages: await convertToModelMessages(uiMessages),
|
||||
stopWhen: stepCountIs(5),
|
||||
experimental_activeTools: isReasoningModel
|
||||
? []
|
||||
|
|
@ -215,32 +218,67 @@ export async function POST(request: Request) {
|
|||
);
|
||||
},
|
||||
generateId: generateUUID,
|
||||
onFinish: async ({ messages }) => {
|
||||
await saveMessages({
|
||||
messages: messages.map((currentMessage) => ({
|
||||
id: currentMessage.id,
|
||||
role: currentMessage.role,
|
||||
parts: currentMessage.parts,
|
||||
createdAt: new Date(),
|
||||
attachments: [],
|
||||
chatId: id,
|
||||
})),
|
||||
});
|
||||
onFinish: async ({ messages: finishedMessages }) => {
|
||||
if (isToolApprovalFlow) {
|
||||
// For tool approval, update existing messages (tool state changed) and save new ones
|
||||
for (const finishedMsg of finishedMessages) {
|
||||
const existingMsg = uiMessages.find((m) => m.id === finishedMsg.id);
|
||||
if (existingMsg) {
|
||||
// Update existing message with new parts (tool state changed)
|
||||
await updateMessage({
|
||||
id: finishedMsg.id,
|
||||
parts: finishedMsg.parts,
|
||||
});
|
||||
} else {
|
||||
// Save new message
|
||||
await saveMessages({
|
||||
messages: [
|
||||
{
|
||||
id: finishedMsg.id,
|
||||
role: finishedMsg.role,
|
||||
parts: finishedMsg.parts,
|
||||
createdAt: new Date(),
|
||||
attachments: [],
|
||||
chatId: id,
|
||||
},
|
||||
],
|
||||
});
|
||||
}
|
||||
}
|
||||
} else if (finishedMessages.length > 0) {
|
||||
// Normal flow - save all finished messages
|
||||
await saveMessages({
|
||||
messages: finishedMessages.map((currentMessage) => ({
|
||||
id: currentMessage.id,
|
||||
role: currentMessage.role,
|
||||
parts: currentMessage.parts,
|
||||
createdAt: new Date(),
|
||||
attachments: [],
|
||||
chatId: id,
|
||||
})),
|
||||
});
|
||||
}
|
||||
},
|
||||
onError: () => {
|
||||
return "Oops, an error occurred!";
|
||||
},
|
||||
});
|
||||
|
||||
// const streamContext = getStreamContext();
|
||||
const streamContext = getStreamContext();
|
||||
|
||||
// if (streamContext) {
|
||||
// return new Response(
|
||||
// await streamContext.resumableStream(streamId, () =>
|
||||
// stream.pipeThrough(new JsonToSseTransformStream())
|
||||
// )
|
||||
// );
|
||||
// }
|
||||
if (streamContext) {
|
||||
try {
|
||||
const resumableStream = await streamContext.resumableStream(
|
||||
streamId,
|
||||
() => stream.pipeThrough(new JsonToSseTransformStream())
|
||||
);
|
||||
if (resumableStream) {
|
||||
return new Response(resumableStream);
|
||||
}
|
||||
} catch (error) {
|
||||
console.error("Failed to create resumable stream:", error);
|
||||
}
|
||||
}
|
||||
|
||||
return new Response(stream.pipeThrough(new JsonToSseTransformStream()));
|
||||
} catch (error) {
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue