🎄 merry christmas: ai sdk v6 beta + tool approval (#1361)

This commit is contained in:
josh 2025-12-19 23:24:24 +00:00 committed by GitHub
parent 6e5b883cf2
commit 4d3ba8d9fe
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
21 changed files with 429 additions and 202 deletions

View file

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