feat: only send user message as part of request (#957)
This commit is contained in:
parent
9279135355
commit
f18af236a0
8 changed files with 90 additions and 54 deletions
|
|
@ -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'),
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
});
|
});
|
||||||
|
|
|
||||||
29
app/(chat)/api/chat/schema.ts
Normal file
29
app/(chat)/api/chat/schema.ts
Normal 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>;
|
||||||
|
|
@ -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));
|
||||||
},
|
},
|
||||||
|
|
|
||||||
|
|
@ -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' }),
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
|
|
|
||||||
|
|
@ -1,12 +1,14 @@
|
||||||
|
import { generateUUID } from '@/lib/utils';
|
||||||
|
|
||||||
export const TEST_PROMPTS = {
|
export const TEST_PROMPTS = {
|
||||||
SKY: {
|
SKY: {
|
||||||
MESSAGES: [
|
MESSAGE: {
|
||||||
{
|
id: generateUUID(),
|
||||||
role: 'user',
|
createdAt: new Date().toISOString(),
|
||||||
content: 'Why is the sky blue?',
|
role: 'user',
|
||||||
parts: [{ type: 'text', text: 'Why is the sky blue?' }],
|
content: '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(),
|
||||||
role: 'user',
|
createdAt: new Date().toISOString(),
|
||||||
content: 'Why is grass green?',
|
role: 'user',
|
||||||
parts: [{ type: 'text', text: 'Why is grass green?' }],
|
content: '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 "',
|
||||||
|
|
|
||||||
|
|
@ -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',
|
||||||
},
|
},
|
||||||
});
|
});
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue