fix(agent): resolve lint issues in final turn render

This commit is contained in:
Anton Bogdanovich 2026-05-09 18:34:12 -07:00
parent 5d929f3a5a
commit d386df4e2d
4 changed files with 423 additions and 158 deletions

View file

@ -33,7 +33,8 @@ func appendTurnActionRecord(
} }
if n := len(records); n > 0 { if n := len(records); n > 0 {
prev := records[n-1] prev := records[n-1]
if prev.Source == rec.Source && prev.Tool == rec.Tool && prev.Text == rec.Text && prev.Error == rec.Error { if prev.Source == rec.Source && prev.Tool == rec.Tool && prev.Text == rec.Text &&
prev.Error == rec.Error {
return records return records
} }
} }
@ -70,10 +71,18 @@ func buildFinalTurnRenderInstruction(exec *turnExecution) string {
b.WriteString("Write the final user-facing reply for this already-completed turn.\n") b.WriteString("Write the final user-facing reply for this already-completed turn.\n")
b.WriteString("Use the same language and general style as the conversation.\n") b.WriteString("Use the same language and general style as the conversation.\n")
b.WriteString("Do not call tools.\n") b.WriteString("Do not call tools.\n")
b.WriteString("Answer the full accumulated user request across this turn, not only the latest follow-up.\n") b.WriteString(
b.WriteString("If a later follow-up clearly corrected, narrowed, or replaced an earlier request, follow the latest clarified intent.\n") "Answer the full accumulated user request across this turn, not only the latest follow-up.\n",
b.WriteString("If later follow-ups added to earlier requests, include the completed additive results together.\n") )
b.WriteString("Use only the facts already present in the conversation and tool results. Do not invent missing results.\n") b.WriteString(
"If a later follow-up clearly corrected, narrowed, or replaced an earlier request, follow the latest clarified intent.\n",
)
b.WriteString(
"If later follow-ups added to earlier requests, include the completed additive results together.\n",
)
b.WriteString(
"Use only the facts already present in the conversation and tool results. Do not invent missing results.\n",
)
b.WriteString("Keep the reply concise and natural.\n") b.WriteString("Keep the reply concise and natural.\n")
if exec == nil || len(exec.actionLog) == 0 { if exec == nil || len(exec.actionLog) == 0 {
@ -96,7 +105,7 @@ func buildFinalTurnRenderInstruction(exec *turnExecution) string {
return b.String() return b.String()
} }
b.WriteString("\nExplicit user-facing outcomes recorded during the turn:\n") b.WriteString("\nExplicit user-facing outcomes recorded during the turn:\n")
b.WriteString(string(raw)) _, _ = b.Write(raw)
return b.String() return b.String()
} }
@ -128,7 +137,7 @@ func tryRenderFinalTurnReply(
}) })
opts := map[string]any{ opts := map[string]any{
"max_tokens": min(ts.agent.MaxTokens, 800), "max_tokens": minInt(ts.agent.MaxTokens, 800),
"temperature": 0.2, "temperature": 0.2,
"prompt_cache_key": ts.agent.ID, "prompt_cache_key": ts.agent.ID,
} }
@ -186,7 +195,7 @@ func shouldFinalizeAfterToolLoopWithRender(al *AgentLoop, exec *turnExecution) b
return !exec.allResponsesHandled return !exec.allResponsesHandled
} }
func min(a, b int) int { func minInt(a, b int) int {
if a < b { if a < b {
return a return a
} }

View file

@ -52,7 +52,10 @@ func (f *fakeMediaChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([
return nil, nil return nil, nil
} }
func (f *fakeMediaChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) { func (f *fakeMediaChannel) SendMedia(
ctx context.Context,
msg bus.OutboundMediaMessage,
) ([]string, error) {
f.sentMedia = append(f.sentMedia, msg) f.sentMedia = append(f.sentMedia, msg)
return nil, nil return nil, nil
} }
@ -75,11 +78,17 @@ func (m *recordingChannelManager) SendMessage(ctx context.Context, msg bus.Outbo
return nil return nil
} }
func (m *recordingChannelManager) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error { func (m *recordingChannelManager) SendMedia(
ctx context.Context,
msg bus.OutboundMediaMessage,
) error {
return nil return nil
} }
func (m *recordingChannelManager) SendPlaceholder(ctx context.Context, channel, chatID string) bool { func (m *recordingChannelManager) SendPlaceholder(
ctx context.Context,
channel, chatID string,
) bool {
return false return false
} }
@ -261,9 +270,11 @@ func TestPublishResponseIfNeeded_DismissesToolFeedbackWhenMessageToolAlreadySent
t.Fatal("expected default agent") t.Fatal("expected default agent")
} }
mt := tools.NewMessageTool() mt := tools.NewMessageTool()
mt.SetSendCallback(func(ctx context.Context, channel, chatID, content, replyToMessageID string) error { mt.SetSendCallback(
return nil func(ctx context.Context, channel, chatID, content, replyToMessageID string) error {
}) return nil
},
)
defaultAgent.Tools.Register(mt) defaultAgent.Tools.Register(mt)
result := mt.Execute( result := mt.Execute(
@ -277,7 +288,13 @@ func TestPublishResponseIfNeeded_DismissesToolFeedbackWhenMessageToolAlreadySent
if result == nil || result.IsError { if result == nil || result.IsError {
t.Fatalf("message tool execute failed: %+v", result) t.Fatalf("message tool execute failed: %+v", result)
} }
al.PublishResponseIfNeeded(context.Background(), "telegram", "-100123", "session-1", "final reply") al.PublishResponseIfNeeded(
context.Background(),
"telegram",
"-100123",
"session-1",
"final reply",
)
if got := cm.dismissed; len(got) != 1 || got[0] != "telegram:-100123" { if got := cm.dismissed; len(got) != 1 || got[0] != "telegram:-100123" {
t.Fatalf("dismissed = %v, want [telegram:-100123]", got) t.Fatalf("dismissed = %v, want [telegram:-100123]", got)
@ -451,7 +468,10 @@ func TestProcessMessage_BtwCommandRunsWithoutPersistingHistory(t *testing.T) {
t.Fatal("provider did not receive any messages") t.Fatal("provider did not receive any messages")
} }
if len(provider.lastMessages) != 4 { if len(provider.lastMessages) != 4 {
t.Fatalf("provider messages len = %d, want 4 (system + prior history + user)", len(provider.lastMessages)) t.Fatalf(
"provider messages len = %d, want 4 (system + prior history + user)",
len(provider.lastMessages),
)
} }
if !reflect.DeepEqual(provider.lastMessages[1:3], initialHistory) { if !reflect.DeepEqual(provider.lastMessages[1:3], initialHistory) {
@ -511,7 +531,10 @@ func TestProcessMessage_BtwCommandIncludesRequestContextAndMedia(t *testing.T) {
if !strings.Contains(systemPrompt, "## Current Session\nChannel: discord\nChat ID: group-1") { if !strings.Contains(systemPrompt, "## Current Session\nChannel: discord\nChat ID: group-1") {
t.Fatalf("system prompt missing current session context:\n%s", systemPrompt) t.Fatalf("system prompt missing current session context:\n%s", systemPrompt)
} }
if !strings.Contains(systemPrompt, "## Current Sender\nCurrent sender: Alice (ID: discord:123)") { if !strings.Contains(
systemPrompt,
"## Current Sender\nCurrent sender: Alice (ID: discord:123)",
) {
t.Fatalf("system prompt missing current sender context:\n%s", systemPrompt) t.Fatalf("system prompt missing current sender context:\n%s", systemPrompt)
} }
@ -587,7 +610,11 @@ func TestProcessMessage_BtwCommandUsesIsolatedProvider(t *testing.T) {
// Verify main session history was NOT modified // Verify main session history was NOT modified
currentHistory := defaultAgent.Sessions.GetHistory(mainSessionKey) currentHistory := defaultAgent.Sessions.GetHistory(mainSessionKey)
if !reflect.DeepEqual(currentHistory, initialHistory) { if !reflect.DeepEqual(currentHistory, initialHistory) {
t.Fatalf("main session history was modified:\ngot %#v\nwant %#v", currentHistory, initialHistory) t.Fatalf(
"main session history was modified:\ngot %#v\nwant %#v",
currentHistory,
initialHistory,
)
} }
} }
@ -1098,7 +1125,9 @@ func TestProcessMessage_MediaToolHandledSkipsFollowUpLLMAndFinalText(t *testing.
store := media.NewFileMediaStore() store := media.NewFileMediaStore()
al.SetMediaStore(store) al.SetMediaStore(store)
telegramChannel := &fakeMediaChannel{fakeChannel: fakeChannel{id: "rid-telegram"}} telegramChannel := &fakeMediaChannel{fakeChannel: fakeChannel{id: "rid-telegram"}}
al.SetChannelManager(newStartedTestChannelManager(t, msgBus, store, "telegram", telegramChannel)) al.SetChannelManager(
newStartedTestChannelManager(t, msgBus, store, "telegram", telegramChannel),
)
imagePath := filepath.Join(tmpDir, "screen.png") imagePath := filepath.Join(tmpDir, "screen.png")
if err := os.WriteFile(imagePath, []byte("fake screenshot"), 0o644); err != nil { if err := os.WriteFile(imagePath, []byte("fake screenshot"), 0o644); err != nil {
@ -1120,7 +1149,10 @@ func TestProcessMessage_MediaToolHandledSkipsFollowUpLLMAndFinalText(t *testing.
t.Fatalf("processMessage() error = %v", err) t.Fatalf("processMessage() error = %v", err)
} }
if response != "" { if response != "" {
t.Fatalf("expected no final response when media tool already handled delivery, got %q", response) t.Fatalf(
"expected no final response when media tool already handled delivery, got %q",
response,
)
} }
if provider.calls != 1 { if provider.calls != 1 {
t.Fatalf("expected exactly 1 LLM call, got %d", provider.calls) t.Fatalf("expected exactly 1 LLM call, got %d", provider.calls)
@ -1133,13 +1165,20 @@ func TestProcessMessage_MediaToolHandledSkipsFollowUpLLMAndFinalText(t *testing.
} }
if len(telegramChannel.sentMedia) != 1 { if len(telegramChannel.sentMedia) != 1 {
t.Fatalf("expected exactly 1 synchronously sent media message, got %d", len(telegramChannel.sentMedia)) t.Fatalf(
"expected exactly 1 synchronously sent media message, got %d",
len(telegramChannel.sentMedia),
)
} }
if telegramChannel.sentMedia[0].Channel != "telegram" || telegramChannel.sentMedia[0].ChatID != "chat1" { if telegramChannel.sentMedia[0].Channel != "telegram" ||
telegramChannel.sentMedia[0].ChatID != "chat1" {
t.Fatalf("unexpected sent media target: %+v", telegramChannel.sentMedia[0]) t.Fatalf("unexpected sent media target: %+v", telegramChannel.sentMedia[0])
} }
if len(telegramChannel.sentMedia[0].Parts) != 1 { if len(telegramChannel.sentMedia[0].Parts) != 1 {
t.Fatalf("expected exactly 1 sent media part, got %d", len(telegramChannel.sentMedia[0].Parts)) t.Fatalf(
"expected exactly 1 sent media part, got %d",
len(telegramChannel.sentMedia[0].Parts),
)
} }
select { select {
@ -1161,22 +1200,29 @@ func TestProcessMessage_MediaToolHandledSkipsFollowUpLLMAndFinalText(t *testing.
if err != nil { if err != nil {
t.Fatalf("resolveMessageRoute() error = %v", err) t.Fatalf("resolveMessageRoute() error = %v", err)
} }
sessionKey := resolveScopeKey(al.allocateRouteSession(route, testInboundMessage(bus.InboundMessage{ sessionKey := resolveScopeKey(
Channel: "telegram", al.allocateRouteSession(route, testInboundMessage(bus.InboundMessage{
ChatID: "chat1", Channel: "telegram",
SenderID: "user1", ChatID: "chat1",
Content: "take a screenshot of the screen and send it to me", SenderID: "user1",
})).SessionKey, "") Content: "take a screenshot of the screen and send it to me",
})).SessionKey,
"",
)
history := defaultAgent.Sessions.GetHistory(sessionKey) history := defaultAgent.Sessions.GetHistory(sessionKey)
if len(history) == 0 { if len(history) == 0 {
t.Fatal("expected session history to be saved") t.Fatal("expected session history to be saved")
} }
last := history[len(history)-1] last := history[len(history)-1]
if last.Role != "assistant" || last.Content != "Requested output delivered via tool attachment." { if last.Role != "assistant" ||
last.Content != "Requested output delivered via tool attachment." {
t.Fatalf("expected handled assistant summary in history, got %+v", last) t.Fatalf("expected handled assistant summary in history, got %+v", last)
} }
if len(last.Attachments) != 1 { if len(last.Attachments) != 1 {
t.Fatalf("expected handled assistant summary attachments in history, got %+v", last.Attachments) t.Fatalf(
"expected handled assistant summary attachments in history, got %+v",
last.Attachments,
)
} }
} }
@ -1200,7 +1246,9 @@ func TestProcessMessage_HandledToolProcessesQueuedSteeringBeforeReturning(t *tes
store := media.NewFileMediaStore() store := media.NewFileMediaStore()
al.SetMediaStore(store) al.SetMediaStore(store)
telegramChannel := &fakeMediaChannel{fakeChannel: fakeChannel{id: "rid-telegram"}} telegramChannel := &fakeMediaChannel{fakeChannel: fakeChannel{id: "rid-telegram"}}
al.SetChannelManager(newStartedTestChannelManager(t, msgBus, store, "telegram", telegramChannel)) al.SetChannelManager(
newStartedTestChannelManager(t, msgBus, store, "telegram", telegramChannel),
)
imagePath := filepath.Join(tmpDir, "screen-steering.png") imagePath := filepath.Join(tmpDir, "screen-steering.png")
if err := os.WriteFile(imagePath, []byte("fake screenshot"), 0o644); err != nil { if err := os.WriteFile(imagePath, []byte("fake screenshot"), 0o644); err != nil {
@ -1229,7 +1277,10 @@ func TestProcessMessage_HandledToolProcessesQueuedSteeringBeforeReturning(t *tes
t.Fatalf("expected 2 LLM calls after queued steering, got %d", provider.calls) t.Fatalf("expected 2 LLM calls after queued steering, got %d", provider.calls)
} }
if len(telegramChannel.sentMedia) != 1 { if len(telegramChannel.sentMedia) != 1 {
t.Fatalf("expected exactly 1 synchronously sent media message, got %d", len(telegramChannel.sentMedia)) t.Fatalf(
"expected exactly 1 synchronously sent media message, got %d",
len(telegramChannel.sentMedia),
)
} }
} }
@ -1248,7 +1299,9 @@ func TestRunAgentLoop_ResponseHandledToolPublishesForUserWhenSendResponseDisable
store := media.NewFileMediaStore() store := media.NewFileMediaStore()
al.SetMediaStore(store) al.SetMediaStore(store)
telegramChannel := &fakeMediaChannel{fakeChannel: fakeChannel{id: "rid-telegram"}} telegramChannel := &fakeMediaChannel{fakeChannel: fakeChannel{id: "rid-telegram"}}
al.SetChannelManager(newStartedTestChannelManager(t, msgBus, store, "telegram", telegramChannel)) al.SetChannelManager(
newStartedTestChannelManager(t, msgBus, store, "telegram", telegramChannel),
)
al.RegisterTool(&handledUserTool{}) al.RegisterTool(&handledUserTool{})
defaultAgent := al.registry.GetDefaultAgent() defaultAgent := al.registry.GetDefaultAgent()
@ -1298,10 +1351,17 @@ func TestRunAgentLoop_ResponseHandledToolPublishesForUserWhenSendResponseDisable
t.Fatalf("unexpected sent text message: %+v", telegramChannel.sentMessages[0]) t.Fatalf("unexpected sent text message: %+v", telegramChannel.sentMessages[0])
} }
if telegramChannel.sentMessages[0].AgentID != defaultAgent.ID { if telegramChannel.sentMessages[0].AgentID != defaultAgent.ID {
t.Fatalf("sent text agent_id = %q, want %q", telegramChannel.sentMessages[0].AgentID, defaultAgent.ID) t.Fatalf(
"sent text agent_id = %q, want %q",
telegramChannel.sentMessages[0].AgentID,
defaultAgent.ID,
)
} }
if telegramChannel.sentMessages[0].SessionKey != "session-1" { if telegramChannel.sentMessages[0].SessionKey != "session-1" {
t.Fatalf("sent text session_key = %q, want session-1", telegramChannel.sentMessages[0].SessionKey) t.Fatalf(
"sent text session_key = %q, want session-1",
telegramChannel.sentMessages[0].SessionKey,
)
} }
if telegramChannel.sentMessages[0].Scope == nil || if telegramChannel.sentMessages[0].Scope == nil ||
telegramChannel.sentMessages[0].Scope.Values["chat"] != "direct:chat1" { telegramChannel.sentMessages[0].Scope.Values["chat"] != "direct:chat1" {
@ -1505,7 +1565,9 @@ func TestProcessMessage_MediaArtifactCanBeForwardedBySendFile(t *testing.T) {
store := media.NewFileMediaStore() store := media.NewFileMediaStore()
al.SetMediaStore(store) al.SetMediaStore(store)
telegramChannel := &fakeMediaChannel{fakeChannel: fakeChannel{id: "rid-telegram"}} telegramChannel := &fakeMediaChannel{fakeChannel: fakeChannel{id: "rid-telegram"}}
al.SetChannelManager(newStartedTestChannelManager(t, msgBus, store, "telegram", telegramChannel)) al.SetChannelManager(
newStartedTestChannelManager(t, msgBus, store, "telegram", telegramChannel),
)
mediaDir := media.TempDir() mediaDir := media.TempDir()
if err := os.MkdirAll(mediaDir, 0o700); err != nil { if err := os.MkdirAll(mediaDir, 0o700); err != nil {
@ -1538,13 +1600,20 @@ func TestProcessMessage_MediaArtifactCanBeForwardedBySendFile(t *testing.T) {
} }
if len(telegramChannel.sentMedia) != 1 { if len(telegramChannel.sentMedia) != 1 {
t.Fatalf("expected exactly 1 synchronously sent media message, got %d", len(telegramChannel.sentMedia)) t.Fatalf(
"expected exactly 1 synchronously sent media message, got %d",
len(telegramChannel.sentMedia),
)
} }
if telegramChannel.sentMedia[0].Channel != "telegram" || telegramChannel.sentMedia[0].ChatID != "chat1" { if telegramChannel.sentMedia[0].Channel != "telegram" ||
telegramChannel.sentMedia[0].ChatID != "chat1" {
t.Fatalf("unexpected sent media target: %+v", telegramChannel.sentMedia[0]) t.Fatalf("unexpected sent media target: %+v", telegramChannel.sentMedia[0])
} }
if len(telegramChannel.sentMedia[0].Parts) != 1 { if len(telegramChannel.sentMedia[0].Parts) != 1 {
t.Fatalf("expected exactly 1 sent media part, got %d", len(telegramChannel.sentMedia[0].Parts)) t.Fatalf(
"expected exactly 1 sent media part, got %d",
len(telegramChannel.sentMedia[0].Parts),
)
} }
select { select {
@ -2014,7 +2083,10 @@ func TestToolFeedbackExplanationFromResponse_UsesExplicitToolCallExtraContent(t
got := toolFeedbackExplanationFromResponse(response, messages) got := toolFeedbackExplanationFromResponse(response, messages)
if got != "Read README.md first to confirm the current project structure." { if got != "Read README.md first to confirm the current project structure." {
t.Fatalf("toolFeedbackExplanationFromResponse() = %q, want explicit tool feedback explanation", got) t.Fatalf(
"toolFeedbackExplanationFromResponse() = %q, want explicit tool feedback explanation",
got,
)
} }
} }
@ -2042,10 +2114,16 @@ func TestToolFeedbackExplanationForToolCall_PrefersToolSpecificExtraContent(t *t
got1 := toolFeedbackExplanationForToolCall(response, response.ToolCalls[0], nil) got1 := toolFeedbackExplanationForToolCall(response, response.ToolCalls[0], nil)
got2 := toolFeedbackExplanationForToolCall(response, response.ToolCalls[1], nil) got2 := toolFeedbackExplanationForToolCall(response, response.ToolCalls[1], nil)
if got1 != "Read README.md first." { if got1 != "Read README.md first." {
t.Fatalf("toolFeedbackExplanationForToolCall() first = %q, want tool-specific explanation", got1) t.Fatalf(
"toolFeedbackExplanationForToolCall() first = %q, want tool-specific explanation",
got1,
)
} }
if got2 != "Update config example after reading it." { if got2 != "Update config example after reading it." {
t.Fatalf("toolFeedbackExplanationForToolCall() second = %q, want tool-specific explanation", got2) t.Fatalf(
"toolFeedbackExplanationForToolCall() second = %q, want tool-specific explanation",
got2,
)
} }
} }
@ -2091,7 +2169,10 @@ func TestToolFeedbackExplanationFromResponse_DoesNotUseReasoningContent(t *testi
got := toolFeedbackExplanationFromResponse(response, messages) got := toolFeedbackExplanationFromResponse(response, messages)
want := utils.ToolFeedbackContinuationHint + ": Inspect README.md and update the config example." want := utils.ToolFeedbackContinuationHint + ": Inspect README.md and update the config example."
if got != want { if got != want {
t.Fatalf("toolFeedbackExplanationFromResponse() = %q, want latest user content fallback", got) t.Fatalf(
"toolFeedbackExplanationFromResponse() = %q, want latest user content fallback",
got,
)
} }
} }
@ -2343,7 +2424,10 @@ func (m *handledMediaWithSteeringTool) Parameters() map[string]any {
} }
} }
func (m *handledMediaWithSteeringTool) Execute(ctx context.Context, args map[string]any) *tools.ToolResult { func (m *handledMediaWithSteeringTool) Execute(
ctx context.Context,
args map[string]any,
) *tools.ToolResult {
if err := m.loop.Steer(providers.Message{Role: "user", Content: "what about this instead?"}); err != nil { if err := m.loop.Steer(providers.Message{Role: "user", Content: "what about this instead?"}); err != nil {
return tools.ErrorResult(err.Error()).WithError(err) return tools.ErrorResult(err.Error()).WithError(err)
} }
@ -2496,7 +2580,11 @@ func newStrictChatCompletionTestServer(
})) }))
} }
func (h testHelper) executeAndGetResponse(tb testing.TB, ctx context.Context, msg bus.InboundMessage) string { func (h testHelper) executeAndGetResponse(
tb testing.TB,
ctx context.Context,
msg bus.InboundMessage,
) string {
// Use a short timeout to avoid hanging // Use a short timeout to avoid hanging
timeoutCtx, cancel := context.WithTimeout(ctx, responseTimeout) timeoutCtx, cancel := context.WithTimeout(ctx, responseTimeout)
defer cancel() defer cancel()
@ -2652,7 +2740,10 @@ func TestProcessMessage_CommandOutcomes(t *testing.T) {
t.Fatalf("unexpected /foo reply: %q", fooResp) t.Fatalf("unexpected /foo reply: %q", fooResp)
} }
if provider.calls != 1 { if provider.calls != 1 {
t.Fatalf("LLM should be called exactly once after /foo passthrough, calls=%d", provider.calls) t.Fatalf(
"LLM should be called exactly once after /foo passthrough, calls=%d",
provider.calls,
)
} }
newResp := helper.executeAndGetResponse(t, context.Background(), bus.InboundMessage{ newResp := helper.executeAndGetResponse(t, context.Background(), bus.InboundMessage{
@ -2857,7 +2948,10 @@ func TestProcessMessage_SwitchModelRejectsUnknownAlias(t *testing.T) {
} }
if provider.calls != 0 { if provider.calls != 0 {
t.Fatalf("LLM should not be called for rejected /switch and /show, calls=%d", provider.calls) t.Fatalf(
"LLM should not be called for rejected /switch and /show, calls=%d",
provider.calls,
)
} }
} }
@ -2875,7 +2969,13 @@ func TestProcessMessage_SwitchModelRoutesSubsequentRequestsToSelectedProvider(t
remoteCalls := 0 remoteCalls := 0
remoteModel := "" remoteModel := ""
remoteServer := newChatCompletionTestServer(t, "remote", "remote reply", &remoteCalls, &remoteModel) remoteServer := newChatCompletionTestServer(
t,
"remote",
"remote reply",
&remoteCalls,
&remoteModel,
)
defer remoteServer.Close() defer remoteServer.Close()
cfg := &config.Config{ cfg := &config.Config{
@ -3055,18 +3155,20 @@ func TestProcessMessage_FallbackUsesPerCandidateProvider(t *testing.T) {
workspace := t.TempDir() workspace := t.TempDir()
primaryCalls := 0 primaryCalls := 0
primaryServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { primaryServer := httptest.NewServer(
primaryCalls++ http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// Return 429 so FallbackChain classifies this as retriable and moves on. primaryCalls++
w.Header().Set("Content-Type", "application/json") // Return 429 so FallbackChain classifies this as retriable and moves on.
w.WriteHeader(http.StatusTooManyRequests) w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(map[string]any{ w.WriteHeader(http.StatusTooManyRequests)
"error": map[string]any{ _ = json.NewEncoder(w).Encode(map[string]any{
"message": "rate limit exceeded", "error": map[string]any{
"type": "rate_limit_error", "message": "rate limit exceeded",
}, "type": "rate_limit_error",
}) },
})) })
}),
)
defer primaryServer.Close() defer primaryServer.Close()
fallbackCalls := 0 fallbackCalls := 0
@ -3139,23 +3241,28 @@ func TestProcessMessage_FallbackUsesActiveProviderWhenCandidateNotRegistered(t *
// Both the primary and the unregistered fallback share this server // Both the primary and the unregistered fallback share this server
// (same api_base) so activeProvider routes both calls here. // (same api_base) so activeProvider routes both calls here.
callCount := 0 callCount := 0
primaryServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { primaryServer := httptest.NewServer(
callCount++ http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json") callCount++
if callCount == 1 { w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusTooManyRequests) if callCount == 1 {
w.WriteHeader(http.StatusTooManyRequests)
_ = json.NewEncoder(w).Encode(map[string]any{
"error": map[string]any{"message": "rate limit", "type": "rate_limit_error"},
})
return
}
// Second call (fallback via activeProvider) succeeds.
_ = json.NewEncoder(w).Encode(map[string]any{ _ = json.NewEncoder(w).Encode(map[string]any{
"error": map[string]any{"message": "rate limit", "type": "rate_limit_error"}, "choices": []map[string]any{
{
"message": map[string]any{"content": "active provider reply"},
"finish_reason": "stop",
},
},
}) })
return }),
} )
// Second call (fallback via activeProvider) succeeds.
_ = json.NewEncoder(w).Encode(map[string]any{
"choices": []map[string]any{
{"message": map[string]any{"content": "active provider reply"}, "finish_reason": "stop"},
},
})
}))
defer primaryServer.Close() defer primaryServer.Close()
cfg := &config.Config{ cfg := &config.Config{
@ -3199,7 +3306,10 @@ func TestProcessMessage_FallbackUsesActiveProviderWhenCandidateNotRegistered(t *
t.Fatalf("response = %q, want %q", resp, "active provider reply") t.Fatalf("response = %q, want %q", resp, "active provider reply")
} }
if callCount < 2 { if callCount < 2 {
t.Fatalf("primary server calls = %d, want >= 2 (one 429 + one success via activeProvider)", callCount) t.Fatalf(
"primary server calls = %d, want >= 2 (one 429 + one success via activeProvider)",
callCount,
)
} }
} }
@ -3338,7 +3448,9 @@ func TestAgentLoop_ContextExhaustionRetry(t *testing.T) {
msgBus := bus.NewMessageBus() msgBus := bus.NewMessageBus()
// Create a provider that fails once with a context error // Create a provider that fails once with a context error
contextErr := fmt.Errorf("InvalidParameter: Total tokens of image and text exceed max message tokens") contextErr := fmt.Errorf(
"InvalidParameter: Total tokens of image and text exceed max message tokens",
)
provider := &failFirstMockProvider{ provider := &failFirstMockProvider{
failures: 1, failures: 1,
failError: contextErr, failError: contextErr,
@ -3482,7 +3594,11 @@ func TestAgentLoop_VisionUnsupportedErrorStripsSessionMedia(t *testing.T) {
t.Fatalf("response = %q, want %q", resp, "ok") t.Fatalf("response = %q, want %q", resp, "ok")
} }
if provider.calls != 2 { if provider.calls != 2 {
t.Fatalf("calls = %d, want %d (fail with media, then retry without media)", provider.calls, 2) t.Fatalf(
"calls = %d, want %d (fail with media, then retry without media)",
provider.calls,
2,
)
} }
if !slices.Equal(provider.mediaSeen, []bool{true, false}) { if !slices.Equal(provider.mediaSeen, []bool{true, false}) {
t.Fatalf("mediaSeen = %v, want %v", provider.mediaSeen, []bool{true, false}) t.Fatalf("mediaSeen = %v, want %v", provider.mediaSeen, []bool{true, false})
@ -3549,7 +3665,13 @@ func TestAgentLoop_EmptyModelResponseUsesAccurateFallback(t *testing.T) {
provider := &simpleMockProvider{response: ""} provider := &simpleMockProvider{response: ""}
al := NewAgentLoop(cfg, msgBus, provider) al := NewAgentLoop(cfg, msgBus, provider)
response, err := al.ProcessDirectWithChannel(context.Background(), "hello", "empty-response", "test", "chat1") response, err := al.ProcessDirectWithChannel(
context.Background(),
"hello",
"empty-response",
"test",
"chat1",
)
if err != nil { if err != nil {
t.Fatalf("ProcessDirectWithChannel failed: %v", err) t.Fatalf("ProcessDirectWithChannel failed: %v", err)
} }
@ -3581,7 +3703,13 @@ func TestAgentLoop_ToolLimitUsesDedicatedFallback(t *testing.T) {
al := NewAgentLoop(cfg, msgBus, provider) al := NewAgentLoop(cfg, msgBus, provider)
al.RegisterTool(&toolLimitTestTool{}) al.RegisterTool(&toolLimitTestTool{})
response, err := al.ProcessDirectWithChannel(context.Background(), "hello", "tool-limit", "test", "chat1") response, err := al.ProcessDirectWithChannel(
context.Background(),
"hello",
"tool-limit",
"test",
"chat1",
)
if err != nil { if err != nil {
t.Fatalf("ProcessDirectWithChannel failed: %v", err) t.Fatalf("ProcessDirectWithChannel failed: %v", err)
} }
@ -3598,11 +3726,13 @@ func TestAgentLoop_ToolLimitUsesDedicatedFallback(t *testing.T) {
ChatType: "direct", ChatType: "direct",
SenderID: "cron", SenderID: "cron",
}) })
history := defaultAgent.Sessions.GetHistory(al.allocateRouteSession(route, testInboundMessage(bus.InboundMessage{ history := defaultAgent.Sessions.GetHistory(
Channel: "test", al.allocateRouteSession(route, testInboundMessage(bus.InboundMessage{
SenderID: "cron", Channel: "test",
ChatID: "chat1", SenderID: "cron",
})).SessionKey) ChatID: "chat1",
})).SessionKey,
)
if len(history) != 4 { if len(history) != 4 {
t.Fatalf("history len = %d, want 4", len(history)) t.Fatalf("history len = %d, want 4", len(history))
} }
@ -3900,7 +4030,9 @@ func TestHandleReasoning(t *testing.T) {
break break
} }
if msg.Content == "should timeout" { if msg.Content == "should timeout" {
t.Fatal("expected reasoning message to be dropped when bus is full, but it was published") t.Fatal(
"expected reasoning message to be dropped when bus is full, but it was published",
)
} }
} }
} }
@ -4015,7 +4147,11 @@ func TestProcessMessage_PicoPublishesReasoningAsThoughtMessage(t *testing.T) {
} }
if thoughtMsg.Channel != "pico" || thoughtMsg.ChatID != "pico:test-session" { if thoughtMsg.Channel != "pico" || thoughtMsg.ChatID != "pico:test-session" {
t.Fatalf("thought message route = %s/%s, want pico/pico:test-session", thoughtMsg.Channel, thoughtMsg.ChatID) t.Fatalf(
"thought message route = %s/%s, want pico/pico:test-session",
thoughtMsg.Channel,
thoughtMsg.ChatID,
)
} }
if thoughtMsg.Context.Raw[metadataKeyMessageKind] != messageKindThought { if thoughtMsg.Context.Raw[metadataKeyMessageKind] != messageKindThought {
t.Fatalf( t.Fatalf(
@ -4057,7 +4193,12 @@ func TestProcessHeartbeat_DoesNotPublishToolFeedback(t *testing.T) {
provider := &toolFeedbackProvider{filePath: heartbeatFile} provider := &toolFeedbackProvider{filePath: heartbeatFile}
al := NewAgentLoop(cfg, msgBus, provider) al := NewAgentLoop(cfg, msgBus, provider)
response, err := al.ProcessHeartbeat(context.Background(), "check heartbeat tasks", "telegram", "chat-1") response, err := al.ProcessHeartbeat(
context.Background(),
"check heartbeat tasks",
"telegram",
"chat-1",
)
if err != nil { if err != nil {
t.Fatalf("ProcessHeartbeat() error = %v", err) t.Fatalf("ProcessHeartbeat() error = %v", err)
} }
@ -4132,10 +4273,16 @@ func TestProcessMessage_PublishesToolFeedbackWhenEnabled(t *testing.T) {
t.Fatalf("tool feedback content = %q, want read_file summary", outbound.Content) t.Fatalf("tool feedback content = %q, want read_file summary", outbound.Content)
} }
if !strings.Contains(outbound.Content, utils.ToolFeedbackContinuationHint) { if !strings.Contains(outbound.Content, utils.ToolFeedbackContinuationHint) {
t.Fatalf("tool feedback content = %q, want continuation hint fallback", outbound.Content) t.Fatalf(
"tool feedback content = %q, want continuation hint fallback",
outbound.Content,
)
} }
if !strings.Contains(outbound.Content, "check tool feedback") { if !strings.Contains(outbound.Content, "check tool feedback") {
t.Fatalf("tool feedback content = %q, want current user intent fallback", outbound.Content) t.Fatalf(
"tool feedback content = %q, want current user intent fallback",
outbound.Content,
)
} }
if !strings.Contains(outbound.Content, "\"path\":") { if !strings.Contains(outbound.Content, "\"path\":") {
t.Fatalf("tool feedback content = %q, want serialized tool arguments", outbound.Content) t.Fatalf("tool feedback content = %q, want serialized tool arguments", outbound.Content)
@ -4144,7 +4291,10 @@ func TestProcessMessage_PublishesToolFeedbackWhenEnabled(t *testing.T) {
t.Fatalf("tool feedback content = %q, want tool argument value", outbound.Content) t.Fatalf("tool feedback content = %q, want tool argument value", outbound.Content)
} }
if strings.Contains(outbound.Content, "Previous turn explanation") { if strings.Contains(outbound.Content, "Previous turn explanation") {
t.Fatalf("tool feedback content = %q, want no previous assistant fallback", outbound.Content) t.Fatalf(
"tool feedback content = %q, want no previous assistant fallback",
outbound.Content,
)
} }
if outbound.AgentID != "main" { if outbound.AgentID != "main" {
t.Fatalf("tool feedback agent_id = %q, want main", outbound.AgentID) t.Fatalf("tool feedback agent_id = %q, want main", outbound.AgentID)
@ -4152,7 +4302,8 @@ func TestProcessMessage_PublishesToolFeedbackWhenEnabled(t *testing.T) {
if outbound.SessionKey == "" { if outbound.SessionKey == "" {
t.Fatal("expected tool feedback to carry session_key") t.Fatal("expected tool feedback to carry session_key")
} }
if outbound.Scope == nil || outbound.Scope.AgentID != "main" || outbound.Scope.Channel != "telegram" { if outbound.Scope == nil || outbound.Scope.AgentID != "main" ||
outbound.Scope.Channel != "telegram" {
t.Fatalf("expected tool feedback scope, got %+v", outbound.Scope) t.Fatalf("expected tool feedback scope, got %+v", outbound.Scope)
} }
case <-time.After(2 * time.Second): case <-time.After(2 * time.Second):
@ -4211,7 +4362,11 @@ func TestProcessMessage_PersistsReasoningContentInSessionHistory(t *testing.T) {
t.Fatalf("last message content = %q, want %q", last.Content, "final answer") t.Fatalf("last message content = %q, want %q", last.Content, "final answer")
} }
if last.ReasoningContent != "thinking trace" { if last.ReasoningContent != "thinking trace" {
t.Fatalf("last message reasoning_content = %q, want %q", last.ReasoningContent, "thinking trace") t.Fatalf(
"last message reasoning_content = %q, want %q",
last.ReasoningContent,
"thinking trace",
)
} }
} }
@ -4268,16 +4423,29 @@ func TestProcessMessage_PersistsReasoningToolResponseAsSingleAssistantRecord(t *
t.Fatal("expected assistant history record with tool_calls") t.Fatal("expected assistant history record with tool_calls")
} }
if assistantWithToolCall.Content != "I'll inspect that file now." { if assistantWithToolCall.Content != "I'll inspect that file now." {
t.Fatalf("assistant content = %q, want %q", assistantWithToolCall.Content, "I'll inspect that file now.") t.Fatalf(
"assistant content = %q, want %q",
assistantWithToolCall.Content,
"I'll inspect that file now.",
)
} }
if assistantWithToolCall.ReasoningContent != "Read the file before answering." { if assistantWithToolCall.ReasoningContent != "Read the file before answering." {
t.Fatalf("assistant reasoning_content = %q, want preserved", assistantWithToolCall.ReasoningContent) t.Fatalf(
"assistant reasoning_content = %q, want preserved",
assistantWithToolCall.ReasoningContent,
)
} }
if len(assistantWithToolCall.ToolCalls) != 1 { if len(assistantWithToolCall.ToolCalls) != 1 {
t.Fatalf("assistant tool calls = %+v, want single read_file tool", assistantWithToolCall.ToolCalls) t.Fatalf(
"assistant tool calls = %+v, want single read_file tool",
assistantWithToolCall.ToolCalls,
)
} }
if got := providers.NormalizeToolCall(assistantWithToolCall.ToolCalls[0]).Name; got != "read_file" { if got := providers.NormalizeToolCall(assistantWithToolCall.ToolCalls[0]).Name; got != "read_file" {
t.Fatalf("assistant tool calls = %+v, want single read_file tool", assistantWithToolCall.ToolCalls) t.Fatalf(
"assistant tool calls = %+v, want single read_file tool",
assistantWithToolCall.ToolCalls,
)
} }
sessionDir := filepath.Join(tmpDir, "sessions") sessionDir := filepath.Join(tmpDir, "sessions")
@ -4317,7 +4485,8 @@ func TestProcessMessage_PersistsReasoningToolResponseAsSingleAssistantRecord(t *
if msg.Role != "assistant" { if msg.Role != "assistant" {
continue continue
} }
if msg.Content == "I'll inspect that file now." || msg.ReasoningContent == "Read the file before answering." { if msg.Content == "I'll inspect that file now." ||
msg.ReasoningContent == "Read the file before answering." {
matchingRecords++ matchingRecords++
toolName := "" toolName := ""
if len(msg.ToolCalls) == 1 { if len(msg.ToolCalls) == 1 {
@ -4327,12 +4496,18 @@ func TestProcessMessage_PersistsReasoningToolResponseAsSingleAssistantRecord(t *
msg.ReasoningContent != "Read the file before answering." || msg.ReasoningContent != "Read the file before answering." ||
len(msg.ToolCalls) != 1 || len(msg.ToolCalls) != 1 ||
toolName != "read_file" { toolName != "read_file" {
t.Fatalf("assistant jsonl record = %+v, want content+reasoning+tool_calls in one line", msg) t.Fatalf(
"assistant jsonl record = %+v, want content+reasoning+tool_calls in one line",
msg,
)
} }
} }
} }
if matchingRecords != 1 { if matchingRecords != 1 {
t.Fatalf("matching assistant jsonl records = %d, want exactly 1 canonical assistant record", matchingRecords) t.Fatalf(
"matching assistant jsonl records = %d, want exactly 1 canonical assistant record",
matchingRecords,
)
} }
} }
@ -4387,10 +4562,16 @@ func TestProcessMessage_DoesNotLeakReasoningContentInToolFeedback(t *testing.T)
t.Fatalf("tool feedback content = %q, want read_file summary", outbound.Content) t.Fatalf("tool feedback content = %q, want read_file summary", outbound.Content)
} }
if !strings.Contains(outbound.Content, utils.ToolFeedbackContinuationHint) { if !strings.Contains(outbound.Content, utils.ToolFeedbackContinuationHint) {
t.Fatalf("tool feedback content = %q, want continuation hint fallback", outbound.Content) t.Fatalf(
"tool feedback content = %q, want continuation hint fallback",
outbound.Content,
)
} }
if !strings.Contains(outbound.Content, "check reasoning fallback") { if !strings.Contains(outbound.Content, "check reasoning fallback") {
t.Fatalf("tool feedback content = %q, want current user intent fallback", outbound.Content) t.Fatalf(
"tool feedback content = %q, want current user intent fallback",
outbound.Content,
)
} }
if !strings.Contains(outbound.Content, "\"path\":") { if !strings.Contains(outbound.Content, "\"path\":") {
t.Fatalf("tool feedback content = %q, want serialized tool arguments", outbound.Content) t.Fatalf("tool feedback content = %q, want serialized tool arguments", outbound.Content)
@ -4399,7 +4580,10 @@ func TestProcessMessage_DoesNotLeakReasoningContentInToolFeedback(t *testing.T)
t.Fatalf("tool feedback content = %q, want tool argument value", outbound.Content) t.Fatalf("tool feedback content = %q, want tool argument value", outbound.Content)
} }
if strings.Contains(outbound.Content, "Read README.md first") { if strings.Contains(outbound.Content, "Read README.md first") {
t.Fatalf("tool feedback content = %q, should not leak hidden reasoning", outbound.Content) t.Fatalf(
"tool feedback content = %q, should not leak hidden reasoning",
outbound.Content,
)
} }
case <-time.After(2 * time.Second): case <-time.After(2 * time.Second):
t.Fatal("expected outbound tool feedback without leaking reasoning") t.Fatal("expected outbound tool feedback without leaking reasoning")
@ -4454,7 +4638,11 @@ func assertToolFeedbackNotPublishedWhenDisabled(t *testing.T, channel string) {
select { select {
case outbound := <-msgBus.OutboundChan(): case outbound := <-msgBus.OutboundChan():
t.Fatalf("expected no outbound tool feedback for %s when disabled, got %+v", channel, outbound) t.Fatalf(
"expected no outbound tool feedback for %s when disabled, got %+v",
channel,
outbound,
)
case <-time.After(200 * time.Millisecond): case <-time.After(200 * time.Millisecond):
} }
} }
@ -4567,13 +4755,20 @@ func TestRun_PicoPublishesAssistantContentDuringToolCallsWithoutFinalDuplicate(t
} }
if outputs[0].Content != "intermediate model text" { if outputs[0].Content != "intermediate model text" {
t.Fatalf("first outbound content = %q, want %q", outputs[0].Content, "intermediate model text") t.Fatalf(
"first outbound content = %q, want %q",
outputs[0].Content,
"intermediate model text",
)
} }
if outputs[1].Context.Raw[metadataKeyMessageKind] != messageKindToolCalls { if outputs[1].Context.Raw[metadataKeyMessageKind] != messageKindToolCalls {
t.Fatalf("second outbound = %+v, want tool_calls message", outputs[1]) t.Fatalf("second outbound = %+v, want tool_calls message", outputs[1])
} }
if !strings.Contains(outputs[1].Context.Raw[metadataKeyToolCalls], "tool_limit_test_tool") { if !strings.Contains(outputs[1].Context.Raw[metadataKeyToolCalls], "tool_limit_test_tool") {
t.Fatalf("second outbound tool_calls = %q, want tool name", outputs[1].Context.Raw[metadataKeyToolCalls]) t.Fatalf(
"second outbound tool_calls = %q, want tool name",
outputs[1].Context.Raw[metadataKeyToolCalls],
)
} }
if outputs[2].Content != "final model text" { if outputs[2].Content != "final model text" {
t.Fatalf("third outbound content = %q, want %q", outputs[2].Content, "final model text") t.Fatalf("third outbound content = %q, want %q", outputs[2].Content, "final model text")
@ -4709,7 +4904,10 @@ func TestRun_PicoToolFeedbackSuppressesDuplicateInterimAssistantContent(t *testi
t.Fatalf("first outbound content = %q, want empty tool_calls content", outputs[0].Content) t.Fatalf("first outbound content = %q, want empty tool_calls content", outputs[0].Content)
} }
if !strings.Contains(outputs[0].Context.Raw[metadataKeyToolCalls], "tool_limit_test_tool") { if !strings.Contains(outputs[0].Context.Raw[metadataKeyToolCalls], "tool_limit_test_tool") {
t.Fatalf("first outbound tool_calls = %q, want tool name", outputs[0].Context.Raw[metadataKeyToolCalls]) t.Fatalf(
"first outbound tool_calls = %q, want tool name",
outputs[0].Context.Raw[metadataKeyToolCalls],
)
} }
if outputs[1].Content != "final model text" { if outputs[1].Content != "final model text" {
t.Fatalf("second outbound content = %q, want %q", outputs[1].Content, "final model text") t.Fatalf("second outbound content = %q, want %q", outputs[1].Content, "final model text")
@ -4858,7 +5056,8 @@ func TestResolveMediaRefs_MultiToolCallPreservesOrdering(t *testing.T) {
if result[3].Role != "user" { if result[3].Role != "user" {
t.Fatalf("result[3] expected user, got %q", result[3].Role) t.Fatalf("result[3] expected user, got %q", result[3].Role)
} }
if len(result[3].Media) != 1 || !strings.HasPrefix(result[3].Media[0], "data:image/png;base64,") { if len(result[3].Media) != 1 ||
!strings.HasPrefix(result[3].Media[0], "data:image/png;base64,") {
t.Fatal("expected synthetic user message to contain base64 image") t.Fatal("expected synthetic user message to contain base64 image")
} }
} }
@ -5335,8 +5534,14 @@ func TestProcessMessage_ContextOverflowRecovery(t *testing.T) {
agent := al.GetRegistry().GetDefaultAgent() agent := al.GetRegistry().GetDefaultAgent()
for i := 0; i < 5; i++ { for i := 0; i < 5; i++ {
agent.Sessions.AddFullMessage(sessionKey, providers.Message{Role: "user", Content: "heavy message"}) agent.Sessions.AddFullMessage(
agent.Sessions.AddFullMessage(sessionKey, providers.Message{Role: "assistant", Content: "response"}) sessionKey,
providers.Message{Role: "user", Content: "heavy message"},
)
agent.Sessions.AddFullMessage(
sessionKey,
providers.Message{Role: "assistant", Content: "response"},
)
} }
response, err := al.processMessage(context.Background(), testInboundMessage(bus.InboundMessage{ response, err := al.processMessage(context.Background(), testInboundMessage(bus.InboundMessage{
@ -5723,7 +5928,10 @@ func (m *activityWithSteeringTool) Parameters() map[string]any {
} }
} }
func (m *activityWithSteeringTool) Execute(ctx context.Context, args map[string]any) *tools.ToolResult { func (m *activityWithSteeringTool) Execute(
ctx context.Context,
args map[string]any,
) *tools.ToolResult {
if err := m.loop.Steer(providers.Message{Role: "user", Content: "и еще 20 приседаний"}); err != nil { if err := m.loop.Steer(providers.Message{Role: "user", Content: "и еще 20 приседаний"}); err != nil {
return tools.ErrorResult(err.Error()).WithError(err) return tools.ErrorResult(err.Error()).WithError(err)
} }
@ -5733,7 +5941,13 @@ func (m *activityWithSteeringTool) Execute(ctx context.Context, args map[string]
} }
} }
func TestProcessMessage_FinalActionSummarySynthesizesAcrossSteering(t *testing.T) { func newFinalTurnRenderTestLoop(
t *testing.T,
provider providers.LLMProvider,
toolFactory func(*AgentLoop) tools.Tool,
) *AgentLoop {
t.Helper()
tmpDir := t.TempDir() tmpDir := t.TempDir()
cfg := &config.Config{ cfg := &config.Config{
Agents: config.AgentsConfig{ Agents: config.AgentsConfig{
@ -5748,9 +5962,16 @@ func TestProcessMessage_FinalActionSummarySynthesizesAcrossSteering(t *testing.T
} }
msgBus := bus.NewMessageBus() msgBus := bus.NewMessageBus()
provider := &activitySummaryWithSteeringProvider{}
al := NewAgentLoop(cfg, msgBus, provider) al := NewAgentLoop(cfg, msgBus, provider)
al.RegisterTool(&activityWithSteeringTool{loop: al}) al.RegisterTool(toolFactory(al))
return al
}
func TestProcessMessage_FinalActionSummarySynthesizesAcrossSteering(t *testing.T) {
provider := &activitySummaryWithSteeringProvider{}
al := newFinalTurnRenderTestLoop(t, provider, func(al *AgentLoop) tools.Tool {
return &activityWithSteeringTool{loop: al}
})
response, err := al.processMessage(context.Background(), testInboundMessage(bus.InboundMessage{ response, err := al.processMessage(context.Background(), testInboundMessage(bus.InboundMessage{
Channel: "telegram", Channel: "telegram",
@ -5869,7 +6090,10 @@ func (t *daySummaryWithSteeringTool) Parameters() map[string]any {
} }
} }
func (t *daySummaryWithSteeringTool) Execute(ctx context.Context, args map[string]any) *tools.ToolResult { func (t *daySummaryWithSteeringTool) Execute(
ctx context.Context,
args map[string]any,
) *tools.ToolResult {
day, _ := args["day"].(string) day, _ := args["day"].(string)
switch day { switch day {
case "today": case "today":
@ -5890,23 +6114,10 @@ func (t *daySummaryWithSteeringTool) Execute(ctx context.Context, args map[strin
} }
func TestProcessMessage_FinalActionSummaryRendersAcrossInformationalSteering(t *testing.T) { func TestProcessMessage_FinalActionSummaryRendersAcrossInformationalSteering(t *testing.T) {
tmpDir := t.TempDir()
cfg := &config.Config{
Agents: config.AgentsConfig{
Defaults: config.AgentDefaults{
Workspace: tmpDir,
ModelName: "test-model",
MaxTokens: 4096,
MaxToolIterations: 10,
FinalTurnRenderMode: "llm",
},
},
}
msgBus := bus.NewMessageBus()
provider := &daySummaryAcrossSteeringProvider{} provider := &daySummaryAcrossSteeringProvider{}
al := NewAgentLoop(cfg, msgBus, provider) al := newFinalTurnRenderTestLoop(t, provider, func(al *AgentLoop) tools.Tool {
al.RegisterTool(&daySummaryWithSteeringTool{loop: al}) return &daySummaryWithSteeringTool{loop: al}
})
response, err := al.processMessage(context.Background(), testInboundMessage(bus.InboundMessage{ response, err := al.processMessage(context.Background(), testInboundMessage(bus.InboundMessage{
Channel: "telegram", Channel: "telegram",

View file

@ -14,7 +14,11 @@ import (
"github.com/sipeed/picoclaw/pkg/providers" "github.com/sipeed/picoclaw/pkg/providers"
) )
func (al *AgentLoop) runTurn(ctx context.Context, ts *turnState, pipeline *Pipeline) (turnResult, error) { func (al *AgentLoop) runTurn(
ctx context.Context,
ts *turnState,
pipeline *Pipeline,
) (turnResult, error) {
turnCtx, turnCancel := context.WithCancel(ctx) turnCtx, turnCancel := context.WithCancel(ctx)
defer turnCancel() defer turnCancel()
ts.setTurnCancel(turnCancel) ts.setTurnCancel(turnCancel)
@ -103,18 +107,26 @@ func (al *AgentLoop) runTurn(ctx context.Context, ts *turnState, pipeline *Pipel
// Check if parent turn has ended (SubTurn support from HEAD) // Check if parent turn has ended (SubTurn support from HEAD)
if ts.parentTurnState != nil && ts.IsParentEnded() { if ts.parentTurnState != nil && ts.IsParentEnded() {
if !ts.critical { if !ts.critical {
logger.InfoCF("agent", "Parent turn ended, non-critical SubTurn exiting gracefully", map[string]any{ logger.InfoCF(
"agent",
"Parent turn ended, non-critical SubTurn exiting gracefully",
map[string]any{
"agent_id": ts.agentID,
"iteration": iteration,
"turn_id": ts.turnID,
},
)
break
}
logger.InfoCF(
"agent",
"Parent turn ended, critical SubTurn continues running",
map[string]any{
"agent_id": ts.agentID, "agent_id": ts.agentID,
"iteration": iteration, "iteration": iteration,
"turn_id": ts.turnID, "turn_id": ts.turnID,
}) },
break )
}
logger.InfoCF("agent", "Parent turn ended, critical SubTurn continues running", map[string]any{
"agent_id": ts.agentID,
"iteration": iteration,
"turn_id": ts.turnID,
})
} }
// Poll for pending SubTurn results (from HEAD) // Poll for pending SubTurn results (from HEAD)
@ -214,24 +226,35 @@ func (al *AgentLoop) runTurn(ctx context.Context, ts *turnState, pipeline *Pipel
messages = exec.messages messages = exec.messages
continue continue
case ToolControlFinalize: case ToolControlFinalize:
finalContent, rendered := tryRenderFinalTurnReply(turnCtx, al, ts, exec, finalContent) renderedContent, rendered := tryRenderFinalTurnReply(
turnCtx,
al,
ts,
exec,
finalContent,
)
if !rendered { if !rendered {
messages = exec.messages messages = exec.messages
continue continue
} }
if steerMsgs := al.dequeueSteeringMessagesForScope(ts.sessionKey); len(steerMsgs) > 0 { if steerMsgs := al.dequeueSteeringMessagesForScope(ts.sessionKey); len(
steerMsgs,
) > 0 {
exec.markSteeringObserved() exec.markSteeringObserved()
logger.InfoCF("agent", "Steering arrived during terminal render; continuing turn", logger.InfoCF(
"agent",
"Steering arrived during terminal render; continuing turn",
map[string]any{ map[string]any{
"agent_id": ts.agent.ID, "agent_id": ts.agent.ID,
"iteration": iteration, "iteration": iteration,
"steering_count": len(steerMsgs), "steering_count": len(steerMsgs),
}) },
)
exec.pendingMessages = append(exec.pendingMessages, steerMsgs...) exec.pendingMessages = append(exec.pendingMessages, steerMsgs...)
messages = exec.messages messages = exec.messages
continue continue
} }
return pipeline.Finalize(ctx, turnCtx, ts, exec, turnStatus, finalContent) return pipeline.Finalize(ctx, turnCtx, ts, exec, turnStatus, renderedContent)
case ToolControlBreak: case ToolControlBreak:
// Hard abort: delegate to abortTurn (sets TurnEndStatusAborted) // Hard abort: delegate to abortTurn (sets TurnEndStatusAborted)
if exec.abortedByHardAbort { if exec.abortedByHardAbort {
@ -323,7 +346,10 @@ func (al *AgentLoop) selectCandidates(
"score": score, "score": score,
"threshold": agent.Router.Threshold(), "threshold": agent.Router.Threshold(),
}) })
return agent.LightCandidates, resolvedCandidateModel(agent.LightCandidates, agent.Router.LightModel()), true return agent.LightCandidates, resolvedCandidateModel(
agent.LightCandidates,
agent.Router.LightModel(),
), true
} }
func (al *AgentLoop) resolveContextManager() ContextManager { func (al *AgentLoop) resolveContextManager() ContextManager {
@ -340,10 +366,14 @@ func (al *AgentLoop) resolveContextManager() ContextManager {
} }
cm, err := factory(al.cfg.Agents.Defaults.ContextManagerConfig, al) cm, err := factory(al.cfg.Agents.Defaults.ContextManagerConfig, al)
if err != nil { if err != nil {
logger.WarnCF("agent", "Failed to create context manager, falling back to legacy", map[string]any{ logger.WarnCF(
"name": name, "agent",
"error": err.Error(), "Failed to create context manager, falling back to legacy",
}) map[string]any{
"name": name,
"error": err.Error(),
},
)
return &legacyContextManager{al: al} return &legacyContextManager{al: al}
} }
return cm return cm
@ -423,7 +453,11 @@ func (al *AgentLoop) askSideQuestion(
forceModel bool, forceModel bool,
callMessages []providers.Message, callMessages []providers.Message,
) (*providers.LLMResponse, error) { ) (*providers.LLMResponse, error) {
provider, providerModel, cleanup, err := al.isolatedSideQuestionProvider(agent, selectedModelName, candidate) provider, providerModel, cleanup, err := al.isolatedSideQuestionProvider(
agent,
selectedModelName,
candidate,
)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@ -443,7 +477,11 @@ func (al *AgentLoop) askSideQuestion(
turnCtx := newTurnContext(nil, nil, nil) turnCtx := newTurnContext(nil, nil, nil)
if opts != nil { if opts != nil {
turnCtx = newTurnContext(opts.Dispatch.InboundContext, opts.Dispatch.RouteResult, opts.Dispatch.SessionScope) turnCtx = newTurnContext(
opts.Dispatch.InboundContext,
opts.Dispatch.RouteResult,
opts.Dispatch.SessionScope,
)
} }
llmModel := activeModel llmModel := activeModel
if al.hooks != nil { if al.hooks != nil {
@ -499,7 +537,8 @@ func (al *AgentLoop) askSideQuestion(
func(ctx context.Context, providerName, model string) (*providers.LLMResponse, error) { func(ctx context.Context, providerName, model string) (*providers.LLMResponse, error) {
candidate := providers.FallbackCandidate{Provider: providerName, Model: model} candidate := providers.FallbackCandidate{Provider: providerName, Model: model}
for _, activeCandidate := range activeCandidates { for _, activeCandidate := range activeCandidates {
if activeCandidate.Provider == providerName && activeCandidate.Model == model { if activeCandidate.Provider == providerName &&
activeCandidate.Model == model {
candidate = activeCandidate candidate = activeCandidate
break break
} }
@ -587,7 +626,9 @@ func (al *AgentLoop) isolatedSideQuestionProvider(
candidate providers.FallbackCandidate, candidate providers.FallbackCandidate,
) (providers.LLMProvider, string, func(), error) { ) (providers.LLMProvider, string, func(), error) {
if agent == nil { if agent == nil {
return nil, "", func() {}, fmt.Errorf("isolatedSideQuestionProvider: no agent available for /btw") return nil, "", func() {}, fmt.Errorf(
"isolatedSideQuestionProvider: no agent available for /btw",
)
} }
modelCfg, err := al.sideQuestionModelConfig(agent, baseModelName, candidate) modelCfg, err := al.sideQuestionModelConfig(agent, baseModelName, candidate)

View file

@ -275,7 +275,7 @@ type AgentDefaults struct {
MaxParallelTurns int `json:"max_parallel_turns,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_PARALLEL_TURNS"` // Max concurrent turns (0 or 1 = sequential) MaxParallelTurns int `json:"max_parallel_turns,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_PARALLEL_TURNS"` // Max concurrent turns (0 or 1 = sequential)
SubTurn SubTurnConfig `json:"subturn" envPrefix:"PICOCLAW_AGENTS_DEFAULTS_SUBTURN_"` SubTurn SubTurnConfig `json:"subturn" envPrefix:"PICOCLAW_AGENTS_DEFAULTS_SUBTURN_"`
ToolFeedback ToolFeedbackConfig `json:"tool_feedback,omitempty"` ToolFeedback ToolFeedbackConfig `json:"tool_feedback,omitempty"`
FinalTurnRenderMode string `json:"final_turn_render_mode,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_FINAL_TURN_RENDER_MODE"` FinalTurnRenderMode string `json:"final_turn_render_mode,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_FINAL_TURN_RENDER_MODE"`
SplitOnMarker bool `json:"split_on_marker" env:"PICOCLAW_AGENTS_DEFAULTS_SPLIT_ON_MARKER"` // split messages on <|[SPLIT]|> marker SplitOnMarker bool `json:"split_on_marker" env:"PICOCLAW_AGENTS_DEFAULTS_SPLIT_ON_MARKER"` // split messages on <|[SPLIT]|> marker
ContextManager string `json:"context_manager,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_CONTEXT_MANAGER"` ContextManager string `json:"context_manager,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_CONTEXT_MANAGER"`
ContextManagerConfig json.RawMessage `json:"context_manager_config,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_CONTEXT_MANAGER_CONFIG"` ContextManagerConfig json.RawMessage `json:"context_manager_config,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_CONTEXT_MANAGER_CONFIG"`
@ -1022,7 +1022,11 @@ func LoadConfig(path string) (*Config, error) {
} }
if e := json.Unmarshal(data, &versionInfo); e != nil { if e := json.Unmarshal(data, &versionInfo); e != nil {
e = wrapJSONError(data, e, "config.json") e = wrapJSONError(data, e, "config.json")
logger.ErrorCF("config", formatDiagnosticLogMessage("Malformed config file", e), map[string]any{"path": path}) logger.ErrorCF(
"config",
formatDiagnosticLogMessage("Malformed config file", e),
map[string]any{"path": path},
)
return nil, e return nil, e
} }
if len(data) <= 10 { if len(data) <= 10 {