Enhance debugging and error tracing in Assistant's streaming process
- Added detailed debug logging throughout the Stream method to trace execution flow and tool call results, improving visibility during runtime. - Implemented error tracing with the new traceAgentFail method to capture failures in agent nodes, enhancing error handling. - Updated MCP tools debugging to provide insights into tool application and validation processes, aiding in troubleshooting. - Refactored output handling in traceAgentOutput for better clarity and consistency in response management.
This commit is contained in:
parent
739a03b4cf
commit
9ea1fceda4
6 changed files with 115 additions and 79 deletions
|
|
@ -64,6 +64,7 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
|
||||||
// ================================================
|
// ================================================
|
||||||
fullMessages, err := ast.WithHistory(ctx, inputMessages, agentNode)
|
fullMessages, err := ast.WithHistory(ctx, inputMessages, agentNode)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
ast.traceAgentFail(agentNode, err)
|
||||||
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
@ -107,10 +108,20 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
|
||||||
// Execute the LLM streaming call
|
// Execute the LLM streaming call
|
||||||
completionResponse, err = ast.executeLLMStream(ctx, completionMessages, completionOptions, agentNode, streamHandler)
|
completionResponse, err = ast.executeLLMStream(ctx, completionMessages, completionOptions, agentNode, streamHandler)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
ast.traceAgentFail(agentNode, err)
|
||||||
// Send error stream_end for root stack
|
// Send error stream_end for root stack
|
||||||
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// === Debug Completion Response ===
|
||||||
|
fmt.Println("--- Debug Completion Response ----------------------")
|
||||||
|
fmt.Printf("completionResponse: %+v\n", completionResponse)
|
||||||
|
if completionResponse != nil {
|
||||||
|
fmt.Printf("ToolCalls: %+v\n", completionResponse.ToolCalls)
|
||||||
|
}
|
||||||
|
fmt.Println("----------------------------------------------------")
|
||||||
|
// === End Debug ===
|
||||||
}
|
}
|
||||||
|
|
||||||
// ================================================
|
// ================================================
|
||||||
|
|
@ -213,6 +224,7 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
|
||||||
var err error
|
var err error
|
||||||
doneResponse, err = ast.Script.Done(ctx, fullMessages, completionResponse, nil)
|
doneResponse, err = ast.Script.Done(ctx, fullMessages, completionResponse, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
ast.traceAgentFail(agentNode, err)
|
||||||
// Send error stream_end for root stack
|
// Send error stream_end for root stack
|
||||||
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
||||||
return nil, err
|
return nil, err
|
||||||
|
|
|
||||||
|
|
@ -97,6 +97,17 @@ func (ast *Assistant) buildCompletionOptions(ctx *context.Context, createRespons
|
||||||
return nil, "", fmt.Errorf("failed to apply MCP tools: %w", err)
|
return nil, "", fmt.Errorf("failed to apply MCP tools: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// === Debug MCP Tools ===
|
||||||
|
fmt.Println("--- Debug MCP Tools after applyMCPTools ---------------")
|
||||||
|
fmt.Printf("options.Tools count: %d\n", len(options.Tools))
|
||||||
|
if len(options.Tools) > 0 {
|
||||||
|
for i, tool := range options.Tools {
|
||||||
|
fmt.Printf("Tool %d: %+v\n", i, tool)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
fmt.Println("-------------------------------------------------------")
|
||||||
|
// === End Debug ===
|
||||||
|
|
||||||
return options, mcpSamplesPrompt, nil
|
return options, mcpSamplesPrompt, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,8 @@
|
||||||
package assistant
|
package assistant
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
"github.com/yaoapp/kun/log"
|
"github.com/yaoapp/kun/log"
|
||||||
"github.com/yaoapp/yao/agent/context"
|
"github.com/yaoapp/yao/agent/context"
|
||||||
"github.com/yaoapp/yao/agent/llm"
|
"github.com/yaoapp/yao/agent/llm"
|
||||||
|
|
@ -18,6 +20,12 @@ func (ast *Assistant) executeLLMStream(
|
||||||
streamHandler message.StreamFunc,
|
streamHandler message.StreamFunc,
|
||||||
) (*context.CompletionResponse, error) {
|
) (*context.CompletionResponse, error) {
|
||||||
|
|
||||||
|
// === Debug LLM Stream Start ===
|
||||||
|
fmt.Println(">>> executeLLMStream: STARTING")
|
||||||
|
fmt.Printf(">>> Messages count: %d\n", len(completionMessages))
|
||||||
|
fmt.Printf(">>> Tools count: %d\n", len(completionOptions.Tools))
|
||||||
|
// === End Debug ===
|
||||||
|
|
||||||
// Get connector object (capabilities were already set above, before stream_start)
|
// Get connector object (capabilities were already set above, before stream_start)
|
||||||
conn, capabilities, err := ast.GetConnector(ctx)
|
conn, capabilities, err := ast.GetConnector(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -44,7 +52,21 @@ func (ast *Assistant) executeLLMStream(
|
||||||
|
|
||||||
// Call the LLM Completion Stream (streamHandler was set earlier)
|
// Call the LLM Completion Stream (streamHandler was set earlier)
|
||||||
log.Trace("[AGENT] Calling LLM Stream: assistant=%s", ast.ID)
|
log.Trace("[AGENT] Calling LLM Stream: assistant=%s", ast.ID)
|
||||||
|
|
||||||
|
// === Debug LLM Stream Call ===
|
||||||
|
fmt.Println(">>> executeLLMStream: CALLING llmInstance.Stream()")
|
||||||
|
// === End Debug ===
|
||||||
|
|
||||||
completionResponse, err := llmInstance.Stream(ctx, completionMessages, completionOptions, streamHandler)
|
completionResponse, err := llmInstance.Stream(ctx, completionMessages, completionOptions, streamHandler)
|
||||||
|
|
||||||
|
// === Debug LLM Stream Return ===
|
||||||
|
fmt.Println(">>> executeLLMStream: llmInstance.Stream() RETURNED")
|
||||||
|
fmt.Printf(">>> err: %v\n", err)
|
||||||
|
if completionResponse != nil {
|
||||||
|
fmt.Printf(">>> ToolCalls: %d\n", len(completionResponse.ToolCalls))
|
||||||
|
}
|
||||||
|
// === End Debug ===
|
||||||
|
|
||||||
log.Trace("[AGENT] LLM Stream returned: assistant=%s, err=%v", ast.ID, err)
|
log.Trace("[AGENT] LLM Stream returned: assistant=%s, err=%v", ast.ID, err)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Trace("[AGENT] Calling sendStreamEndOnError")
|
log.Trace("[AGENT] Calling sendStreamEndOnError")
|
||||||
|
|
|
||||||
|
|
@ -207,19 +207,33 @@ func (ast *Assistant) executeToolCalls(ctx *agentContext.Context, toolCalls []ag
|
||||||
return nil, false
|
return nil, false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// === Debug ===
|
||||||
|
fmt.Printf(">>> executeToolCalls: START (attempt %d, toolCalls count: %d)\n", attempt, len(toolCalls))
|
||||||
|
// === End Debug ===
|
||||||
|
|
||||||
log.Trace("[Assistant MCP] Executing %d tool calls (attempt %d)", len(toolCalls), attempt)
|
log.Trace("[Assistant MCP] Executing %d tool calls (attempt %d)", len(toolCalls), attempt)
|
||||||
|
|
||||||
// Single tool call
|
// Single tool call
|
||||||
if len(toolCalls) == 1 {
|
if len(toolCalls) == 1 {
|
||||||
return ast.executeSingleToolCall(ctx, toolCalls[0])
|
fmt.Println(">>> executeToolCalls: Calling executeSingleToolCall")
|
||||||
|
results, hasErrors := ast.executeSingleToolCall(ctx, toolCalls[0])
|
||||||
|
fmt.Printf(">>> executeToolCalls: executeSingleToolCall RETURNED (hasErrors: %v)\n", hasErrors)
|
||||||
|
return results, hasErrors
|
||||||
}
|
}
|
||||||
|
|
||||||
// Multiple tool calls - try parallel first
|
// Multiple tool calls - try parallel first
|
||||||
return ast.executeMultipleToolCallsParallel(ctx, toolCalls)
|
fmt.Println(">>> executeToolCalls: Calling executeMultipleToolCallsParallel")
|
||||||
|
results, hasErrors := ast.executeMultipleToolCallsParallel(ctx, toolCalls)
|
||||||
|
fmt.Printf(">>> executeToolCalls: executeMultipleToolCallsParallel RETURNED (hasErrors: %v)\n", hasErrors)
|
||||||
|
return results, hasErrors
|
||||||
}
|
}
|
||||||
|
|
||||||
// executeSingleToolCall executes a single tool call with trace logging
|
// executeSingleToolCall executes a single tool call with trace logging
|
||||||
func (ast *Assistant) executeSingleToolCall(ctx *agentContext.Context, toolCall agentContext.ToolCall) ([]ToolCallResult, bool) {
|
func (ast *Assistant) executeSingleToolCall(ctx *agentContext.Context, toolCall agentContext.ToolCall) ([]ToolCallResult, bool) {
|
||||||
|
// === Debug ===
|
||||||
|
fmt.Printf(">>> executeSingleToolCall: START (tool: %s)\n", toolCall.Function.Name)
|
||||||
|
// === End Debug ===
|
||||||
|
|
||||||
trace, _ := ctx.Trace()
|
trace, _ := ctx.Trace()
|
||||||
|
|
||||||
// Use the agent context for cancellation and timeout control
|
// Use the agent context for cancellation and timeout control
|
||||||
|
|
@ -234,6 +248,7 @@ func (ast *Assistant) executeSingleToolCall(ctx *agentContext.Context, toolCall
|
||||||
}
|
}
|
||||||
|
|
||||||
// Parse tool name
|
// Parse tool name
|
||||||
|
fmt.Println(">>> executeSingleToolCall: Parsing tool name")
|
||||||
serverID, toolName, ok := ParseMCPToolName(toolCall.Function.Name)
|
serverID, toolName, ok := ParseMCPToolName(toolCall.Function.Name)
|
||||||
if !ok {
|
if !ok {
|
||||||
result.Error = fmt.Errorf("invalid MCP tool name format: %s", toolCall.Function.Name)
|
result.Error = fmt.Errorf("invalid MCP tool name format: %s", toolCall.Function.Name)
|
||||||
|
|
@ -319,22 +334,30 @@ func (ast *Assistant) executeSingleToolCall(ctx *agentContext.Context, toolCall
|
||||||
|
|
||||||
// Validate arguments against tool schema if available
|
// Validate arguments against tool schema if available
|
||||||
if toolSchema != nil {
|
if toolSchema != nil {
|
||||||
|
fmt.Println(">>> executeSingleToolCall: Validating arguments against schema")
|
||||||
if err := gouJson.Validate(args, toolSchema); err != nil {
|
if err := gouJson.Validate(args, toolSchema); err != nil {
|
||||||
|
fmt.Printf(">>> executeSingleToolCall: Validation FAILED: %v\n", err)
|
||||||
result.Error = fmt.Errorf("argument validation failed: %w", err)
|
result.Error = fmt.Errorf("argument validation failed: %w", err)
|
||||||
result.Content = result.Error.Error()
|
result.Content = result.Error.Error()
|
||||||
result.IsRetryableError = true // Validation error is retryable by LLM
|
result.IsRetryableError = true // Validation error is retryable by LLM
|
||||||
log.Error("[Assistant MCP] %v", result.Error)
|
log.Error("[Assistant MCP] %v", result.Error)
|
||||||
if toolNode != nil {
|
if toolNode != nil {
|
||||||
|
fmt.Println(">>> executeSingleToolCall: Failing toolNode due to validation error")
|
||||||
toolNode.Fail(result.Error)
|
toolNode.Fail(result.Error)
|
||||||
|
fmt.Println(">>> executeSingleToolCall: toolNode.Fail() finished")
|
||||||
}
|
}
|
||||||
|
fmt.Println(">>> executeSingleToolCall: RETURNING with validation error")
|
||||||
return []ToolCallResult{result}, true
|
return []ToolCallResult{result}, true
|
||||||
}
|
}
|
||||||
|
fmt.Println(">>> executeSingleToolCall: Validation PASSED")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Call the tool
|
// Call the tool
|
||||||
log.Trace("[Assistant MCP] Calling tool: %s (server: %s)", toolName, serverID)
|
log.Trace("[Assistant MCP] Calling tool: %s (server: %s)", toolName, serverID)
|
||||||
|
fmt.Printf(">>> executeSingleToolCall: CALLING client.CallTool (tool: %s, server: %s)\n", toolName, serverID)
|
||||||
callResult, err := client.CallTool(mcpCtx, toolName, args)
|
callResult, err := client.CallTool(mcpCtx, toolName, args)
|
||||||
|
fmt.Printf(">>> executeSingleToolCall: client.CallTool RETURNED (err: %v)\n", err)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
result.Error = fmt.Errorf("tool call failed: %w", err)
|
result.Error = fmt.Errorf("tool call failed: %w", err)
|
||||||
result.Content = result.Error.Error()
|
result.Content = result.Error.Error()
|
||||||
|
|
@ -344,6 +367,7 @@ func (ast *Assistant) executeSingleToolCall(ctx *agentContext.Context, toolCall
|
||||||
if toolNode != nil {
|
if toolNode != nil {
|
||||||
toolNode.Fail(result.Error)
|
toolNode.Fail(result.Error)
|
||||||
}
|
}
|
||||||
|
fmt.Println(">>> executeSingleToolCall: RETURNING with error")
|
||||||
return []ToolCallResult{result}, true
|
return []ToolCallResult{result}, true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -370,11 +394,14 @@ func (ast *Assistant) executeSingleToolCall(ctx *agentContext.Context, toolCall
|
||||||
log.Trace("[Assistant MCP] Tool call succeeded: %s", toolName)
|
log.Trace("[Assistant MCP] Tool call succeeded: %s", toolName)
|
||||||
|
|
||||||
if toolNode != nil {
|
if toolNode != nil {
|
||||||
|
fmt.Println(">>> executeSingleToolCall: Completing toolNode")
|
||||||
toolNode.Complete(map[string]any{
|
toolNode.Complete(map[string]any{
|
||||||
"result": callResult,
|
"result": callResult,
|
||||||
})
|
})
|
||||||
|
fmt.Println(">>> executeSingleToolCall: toolNode.Complete() finished")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fmt.Println(">>> executeSingleToolCall: RETURNING success")
|
||||||
return []ToolCallResult{result}, false
|
return []ToolCallResult{result}, false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -558,7 +585,16 @@ func (ast *Assistant) executeServerToolsParallelWithTrace(mcpCtx context.Context
|
||||||
// Create parallel trace nodes
|
// Create parallel trace nodes
|
||||||
var toolNodes []types.Node
|
var toolNodes []types.Node
|
||||||
if trace != nil && len(parallelInputs) > 0 {
|
if trace != nil && len(parallelInputs) > 0 {
|
||||||
toolNodes, _ = trace.Parallel(parallelInputs)
|
fmt.Printf(">>> executeServerToolsParallelWithTrace: Creating %d parallel trace nodes\n", len(parallelInputs))
|
||||||
|
var err error
|
||||||
|
toolNodes, err = trace.Parallel(parallelInputs)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf(">>> executeServerToolsParallelWithTrace: trace.Parallel() FAILED: %v\n", err)
|
||||||
|
} else {
|
||||||
|
fmt.Printf(">>> executeServerToolsParallelWithTrace: Created %d trace nodes\n", len(toolNodes))
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
fmt.Printf(">>> executeServerToolsParallelWithTrace: NOT creating trace nodes (trace: %v, inputs: %d)\n", trace != nil, len(parallelInputs))
|
||||||
}
|
}
|
||||||
|
|
||||||
// Call tools in parallel
|
// Call tools in parallel
|
||||||
|
|
@ -617,9 +653,13 @@ func (ast *Assistant) executeServerToolsParallelWithTrace(mcpCtx context.Context
|
||||||
} else {
|
} else {
|
||||||
// Success
|
// Success
|
||||||
if toolNode != nil {
|
if toolNode != nil {
|
||||||
|
fmt.Printf(">>> executeServerToolsParallelWithTrace: Completing toolNode %d\n", i)
|
||||||
toolNode.Complete(map[string]any{
|
toolNode.Complete(map[string]any{
|
||||||
"result": mcpResult.Content,
|
"result": mcpResult.Content,
|
||||||
})
|
})
|
||||||
|
fmt.Printf(">>> executeServerToolsParallelWithTrace: toolNode %d completed\n", i)
|
||||||
|
} else {
|
||||||
|
fmt.Printf(">>> executeServerToolsParallelWithTrace: toolNode %d is nil!\n", i)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -89,11 +89,13 @@ func (ast *Assistant) traceAgentOutput(agentNode types.Node, createResponse *con
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
agentNode.SetOutput(context.Response{
|
output := context.Response{
|
||||||
Create: createResponse,
|
Create: createResponse,
|
||||||
Done: doneResponse,
|
Done: doneResponse,
|
||||||
Completion: completionResponse,
|
Completion: completionResponse,
|
||||||
})
|
}
|
||||||
|
|
||||||
|
agentNode.Complete(output)
|
||||||
}
|
}
|
||||||
|
|
||||||
// traceAgentFail marks the agent trace node as failed
|
// traceAgentFail marks the agent trace node as failed
|
||||||
|
|
|
||||||
|
|
@ -213,7 +213,6 @@ func (p *Provider) Stream(ctx *context.Context, messages []context.Message, opti
|
||||||
}
|
}
|
||||||
|
|
||||||
maxRetries := 3
|
maxRetries := 3
|
||||||
maxValidationRetries := 3
|
|
||||||
var lastErr error
|
var lastErr error
|
||||||
|
|
||||||
// Get Go context for cancellation support
|
// Get Go context for cancellation support
|
||||||
|
|
@ -305,43 +304,14 @@ func (p *Provider) Stream(ctx *context.Context, messages []context.Message, opti
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check if error is tool call validation failure
|
// Note: Tool call validation errors should not reach here anymore
|
||||||
|
// because we now pass through validation failures to Agent layer
|
||||||
|
// This check is kept for safety but should not trigger
|
||||||
if isToolCallValidationError(err) {
|
if isToolCallValidationError(err) {
|
||||||
// Handle tool call validation retry with feedback to LLM
|
if trace != nil {
|
||||||
validationRetryMessages := currentMessages
|
trace.Debug("Tool call validation error (unexpected, should be handled differently)", map[string]any{
|
||||||
for validationAttempt := 0; validationAttempt < maxValidationRetries; validationAttempt++ {
|
"error": err.Error(),
|
||||||
if trace != nil {
|
|
||||||
trace.Warn("Tool call validation failed", map[string]any{
|
|
||||||
"attempt": validationAttempt + 1,
|
|
||||||
"max_retries": maxValidationRetries,
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// Add error feedback to conversation history
|
|
||||||
validationRetryMessages = append(validationRetryMessages, context.Message{
|
|
||||||
Role: context.RoleSystem,
|
|
||||||
Content: fmt.Sprintf("Tool call validation error: %v. Please correct the tool call arguments to match the required schema.", err),
|
|
||||||
})
|
})
|
||||||
|
|
||||||
// Retry with feedback
|
|
||||||
response, err = p.streamWithRetry(ctx, validationRetryMessages, options, handler)
|
|
||||||
if err == nil {
|
|
||||||
return response, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Check if still validation error
|
|
||||||
if !isToolCallValidationError(err) {
|
|
||||||
// Different error type, break out of validation retry loop
|
|
||||||
lastErr = err
|
|
||||||
break
|
|
||||||
}
|
|
||||||
lastErr = err
|
|
||||||
}
|
|
||||||
|
|
||||||
// If we exhausted validation retries, return the error
|
|
||||||
if isToolCallValidationError(lastErr) {
|
|
||||||
return nil, fmt.Errorf("tool call validation failed after %d retries: %w", maxValidationRetries, lastErr)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -803,12 +773,21 @@ func (p *Provider) streamWithRetry(ctx *context.Context, messages []context.Mess
|
||||||
response.ToolCalls = toolCalls
|
response.ToolCalls = toolCalls
|
||||||
|
|
||||||
// Validate tool call results if schema is provided
|
// Validate tool call results if schema is provided
|
||||||
|
// Note: If validation fails, we log the error but DO NOT return error
|
||||||
|
// Instead, we let the response through so Agent layer can handle it
|
||||||
|
// Agent layer will re-validate and provide better error feedback to LLM
|
||||||
if err := p.validateToolCallResults(options, toolCalls); err != nil {
|
if err := p.validateToolCallResults(options, toolCalls); err != nil {
|
||||||
|
// Log validation error
|
||||||
|
if trace, _ := ctx.Trace(); trace != nil {
|
||||||
|
trace.Warn("Tool call validation failed at LLM layer, passing to Agent layer for handling", map[string]any{
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
}
|
||||||
// End current message
|
// End current message
|
||||||
messageTracker.endMessage(handler)
|
messageTracker.endMessage(handler)
|
||||||
|
|
||||||
// Tool call validation failed, need to retry with error feedback
|
// Continue and return response (don't return error)
|
||||||
return nil, fmt.Errorf("tool call validation failed: %w", err)
|
// Agent layer will handle validation and retry
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -829,7 +808,6 @@ func (p *Provider) Post(ctx *context.Context, messages []context.Message, option
|
||||||
}
|
}
|
||||||
|
|
||||||
maxRetries := 3
|
maxRetries := 3
|
||||||
maxValidationRetries := 3
|
|
||||||
var lastErr error
|
var lastErr error
|
||||||
|
|
||||||
// Get Go context for cancellation support
|
// Get Go context for cancellation support
|
||||||
|
|
@ -880,43 +858,14 @@ func (p *Provider) Post(ctx *context.Context, messages []context.Message, option
|
||||||
}
|
}
|
||||||
lastErr = err
|
lastErr = err
|
||||||
|
|
||||||
// Check if error is tool call validation failure
|
// Note: Tool call validation errors should not reach here anymore
|
||||||
|
// because we now pass through validation failures to Agent layer
|
||||||
|
// This check is kept for safety but should not trigger
|
||||||
if isToolCallValidationError(err) {
|
if isToolCallValidationError(err) {
|
||||||
// Handle tool call validation retry with feedback to LLM
|
if trace != nil {
|
||||||
validationRetryMessages := currentMessages
|
trace.Debug("Tool call validation error in Post (unexpected, should be handled differently)", map[string]any{
|
||||||
for validationAttempt := 0; validationAttempt < maxValidationRetries; validationAttempt++ {
|
"error": err.Error(),
|
||||||
if trace != nil {
|
|
||||||
trace.Warn("Tool call validation failed", map[string]any{
|
|
||||||
"attempt": validationAttempt + 1,
|
|
||||||
"max_retries": maxValidationRetries,
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// Add error feedback to conversation history
|
|
||||||
validationRetryMessages = append(validationRetryMessages, context.Message{
|
|
||||||
Role: context.RoleSystem,
|
|
||||||
Content: fmt.Sprintf("Tool call validation error: %v. Please correct the tool call arguments to match the required schema.", err),
|
|
||||||
})
|
})
|
||||||
|
|
||||||
// Retry with feedback
|
|
||||||
response, err = p.postWithRetry(ctx, validationRetryMessages, options)
|
|
||||||
if err == nil {
|
|
||||||
return response, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Check if still validation error
|
|
||||||
if !isToolCallValidationError(err) {
|
|
||||||
// Different error type, break out of validation retry loop
|
|
||||||
lastErr = err
|
|
||||||
break
|
|
||||||
}
|
|
||||||
lastErr = err
|
|
||||||
}
|
|
||||||
|
|
||||||
// If we exhausted validation retries, return the error
|
|
||||||
if isToolCallValidationError(lastErr) {
|
|
||||||
return nil, fmt.Errorf("tool call validation failed after %d retries: %w", maxValidationRetries, lastErr)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue