feat: only send user message as part of request (#957)

This commit is contained in:
Jeremy 2025-04-26 01:09:01 -07:00 committed by GitHub
parent 9279135355
commit f18af236a0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 90 additions and 54 deletions

View file

@ -1,14 +1,13 @@
'use server';
import { generateText, Message } from 'ai';
import { generateText, type UIMessage } from 'ai';
import { cookies } from 'next/headers';
import {
deleteMessagesByChatIdAfterTimestamp,
getMessageById,
updateChatVisiblityById,
} from '@/lib/db/queries';
import { VisibilityType } from '@/components/visibility-selector';
import type { VisibilityType } from '@/components/visibility-selector';
import { myProvider } from '@/lib/ai/providers';
export async function saveChatModelAsCookie(model: string) {
@ -19,7 +18,7 @@ export async function saveChatModelAsCookie(model: string) {
export async function generateTitleFromUserMessage({
message,
}: {
message: Message;
message: UIMessage;
}) {
const { text: title } = await generateText({
model: myProvider.languageModel('title-model'),

View file

@ -1,5 +1,5 @@
import {
type UIMessage,
appendClientMessage,
appendResponseMessages,
createDataStreamResponse,
smoothStream,
@ -11,14 +11,11 @@ import {
deleteChatById,
getChatById,
getMessageCountByUserId,
getMessagesByChatId,
saveChat,
saveMessages,
} from '@/lib/db/queries';
import {
generateUUID,
getMostRecentUserMessage,
getTrailingMessageId,
} from '@/lib/utils';
import { generateUUID, getTrailingMessageId } from '@/lib/utils';
import { generateTitleFromUserMessage } from '../../actions';
import { createDocument } from '@/lib/ai/tools/create-document';
import { updateDocument } from '@/lib/ai/tools/update-document';
@ -27,24 +24,26 @@ import { getWeather } from '@/lib/ai/tools/get-weather';
import { isProductionEnvironment } from '@/lib/constants';
import { myProvider } from '@/lib/ai/providers';
import { entitlementsByUserType } from '@/lib/ai/entitlements';
import { postRequestBodySchema, type PostRequestBody } from './schema';
export const maxDuration = 60;
export async function POST(request: Request) {
let requestBody: PostRequestBody;
try {
const {
id,
messages,
selectedChatModel,
}: {
id: string;
messages: Array<UIMessage>;
selectedChatModel: string;
} = await request.json();
const json = await request.json();
requestBody = postRequestBodySchema.parse(json);
} catch (_) {
return new Response('Invalid request body', { status: 400 });
}
try {
const { id, message, selectedChatModel } = requestBody;
const session = await auth();
if (!session?.user?.id) {
if (!session?.user) {
return new Response('Unauthorized', { status: 401 });
}
@ -64,17 +63,11 @@ export async function POST(request: Request) {
);
}
const userMessage = getMostRecentUserMessage(messages);
if (!userMessage) {
return new Response('No user message found', { status: 400 });
}
const chat = await getChatById({ id });
if (!chat) {
const title = await generateTitleFromUserMessage({
message: userMessage,
message,
});
await saveChat({ id, userId: session.user.id, title });
@ -84,14 +77,22 @@ 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,
});
await saveMessages({
messages: [
{
chatId: id,
id: userMessage.id,
id: message.id,
role: 'user',
parts: userMessage.parts,
attachments: userMessage.experimental_attachments ?? [],
parts: message.parts,
attachments: message.experimental_attachments ?? [],
createdAt: new Date(),
},
],
@ -138,7 +139,7 @@ export async function POST(request: Request) {
}
const [, assistantMessage] = appendResponseMessages({
messages: [userMessage],
messages: [message],
responseMessages: response.messages,
});
@ -176,7 +177,7 @@ export async function POST(request: Request) {
return 'Oops, an error occurred!';
},
});
} catch (error) {
} catch (_) {
return new Response('An error occurred while processing your request!', {
status: 500,
});

View file

@ -0,0 +1,29 @@
import { z } from 'zod';
const textPartSchema = z.object({
text: z.string().min(1).max(2000),
type: z.enum(['text']),
});
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(),
}),
selectedChatModel: z.enum(['chat-model', 'chat-model-reasoning']),
});
export type PostRequestBody = z.infer<typeof postRequestBodySchema>;

View file

@ -46,11 +46,15 @@ export function Chat({
reload,
} = useChat({
id,
body: { id, selectedChatModel: selectedChatModel },
initialMessages,
experimental_throttle: 100,
sendExtraMessageFields: true,
generateId: generateUUID,
experimental_prepareRequestBody: (body) => ({
id,
message: body.messages.at(-1),
selectedChatModel,
}),
onFinish: () => {
mutate(unstable_serialize(getChatHistoryPaginationKey));
},

View file

@ -23,7 +23,7 @@ export const myProvider = isTestEnvironment
})
: customProvider({
languageModels: {
'chat-model': xai('grok-2-1212'),
'chat-model': xai('grok-2-vision-1212'),
'chat-model-reasoning': wrapLanguageModel({
model: xai('grok-3-mini-beta'),
middleware: extractReasoningMiddleware({ tagName: 'think' }),

View file

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

View file

@ -1,12 +1,14 @@
import { generateUUID } from '@/lib/utils';
export const TEST_PROMPTS = {
SKY: {
MESSAGES: [
{
role: 'user',
content: 'Why is the sky blue?',
parts: [{ type: 'text', text: 'Why is the sky blue?' }],
},
],
MESSAGE: {
id: generateUUID(),
createdAt: new Date().toISOString(),
role: 'user',
content: 'Why is the sky blue?',
parts: [{ type: 'text', text: 'Why is the sky blue?' }],
},
OUTPUT_STREAM: [
'0:"It\'s "',
'0:"just "',
@ -17,13 +19,14 @@ export const TEST_PROMPTS = {
],
},
GRASS: {
MESSAGES: [
{
role: 'user',
content: 'Why is grass green?',
parts: [{ type: 'text', text: 'Why is grass green?' }],
},
],
MESSAGE: {
id: generateUUID(),
createdAt: new Date().toISOString(),
role: 'user',
content: 'Why is grass green?',
parts: [{ type: 'text', text: 'Why is grass green?' }],
},
OUTPUT_STREAM: [
'0:"It\'s "',
'0:"just "',

View file

@ -10,12 +10,12 @@ test.describe
adaContext,
}) => {
const response = await adaContext.request.post('/api/chat', {
data: {},
data: JSON.stringify({}),
});
expect(response.status()).toBe(500);
expect(response.status()).toBe(400);
const text = await response.text();
expect(text).toEqual('An error occurred while processing your request!');
expect(text).toEqual('Invalid request body');
});
test('Ada can invoke chat generation', async ({ adaContext }) => {
@ -24,7 +24,7 @@ test.describe
const response = await adaContext.request.post('/api/chat', {
data: {
id: chatId,
messages: TEST_PROMPTS.SKY.MESSAGES,
message: TEST_PROMPTS.SKY.MESSAGE,
selectedChatModel: 'chat-model',
},
});
@ -47,7 +47,7 @@ test.describe
const response = await babbageContext.request.post('/api/chat', {
data: {
id: chatId,
messages: TEST_PROMPTS.GRASS.MESSAGES,
message: TEST_PROMPTS.GRASS.MESSAGE,
selectedChatModel: 'chat-model',
},
});