feat(tamagotchi): show tool calls (#184)

This commit is contained in:
LemonNeko
2025-05-29 20:25:34 +08:00
committed by GitHub
parent 37131fd79d
commit dc8bd7e87c
6 changed files with 165 additions and 35 deletions
@@ -18,18 +18,16 @@ const { onBeforeMessageComposed, onTokenLiteral } = useChatStore()
onBeforeMessageComposed(async () => {
// Scroll down to the new sent message
nextTick().then(() => {
bounding.update()
chatHistoryContainerY.value = bounding.height.value
})
await nextTick()
bounding.update()
chatHistoryContainerY.value = bounding.height.value
})
onTokenLiteral(async () => {
// Scroll down to the new responding message
nextTick().then(() => {
bounding.update()
chatHistoryContainerY.value = bounding.height.value
})
await nextTick()
bounding.update()
chatHistoryContainerY.value = bounding.height.value
})
</script>
@@ -71,8 +69,19 @@ onTokenLiteral(async () => {
</div>
<div
v-if="message.content" class="markdown-content" text="xs primary-400"
v-html="process(message.content as string)"
/>
>
<div v-for="(slice, sliceIndex) in message.slices" :key="sliceIndex">
<div v-if="slice.type === 'tool-call'">
<div
p="1" border="1 solid primary-200" rounded-lg m="y-1" bg="primary-100"
>
Called {{ slice.toolCall.function.name }}
</div>
</div>
<div v-else-if="slice.type === 'tool-call-result'" /> <!-- this line should be unreachable -->
<div v-else v-html="process(slice.text)" />
</div>
</div>
<div v-else i-eos-icons:three-dots-loading />
</div>
</div>
+81 -23
View File
@@ -1,12 +1,14 @@
import type { ChatProvider } from '@xsai-ext/shared-providers'
import type { AssistantMessage, Message, SystemMessage } from '@xsai/shared-chat'
import type { Message, SystemMessage } from '@xsai/shared-chat'
import type { ChatAssistantMessage, ChatMessage, ChatSlices } from '../types/chat'
import { defineStore, storeToRefs } from 'pinia'
import { ref, toRaw } from 'vue'
import { useQueue } from '../composables'
import { useLlmmarkerParser } from '../composables/llmmarkerParser'
import { useLLM } from '../stores/llm'
import { asyncIteratorFromReadableStream } from '../utils/iterator'
import { asyncIteratorFromReadableStream } from '../utils'
import { useAiriCardStore } from './modules'
export interface ErrorMessage {
@@ -61,14 +63,14 @@ export const useChatStore = defineStore('chat', () => {
onAssistantResponseEndHooks.value.push(cb)
}
const messages = ref<Array<Message | ErrorMessage>>([
const messages = ref<Array<ChatMessage | ErrorMessage>>([
{
role: 'system',
content: systemPrompt.value, // TODO: compose, replace {{ user }} tag, etc
} satisfies SystemMessage,
])
const streamingMessage = ref<AssistantMessage>({ role: 'assistant', content: '' })
const streamingMessage = ref<ChatAssistantMessage>({ role: 'assistant', content: '', slices: [], tool_results: [] })
async function send(sendingMessage: string, options: { model: string, chatProvider: ChatProvider, providerConfig?: Record<string, unknown> }) {
try {
@@ -81,10 +83,63 @@ export const useChatStore = defineStore('chat', () => {
await hook(sendingMessage)
}
streamingMessage.value = { role: 'assistant', content: '' }
const parser = useLlmmarkerParser({
onLiteral: async (literal) => {
for (const hook of onTokenLiteralHooks.value) {
await hook(literal)
}
streamingMessage.value.content += literal
// merge text slices for markdown
const lastSlice = streamingMessage.value.slices.at(-1)
if (lastSlice?.type === 'text') {
lastSlice.text += literal
return
}
streamingMessage.value.slices.push({
type: 'text',
text: literal,
})
},
onSpecial: async (special) => {
for (const hook of onTokenSpecialHooks.value) {
await hook(special)
}
},
})
const slicesQueue = useQueue<ChatSlices>({
handlers: [
async (ctx) => { // FIXME: it still looks dirty
if (ctx.data.type === 'text') {
await parser.consume(ctx.data.text)
return
}
if (ctx.data.type === 'tool-call') {
streamingMessage.value.slices.push(ctx.data)
return
}
if (ctx.data.type === 'tool-call-result') {
streamingMessage.value.tool_results.push(ctx.data)
}
},
],
})
streamingMessage.value = { role: 'assistant', content: '', slices: [], tool_results: [] }
messages.value.push({ role: 'user', content: sendingMessage })
messages.value.push(streamingMessage.value)
const newMessages = messages.value.slice(0, messages.value.length - 1).map(msg => toRaw(msg))
const newMessages = messages.value.slice(0, messages.value.length - 1).map((msg) => {
if (msg.role === 'assistant') {
const { slices: _, ...rest } = msg // exclude slices
return toRaw(rest)
}
return toRaw(msg)
})
for (const hook of onAfterMessageComposedHooks.value) {
await hook(sendingMessage)
@@ -95,7 +150,22 @@ export const useChatStore = defineStore('chat', () => {
}
const headers = (options.providerConfig?.headers || {}) as Record<string, string>
const res = await stream(options.model, options.chatProvider, newMessages as Message[], { headers })
const res = await stream(options.model, options.chatProvider, newMessages as Message[], {
headers,
onToolCall(toolCall) {
slicesQueue.add({
type: 'tool-call',
toolCall,
})
},
onToolCallResult(toolCallResult) {
slicesQueue.add({
type: 'tool-call-result',
id: toolCallResult.id,
result: toolCallResult.result,
})
},
})
for (const hook of onAfterSendHooks.value) {
await hook(sendingMessage)
@@ -103,24 +173,12 @@ export const useChatStore = defineStore('chat', () => {
let fullText = ''
const parser = useLlmmarkerParser({
onLiteral: async (literal) => {
for (const hook of onTokenLiteralHooks.value) {
await hook(literal)
}
streamingMessage.value.content += literal
},
onSpecial: async (special) => {
for (const hook of onTokenSpecialHooks.value) {
await hook(special)
}
},
})
for await (const textPart of asyncIteratorFromReadableStream(res.textStream, async v => v)) {
slicesQueue.add({
type: 'text',
text: textPart,
})
fullText += textPart
await parser.consume(textPart)
}
await parser.end()
+16 -2
View File
@@ -1,15 +1,20 @@
import type { ChatProvider } from '@xsai-ext/shared-providers'
import type { Message } from '@xsai/shared-chat'
import type { Message, ToolCall, ToolMessagePart } from '@xsai/shared-chat'
import { listModels } from '@xsai/model'
import { streamText } from '@xsai/stream-text'
import { defineStore } from 'pinia'
import { mcp } from '../tools'
import { debug, mcp } from '../tools'
export const useLLM = defineStore('llm', () => {
async function stream(model: string, chatProvider: ChatProvider, messages: Message[], options?: {
headers?: Record<string, string>
onToolCall?: (toolCall: ToolCall) => void
onToolCallResult?: (toolCallResult: {
id: string
result?: string | ToolMessagePart[]
}) => void
}) {
const headers = options?.headers
@@ -20,7 +25,16 @@ export const useLLM = defineStore('llm', () => {
headers,
tools: [
...await mcp(),
...await debug(),
],
onEvent(event) {
if (event.type === 'tool-call') {
options?.onToolCall?.(event.toolCall)
}
else if (event.type === 'tool-call-result') {
options?.onToolCallResult?.({ id: event.id, result: event.result })
}
},
})
}
+19
View File
@@ -0,0 +1,19 @@
import { tool } from '@xsai/tool'
import { z } from 'zod'
const tools = [
tool({
name: 'debug_random_number',
description: 'Generate a random number between 0 and 1',
execute: async () => {
return new Promise((resolve) => {
setTimeout(() => {
resolve(Math.random().toString())
}, 1000)
})
},
parameters: z.object({}),
}),
]
export const debug = async () => Promise.all(tools)
+1
View File
@@ -1 +1,2 @@
export * from './debug'
export * from './mcp'
+29
View File
@@ -0,0 +1,29 @@
import type { AssistantMessage, SystemMessage, ToolCall, ToolMessage, ToolMessagePart, UserMessage } from '@xsai/shared-chat'
export interface ChatSlicesText {
type: 'text'
text: string
}
export interface ChatSlicesToolCall {
type: 'tool-call'
toolCall: ToolCall
}
export interface ChatSlicesToolCallResult {
type: 'tool-call-result'
id: string
result?: string | ToolMessagePart[]
}
export type ChatSlices = ChatSlicesText | ChatSlicesToolCall | ChatSlicesToolCallResult
export interface ChatAssistantMessage extends AssistantMessage {
slices: ChatSlices[]
tool_results: {
id: string
result?: string | ToolMessagePart[]
}[]
}
export type ChatMessage = ChatAssistantMessage | SystemMessage | ToolMessage | UserMessage