fix: reviewer comments improve context handling for tool execution and ensure defaults for non-conversation callers

Signed-off-by: Boris Bliznioukov <blib@mail.com>
This commit is contained in:
Boris Bliznioukov 2026-03-04 13:28:41 +01:00
parent 5d9570bd73
commit b900af334f
No known key found for this signature in database
6 changed files with 38 additions and 17 deletions

View file

@ -480,8 +480,8 @@ func (al *AgentLoop) processMessage(ctx context.Context, msg bus.InboundMessage)
// Reset message-tool state for this round so we don't skip publishing due to a previous round. // Reset message-tool state for this round so we don't skip publishing due to a previous round.
if tool, ok := agent.Tools.Get("message"); ok { if tool, ok := agent.Tools.Get("message"); ok {
if mt, ok := tool.(*tools.MessageTool); ok { if resetter, ok := tool.(interface{ ResetSentInRound() }); ok {
mt.ResetSentInRound() resetter.ResetSentInRound()
} }
} }

View file

@ -14,30 +14,33 @@ type Tool interface {
// //
// Carried via context.Value so that concurrent tool calls each receive // Carried via context.Value so that concurrent tool calls each receive
// their own immutable copy — no mutable state on singleton tool instances. // their own immutable copy — no mutable state on singleton tool instances.
//
// Keys are unexported pointer-typed vars — guaranteed collision-free,
// and only accessible through the helper functions below.
type ToolContextKey string type toolCtxKey struct{ name string }
const ( var (
CtxKeyChannel ToolContextKey = "channel" ctxKeyChannel = &toolCtxKey{"channel"}
CtxKeyChatID ToolContextKey = "chatID" ctxKeyChatID = &toolCtxKey{"chatID"}
) )
// WithToolContext returns a child context carrying channel and chatID. // WithToolContext returns a child context carrying channel and chatID.
func WithToolContext(ctx context.Context, channel, chatID string) context.Context { func WithToolContext(ctx context.Context, channel, chatID string) context.Context {
ctx = context.WithValue(ctx, CtxKeyChannel, channel) ctx = context.WithValue(ctx, ctxKeyChannel, channel)
ctx = context.WithValue(ctx, CtxKeyChatID, chatID) ctx = context.WithValue(ctx, ctxKeyChatID, chatID)
return ctx return ctx
} }
// ToolChannel extracts the channel from ctx, or "" if unset. // ToolChannel extracts the channel from ctx, or "" if unset.
func ToolChannel(ctx context.Context) string { func ToolChannel(ctx context.Context) string {
v, _ := ctx.Value(CtxKeyChannel).(string) v, _ := ctx.Value(ctxKeyChannel).(string)
return v return v
} }
// ToolChatID extracts the chatID from ctx, or "" if unset. // ToolChatID extracts the chatID from ctx, or "" if unset.
func ToolChatID(ctx context.Context) string { func ToolChatID(ctx context.Context) string {
v, _ := ctx.Value(CtxKeyChatID).(string) v, _ := ctx.Value(ctxKeyChatID).(string)
return v return v
} }

View file

@ -71,10 +71,8 @@ func (r *ToolRegistry) ExecuteWithContext(
} }
// Inject channel/chatID into ctx so tools read them via ToolChannel(ctx)/ToolChatID(ctx). // Inject channel/chatID into ctx so tools read them via ToolChannel(ctx)/ToolChatID(ctx).
// Immutable per-call — no shared mutable state on tool instances. // Always inject — tools validate what they require.
if channel != "" && chatID != "" {
ctx = WithToolContext(ctx, channel, chatID) ctx = WithToolContext(ctx, channel, chatID)
}
// If tool implements AsyncExecutor and callback is provided, use ExecuteAsync. // If tool implements AsyncExecutor and callback is provided, use ExecuteAsync.
// The callback is a call parameter, not mutable state on the tool instance. // The callback is a call parameter, not mutable state on the tool instance.

View file

@ -156,7 +156,7 @@ func TestToolRegistry_ExecuteWithContext_InjectsToolContext(t *testing.T) {
} }
} }
func TestToolRegistry_ExecuteWithContext_SkipsEmptyContext(t *testing.T) { func TestToolRegistry_ExecuteWithContext_EmptyContext(t *testing.T) {
r := NewToolRegistry() r := NewToolRegistry()
ct := &mockContextAwareTool{ ct := &mockContextAwareTool{
mockRegistryTool: *newMockTool("ctx_tool", "needs context"), mockRegistryTool: *newMockTool("ctx_tool", "needs context"),
@ -168,6 +168,7 @@ func TestToolRegistry_ExecuteWithContext_SkipsEmptyContext(t *testing.T) {
if ct.lastCtx == nil { if ct.lastCtx == nil {
t.Fatal("expected Execute to be called") t.Fatal("expected Execute to be called")
} }
// Empty values are still injected; tools decide what to do with them.
if got := ToolChannel(ct.lastCtx); got != "" { if got := ToolChannel(ct.lastCtx); got != "" {
t.Errorf("expected empty channel, got %q", got) t.Errorf("expected empty channel, got %q", got)
} }

View file

@ -83,9 +83,17 @@ func (t *SpawnTool) execute(ctx context.Context, args map[string]any, cb AsyncCa
return ErrorResult("Subagent manager not configured") return ErrorResult("Subagent manager not configured")
} }
// Read channel/chatID from context (injected by registry) // Read channel/chatID from context (injected by registry).
// Fall back to "cli"/"direct" for non-conversation callers (e.g., CLI, tests)
// to preserve the same defaults as the original NewSpawnTool constructor.
channel := ToolChannel(ctx) channel := ToolChannel(ctx)
if channel == "" {
channel = "cli"
}
chatID := ToolChatID(ctx) chatID := ToolChatID(ctx)
if chatID == "" {
chatID = "direct"
}
// Pass callback to manager for async completion notification // Pass callback to manager for async completion notification
result, err := t.manager.Spawn(ctx, task, label, agentID, channel, chatID, cb) result, err := t.manager.Spawn(ctx, task, label, agentID, channel, chatID, cb)

View file

@ -332,13 +332,24 @@ func (t *SubagentTool) Execute(ctx context.Context, args map[string]any) *ToolRe
} }
} }
// Fall back to "cli"/"direct" for non-conversation callers (e.g., CLI, tests)
// to preserve the same defaults as the original NewSubagentTool constructor.
channel := ToolChannel(ctx)
if channel == "" {
channel = "cli"
}
chatID := ToolChatID(ctx)
if chatID == "" {
chatID = "direct"
}
loopResult, err := RunToolLoop(ctx, ToolLoopConfig{ loopResult, err := RunToolLoop(ctx, ToolLoopConfig{
Provider: sm.provider, Provider: sm.provider,
Model: sm.defaultModel, Model: sm.defaultModel,
Tools: tools, Tools: tools,
MaxIterations: maxIter, MaxIterations: maxIter,
LLMOptions: llmOptions, LLMOptions: llmOptions,
}, messages, ToolChannel(ctx), ToolChatID(ctx)) }, messages, channel, chatID)
if err != nil { if err != nil {
return ErrorResult(fmt.Sprintf("Subagent execution failed: %v", err)).WithError(err) return ErrorResult(fmt.Sprintf("Subagent execution failed: %v", err)).WithError(err)
} }