feat: add error handling for tool responses in agent loop
This commit is contained in:
parent
c4c259c748
commit
4b62382b31
2 changed files with 123 additions and 4 deletions
|
|
@ -1183,6 +1183,72 @@ func TestRunAgentLoop_ResponseHandledToolPublishesForUserWhenSendResponseDisable
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRunAgentLoop_ErrorToolPublishesForUserWhenSendResponseDisabled(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.Agents.Defaults.Workspace = tmpDir
|
||||||
|
cfg.Agents.Defaults.ModelName = "test-model"
|
||||||
|
cfg.Agents.Defaults.MaxTokens = 4096
|
||||||
|
cfg.Agents.Defaults.MaxToolIterations = 10
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
provider := &errorToolProvider{}
|
||||||
|
al := NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
|
||||||
|
store := media.NewFileMediaStore()
|
||||||
|
al.SetMediaStore(store)
|
||||||
|
telegramChannel := &fakeMediaChannel{fakeChannel: fakeChannel{id: "rid-telegram"}}
|
||||||
|
al.SetChannelManager(newStartedTestChannelManager(t, msgBus, store, "telegram", telegramChannel))
|
||||||
|
al.RegisterTool(&errorTool{})
|
||||||
|
|
||||||
|
defaultAgent := al.registry.GetDefaultAgent()
|
||||||
|
if defaultAgent == nil {
|
||||||
|
t.Fatal("expected default agent")
|
||||||
|
}
|
||||||
|
|
||||||
|
response, err := al.runAgentLoop(context.Background(), defaultAgent, processOptions{
|
||||||
|
Dispatch: DispatchRequest{
|
||||||
|
SessionKey: "session-error-tool",
|
||||||
|
UserMessage: "please run failing tool",
|
||||||
|
SessionScope: &session.SessionScope{
|
||||||
|
Version: session.ScopeVersionV1,
|
||||||
|
AgentID: defaultAgent.ID,
|
||||||
|
Channel: "telegram",
|
||||||
|
Dimensions: []string{"chat"},
|
||||||
|
Values: map[string]string{
|
||||||
|
"chat": "direct:chat1",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
InboundContext: &bus.InboundContext{
|
||||||
|
Channel: "telegram",
|
||||||
|
ChatID: "chat1",
|
||||||
|
ChatType: "direct",
|
||||||
|
SenderID: "user1",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
DefaultResponse: defaultResponse,
|
||||||
|
EnableSummary: false,
|
||||||
|
SendResponse: false,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("runAgentLoop() error = %v", err)
|
||||||
|
}
|
||||||
|
if response != defaultResponse {
|
||||||
|
t.Fatalf("expected default final response, got %q", response)
|
||||||
|
}
|
||||||
|
|
||||||
|
deadline := time.Now().Add(2 * time.Second)
|
||||||
|
for len(telegramChannel.sentMessages) == 0 && time.Now().Before(deadline) {
|
||||||
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
}
|
||||||
|
if len(telegramChannel.sentMessages) != 1 {
|
||||||
|
t.Fatalf("expected exactly 1 sent text message, got %d", len(telegramChannel.sentMessages))
|
||||||
|
}
|
||||||
|
if telegramChannel.sentMessages[0].Content != "simulated tool failure" {
|
||||||
|
t.Fatalf("unexpected sent text message: %+v", telegramChannel.sentMessages[0])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestAppendEventContextFields_IncludesInboundRouteAndScope(t *testing.T) {
|
func TestAppendEventContextFields_IncludesInboundRouteAndScope(t *testing.T) {
|
||||||
fields := map[string]any{}
|
fields := map[string]any{}
|
||||||
|
|
||||||
|
|
@ -1609,6 +1675,10 @@ type handledUserProvider struct {
|
||||||
calls int
|
calls int
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type errorToolProvider struct {
|
||||||
|
calls int
|
||||||
|
}
|
||||||
|
|
||||||
func (m *handledUserProvider) Chat(
|
func (m *handledUserProvider) Chat(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
messages []providers.Message,
|
messages []providers.Message,
|
||||||
|
|
@ -1635,6 +1705,32 @@ func (m *handledUserProvider) GetDefaultModel() string {
|
||||||
return "handled-user-model"
|
return "handled-user-model"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (m *errorToolProvider) Chat(
|
||||||
|
ctx context.Context,
|
||||||
|
messages []providers.Message,
|
||||||
|
tools []providers.ToolDefinition,
|
||||||
|
model string,
|
||||||
|
opts map[string]any,
|
||||||
|
) (*providers.LLMResponse, error) {
|
||||||
|
m.calls++
|
||||||
|
if m.calls == 1 {
|
||||||
|
return &providers.LLMResponse{
|
||||||
|
Content: "Running the requested tool.",
|
||||||
|
ToolCalls: []providers.ToolCall{{
|
||||||
|
ID: "call_error_tool",
|
||||||
|
Type: "function",
|
||||||
|
Name: "error_tool",
|
||||||
|
Arguments: map[string]any{},
|
||||||
|
}},
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
return &providers.LLMResponse{}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *errorToolProvider) GetDefaultModel() string {
|
||||||
|
return "error-tool-model"
|
||||||
|
}
|
||||||
|
|
||||||
type messageToolProvider struct {
|
type messageToolProvider struct {
|
||||||
calls int
|
calls int
|
||||||
}
|
}
|
||||||
|
|
@ -2021,6 +2117,8 @@ func (m *handledMediaTool) Execute(ctx context.Context, args map[string]any) *to
|
||||||
|
|
||||||
type handledUserTool struct{}
|
type handledUserTool struct{}
|
||||||
|
|
||||||
|
type errorTool struct{}
|
||||||
|
|
||||||
func (m *handledUserTool) Name() string { return "handled_user_tool" }
|
func (m *handledUserTool) Name() string { return "handled_user_tool" }
|
||||||
func (m *handledUserTool) Description() string {
|
func (m *handledUserTool) Description() string {
|
||||||
return "Returns a user-visible result and marks delivery as handled"
|
return "Returns a user-visible result and marks delivery as handled"
|
||||||
|
|
@ -2037,6 +2135,22 @@ func (m *handledUserTool) Execute(ctx context.Context, args map[string]any) *too
|
||||||
return tools.UserResult("Handled user output from tool.").WithResponseHandled()
|
return tools.UserResult("Handled user output from tool.").WithResponseHandled()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (m *errorTool) Name() string { return "error_tool" }
|
||||||
|
func (m *errorTool) Description() string {
|
||||||
|
return "Returns a tool error for testing delivery behavior"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *errorTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *errorTool) Execute(ctx context.Context, args map[string]any) *tools.ToolResult {
|
||||||
|
return tools.ErrorResult("simulated tool failure")
|
||||||
|
}
|
||||||
|
|
||||||
type handledMediaWithSteeringProvider struct {
|
type handledMediaWithSteeringProvider struct {
|
||||||
calls int
|
calls int
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -1263,15 +1263,20 @@ turnLoop:
|
||||||
allResponsesHandled = false
|
allResponsesHandled = false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
forUserContent := toolResult.ForUser
|
||||||
|
if toolResult.IsError && forUserContent == "" {
|
||||||
|
forUserContent = toolResult.ContentForLLM()
|
||||||
|
}
|
||||||
|
|
||||||
shouldSendForUser := !toolResult.Silent &&
|
shouldSendForUser := !toolResult.Silent &&
|
||||||
toolResult.ForUser != "" &&
|
forUserContent != "" &&
|
||||||
(ts.opts.SendResponse || toolResult.ResponseHandled)
|
(ts.opts.SendResponse || toolResult.ResponseHandled || toolResult.IsError)
|
||||||
if shouldSendForUser {
|
if shouldSendForUser {
|
||||||
al.bus.PublishOutbound(ctx, outboundMessageForTurn(ts, toolResult.ForUser))
|
al.bus.PublishOutbound(ctx, outboundMessageForTurn(ts, forUserContent))
|
||||||
logger.DebugCF("agent", "Sent tool result to user",
|
logger.DebugCF("agent", "Sent tool result to user",
|
||||||
map[string]any{
|
map[string]any{
|
||||||
"tool": toolName,
|
"tool": toolName,
|
||||||
"content_len": len(toolResult.ForUser),
|
"content_len": len(forUserContent),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
contentForLLM := toolResult.ContentForLLM()
|
contentForLLM := toolResult.ContentForLLM()
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue