Fix generating title from messages (#442)

This commit is contained in:
Jeremy 2024-10-15 10:50:50 +05:30 committed by GitHub
parent 2705d83d6c
commit 00b125378c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 103 additions and 77 deletions

View file

@ -1,84 +1,11 @@
import { CoreMessage, CoreToolMessage, Message, ToolInvocation } from "ai"; import { CoreMessage } from "ai";
import { notFound } from "next/navigation"; import { notFound } from "next/navigation";
import { auth } from "@/app/(auth)/auth"; import { auth } from "@/app/(auth)/auth";
import { Chat as PreviewChat } from "@/components/custom/chat"; import { Chat as PreviewChat } from "@/components/custom/chat";
import { getChatById } from "@/db/queries"; import { getChatById } from "@/db/queries";
import { Chat } from "@/db/schema"; import { Chat } from "@/db/schema";
import { generateUUID } from "@/lib/utils"; import { convertToUIMessages, generateUUID } from "@/lib/utils";
function addToolMessageToChat({
toolMessage,
messages,
}: {
toolMessage: CoreToolMessage;
messages: Array<Message>;
}): Array<Message> {
return messages.map((message) => {
if (message.toolInvocations) {
return {
...message,
toolInvocations: message.toolInvocations.map((toolInvocation) => {
const toolResult = toolMessage.content.find(
(tool) => tool.toolCallId === toolInvocation.toolCallId,
);
if (toolResult) {
return {
...toolInvocation,
state: "result",
result: toolResult.result,
};
}
return toolInvocation;
}),
};
}
return message;
});
}
function convertToUIMessages(messages: Array<CoreMessage>): Array<Message> {
return messages.reduce((chatMessages: Array<Message>, message) => {
if (message.role === "tool") {
return addToolMessageToChat({
toolMessage: message as CoreToolMessage,
messages: chatMessages,
});
}
let textContent = "";
let toolInvocations: Array<ToolInvocation> = [];
if (typeof message.content === "string") {
textContent = message.content;
} else if (Array.isArray(message.content)) {
for (const content of message.content) {
if (content.type === "text") {
textContent += content.text;
} else if (content.type === "tool-call") {
toolInvocations.push({
state: "call",
toolCallId: content.toolCallId,
toolName: content.toolName,
args: content.args,
});
}
}
}
chatMessages.push({
id: generateUUID(),
role: message.role,
content: textContent,
toolInvocations,
});
return chatMessages;
}, []);
}
export default async function Page({ params }: { params: any }) { export default async function Page({ params }: { params: any }) {
const { id } = params; const { id } = params;

View file

@ -103,6 +103,10 @@
} }
.skeleton { .skeleton {
* {
pointer-events: none !important;
}
*[class^="text-"] { *[class^="text-"] {
color: transparent; color: transparent;
@apply rounded-md bg-foreground/20 select-none animate-pulse; @apply rounded-md bg-foreground/20 select-none animate-pulse;

View file

@ -10,7 +10,7 @@ import { toast } from "sonner";
import useSWR from "swr"; import useSWR from "swr";
import { Chat } from "@/db/schema"; import { Chat } from "@/db/schema";
import { fetcher } from "@/lib/utils"; import { fetcher, getTitleFromChat } from "@/lib/utils";
import { import {
InfoIcon, InfoIcon,
@ -183,7 +183,7 @@ export const History = ({ user }: { user: User | undefined }) => {
href={`/chat/${chat.id}`} href={`/chat/${chat.id}`}
className="text-ellipsis overflow-hidden text-left py-2 pl-2 rounded-lg outline-zinc-900" className="text-ellipsis overflow-hidden text-left py-2 pl-2 rounded-lg outline-zinc-900"
> >
{chat.messages[0].content as string} {getTitleFromChat(chat)}
</Link> </Link>
</Button> </Button>

View file

@ -1,6 +1,15 @@
import {
CoreMessage,
CoreToolMessage,
generateId,
Message,
ToolInvocation,
} from "ai";
import { clsx, type ClassValue } from "clsx"; import { clsx, type ClassValue } from "clsx";
import { twMerge } from "tailwind-merge"; import { twMerge } from "tailwind-merge";
import { Chat } from "@/db/schema";
export function cn(...inputs: ClassValue[]) { export function cn(...inputs: ClassValue[]) {
return twMerge(clsx(inputs)); return twMerge(clsx(inputs));
} }
@ -41,3 +50,89 @@ export function generateUUID(): string {
return v.toString(16); return v.toString(16);
}); });
} }
function addToolMessageToChat({
toolMessage,
messages,
}: {
toolMessage: CoreToolMessage;
messages: Array<Message>;
}): Array<Message> {
return messages.map((message) => {
if (message.toolInvocations) {
return {
...message,
toolInvocations: message.toolInvocations.map((toolInvocation) => {
const toolResult = toolMessage.content.find(
(tool) => tool.toolCallId === toolInvocation.toolCallId,
);
if (toolResult) {
return {
...toolInvocation,
state: "result",
result: toolResult.result,
};
}
return toolInvocation;
}),
};
}
return message;
});
}
export function convertToUIMessages(
messages: Array<CoreMessage>,
): Array<Message> {
return messages.reduce((chatMessages: Array<Message>, message) => {
if (message.role === "tool") {
return addToolMessageToChat({
toolMessage: message as CoreToolMessage,
messages: chatMessages,
});
}
let textContent = "";
let toolInvocations: Array<ToolInvocation> = [];
if (typeof message.content === "string") {
textContent = message.content;
} else if (Array.isArray(message.content)) {
for (const content of message.content) {
if (content.type === "text") {
textContent += content.text;
} else if (content.type === "tool-call") {
toolInvocations.push({
state: "call",
toolCallId: content.toolCallId,
toolName: content.toolName,
args: content.args,
});
}
}
}
chatMessages.push({
id: generateId(),
role: message.role,
content: textContent,
toolInvocations,
});
return chatMessages;
}, []);
}
export function getTitleFromChat(chat: Chat) {
const messages = convertToUIMessages(chat.messages as Array<CoreMessage>);
const firstMessage = messages[0];
if (!firstMessage) {
return "Untitled";
}
return firstMessage.content;
}