feat: add reasoning model (#750)

Co-authored-by: Matt Apperson <me@mattapperson.com>
This commit is contained in:
Jeremy 2025-02-03 20:33:15 +05:30 committed by GitHub
parent 76804269c4
commit c61d4f91d4
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
19 changed files with 342 additions and 209 deletions

View file

@ -1,19 +1,19 @@
'use server';
import { type CoreUserMessage, generateText, Message } from 'ai';
import { generateText, Message } from 'ai';
import { cookies } from 'next/headers';
import { customModel } from '@/lib/ai';
import {
deleteMessagesByChatIdAfterTimestamp,
getMessageById,
updateChatVisiblityById,
} from '@/lib/db/queries';
import { VisibilityType } from '@/components/visibility-selector';
import { myProvider } from '@/lib/ai/models';
export async function saveModelId(model: string) {
export async function saveChatModelAsCookie(model: string) {
const cookieStore = await cookies();
cookieStore.set('model-id', model);
cookieStore.set('chat-model', model);
}
export async function generateTitleFromUserMessage({
@ -22,7 +22,7 @@ export async function generateTitleFromUserMessage({
message: Message;
}) {
const { text: title } = await generateText({
model: customModel('gpt-4o-mini'),
model: myProvider.languageModel('title-model'),
system: `\n
- you will generate a short title based on the first message a user begins a conversation with
- ensure it is not more than 80 characters long

View file

@ -3,11 +3,11 @@ import {
createDataStreamResponse,
smoothStream,
streamText,
wrapLanguageModel,
} from 'ai';
import { auth } from '@/app/(auth)/auth';
import { customModel } from '@/lib/ai';
import { models } from '@/lib/ai/models';
import { myProvider } from '@/lib/ai/models';
import { systemPrompt } from '@/lib/ai/prompts';
import {
deleteChatById,
@ -48,8 +48,8 @@ export async function POST(request: Request) {
const {
id,
messages,
modelId,
}: { id: string; messages: Array<Message>; modelId: string } =
selectedChatModel,
}: { id: string; messages: Array<Message>; selectedChatModel: string } =
await request.json();
const session = await auth();
@ -58,12 +58,6 @@ export async function POST(request: Request) {
return new Response('Unauthorized', { status: 401 });
}
const model = models.find((model) => model.id === modelId);
if (!model) {
return new Response('Model not found', { status: 404 });
}
const userMessage = getMostRecentUserMessage(messages);
if (!userMessage) {
@ -84,7 +78,7 @@ export async function POST(request: Request) {
return createDataStreamResponse({
execute: (dataStream) => {
const result = streamText({
model: customModel(model.apiIdentifier),
model: myProvider.languageModel(selectedChatModel),
system: systemPrompt,
messages,
maxSteps: 5,
@ -93,32 +87,31 @@ export async function POST(request: Request) {
experimental_generateMessageId: generateUUID,
tools: {
getWeather,
createDocument: createDocument({ session, dataStream, model }),
updateDocument: updateDocument({ session, dataStream, model }),
createDocument: createDocument({ session, dataStream }),
updateDocument: updateDocument({ session, dataStream }),
requestSuggestions: requestSuggestions({
session,
dataStream,
model,
}),
},
onFinish: async ({ response }) => {
onFinish: async ({ response, reasoning }) => {
if (session.user?.id) {
try {
const responseMessagesWithoutIncompleteToolCalls =
sanitizeResponseMessages(response.messages);
const sanitizedResponseMessages = sanitizeResponseMessages({
messages: response.messages,
reasoning,
});
await saveMessages({
messages: responseMessagesWithoutIncompleteToolCalls.map(
(message) => {
return {
id: message.id,
chatId: id,
role: message.role,
content: message.content,
createdAt: new Date(),
};
},
),
messages: sanitizedResponseMessages.map((message) => {
return {
id: message.id,
chatId: id,
role: message.role,
content: message.content,
createdAt: new Date(),
};
}),
});
} catch (error) {
console.error('Failed to save chat');
@ -131,7 +124,12 @@ export async function POST(request: Request) {
},
});
result.mergeIntoDataStream(dataStream);
result.mergeIntoDataStream(dataStream, {
sendReasoning: true,
});
},
onError: (error) => {
return 'Oops, an error occured!';
},
});
}

View file

@ -3,10 +3,10 @@ import { notFound } from 'next/navigation';
import { auth } from '@/app/(auth)/auth';
import { Chat } from '@/components/chat';
import { DEFAULT_MODEL_NAME, models } from '@/lib/ai/models';
import { getChatById, getMessagesByChatId } from '@/lib/db/queries';
import { convertToUIMessages } from '@/lib/utils';
import { DataStreamHandler } from '@/components/data-stream-handler';
import { DEFAULT_CHAT_MODEL } from '@/lib/ai/models';
export default async function Page(props: { params: Promise<{ id: string }> }) {
const params = await props.params;
@ -34,17 +34,29 @@ export default async function Page(props: { params: Promise<{ id: string }> }) {
});
const cookieStore = await cookies();
const modelIdFromCookie = cookieStore.get('model-id')?.value;
const selectedModelId =
models.find((model) => model.id === modelIdFromCookie)?.id ||
DEFAULT_MODEL_NAME;
const chatModelFromCookie = cookieStore.get('chat-model');
if (!chatModelFromCookie) {
return (
<>
<Chat
id={chat.id}
initialMessages={convertToUIMessages(messagesFromDb)}
selectedChatModel={DEFAULT_CHAT_MODEL}
selectedVisibilityType={chat.visibility}
isReadonly={session?.user?.id !== chat.userId}
/>
<DataStreamHandler id={id} />
</>
);
}
return (
<>
<Chat
id={chat.id}
initialMessages={convertToUIMessages(messagesFromDb)}
selectedModelId={selectedModelId}
selectedChatModel={chatModelFromCookie.value}
selectedVisibilityType={chat.visibility}
isReadonly={session?.user?.id !== chat.userId}
/>

View file

@ -1,7 +1,7 @@
import { cookies } from 'next/headers';
import { Chat } from '@/components/chat';
import { DEFAULT_MODEL_NAME, models } from '@/lib/ai/models';
import { DEFAULT_CHAT_MODEL } from '@/lib/ai/models';
import { generateUUID } from '@/lib/utils';
import { DataStreamHandler } from '@/components/data-stream-handler';
@ -9,11 +9,23 @@ export default async function Page() {
const id = generateUUID();
const cookieStore = await cookies();
const modelIdFromCookie = cookieStore.get('model-id')?.value;
const modelIdFromCookie = cookieStore.get('chat-model');
const selectedModelId =
models.find((model) => model.id === modelIdFromCookie)?.id ||
DEFAULT_MODEL_NAME;
if (!modelIdFromCookie) {
return (
<>
<Chat
key={id}
id={id}
initialMessages={[]}
selectedChatModel={DEFAULT_CHAT_MODEL}
selectedVisibilityType="private"
isReadonly={false}
/>
<DataStreamHandler id={id} />
</>
);
}
return (
<>
@ -21,7 +33,7 @@ export default async function Page() {
key={id}
id={id}
initialMessages={[]}
selectedModelId={selectedModelId}
selectedChatModel={modelIdFromCookie.value}
selectedVisibilityType="private"
isReadonly={false}
/>