feat: add reasoning model (#750)
Co-authored-by: Matt Apperson <me@mattapperson.com>
This commit is contained in:
parent
76804269c4
commit
c61d4f91d4
19 changed files with 342 additions and 209 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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!';
|
||||
},
|
||||
});
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
/>
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
/>
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue