Implement tool call retry logic in Assistant's streaming process
- Enhanced the Stream method to include a retry mechanism for tool calls, allowing for up to three attempts in case of errors. - Introduced detailed logging for each retry attempt, improving traceability and debugging capabilities. - Added a new method, buildToolRetryMessages, to construct messages for LLM retries, ensuring proper context is maintained. - Updated the executeLLMForToolRetry method to handle streaming output during retries, providing real-time feedback to users. - Improved error handling to differentiate between retryable and non-retryable errors, enhancing the robustness of tool call executions.
This commit is contained in:
parent
151d7f19e0
commit
a686483f15
5 changed files with 904 additions and 12 deletions
|
|
@ -92,9 +92,11 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
|
||||||
// ================================================
|
// ================================================
|
||||||
// LLM Call Stream ( Optional )
|
// LLM Call Stream ( Optional )
|
||||||
var completionResponse *context.CompletionResponse
|
var completionResponse *context.CompletionResponse
|
||||||
|
var completionMessages []context.Message
|
||||||
|
var completionOptions *context.CompletionOptions
|
||||||
if ast.Prompts != nil || ast.MCP != nil {
|
if ast.Prompts != nil || ast.MCP != nil {
|
||||||
// Build the LLM request first
|
// Build the LLM request first
|
||||||
completionMessages, completionOptions, err := ast.BuildRequest(ctx, inputMessages, createResponse)
|
completionMessages, completionOptions, err = ast.BuildRequest(ctx, inputMessages, createResponse)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
ast.traceAgentFail(agentNode, err)
|
ast.traceAgentFail(agentNode, err)
|
||||||
// Send error stream_end for root stack
|
// Send error stream_end for root stack
|
||||||
|
|
@ -112,13 +114,97 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
|
||||||
}
|
}
|
||||||
|
|
||||||
// ================================================
|
// ================================================
|
||||||
// Execute tool calls
|
// Execute tool calls with retry
|
||||||
// ================================================
|
// ================================================
|
||||||
if completionResponse != nil && completionResponse.ToolCalls != nil {
|
if completionResponse != nil && completionResponse.ToolCalls != nil {
|
||||||
|
|
||||||
fmt.Println("--- completionResponse ToolCalls --------------------------------")
|
// === Debug Tool Calls ===
|
||||||
|
fmt.Println("--- Debug Tool Calls --------------------------------")
|
||||||
utils.Dump(completionResponse.ToolCalls)
|
utils.Dump(completionResponse.ToolCalls)
|
||||||
fmt.Println("--------------------------------")
|
|
||||||
|
// === End Debug Tool Calls ===
|
||||||
|
|
||||||
|
maxToolRetries := 3
|
||||||
|
currentMessages := completionMessages
|
||||||
|
currentResponse := completionResponse
|
||||||
|
|
||||||
|
for attempt := 0; attempt < maxToolRetries; attempt++ {
|
||||||
|
|
||||||
|
fmt.Println("attempt", attempt)
|
||||||
|
// Execute all tool calls
|
||||||
|
toolResults, hasErrors := ast.executeToolCalls(ctx, currentResponse.ToolCalls, attempt)
|
||||||
|
// If all successful, break out
|
||||||
|
if !hasErrors {
|
||||||
|
log.Trace("[AGENT] All tool calls succeeded (attempt %d)", attempt)
|
||||||
|
for _, result := range toolResults {
|
||||||
|
fmt.Println("--")
|
||||||
|
fmt.Printf("Result :%s %s %s\n", result.ToolCallID, result.ServerID(), result.ToolName())
|
||||||
|
res, err := result.ParsedContent()
|
||||||
|
if err != nil {
|
||||||
|
fmt.Println("Error: ", err)
|
||||||
|
}
|
||||||
|
utils.Dump(res)
|
||||||
|
fmt.Println("--")
|
||||||
|
}
|
||||||
|
fmt.Println("--------------------------------")
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if any errors are retryable (parameter/validation issues)
|
||||||
|
hasRetryableErrors := false
|
||||||
|
for _, result := range toolResults {
|
||||||
|
if result.Error != nil && result.IsRetryableError {
|
||||||
|
hasRetryableErrors = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// If no retryable errors, don't retry (MCP internal issues)
|
||||||
|
if !hasRetryableErrors {
|
||||||
|
err := fmt.Errorf("tool calls failed with non-retryable errors (MCP internal issues)")
|
||||||
|
log.Error("[AGENT] %v", err)
|
||||||
|
ast.traceAgentFail(agentNode, err)
|
||||||
|
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// If it's the last attempt, return error
|
||||||
|
if attempt == maxToolRetries-1 {
|
||||||
|
err := fmt.Errorf("tool calls failed after %d attempts", maxToolRetries)
|
||||||
|
log.Error("[AGENT] %v", err)
|
||||||
|
ast.traceAgentFail(agentNode, err)
|
||||||
|
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build retry messages with tool call results (including errors)
|
||||||
|
retryMessages := ast.buildToolRetryMessages(currentMessages, currentResponse, toolResults)
|
||||||
|
|
||||||
|
// Retry LLM call (streaming to keep user informed)
|
||||||
|
log.Trace("[AGENT] Retrying LLM for tool call correction (attempt %d/%d)", attempt+1, maxToolRetries-1)
|
||||||
|
currentResponse, err = ast.executeLLMForToolRetry(ctx, retryMessages, completionOptions, agentNode, streamHandler)
|
||||||
|
if err != nil {
|
||||||
|
log.Error("[AGENT] LLM retry failed: %v", err)
|
||||||
|
ast.traceAgentFail(agentNode, err)
|
||||||
|
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// If LLM didn't return tool calls, it might have given up
|
||||||
|
if currentResponse.ToolCalls == nil {
|
||||||
|
err := fmt.Errorf("LLM did not return tool calls in retry attempt %d", attempt+1)
|
||||||
|
log.Error("[AGENT] %v", err)
|
||||||
|
ast.traceAgentFail(agentNode, err)
|
||||||
|
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update messages for next iteration
|
||||||
|
currentMessages = retryMessages
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update completionResponse with the final successful response
|
||||||
|
completionResponse = currentResponse
|
||||||
}
|
}
|
||||||
|
|
||||||
// Request Done hook ( Optional )
|
// Request Done hook ( Optional )
|
||||||
|
|
@ -403,3 +489,51 @@ func (ast *Assistant) initializeCapabilities(ctx *context.Context) error {
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// buildToolRetryMessages builds messages for LLM retry with tool call results
|
||||||
|
// Format follows OpenAI's tool call response pattern:
|
||||||
|
// 1. Assistant message with tool calls
|
||||||
|
// 2. Tool messages with results (one per tool call)
|
||||||
|
// 3. System message explaining the retry
|
||||||
|
func (ast *Assistant) buildToolRetryMessages(
|
||||||
|
previousMessages []context.Message,
|
||||||
|
completionResponse *context.CompletionResponse,
|
||||||
|
toolResults []ToolCallResult,
|
||||||
|
) []context.Message {
|
||||||
|
retryMessages := make([]context.Message, 0, len(previousMessages)+len(toolResults)+2)
|
||||||
|
|
||||||
|
// Add all previous messages
|
||||||
|
retryMessages = append(retryMessages, previousMessages...)
|
||||||
|
|
||||||
|
// Add assistant message with tool calls
|
||||||
|
assistantMsg := context.Message{
|
||||||
|
Role: context.RoleAssistant,
|
||||||
|
Content: completionResponse.Content,
|
||||||
|
ToolCalls: completionResponse.ToolCalls,
|
||||||
|
}
|
||||||
|
retryMessages = append(retryMessages, assistantMsg)
|
||||||
|
|
||||||
|
// Add tool result messages (one per tool call)
|
||||||
|
for _, result := range toolResults {
|
||||||
|
toolMsg := context.Message{
|
||||||
|
Role: context.RoleTool,
|
||||||
|
Content: result.Content,
|
||||||
|
ToolCallID: &result.ToolCallID,
|
||||||
|
}
|
||||||
|
// Add tool name if available
|
||||||
|
if result.Name != "" {
|
||||||
|
name := result.Name
|
||||||
|
toolMsg.Name = &name
|
||||||
|
}
|
||||||
|
retryMessages = append(retryMessages, toolMsg)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add system message explaining the retry (optional, helps LLM understand context)
|
||||||
|
systemMsg := context.Message{
|
||||||
|
Role: context.RoleSystem,
|
||||||
|
Content: i18n.Tr(ast.ID, "en", "assistant.agent.tool_retry_prompt"),
|
||||||
|
}
|
||||||
|
retryMessages = append(retryMessages, systemMsg)
|
||||||
|
|
||||||
|
return retryMessages
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -56,3 +56,49 @@ func (ast *Assistant) executeLLMStream(
|
||||||
|
|
||||||
return completionResponse, nil
|
return completionResponse, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// executeLLMForToolRetry executes LLM call for tool retry with streaming output
|
||||||
|
// This is used when retrying tool calls - we still want to show LLM's response to users
|
||||||
|
// Returns completionResponse and error
|
||||||
|
func (ast *Assistant) executeLLMForToolRetry(
|
||||||
|
ctx *context.Context,
|
||||||
|
completionMessages []context.Message,
|
||||||
|
completionOptions *context.CompletionOptions,
|
||||||
|
agentNode types.Node,
|
||||||
|
streamHandler message.StreamFunc,
|
||||||
|
) (*context.CompletionResponse, error) {
|
||||||
|
|
||||||
|
// Get connector object
|
||||||
|
conn, capabilities, err := ast.GetConnector(ctx)
|
||||||
|
if err != nil {
|
||||||
|
ast.traceAgentFail(agentNode, err)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set capabilities in options if not already set
|
||||||
|
if completionOptions.Capabilities == nil && capabilities != nil {
|
||||||
|
completionOptions.Capabilities = capabilities
|
||||||
|
}
|
||||||
|
|
||||||
|
// Trace Add LLM retry request
|
||||||
|
ast.traceLLMRetryRequest(ctx, conn.ID(), completionMessages, completionOptions)
|
||||||
|
|
||||||
|
// Create LLM instance with connector and options
|
||||||
|
llmInstance, err := llm.New(conn, completionOptions)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Call the LLM Completion Stream (still streaming for tool retry)
|
||||||
|
log.Trace("[AGENT] Calling LLM Stream for tool retry: assistant=%s", ast.ID)
|
||||||
|
completionResponse, err := llmInstance.Stream(ctx, completionMessages, completionOptions, streamHandler)
|
||||||
|
log.Trace("[AGENT] LLM tool retry stream returned: assistant=%s, err=%v", ast.ID, err)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Mark LLM Request Complete
|
||||||
|
ast.traceLLMComplete(ctx, completionResponse)
|
||||||
|
|
||||||
|
return completionResponse, nil
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -5,10 +5,13 @@ import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
jsoniter "github.com/json-iterator/go"
|
||||||
|
gouJson "github.com/yaoapp/gou/json"
|
||||||
"github.com/yaoapp/gou/mcp"
|
"github.com/yaoapp/gou/mcp"
|
||||||
mcpTypes "github.com/yaoapp/gou/mcp/types"
|
mcpTypes "github.com/yaoapp/gou/mcp/types"
|
||||||
"github.com/yaoapp/kun/log"
|
"github.com/yaoapp/kun/log"
|
||||||
agentContext "github.com/yaoapp/yao/agent/context"
|
agentContext "github.com/yaoapp/yao/agent/context"
|
||||||
|
"github.com/yaoapp/yao/trace/types"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
|
|
@ -67,13 +70,6 @@ func ParseMCPToolName(formattedName string) (string, string, bool) {
|
||||||
return serverID, toolName, true
|
return serverID, toolName, true
|
||||||
}
|
}
|
||||||
|
|
||||||
// MCPTool represents a simplified MCP tool for building LLM requests
|
|
||||||
type MCPTool struct {
|
|
||||||
Name string
|
|
||||||
Description string
|
|
||||||
Parameters interface{}
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildMCPTools builds tool definitions and samples system prompt from MCP servers
|
// buildMCPTools builds tool definitions and samples system prompt from MCP servers
|
||||||
// Returns (tools, samplesPrompt, error)
|
// Returns (tools, samplesPrompt, error)
|
||||||
func (ast *Assistant) buildMCPTools(ctx *agentContext.Context) ([]MCPTool, string, error) {
|
func (ast *Assistant) buildMCPTools(ctx *agentContext.Context) ([]MCPTool, string, error) {
|
||||||
|
|
@ -81,7 +77,12 @@ func (ast *Assistant) buildMCPTools(ctx *agentContext.Context) ([]MCPTool, strin
|
||||||
return nil, "", nil
|
return nil, "", nil
|
||||||
}
|
}
|
||||||
|
|
||||||
mcpCtx := context.Background()
|
// Use the agent context for cancellation and timeout control
|
||||||
|
mcpCtx := ctx.Context
|
||||||
|
if mcpCtx == nil {
|
||||||
|
mcpCtx = context.Background()
|
||||||
|
}
|
||||||
|
|
||||||
allTools := make([]MCPTool, 0)
|
allTools := make([]MCPTool, 0)
|
||||||
samplesBuilder := strings.Builder{}
|
samplesBuilder := strings.Builder{}
|
||||||
hasSamples := false
|
hasSamples := false
|
||||||
|
|
@ -195,3 +196,599 @@ func (ast *Assistant) buildMCPTools(ctx *agentContext.Context) ([]MCPTool, strin
|
||||||
log.Trace("[Assistant MCP] Total MCP tools loaded: %d", len(allTools))
|
log.Trace("[Assistant MCP] Total MCP tools loaded: %d", len(allTools))
|
||||||
return allTools, samplesPrompt, nil
|
return allTools, samplesPrompt, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ToolCallResult represents the result of a tool call execution
|
||||||
|
// executeToolCalls executes tool calls with intelligent strategy and trace logging:
|
||||||
|
// - Single tool: use CallTool, single trace node
|
||||||
|
// - Multiple tools: use CallToolsParallel with parallel trace nodes, fallback to sequential on certain errors
|
||||||
|
// Returns (results, hasErrors)
|
||||||
|
func (ast *Assistant) executeToolCalls(ctx *agentContext.Context, toolCalls []agentContext.ToolCall, attempt int) ([]ToolCallResult, bool) {
|
||||||
|
if len(toolCalls) == 0 {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Trace("[Assistant MCP] Executing %d tool calls (attempt %d)", len(toolCalls), attempt)
|
||||||
|
|
||||||
|
// Single tool call
|
||||||
|
if len(toolCalls) == 1 {
|
||||||
|
return ast.executeSingleToolCall(ctx, toolCalls[0])
|
||||||
|
}
|
||||||
|
|
||||||
|
// Multiple tool calls - try parallel first
|
||||||
|
return ast.executeMultipleToolCallsParallel(ctx, toolCalls)
|
||||||
|
}
|
||||||
|
|
||||||
|
// executeSingleToolCall executes a single tool call with trace logging
|
||||||
|
func (ast *Assistant) executeSingleToolCall(ctx *agentContext.Context, toolCall agentContext.ToolCall) ([]ToolCallResult, bool) {
|
||||||
|
trace, _ := ctx.Trace()
|
||||||
|
|
||||||
|
// Use the agent context for cancellation and timeout control
|
||||||
|
mcpCtx := ctx.Context
|
||||||
|
if mcpCtx == nil {
|
||||||
|
mcpCtx = context.Background()
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ToolCallResult{
|
||||||
|
ToolCallID: toolCall.ID,
|
||||||
|
Name: toolCall.Function.Name,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse tool name
|
||||||
|
serverID, toolName, ok := ParseMCPToolName(toolCall.Function.Name)
|
||||||
|
if !ok {
|
||||||
|
result.Error = fmt.Errorf("invalid MCP tool name format: %s", toolCall.Function.Name)
|
||||||
|
result.Content = result.Error.Error()
|
||||||
|
log.Error("[Assistant MCP] %v", result.Error)
|
||||||
|
return []ToolCallResult{result}, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get MCP client
|
||||||
|
client, err := mcp.Select(serverID)
|
||||||
|
if err != nil {
|
||||||
|
result.Error = fmt.Errorf("failed to select MCP client '%s': %w", serverID, err)
|
||||||
|
result.Content = result.Error.Error()
|
||||||
|
result.IsRetryableError = false // MCP client selection error is not retryable
|
||||||
|
log.Error("[Assistant MCP] %v", result.Error)
|
||||||
|
return []ToolCallResult{result}, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get tool info for description and schema
|
||||||
|
toolsResponse, err := client.ListTools(mcpCtx, "")
|
||||||
|
var toolDescription string
|
||||||
|
var toolSchema interface{}
|
||||||
|
if err == nil {
|
||||||
|
for _, t := range toolsResponse.Tools {
|
||||||
|
if t.Name == toolName {
|
||||||
|
toolDescription = t.Description
|
||||||
|
toolSchema = t.InputSchema
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if toolDescription == "" {
|
||||||
|
toolDescription = fmt.Sprintf("MCP tool '%s'", toolName)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add trace node for this tool call
|
||||||
|
var toolNode types.Node
|
||||||
|
if trace != nil {
|
||||||
|
toolNode, _ = trace.Add(
|
||||||
|
map[string]any{
|
||||||
|
"tool_call_id": toolCall.ID,
|
||||||
|
"server": serverID,
|
||||||
|
"tool": toolName,
|
||||||
|
"arguments": toolCall.Function.Arguments,
|
||||||
|
},
|
||||||
|
types.TraceNodeOption{
|
||||||
|
Label: toolDescription,
|
||||||
|
Type: "mcp_tool",
|
||||||
|
Icon: "build",
|
||||||
|
Description: fmt.Sprintf("Calling '%s' on server '%s'", toolName, serverID),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse arguments with repair support for better tolerance
|
||||||
|
var args map[string]interface{}
|
||||||
|
if toolCall.Function.Arguments != "" {
|
||||||
|
parsed, err := gouJson.Parse(toolCall.Function.Arguments)
|
||||||
|
if err != nil {
|
||||||
|
result.Error = fmt.Errorf("failed to parse arguments: %w", err)
|
||||||
|
result.Content = result.Error.Error()
|
||||||
|
result.IsRetryableError = true // Argument parsing error is retryable by LLM
|
||||||
|
log.Error("[Assistant MCP] %v", result.Error)
|
||||||
|
if toolNode != nil {
|
||||||
|
toolNode.Fail(result.Error)
|
||||||
|
}
|
||||||
|
return []ToolCallResult{result}, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Convert to map
|
||||||
|
if argsMap, ok := parsed.(map[string]interface{}); ok {
|
||||||
|
args = argsMap
|
||||||
|
} else {
|
||||||
|
result.Error = fmt.Errorf("arguments must be an object, got %T", parsed)
|
||||||
|
result.Content = result.Error.Error()
|
||||||
|
result.IsRetryableError = true // Type error is retryable by LLM
|
||||||
|
log.Error("[Assistant MCP] %v", result.Error)
|
||||||
|
if toolNode != nil {
|
||||||
|
toolNode.Fail(result.Error)
|
||||||
|
}
|
||||||
|
return []ToolCallResult{result}, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate arguments against tool schema if available
|
||||||
|
if toolSchema != nil {
|
||||||
|
if err := gouJson.Validate(args, toolSchema); err != nil {
|
||||||
|
result.Error = fmt.Errorf("argument validation failed: %w", err)
|
||||||
|
result.Content = result.Error.Error()
|
||||||
|
result.IsRetryableError = true // Validation error is retryable by LLM
|
||||||
|
log.Error("[Assistant MCP] %v", result.Error)
|
||||||
|
if toolNode != nil {
|
||||||
|
toolNode.Fail(result.Error)
|
||||||
|
}
|
||||||
|
return []ToolCallResult{result}, true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Call the tool
|
||||||
|
log.Trace("[Assistant MCP] Calling tool: %s (server: %s)", toolName, serverID)
|
||||||
|
callResult, err := client.CallTool(mcpCtx, toolName, args)
|
||||||
|
if err != nil {
|
||||||
|
result.Error = fmt.Errorf("tool call failed: %w", err)
|
||||||
|
result.Content = result.Error.Error()
|
||||||
|
// Check if error is retryable (parameter/validation errors)
|
||||||
|
result.IsRetryableError = isRetryableToolError(err)
|
||||||
|
log.Error("[Assistant MCP] Tool call failed: %v (retryable: %v)", err, result.IsRetryableError)
|
||||||
|
if toolNode != nil {
|
||||||
|
toolNode.Fail(result.Error)
|
||||||
|
}
|
||||||
|
return []ToolCallResult{result}, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if result is an error
|
||||||
|
if callResult.IsError {
|
||||||
|
result.Error = fmt.Errorf("MCP tool error")
|
||||||
|
result.IsRetryableError = false // MCP internal error is not retryable
|
||||||
|
}
|
||||||
|
|
||||||
|
// Serialize the Content field only ([]ToolContent)
|
||||||
|
contentBytes, err := jsoniter.Marshal(callResult.Content)
|
||||||
|
if err != nil {
|
||||||
|
result.Error = fmt.Errorf("failed to serialize result: %w", err)
|
||||||
|
result.Content = result.Error.Error()
|
||||||
|
result.IsRetryableError = false
|
||||||
|
log.Error("[Assistant MCP] %v", result.Error)
|
||||||
|
if toolNode != nil {
|
||||||
|
toolNode.Fail(result.Error)
|
||||||
|
}
|
||||||
|
return []ToolCallResult{result}, true
|
||||||
|
}
|
||||||
|
|
||||||
|
result.Content = string(contentBytes)
|
||||||
|
log.Trace("[Assistant MCP] Tool call succeeded: %s", toolName)
|
||||||
|
|
||||||
|
if toolNode != nil {
|
||||||
|
toolNode.Complete(map[string]any{
|
||||||
|
"result": callResult,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
return []ToolCallResult{result}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
// executeMultipleToolCallsParallel executes multiple tool calls in parallel with trace logging
|
||||||
|
func (ast *Assistant) executeMultipleToolCallsParallel(ctx *agentContext.Context, toolCalls []agentContext.ToolCall) ([]ToolCallResult, bool) {
|
||||||
|
trace, _ := ctx.Trace()
|
||||||
|
|
||||||
|
// Use the agent context for cancellation and timeout control
|
||||||
|
mcpCtx := ctx.Context
|
||||||
|
if mcpCtx == nil {
|
||||||
|
mcpCtx = context.Background()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Group tool calls by server
|
||||||
|
serverGroups := make(map[string][]agentContext.ToolCall)
|
||||||
|
for _, tc := range toolCalls {
|
||||||
|
serverID, _, ok := ParseMCPToolName(tc.Function.Name)
|
||||||
|
if !ok {
|
||||||
|
log.Warn("[Assistant MCP] Invalid tool name format: %s", tc.Function.Name)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
serverGroups[serverID] = append(serverGroups[serverID], tc)
|
||||||
|
}
|
||||||
|
|
||||||
|
results := make([]ToolCallResult, 0, len(toolCalls))
|
||||||
|
hasErrors := false
|
||||||
|
|
||||||
|
// Process each server's tools
|
||||||
|
for serverID, calls := range serverGroups {
|
||||||
|
client, err := mcp.Select(serverID)
|
||||||
|
if err != nil {
|
||||||
|
log.Error("[Assistant MCP] Failed to select MCP client '%s': %v", serverID, err)
|
||||||
|
// Add error results for all calls to this server
|
||||||
|
for _, tc := range calls {
|
||||||
|
results = append(results, ToolCallResult{
|
||||||
|
ToolCallID: tc.ID,
|
||||||
|
Name: tc.Function.Name,
|
||||||
|
Content: fmt.Sprintf("Failed to select MCP client: %v", err),
|
||||||
|
Error: err,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
hasErrors = true
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try parallel execution
|
||||||
|
serverResults, serverHasErrors := ast.executeServerToolsParallelWithTrace(
|
||||||
|
mcpCtx, trace, client, serverID, calls,
|
||||||
|
)
|
||||||
|
|
||||||
|
// If parallel execution failed with retryable error, try sequential
|
||||||
|
if serverHasErrors && ast.shouldRetrySequential(serverResults) {
|
||||||
|
log.Warn("[Assistant MCP] Parallel execution had parameter errors for server '%s', retrying sequentially", serverID)
|
||||||
|
serverResults, serverHasErrors = ast.executeServerToolsSequentialWithTrace(
|
||||||
|
mcpCtx, trace, client, serverID, calls,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
results = append(results, serverResults...)
|
||||||
|
if serverHasErrors {
|
||||||
|
hasErrors = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return results, hasErrors
|
||||||
|
}
|
||||||
|
|
||||||
|
// isRetryableToolError checks if an error is retryable by LLM (parameter/validation errors)
|
||||||
|
// Returns true for errors that LLM can potentially fix by adjusting parameters
|
||||||
|
// Returns false for MCP internal errors (network, auth, service unavailable, etc.)
|
||||||
|
func isRetryableToolError(err error) bool {
|
||||||
|
if err == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
errMsg := strings.ToLower(err.Error())
|
||||||
|
|
||||||
|
// These are NOT retryable (MCP internal issues)
|
||||||
|
nonRetryablePatterns := []string{
|
||||||
|
"network",
|
||||||
|
"timeout",
|
||||||
|
"connection",
|
||||||
|
"unauthorized",
|
||||||
|
"forbidden",
|
||||||
|
"unavailable",
|
||||||
|
"failed to select",
|
||||||
|
"context canceled",
|
||||||
|
"context deadline",
|
||||||
|
"server error",
|
||||||
|
"internal error",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, pattern := range nonRetryablePatterns {
|
||||||
|
if strings.Contains(errMsg, pattern) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// These ARE retryable (parameter/validation issues LLM can fix)
|
||||||
|
retryablePatterns := []string{
|
||||||
|
"invalid",
|
||||||
|
"required",
|
||||||
|
"missing",
|
||||||
|
"validation",
|
||||||
|
"schema",
|
||||||
|
"type",
|
||||||
|
"format",
|
||||||
|
"parse",
|
||||||
|
"argument",
|
||||||
|
"parameter",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, pattern := range retryablePatterns {
|
||||||
|
if strings.Contains(errMsg, pattern) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Default: assume it's retryable unless proven otherwise
|
||||||
|
// This allows LLM to attempt fixes for unknown error types
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// shouldRetrySequential checks if errors are retryable (parameter issues, not network/service issues)
|
||||||
|
func (ast *Assistant) shouldRetrySequential(results []ToolCallResult) bool {
|
||||||
|
// Check if any result has a retryable error
|
||||||
|
hasRetryable := false
|
||||||
|
for _, result := range results {
|
||||||
|
if result.Error != nil && result.IsRetryableError {
|
||||||
|
hasRetryable = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return hasRetryable
|
||||||
|
}
|
||||||
|
|
||||||
|
// executeServerToolsParallelWithTrace executes tools for a single server in parallel with trace
|
||||||
|
func (ast *Assistant) executeServerToolsParallelWithTrace(mcpCtx context.Context, trace types.Manager, client mcp.Client, serverID string, toolCalls []agentContext.ToolCall) ([]ToolCallResult, bool) {
|
||||||
|
// Prepare parallel trace inputs
|
||||||
|
var parallelInputs []types.TraceParallelInput
|
||||||
|
mcpCalls := make([]mcpTypes.ToolCall, 0, len(toolCalls))
|
||||||
|
callMap := make(map[string]agentContext.ToolCall)
|
||||||
|
|
||||||
|
for _, tc := range toolCalls {
|
||||||
|
_, toolName, ok := ParseMCPToolName(tc.Function.Name)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
var args map[string]interface{}
|
||||||
|
if tc.Function.Arguments != "" {
|
||||||
|
if err := jsoniter.UnmarshalFromString(tc.Function.Arguments, &args); err != nil {
|
||||||
|
log.Error("[Assistant MCP] Failed to parse arguments for %s: %v", toolName, err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
mcpCalls = append(mcpCalls, mcpTypes.ToolCall{
|
||||||
|
Name: toolName,
|
||||||
|
Arguments: args,
|
||||||
|
})
|
||||||
|
callMap[toolName] = tc
|
||||||
|
|
||||||
|
// Add trace input for this tool
|
||||||
|
parallelInputs = append(parallelInputs, types.TraceParallelInput{
|
||||||
|
Input: map[string]any{
|
||||||
|
"tool_call_id": tc.ID,
|
||||||
|
"server": serverID,
|
||||||
|
"tool": toolName,
|
||||||
|
"arguments": tc.Function.Arguments,
|
||||||
|
},
|
||||||
|
Option: types.TraceNodeOption{
|
||||||
|
Label: fmt.Sprintf("Tool: %s", toolName),
|
||||||
|
Type: "mcp_tool",
|
||||||
|
Icon: "build",
|
||||||
|
Description: fmt.Sprintf("Calling MCP tool '%s' on server '%s'", toolName, serverID),
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create parallel trace nodes
|
||||||
|
var toolNodes []types.Node
|
||||||
|
if trace != nil && len(parallelInputs) > 0 {
|
||||||
|
toolNodes, _ = trace.Parallel(parallelInputs)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Call tools in parallel
|
||||||
|
log.Trace("[Assistant MCP] Calling %d tools in parallel on server '%s'", len(mcpCalls), serverID)
|
||||||
|
mcpResponse, err := client.CallToolsParallel(mcpCtx, mcpCalls)
|
||||||
|
if err != nil {
|
||||||
|
log.Error("[Assistant MCP] Parallel call failed: %v", err)
|
||||||
|
// Mark all trace nodes as failed
|
||||||
|
for _, node := range toolNodes {
|
||||||
|
if node != nil {
|
||||||
|
node.Fail(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Process results
|
||||||
|
results := make([]ToolCallResult, 0, len(mcpResponse.Results))
|
||||||
|
hasErrors := false
|
||||||
|
|
||||||
|
for i, mcpResult := range mcpResponse.Results {
|
||||||
|
toolName := mcpCalls[i].Name
|
||||||
|
originalCall := callMap[toolName]
|
||||||
|
var toolNode types.Node
|
||||||
|
if i < len(toolNodes) {
|
||||||
|
toolNode = toolNodes[i]
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ToolCallResult{
|
||||||
|
ToolCallID: originalCall.ID,
|
||||||
|
Name: originalCall.Function.Name,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Serialize content
|
||||||
|
contentBytes, err := jsoniter.Marshal(mcpResult.Content)
|
||||||
|
if err != nil {
|
||||||
|
result.Error = fmt.Errorf("failed to serialize result: %w", err)
|
||||||
|
result.Content = result.Error.Error()
|
||||||
|
result.IsRetryableError = false // Serialization error is not retryable
|
||||||
|
hasErrors = true
|
||||||
|
if toolNode != nil {
|
||||||
|
toolNode.Fail(result.Error)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
result.Content = string(contentBytes)
|
||||||
|
|
||||||
|
// Check if it's an error result
|
||||||
|
if mcpResult.IsError {
|
||||||
|
result.Error = fmt.Errorf("tool call error: %s", result.Content)
|
||||||
|
result.IsRetryableError = isRetryableToolError(result.Error)
|
||||||
|
hasErrors = true
|
||||||
|
log.Error("[Assistant MCP] Tool call failed: %s - %s (retryable: %v)", toolName, result.Content, result.IsRetryableError)
|
||||||
|
if toolNode != nil {
|
||||||
|
toolNode.Fail(result.Error)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// Success
|
||||||
|
if toolNode != nil {
|
||||||
|
toolNode.Complete(map[string]any{
|
||||||
|
"result": mcpResult.Content,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
results = append(results, result)
|
||||||
|
}
|
||||||
|
|
||||||
|
return results, hasErrors
|
||||||
|
}
|
||||||
|
|
||||||
|
// executeServerToolsSequentialWithTrace executes tools for a single server sequentially with trace
|
||||||
|
func (ast *Assistant) executeServerToolsSequentialWithTrace(mcpCtx context.Context, trace types.Manager, client mcp.Client, serverID string, toolCalls []agentContext.ToolCall) ([]ToolCallResult, bool) {
|
||||||
|
results := make([]ToolCallResult, 0, len(toolCalls))
|
||||||
|
hasErrors := false
|
||||||
|
|
||||||
|
log.Trace("[Assistant MCP] Calling %d tools sequentially on server '%s'", len(toolCalls), serverID)
|
||||||
|
|
||||||
|
for _, tc := range toolCalls {
|
||||||
|
_, toolName, ok := ParseMCPToolName(tc.Function.Name)
|
||||||
|
if !ok {
|
||||||
|
results = append(results, ToolCallResult{
|
||||||
|
ToolCallID: tc.ID,
|
||||||
|
Name: tc.Function.Name,
|
||||||
|
Content: fmt.Sprintf("Invalid tool name format: %s", tc.Function.Name),
|
||||||
|
Error: fmt.Errorf("invalid tool name format"),
|
||||||
|
})
|
||||||
|
hasErrors = true
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get tool schema for validation
|
||||||
|
toolsResponse, err := client.ListTools(mcpCtx, "")
|
||||||
|
var toolSchema interface{}
|
||||||
|
if err == nil {
|
||||||
|
for _, t := range toolsResponse.Tools {
|
||||||
|
if t.Name == toolName {
|
||||||
|
toolSchema = t.InputSchema
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add trace node for this tool call
|
||||||
|
var toolNode types.Node
|
||||||
|
if trace != nil {
|
||||||
|
toolNode, _ = trace.Add(
|
||||||
|
map[string]any{
|
||||||
|
"tool_call_id": tc.ID,
|
||||||
|
"server": serverID,
|
||||||
|
"tool": toolName,
|
||||||
|
"arguments": tc.Function.Arguments,
|
||||||
|
},
|
||||||
|
types.TraceNodeOption{
|
||||||
|
Label: fmt.Sprintf("Tool: %s (sequential retry)", toolName),
|
||||||
|
Type: "mcp_tool",
|
||||||
|
Icon: "build",
|
||||||
|
Description: fmt.Sprintf("Retrying MCP tool '%s' on server '%s' sequentially", toolName, serverID),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse arguments with repair support
|
||||||
|
var args map[string]interface{}
|
||||||
|
if tc.Function.Arguments != "" {
|
||||||
|
parsed, err := gouJson.Parse(tc.Function.Arguments)
|
||||||
|
if err != nil {
|
||||||
|
result := ToolCallResult{
|
||||||
|
ToolCallID: tc.ID,
|
||||||
|
Name: tc.Function.Name,
|
||||||
|
Content: fmt.Sprintf("Failed to parse arguments: %v", err),
|
||||||
|
Error: err,
|
||||||
|
IsRetryableError: true, // Parsing error is retryable
|
||||||
|
}
|
||||||
|
results = append(results, result)
|
||||||
|
hasErrors = true
|
||||||
|
if toolNode != nil {
|
||||||
|
toolNode.Fail(err)
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Convert to map
|
||||||
|
if argsMap, ok := parsed.(map[string]interface{}); ok {
|
||||||
|
args = argsMap
|
||||||
|
} else {
|
||||||
|
err := fmt.Errorf("arguments must be an object, got %T", parsed)
|
||||||
|
result := ToolCallResult{
|
||||||
|
ToolCallID: tc.ID,
|
||||||
|
Name: tc.Function.Name,
|
||||||
|
Content: err.Error(),
|
||||||
|
Error: err,
|
||||||
|
IsRetryableError: true, // Type error is retryable
|
||||||
|
}
|
||||||
|
results = append(results, result)
|
||||||
|
hasErrors = true
|
||||||
|
if toolNode != nil {
|
||||||
|
toolNode.Fail(err)
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate arguments against tool schema if available
|
||||||
|
if toolSchema != nil {
|
||||||
|
if err := gouJson.Validate(args, toolSchema); err != nil {
|
||||||
|
result := ToolCallResult{
|
||||||
|
ToolCallID: tc.ID,
|
||||||
|
Name: tc.Function.Name,
|
||||||
|
Content: fmt.Sprintf("Argument validation failed: %v", err),
|
||||||
|
Error: err,
|
||||||
|
IsRetryableError: true, // Validation error is retryable
|
||||||
|
}
|
||||||
|
results = append(results, result)
|
||||||
|
hasErrors = true
|
||||||
|
if toolNode != nil {
|
||||||
|
toolNode.Fail(err)
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Call single tool
|
||||||
|
log.Trace("[Assistant MCP] Calling tool: %s", toolName)
|
||||||
|
mcpResult, err := client.CallTool(mcpCtx, toolName, args)
|
||||||
|
|
||||||
|
result := ToolCallResult{
|
||||||
|
ToolCallID: tc.ID,
|
||||||
|
Name: tc.Function.Name,
|
||||||
|
}
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
result.Error = err
|
||||||
|
result.Content = fmt.Sprintf("Tool call failed: %v", err)
|
||||||
|
result.IsRetryableError = isRetryableToolError(err)
|
||||||
|
hasErrors = true
|
||||||
|
log.Error("[Assistant MCP] Tool call failed: %s - %v (retryable: %v)", toolName, err, result.IsRetryableError)
|
||||||
|
if toolNode != nil {
|
||||||
|
toolNode.Fail(err)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// Check if result is an error
|
||||||
|
if mcpResult.IsError {
|
||||||
|
result.Error = fmt.Errorf("MCP tool error")
|
||||||
|
result.IsRetryableError = false // MCP internal error is not retryable
|
||||||
|
hasErrors = true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Serialize the Content field only ([]ToolContent)
|
||||||
|
contentBytes, err := jsoniter.Marshal(mcpResult.Content)
|
||||||
|
if err != nil {
|
||||||
|
result.Error = err
|
||||||
|
result.Content = fmt.Sprintf("Failed to serialize result: %v", err)
|
||||||
|
result.IsRetryableError = false // Serialization error is not retryable
|
||||||
|
hasErrors = true
|
||||||
|
if toolNode != nil {
|
||||||
|
toolNode.Fail(err)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
result.Content = string(contentBytes)
|
||||||
|
if toolNode != nil {
|
||||||
|
toolNode.Complete(map[string]any{
|
||||||
|
"result": mcpResult.Content,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
results = append(results, result)
|
||||||
|
}
|
||||||
|
|
||||||
|
return results, hasErrors
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -104,3 +104,21 @@ func (ast *Assistant) traceAgentFail(agentNode types.Node, err error) {
|
||||||
|
|
||||||
agentNode.Fail(err)
|
agentNode.Fail(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// traceLLMRetryRequest adds a LLM retry trace node to the trace
|
||||||
|
func (ast *Assistant) traceLLMRetryRequest(ctx *context.Context, connID string, completionMessages []context.Message, completionOptions *context.CompletionOptions) {
|
||||||
|
trace, _ := ctx.Trace()
|
||||||
|
if trace == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
trace.Add(
|
||||||
|
map[string]any{"messages": completionMessages, "options": completionOptions},
|
||||||
|
types.TraceNodeOption{
|
||||||
|
Label: fmt.Sprintf("LLM %s (Tool Retry)", connID),
|
||||||
|
Type: "llm_retry",
|
||||||
|
Icon: "refresh",
|
||||||
|
Description: fmt.Sprintf("LLM %s is retrying with tool call error feedback", connID),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -5,6 +5,7 @@ import (
|
||||||
"io"
|
"io"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
jsoniter "github.com/json-iterator/go"
|
||||||
"github.com/yaoapp/yao/agent/assistant/hook"
|
"github.com/yaoapp/yao/agent/assistant/hook"
|
||||||
chatctx "github.com/yaoapp/yao/agent/context"
|
chatctx "github.com/yaoapp/yao/agent/context"
|
||||||
"github.com/yaoapp/yao/agent/message"
|
"github.com/yaoapp/yao/agent/message"
|
||||||
|
|
@ -164,3 +165,99 @@ type FileResponse struct {
|
||||||
ContentType string
|
ContentType string
|
||||||
Extension string
|
Extension string
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// MCPTool represents a simplified MCP tool for building LLM requests
|
||||||
|
// This is an internal representation used when collecting tools from MCP servers
|
||||||
|
// and preparing them for the LLM's tool calling interface
|
||||||
|
type MCPTool struct {
|
||||||
|
Name string // Formatted tool name with server prefix (e.g., "server_id__tool_name")
|
||||||
|
Description string // Tool description from MCP server
|
||||||
|
Parameters interface{} // JSON Schema for tool parameters (from MCP InputSchema)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ToolCallResult represents the result of a tool call execution
|
||||||
|
// Used to track the outcome of MCP tool invocations during agent execution
|
||||||
|
type ToolCallResult struct {
|
||||||
|
ToolCallID string // Tool call ID from the LLM (matches the ID in the LLM's tool_calls response)
|
||||||
|
Name string // Tool name (formatted with server prefix, e.g., "server_id__tool_name")
|
||||||
|
Content string // Result content (JSON string of the tool's output or error message)
|
||||||
|
Error error // Error if the call failed (nil if successful)
|
||||||
|
IsRetryableError bool // Whether the error should be sent to LLM for retry
|
||||||
|
// true: parameter/validation errors that LLM can fix (e.g., "missing required field")
|
||||||
|
// false: MCP internal errors that LLM cannot fix (e.g., "network error", "service unavailable")
|
||||||
|
}
|
||||||
|
|
||||||
|
// ServerID extracts the MCP server ID from the formatted tool name
|
||||||
|
// Example: "echo__ping" -> "echo"
|
||||||
|
func (r *ToolCallResult) ServerID() string {
|
||||||
|
serverID, _, _ := ParseMCPToolName(r.Name)
|
||||||
|
return serverID
|
||||||
|
}
|
||||||
|
|
||||||
|
// ToolName extracts the original tool name without server prefix
|
||||||
|
// Example: "echo__ping" -> "ping"
|
||||||
|
func (r *ToolCallResult) ToolName() string {
|
||||||
|
_, toolName, _ := ParseMCPToolName(r.Name)
|
||||||
|
return toolName
|
||||||
|
}
|
||||||
|
|
||||||
|
// ParsedContent extracts the actual tool return value from MCP ToolContent array
|
||||||
|
// According to MCP protocol:
|
||||||
|
// - Content is []ToolContent array
|
||||||
|
// - For "text" type, the actual value is in Text field (usually JSON string)
|
||||||
|
// - For "image" type, returns the Data field
|
||||||
|
// - For "resource" type, returns the Resource object
|
||||||
|
// If there are multiple content items, returns an array of parsed values
|
||||||
|
func (r *ToolCallResult) ParsedContent() (interface{}, error) {
|
||||||
|
if r.Content == "" {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse Content as []ToolContent
|
||||||
|
var toolContents []map[string]interface{}
|
||||||
|
if err := jsoniter.UnmarshalFromString(r.Content, &toolContents); err != nil {
|
||||||
|
// If parsing fails, return the string content directly (error message)
|
||||||
|
return r.Content, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract actual values from ToolContent items
|
||||||
|
var results []interface{}
|
||||||
|
for _, tc := range toolContents {
|
||||||
|
contentType, _ := tc["type"].(string)
|
||||||
|
|
||||||
|
switch contentType {
|
||||||
|
case "text":
|
||||||
|
// For text type, parse the Text field (usually JSON)
|
||||||
|
if textStr, ok := tc["text"].(string); ok {
|
||||||
|
// Try to parse as JSON
|
||||||
|
var parsed interface{}
|
||||||
|
if err := jsoniter.UnmarshalFromString(textStr, &parsed); err == nil {
|
||||||
|
results = append(results, parsed)
|
||||||
|
} else {
|
||||||
|
// If not JSON, return as plain string
|
||||||
|
results = append(results, textStr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case "image":
|
||||||
|
// For image type, return the data and mimeType
|
||||||
|
results = append(results, map[string]interface{}{
|
||||||
|
"type": "image",
|
||||||
|
"data": tc["data"],
|
||||||
|
"mimeType": tc["mimeType"],
|
||||||
|
})
|
||||||
|
case "resource":
|
||||||
|
// For resource type, return the resource object
|
||||||
|
results = append(results, tc["resource"])
|
||||||
|
default:
|
||||||
|
// Unknown type, return as-is
|
||||||
|
results = append(results, tc)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// If only one result, return it directly (not as array)
|
||||||
|
if len(results) == 1 {
|
||||||
|
return results[0], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return results, nil
|
||||||
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue