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 }