From dc8bd7e87c0c7b6638f197389bc7e9797f2f50e9 Mon Sep 17 00:00:00 2001 From: LemonNeko Date: Thu, 29 May 2025 20:25:34 +0800 Subject: [PATCH] feat(tamagotchi): show tool calls (#184) --- .../src/components/ChatHistory.vue | 29 +++-- packages/stage-ui/src/stores/chat.ts | 104 ++++++++++++++---- packages/stage-ui/src/stores/llm.ts | 18 ++- packages/stage-ui/src/tools/debug.ts | 19 ++++ packages/stage-ui/src/tools/index.ts | 1 + packages/stage-ui/src/types/chat.ts | 29 +++++ 6 files changed, 165 insertions(+), 35 deletions(-) create mode 100644 packages/stage-ui/src/tools/debug.ts create mode 100644 packages/stage-ui/src/types/chat.ts diff --git a/apps/stage-tamagotchi/src/components/ChatHistory.vue b/apps/stage-tamagotchi/src/components/ChatHistory.vue index 8d1b79b5b..39606b900 100644 --- a/apps/stage-tamagotchi/src/components/ChatHistory.vue +++ b/apps/stage-tamagotchi/src/components/ChatHistory.vue @@ -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 }) @@ -71,8 +69,19 @@ onTokenLiteral(async () => {
+ > +
+
+
+ Called {{ slice.toolCall.function.name }} +
+
+
+
+
+
diff --git a/packages/stage-ui/src/stores/chat.ts b/packages/stage-ui/src/stores/chat.ts index 28d85e9fe..20d8ef73d 100644 --- a/packages/stage-ui/src/stores/chat.ts +++ b/packages/stage-ui/src/stores/chat.ts @@ -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>([ + const messages = ref>([ { role: 'system', content: systemPrompt.value, // TODO: compose, replace {{ user }} tag, etc } satisfies SystemMessage, ]) - const streamingMessage = ref({ role: 'assistant', content: '' }) + const streamingMessage = ref({ role: 'assistant', content: '', slices: [], tool_results: [] }) async function send(sendingMessage: string, options: { model: string, chatProvider: ChatProvider, providerConfig?: Record }) { 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({ + 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 - 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() diff --git a/packages/stage-ui/src/stores/llm.ts b/packages/stage-ui/src/stores/llm.ts index 9df58cb60..d8add23b6 100644 --- a/packages/stage-ui/src/stores/llm.ts +++ b/packages/stage-ui/src/stores/llm.ts @@ -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 + 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 }) + } + }, }) } diff --git a/packages/stage-ui/src/tools/debug.ts b/packages/stage-ui/src/tools/debug.ts new file mode 100644 index 000000000..7286dbe61 --- /dev/null +++ b/packages/stage-ui/src/tools/debug.ts @@ -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) diff --git a/packages/stage-ui/src/tools/index.ts b/packages/stage-ui/src/tools/index.ts index 75e700c28..d1fc2cfba 100644 --- a/packages/stage-ui/src/tools/index.ts +++ b/packages/stage-ui/src/tools/index.ts @@ -1 +1,2 @@ +export * from './debug' export * from './mcp' diff --git a/packages/stage-ui/src/types/chat.ts b/packages/stage-ui/src/types/chat.ts new file mode 100644 index 000000000..cef39f034 --- /dev/null +++ b/packages/stage-ui/src/types/chat.ts @@ -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