Files
aichat/toolrouter/loop.go
T
2026-06-17 11:29:51 +08:00

126 lines
5.7 KiB
Go

package toolrouter
import (
"context"
"fmt"
"strings"
"time"
searchagent "aichat/agents/search"
sqlquery "aichat/agents/sql"
"aichat/llm"
"aichat/message"
"aichat/stream"
"aichat/utils"
"github.com/volcengine/volcengine-go-sdk/service/arkruntime/model"
)
const maxAgentToolIterations = 6
func RunAgentToolLoop(ctx context.Context, state *State, profile *llm.Profile, chatMessages []message.ChatMessage, searchState *searchagent.State, sqlState *sqlquery.State, emit stream.EmitFunc) ([]*model.ChatCompletionMessage, error) {
finalMessages, err := message.BuildArkMessages(chatMessages)
if err != nil {
return nil, err
}
routerProfile := profile
if state != nil {
routerProfile = state.RouterProfile(profile)
}
tools := availableAgentTools(state, routerProfile, searchState, sqlState, emit)
if len(tools) == 0 {
return finalMessages, nil
}
decisionMessages := append([]*model.ChatCompletionMessage(nil), finalMessages...)
if message.HasImageMessage(chatMessages) {
decisionMessages, err = message.BuildToolDecisionMessages(chatMessages)
if err != nil {
return nil, err
}
if emit != nil {
emit(stream.Frame{Type: "trace", Tool: "agent_tools", Stage: "prepare", Status: "success", Message: "检测到图片输入,工具判断阶段将使用纯文本上下文"})
}
}
toolByName := make(map[string]AgentTool, len(tools))
definitions := make([]*model.Tool, 0, len(tools))
availableNames := make([]string, 0, len(tools))
toolDescriptions := make([]string, 0, len(tools))
for _, tool := range tools {
toolByName[tool.name] = tool
definitions = append(definitions, tool.definition)
availableNames = append(availableNames, tool.name)
if tool.definition != nil && tool.definition.Function != nil {
toolDescriptions = append(toolDescriptions, fmt.Sprintf("%s: %s", tool.name, tool.definition.Function.Description))
}
}
if emit != nil {
emit(stream.Frame{Type: "trace", Tool: "agent_tools", Stage: "prepare", Status: "success", Message: "已准备可用工具", Data: map[string]any{"tools": availableNames, "tool_descriptions": toolDescriptions}})
}
if state == nil || state.cfg == nil {
return finalMessages, nil
}
if prompt := strings.TrimSpace(state.cfg.SystemPrompt); prompt != "" {
systemMessage := &model.ChatCompletionMessage{Role: model.ChatMessageRoleSystem, Content: message.StringContent(prompt)}
finalMessages = append([]*model.ChatCompletionMessage{systemMessage}, finalMessages...)
decisionMessages = append([]*model.ChatCompletionMessage{systemMessage}, decisionMessages...)
}
for i := 0; i < maxAgentToolIterations; i++ {
if emit != nil {
emit(stream.Frame{Type: "trace", Tool: "agent_tools", Stage: "request", Status: "running", Message: fmt.Sprintf("正在进行第 %d 轮工具判断", i+1), Data: map[string]any{"iteration": i + 1, "max_iterations": maxAgentToolIterations, "tools": availableNames}})
}
resp, err := state.complete(ctx, routerProfile, model.CreateChatCompletionRequest{
Model: routerProfile.Config.Model,
Messages: decisionMessages,
MaxTokens: utils.IntPtr(state.cfg.MaxTokens),
Tools: definitions,
ToolChoice: model.ToolChoiceStringTypeAuto,
ParallelToolCalls: utils.BoolPtr(false),
}, time.Duration(state.cfg.Timeout)*time.Second)
if err != nil {
return finalMessages, err
}
if tracker := stream.TrackerFromContext(ctx); tracker != nil {
tracker.AddTool(resp.Usage.PromptTokens, resp.Usage.CompletionTokens)
}
if len(resp.Choices) == 0 {
return finalMessages, nil
}
choice := resp.Choices[0]
decisionPreview := message.ChatMessageContentString(choice.Message.Content)
if emit != nil {
emit(stream.Frame{Type: "trace", Tool: "agent_tools", Stage: "decision", Status: "success", Message: "工具判断响应已返回", Data: map[string]any{"iteration": i + 1, "finish_reason": string(choice.FinishReason), "content_preview": utils.TruncateString(decisionPreview, 800)}})
}
calls := choice.Message.ToolCalls
if len(calls) == 0 && choice.Message.FunctionCall != nil {
calls = []*model.ToolCall{{ID: "legacy_function_call", Type: model.ToolTypeFunction, Function: *choice.Message.FunctionCall}}
}
if len(calls) == 0 {
if emit != nil {
emit(stream.Frame{Type: "trace", Tool: "agent_tools", Stage: "request", Status: "success", Message: "模型未请求工具,进入回答生成"})
}
return finalMessages, nil
}
callNames := make([]string, 0, len(calls))
for _, call := range calls {
if call != nil {
callNames = append(callNames, call.Function.Name)
}
}
if emit != nil {
emit(stream.Frame{Type: "trace", Tool: "agent_tools", Stage: "tool_calls", Status: "running", Message: fmt.Sprintf("模型请求调用 %d 个工具", len(calls)), Data: map[string]any{"tools": callNames, "iteration": i + 1}})
}
assistantMessage := &model.ChatCompletionMessage{Role: model.ChatMessageRoleAssistant, ToolCalls: calls, Content: choice.Message.Content}
finalMessages = append(finalMessages, assistantMessage)
decisionMessages = append(decisionMessages, assistantMessage)
for _, call := range calls {
result := ExecuteAgentToolCall(ctx, call, toolByName, emit)
toolMessage := &model.ChatCompletionMessage{Role: model.ChatMessageRoleTool, ToolCallID: call.ID, Content: message.StringContent(result)}
finalMessages = append(finalMessages, toolMessage)
decisionMessages = append(decisionMessages, toolMessage)
}
}
limitMessage := &model.ChatCompletionMessage{Role: model.ChatMessageRoleSystem, Content: message.StringContent("工具调用轮数已达到上限。请基于已有工具结果回答,并说明可能未完成全部工具调用。")}
finalMessages = append(finalMessages, limitMessage)
return finalMessages, nil
}