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

View file

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

View file

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

View file

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

View file

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

View file

@ -10,12 +10,12 @@ test.describe
adaContext, adaContext,
}) => { }) => {
const response = await adaContext.request.post('/api/chat', { 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(); 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 }) => { test('Ada can invoke chat generation', async ({ adaContext }) => {
@ -24,7 +24,7 @@ test.describe
const response = await adaContext.request.post('/api/chat', { const response = await adaContext.request.post('/api/chat', {
data: { data: {
id: chatId, id: chatId,
messages: TEST_PROMPTS.SKY.MESSAGES, message: TEST_PROMPTS.SKY.MESSAGE,
selectedChatModel: 'chat-model', selectedChatModel: 'chat-model',
}, },
}); });
@ -47,7 +47,7 @@ test.describe
const response = await babbageContext.request.post('/api/chat', { const response = await babbageContext.request.post('/api/chat', {
data: { data: {
id: chatId, id: chatId,
messages: TEST_PROMPTS.GRASS.MESSAGES, message: TEST_PROMPTS.GRASS.MESSAGE,
selectedChatModel: 'chat-model', selectedChatModel: 'chat-model',
}, },
}); });