chore: revert unrelated golines formatting

This commit is contained in:
afjcjsbx 2026-03-29 14:06:19 +02:00
parent 3b173c0bee
commit 07748bf076
57 changed files with 297 additions and 1278 deletions

View file

@ -500,11 +500,8 @@ func TestEstimateMessageTokens_ReasoningContent(t *testing.T) {
reasoningTokens := estimateMessageTokens(withReasoning) reasoningTokens := estimateMessageTokens(withReasoning)
if reasoningTokens <= plainTokens { if reasoningTokens <= plainTokens {
t.Errorf( t.Errorf("message with ReasoningContent (%d tokens) should exceed plain message (%d tokens)",
"message with ReasoningContent (%d tokens) should exceed plain message (%d tokens)", reasoningTokens, plainTokens)
reasoningTokens,
plainTokens,
)
} }
} }
@ -767,11 +764,7 @@ func TestEstimateMessageTokens_WithReasoningAndMedia(t *testing.T) {
tokensNoReasoning := estimateMessageTokens(msgNoReasoning) tokensNoReasoning := estimateMessageTokens(msgNoReasoning)
if tokens <= tokensNoReasoning { if tokens <= tokensNoReasoning {
t.Errorf( t.Errorf("reasoning content should add tokens: with=%d, without=%d", tokens, tokensNoReasoning)
"reasoning content should add tokens: with=%d, without=%d",
tokens,
tokensNoReasoning,
)
} }
} }

View file

@ -82,16 +82,7 @@ func TestSingleSystemMessage(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
msgs := cb.BuildMessages( msgs := cb.BuildMessages(tt.history, tt.summary, tt.message, nil, "test", "chat1", "", "")
tt.history,
tt.summary,
tt.message,
nil,
"test",
"chat1",
"",
"",
)
systemCount := 0 systemCount := 0
for _, m := range msgs { for _, m := range msgs {
@ -177,16 +168,7 @@ func TestBuildMessages_CurrentSenderDynamicContext(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
msgs := cb.BuildMessages( msgs := cb.BuildMessages(nil, "", "hello", nil, "discord", "chat1", tt.senderID, tt.senderDisplayName)
nil,
"",
"hello",
nil,
"discord",
"chat1",
tt.senderID,
tt.senderDisplayName,
)
sys := msgs[0].Content sys := msgs[0].Content
if tt.wantSection { if tt.wantSection {
@ -400,10 +382,7 @@ func TestNewFileCreationInvalidatesCache(t *testing.T) {
// Cache should auto-invalidate because file went from absent -> present // Cache should auto-invalidate because file went from absent -> present
sp2 := cb.BuildSystemPromptWithCache() sp2 := cb.BuildSystemPromptWithCache()
if !strings.Contains(sp2, tt.checkField) { if !strings.Contains(sp2, tt.checkField) {
t.Errorf( t.Errorf("cache not invalidated on new file creation: expected %q in prompt", tt.checkField)
"cache not invalidated on new file creation: expected %q in prompt",
tt.checkField,
)
} }
}) })
} }

View file

@ -151,19 +151,7 @@ func TestSanitizeHistoryForProvider_MultiToolCallsThenNewRound(t *testing.T) {
if len(result) != 9 { if len(result) != 9 {
t.Fatalf("expected 9 messages, got %d: %+v", len(result), roles(result)) t.Fatalf("expected 9 messages, got %d: %+v", len(result), roles(result))
} }
assertRoles( assertRoles(t, result, "user", "assistant", "tool", "tool", "assistant", "user", "assistant", "tool", "assistant")
t,
result,
"user",
"assistant",
"tool",
"tool",
"assistant",
"user",
"assistant",
"tool",
"assistant",
)
} }
func TestSanitizeHistoryForProvider_ConsecutiveMultiToolRounds(t *testing.T) { func TestSanitizeHistoryForProvider_ConsecutiveMultiToolRounds(t *testing.T) {
@ -182,18 +170,7 @@ func TestSanitizeHistoryForProvider_ConsecutiveMultiToolRounds(t *testing.T) {
if len(result) != 8 { if len(result) != 8 {
t.Fatalf("expected 8 messages, got %d: %+v", len(result), roles(result)) t.Fatalf("expected 8 messages, got %d: %+v", len(result), roles(result))
} }
assertRoles( assertRoles(t, result, "user", "assistant", "tool", "tool", "assistant", "tool", "tool", "assistant")
t,
result,
"user",
"assistant",
"tool",
"tool",
"assistant",
"tool",
"tool",
"assistant",
)
} }
func TestSanitizeHistoryForProvider_PlainConversation(t *testing.T) { func TestSanitizeHistoryForProvider_PlainConversation(t *testing.T) {
@ -327,17 +304,5 @@ func TestSanitizeHistoryForProvider_PartialToolResultsInMiddle(t *testing.T) {
if len(result) != 9 { if len(result) != 9 {
t.Fatalf("expected 9 messages, got %d: %+v", len(result), roles(result)) t.Fatalf("expected 9 messages, got %d: %+v", len(result), roles(result))
} }
assertRoles( assertRoles(t, result, "user", "assistant", "tool", "assistant", "user", "user", "assistant", "tool", "assistant")
t,
result,
"user",
"assistant",
"tool",
"assistant",
"user",
"user",
"assistant",
"tool",
"assistant",
)
} }

View file

@ -61,12 +61,8 @@ Act directly and use tools first.
if len(definition.Agent.Frontmatter.Skills) != 2 { if len(definition.Agent.Frontmatter.Skills) != 2 {
t.Fatalf("expected skills to be parsed, got %v", definition.Agent.Frontmatter.Skills) t.Fatalf("expected skills to be parsed, got %v", definition.Agent.Frontmatter.Skills)
} }
if len(definition.Agent.Frontmatter.MCPServers) != 1 || if len(definition.Agent.Frontmatter.MCPServers) != 1 || definition.Agent.Frontmatter.MCPServers[0] != "github" {
definition.Agent.Frontmatter.MCPServers[0] != "github" { t.Fatalf("expected mcpServers to be parsed, got %v", definition.Agent.Frontmatter.MCPServers)
t.Fatalf(
"expected mcpServers to be parsed, got %v",
definition.Agent.Frontmatter.MCPServers,
)
} }
if definition.Agent.Frontmatter.Fields["metadata"] == nil { if definition.Agent.Frontmatter.Fields["metadata"] == nil {
t.Fatal("expected arbitrary frontmatter fields to remain available") t.Fatal("expected arbitrary frontmatter fields to remain available")
@ -100,10 +96,7 @@ func TestLoadAgentDefinitionFallsBackToLegacyAgentsMarkdown(t *testing.T) {
t.Fatal("expected AGENTS.md to be loaded") t.Fatal("expected AGENTS.md to be loaded")
} }
if definition.Agent.RawFrontmatter != "" { if definition.Agent.RawFrontmatter != "" {
t.Fatalf( t.Fatalf("legacy AGENTS.md should not have frontmatter, got %q", definition.Agent.RawFrontmatter)
"legacy AGENTS.md should not have frontmatter, got %q",
definition.Agent.RawFrontmatter,
)
} }
if !strings.Contains(definition.Agent.Body, "Keep compatibility") { if !strings.Contains(definition.Agent.Body, "Keep compatibility") {
t.Fatalf("expected legacy body to be preserved, got %q", definition.Agent.Body) t.Fatalf("expected legacy body to be preserved, got %q", definition.Agent.Body)
@ -166,10 +159,7 @@ Keep going.
len(definition.Agent.Frontmatter.Skills) != 0 || len(definition.Agent.Frontmatter.Skills) != 0 ||
len(definition.Agent.Frontmatter.MCPServers) != 0 || len(definition.Agent.Frontmatter.MCPServers) != 0 ||
len(definition.Agent.Frontmatter.Fields) != 0 { len(definition.Agent.Frontmatter.Fields) != 0 {
t.Fatalf( t.Fatalf("expected invalid frontmatter to decode as empty struct, got %+v", definition.Agent.Frontmatter)
"expected invalid frontmatter to decode as empty struct, got %+v",
definition.Agent.Frontmatter,
)
} }
} }

View file

@ -275,13 +275,7 @@ func TestAgentLoop_EmitsSteeringAndSkippedToolEvents(t *testing.T) {
resultCh := make(chan string, 1) resultCh := make(chan string, 1)
go func() { go func() {
resp, _ := al.ProcessDirectWithChannel( resp, _ := al.ProcessDirectWithChannel(context.Background(), "do something", "test-session", "test", "chat1")
context.Background(),
"do something",
"test-session",
"test",
"chat1",
)
resultCh <- resp resultCh <- resp
}() }()
@ -344,11 +338,7 @@ func TestAgentLoop_EmitsSteeringAndSkippedToolEvents(t *testing.T) {
t.Fatalf("expected steering interrupt kind, got %q", interruptPayload.Kind) t.Fatalf("expected steering interrupt kind, got %q", interruptPayload.Kind)
} }
if interruptPayload.ContentLen != len("change course") { if interruptPayload.ContentLen != len("change course") {
t.Fatalf( t.Fatalf("expected interrupt content len %d, got %d", len("change course"), interruptPayload.ContentLen)
"expected interrupt content len %d, got %d",
len("change course"),
interruptPayload.ContentLen,
)
} }
} }
@ -370,9 +360,7 @@ func TestAgentLoop_EmitsContextCompressEventOnRetry(t *testing.T) {
}, },
} }
contextErr := stringError( contextErr := stringError("InvalidParameter: Total tokens of image and text exceed max message tokens")
"InvalidParameter: Total tokens of image and text exceed max message tokens",
)
provider := &failFirstMockProvider{ provider := &failFirstMockProvider{
failures: 1, failures: 1,
failError: contextErr, failError: contextErr,
@ -615,12 +603,7 @@ func collectEventStream(ch <-chan Event) []Event {
} }
} }
func waitForEvent( func waitForEvent(t *testing.T, ch <-chan Event, timeout time.Duration, match func(Event) bool) Event {
t *testing.T,
ch <-chan Event,
timeout time.Duration,
match func(Event) bool,
) Event {
t.Helper() t.Helper()
timer := time.NewTimer(timeout) timer := time.NewTimer(timeout)

View file

@ -40,11 +40,7 @@ func (h *builtinAutoHook) AfterLLM(
return next, HookDecision{Action: HookActionModify}, nil return next, HookDecision{Action: HookActionModify}, nil
} }
func newConfiguredHookLoop( func newConfiguredHookLoop(t *testing.T, provider *llmHookTestProvider, hooks config.HooksConfig) *AgentLoop {
t *testing.T,
provider *llmHookTestProvider,
hooks config.HooksConfig,
) *AgentLoop {
t.Helper() t.Helper()
cfg := &config.Config{ cfg := &config.Config{
@ -106,13 +102,7 @@ func TestAgentLoop_ProcessDirectWithChannel_AutoMountsBuiltinHook(t *testing.T)
}) })
defer al.Close() defer al.Close()
resp, err := al.ProcessDirectWithChannel( resp, err := al.ProcessDirectWithChannel(context.Background(), "hello", "session-1", "cli", "direct")
context.Background(),
"hello",
"session-1",
"cli",
"direct",
)
if err != nil { if err != nil {
t.Fatalf("ProcessDirectWithChannel failed: %v", err) t.Fatalf("ProcessDirectWithChannel failed: %v", err)
} }
@ -150,13 +140,7 @@ func TestAgentLoop_ProcessDirectWithChannel_AutoMountsProcessHook(t *testing.T)
}) })
defer al.Close() defer al.Close()
resp, err := al.ProcessDirectWithChannel( resp, err := al.ProcessDirectWithChannel(context.Background(), "hello", "session-1", "cli", "direct")
context.Background(),
"hello",
"session-1",
"cli",
"direct",
)
if err != nil { if err != nil {
t.Fatalf("ProcessDirectWithChannel failed: %v", err) t.Fatalf("ProcessDirectWithChannel failed: %v", err)
} }
@ -188,13 +172,7 @@ func TestAgentLoop_ProcessDirectWithChannel_InvalidConfiguredHookFails(t *testin
}) })
defer al.Close() defer al.Close()
_, err := al.ProcessDirectWithChannel( _, err := al.ProcessDirectWithChannel(context.Background(), "hello", "session-1", "cli", "direct")
context.Background(),
"hello",
"session-1",
"cli",
"direct",
)
if err == nil { if err == nil {
t.Fatal("expected invalid configured hook error") t.Fatal("expected invalid configured hook error")
} }

View file

@ -98,11 +98,7 @@ type processHookAfterToolResponse struct {
Result *ToolResultHookResponse `json:"result,omitempty"` Result *ToolResultHookResponse `json:"result,omitempty"`
} }
func NewProcessHook( func NewProcessHook(ctx context.Context, name string, opts ProcessHookOptions) (*ProcessHook, error) {
ctx context.Context,
name string,
opts ProcessHookOptions,
) (*ProcessHook, error) {
if len(opts.Command) == 0 { if len(opts.Command) == 0 {
return nil, fmt.Errorf("process hook command is required") return nil, fmt.Errorf("process hook command is required")
} }
@ -266,10 +262,7 @@ func (ph *ProcessHook) AfterTool(
return resp.Result, HookDecision{Action: resp.Action, Reason: resp.Reason}, nil return resp.Result, HookDecision{Action: resp.Action, Reason: resp.Reason}, nil
} }
func (ph *ProcessHook) ApproveTool( func (ph *ProcessHook) ApproveTool(ctx context.Context, req *ToolApprovalRequest) (ApprovalDecision, error) {
ctx context.Context,
req *ToolApprovalRequest,
) (ApprovalDecision, error) {
if ph == nil || !ph.opts.ApproveTool { if ph == nil || !ph.opts.ApproveTool {
return ApprovalDecision{Approved: true}, nil return ApprovalDecision{Approved: true}, nil
} }
@ -480,11 +473,7 @@ func (ph *ProcessHook) removePending(id uint64) {
} }
} }
func (al *AgentLoop) MountProcessHook( func (al *AgentLoop) MountProcessHook(ctx context.Context, name string, opts ProcessHookOptions) error {
ctx context.Context,
name string,
opts ProcessHookOptions,
) error {
if al == nil { if al == nil {
return fmt.Errorf("agent loop is nil") return fmt.Errorf("agent loop is nil")
} }

View file

@ -79,14 +79,8 @@ type LLMInterceptor interface {
} }
type ToolInterceptor interface { type ToolInterceptor interface {
BeforeTool( BeforeTool(ctx context.Context, call *ToolCallHookRequest) (*ToolCallHookRequest, HookDecision, error)
ctx context.Context, AfterTool(ctx context.Context, result *ToolResultHookResponse) (*ToolResultHookResponse, HookDecision, error)
call *ToolCallHookRequest,
) (*ToolCallHookRequest, HookDecision, error)
AfterTool(
ctx context.Context,
result *ToolResultHookResponse,
) (*ToolResultHookResponse, HookDecision, error)
} }
type ToolApprover interface { type ToolApprover interface {
@ -301,10 +295,7 @@ func (hm *HookManager) dispatchEvents() {
} }
} }
func (hm *HookManager) BeforeLLM( func (hm *HookManager) BeforeLLM(ctx context.Context, req *LLMHookRequest) (*LLMHookRequest, HookDecision) {
ctx context.Context,
req *LLMHookRequest,
) (*LLMHookRequest, HookDecision) {
if hm == nil || req == nil { if hm == nil || req == nil {
return req, HookDecision{Action: HookActionContinue} return req, HookDecision{Action: HookActionContinue}
} }
@ -335,10 +326,7 @@ func (hm *HookManager) BeforeLLM(
return current, HookDecision{Action: HookActionContinue} return current, HookDecision{Action: HookActionContinue}
} }
func (hm *HookManager) AfterLLM( func (hm *HookManager) AfterLLM(ctx context.Context, resp *LLMHookResponse) (*LLMHookResponse, HookDecision) {
ctx context.Context,
resp *LLMHookResponse,
) (*LLMHookResponse, HookDecision) {
if hm == nil || resp == nil { if hm == nil || resp == nil {
return resp, HookDecision{Action: HookActionContinue} return resp, HookDecision{Action: HookActionContinue}
} }

View file

@ -293,10 +293,7 @@ func TestAgentLoop_Hooks_ToolInterceptorCanRewrite(t *testing.T) {
type denyApprovalHook struct{} type denyApprovalHook struct{}
func (h *denyApprovalHook) ApproveTool( func (h *denyApprovalHook) ApproveTool(ctx context.Context, req *ToolApprovalRequest) (ApprovalDecision, error) {
ctx context.Context,
req *ToolApprovalRequest,
) (ApprovalDecision, error) {
return ApprovalDecision{ return ApprovalDecision{
Approved: false, Approved: false,
Reason: "blocked", Reason: "blocked",

View file

@ -156,11 +156,7 @@ func TestNewAgentInstance_ResolveCandidatesFromModelListAlias(t *testing.T) {
t.Fatalf("len(Candidates) = %d, want 1", len(agent.Candidates)) t.Fatalf("len(Candidates) = %d, want 1", len(agent.Candidates))
} }
if agent.Candidates[0].Provider != tt.wantProvider { if agent.Candidates[0].Provider != tt.wantProvider {
t.Fatalf( t.Fatalf("candidate provider = %q, want %q", agent.Candidates[0].Provider, tt.wantProvider)
"candidate provider = %q, want %q",
agent.Candidates[0].Provider,
tt.wantProvider,
)
} }
if agent.Candidates[0].Model != tt.wantModel { if agent.Candidates[0].Model != tt.wantModel {
t.Fatalf("candidate model = %q, want %q", agent.Candidates[0].Model, tt.wantModel) t.Fatalf("candidate model = %q, want %q", agent.Candidates[0].Model, tt.wantModel)

View file

@ -192,11 +192,7 @@ func registerSharedTools(
Proxy: cfg.Tools.Web.Proxy, Proxy: cfg.Tools.Web.Proxy,
}) })
if err != nil { if err != nil {
logger.ErrorCF( logger.ErrorCF("agent", "Failed to create web search tool", map[string]any{"error": err.Error()})
"agent",
"Failed to create web search tool",
map[string]any{"error": err.Error()},
)
} else if searchTool != nil { } else if searchTool != nil {
agent.Tools.Register(searchTool) agent.Tools.Register(searchTool)
} }
@ -209,11 +205,7 @@ func registerSharedTools(
cfg.Tools.Web.FetchLimitBytes, cfg.Tools.Web.FetchLimitBytes,
cfg.Tools.Web.PrivateHostWhitelist) cfg.Tools.Web.PrivateHostWhitelist)
if err != nil { if err != nil {
logger.ErrorCF( logger.ErrorCF("agent", "Failed to create web fetch tool", map[string]any{"error": err.Error()})
"agent",
"Failed to create web fetch tool",
map[string]any{"error": err.Error()},
)
} else { } else {
agent.Tools.Register(fetchTool) agent.Tools.Register(fetchTool)
} }
@ -483,12 +475,7 @@ func (al *AgentLoop) Run(ctx context.Context) error {
"queue_depth": al.pendingSteeringCountForScope(target.SessionKey), "queue_depth": al.pendingSteeringCountForScope(target.SessionKey),
}) })
continued, continueErr := al.Continue( continued, continueErr := al.Continue(ctx, target.SessionKey, target.Channel, target.ChatID)
ctx,
target.SessionKey,
target.Channel,
target.ChatID,
)
if continueErr != nil { if continueErr != nil {
logger.WarnCF("agent", "Failed to continue queued steering", logger.WarnCF("agent", "Failed to continue queued steering",
map[string]any{ map[string]any{
@ -516,22 +503,14 @@ func (al *AgentLoop) Run(ctx context.Context) error {
"queue_depth": al.pendingSteeringCountForScope(target.SessionKey), "queue_depth": al.pendingSteeringCountForScope(target.SessionKey),
}) })
continued, continueErr := al.Continue( continued, continueErr := al.Continue(ctx, target.SessionKey, target.Channel, target.ChatID)
ctx,
target.SessionKey,
target.Channel,
target.ChatID,
)
if continueErr != nil { if continueErr != nil {
logger.WarnCF( logger.WarnCF("agent", "Failed to continue queued steering after shutdown drain",
"agent",
"Failed to continue queued steering after shutdown drain",
map[string]any{ map[string]any{
"channel": target.Channel, "channel": target.Channel,
"chat_id": target.ChatID, "chat_id": target.ChatID,
"error": continueErr.Error(), "error": continueErr.Error(),
}, })
)
return return
} }
if continued == "" { if continued == "" {
@ -586,15 +565,11 @@ func (al *AgentLoop) drainBusToSteering(ctx context.Context, activeScope, active
msgScope, _, scopeOK := al.resolveSteeringTarget(msg) msgScope, _, scopeOK := al.resolveSteeringTarget(msg)
if !scopeOK || msgScope != activeScope { if !scopeOK || msgScope != activeScope {
if err := al.requeueInboundMessage(msg); err != nil { if err := al.requeueInboundMessage(msg); err != nil {
logger.WarnCF( logger.WarnCF("agent", "Failed to requeue non-steering inbound message", map[string]any{
"agent", "error": err.Error(),
"Failed to requeue non-steering inbound message", "channel": msg.Channel,
map[string]any{ "sender_id": msg.SenderID,
"error": err.Error(), })
"channel": msg.Channel,
"sender_id": msg.SenderID,
},
)
} }
continue continue
} }
@ -628,10 +603,7 @@ func (al *AgentLoop) Stop() {
al.running.Store(false) al.running.Store(false)
} }
func (al *AgentLoop) PublishResponseIfNeeded( func (al *AgentLoop) PublishResponseIfNeeded(ctx context.Context, channel, chatID, response string) {
ctx context.Context,
channel, chatID, response string,
) {
if response == "" { if response == "" {
return return
} }
@ -1081,10 +1053,7 @@ var audioAnnotationRe = regexp.MustCompile(`\[(voice|audio)(?::[^\]]*)?\]`)
// transcribeAudioInMessage resolves audio media refs, transcribes them, and // transcribeAudioInMessage resolves audio media refs, transcribes them, and
// replaces audio annotations in msg.Content with the transcribed text. // replaces audio annotations in msg.Content with the transcribed text.
// Returns the (possibly modified) message and true if audio was transcribed. // Returns the (possibly modified) message and true if audio was transcribed.
func (al *AgentLoop) transcribeAudioInMessage( func (al *AgentLoop) transcribeAudioInMessage(ctx context.Context, msg bus.InboundMessage) (bus.InboundMessage, bool) {
ctx context.Context,
msg bus.InboundMessage,
) (bus.InboundMessage, bool) {
if al.transcriber == nil || al.mediaStore == nil || len(msg.Media) == 0 { if al.transcriber == nil || al.mediaStore == nil || len(msg.Media) == 0 {
return msg, false return msg, false
} }
@ -1094,11 +1063,7 @@ func (al *AgentLoop) transcribeAudioInMessage(
for _, ref := range msg.Media { for _, ref := range msg.Media {
path, meta, err := al.mediaStore.ResolveWithMeta(ref) path, meta, err := al.mediaStore.ResolveWithMeta(ref)
if err != nil { if err != nil {
logger.WarnCF( logger.WarnCF("voice", "Failed to resolve media ref", map[string]any{"ref": ref, "error": err})
"voice",
"Failed to resolve media ref",
map[string]any{"ref": ref, "error": err},
)
continue continue
} }
if !utils.IsAudioFile(meta.Filename, meta.ContentType) { if !utils.IsAudioFile(meta.Filename, meta.ContentType) {
@ -1176,11 +1141,7 @@ func (al *AgentLoop) sendTranscriptionFeedback(
ReplyToMessageID: messageID, ReplyToMessageID: messageID,
}) })
if err != nil { if err != nil {
logger.WarnCF( logger.WarnCF("voice", "Failed to send transcription feedback", map[string]any{"error": err.Error()})
"voice",
"Failed to send transcription feedback",
map[string]any{"error": err.Error()},
)
} }
} }
@ -1381,9 +1342,7 @@ func (al *AgentLoop) processMessage(ctx context.Context, msg bus.InboundMessage)
return al.runAgentLoop(ctx, agent, opts) return al.runAgentLoop(ctx, agent, opts)
} }
func (al *AgentLoop) resolveMessageRoute( func (al *AgentLoop) resolveMessageRoute(msg bus.InboundMessage) (routing.ResolvedRoute, *AgentInstance, error) {
msg bus.InboundMessage,
) (routing.ResolvedRoute, *AgentInstance, error) {
registry := al.GetRegistry() registry := al.GetRegistry()
route := registry.ResolveRoute(routing.RouteInput{ route := registry.ResolveRoute(routing.RouteInput{
Channel: msg.Channel, Channel: msg.Channel,
@ -1399,10 +1358,7 @@ func (al *AgentLoop) resolveMessageRoute(
agent = registry.GetDefaultAgent() agent = registry.GetDefaultAgent()
} }
if agent == nil { if agent == nil {
return routing.ResolvedRoute{}, nil, fmt.Errorf( return routing.ResolvedRoute{}, nil, fmt.Errorf("no agent available for route (agent_id=%s)", route.AgentID)
"no agent available for route (agent_id=%s)",
route.AgentID,
)
} }
return route, agent, nil return route, agent, nil
@ -1727,11 +1683,7 @@ func (al *AgentLoop) runTurn(ctx context.Context, ts *turnState) (turnResult, er
ts.recordPersistedMessage(rootMsg) ts.recordPersistedMessage(rootMsg)
} }
activeCandidates, activeModel, usedLight := al.selectCandidates( activeCandidates, activeModel, usedLight := al.selectCandidates(ts.agent, ts.userMessage, messages)
ts.agent,
ts.userMessage,
messages,
)
activeProvider := ts.agent.Provider activeProvider := ts.agent.Provider
if usedLight && ts.agent.LightProvider != nil { if usedLight && ts.agent.LightProvider != nil {
activeProvider = ts.agent.LightProvider activeProvider = ts.agent.LightProvider
@ -2704,15 +2656,12 @@ turnLoop:
} }
if steerMsgs := al.dequeueSteeringMessagesForScope(ts.sessionKey); len(steerMsgs) > 0 { if steerMsgs := al.dequeueSteeringMessagesForScope(ts.sessionKey); len(steerMsgs) > 0 {
logger.InfoCF( logger.InfoCF("agent", "Steering arrived after turn completion; continuing turn before finalizing",
"agent",
"Steering arrived after turn completion; continuing turn before finalizing",
map[string]any{ map[string]any{
"agent_id": ts.agent.ID, "agent_id": ts.agent.ID,
"steering_count": len(steerMsgs), "steering_count": len(steerMsgs),
"session_key": ts.sessionKey, "session_key": ts.sessionKey,
}, })
)
pendingMessages = append(pendingMessages, steerMsgs...) pendingMessages = append(pendingMessages, steerMsgs...)
finalContent = "" finalContent = ""
goto turnLoop goto turnLoop
@ -2828,18 +2777,11 @@ func (al *AgentLoop) selectCandidates(
"score": score, "score": score,
"threshold": agent.Router.Threshold(), "threshold": agent.Router.Threshold(),
}) })
return agent.LightCandidates, resolvedCandidateModel( return agent.LightCandidates, resolvedCandidateModel(agent.LightCandidates, agent.Router.LightModel()), true
agent.LightCandidates,
agent.Router.LightModel(),
), true
} }
// maybeSummarize triggers summarization if the session history exceeds thresholds. // maybeSummarize triggers summarization if the session history exceeds thresholds.
func (al *AgentLoop) maybeSummarize( func (al *AgentLoop) maybeSummarize(agent *AgentInstance, sessionKey string, turnScope turnEventScope) {
agent *AgentInstance,
sessionKey string,
turnScope turnEventScope,
) {
newHistory := agent.Sessions.GetHistory(sessionKey) newHistory := agent.Sessions.GetHistory(sessionKey)
tokenEstimate := al.estimateTokens(newHistory) tokenEstimate := al.estimateTokens(newHistory)
threshold := agent.ContextWindow * agent.SummarizeTokenPercent / 100 threshold := agent.ContextWindow * agent.SummarizeTokenPercent / 100
@ -2873,10 +2815,7 @@ type compressionResult struct {
// prompt is built dynamically by BuildMessages and is NOT stored here. // prompt is built dynamically by BuildMessages and is NOT stored here.
// The compression note is recorded in the session summary so that // The compression note is recorded in the session summary so that
// BuildMessages can include it in the next system prompt. // BuildMessages can include it in the next system prompt.
func (al *AgentLoop) forceCompression( func (al *AgentLoop) forceCompression(agent *AgentInstance, sessionKey string) (compressionResult, bool) {
agent *AgentInstance,
sessionKey string,
) (compressionResult, bool) {
history := agent.Sessions.GetHistory(sessionKey) history := agent.Sessions.GetHistory(sessionKey)
if len(history) <= 2 { if len(history) <= 2 {
return compressionResult{}, false return compressionResult{}, false
@ -3029,11 +2968,7 @@ func formatToolsForLog(toolDefs []providers.ToolDefinition) string {
} }
// summarizeSession summarizes the conversation history for a session. // summarizeSession summarizes the conversation history for a session.
func (al *AgentLoop) summarizeSession( func (al *AgentLoop) summarizeSession(agent *AgentInstance, sessionKey string, turnScope turnEventScope) {
agent *AgentInstance,
sessionKey string,
turnScope turnEventScope,
) {
ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second)
defer cancel() defer cancel()
@ -3385,10 +3320,7 @@ func (al *AgentLoop) applyExplicitSkillCommand(
skillName, ok := agent.ContextBuilder.ResolveSkillName(arg) skillName, ok := agent.ContextBuilder.ResolveSkillName(arg)
if !ok { if !ok {
return true, true, fmt.Sprintf( return true, true, fmt.Sprintf("Unknown skill: %s\nUse /list skills to see installed skills.", arg)
"Unknown skill: %s\nUse /list skills to see installed skills.",
arg,
)
} }
if len(parts) < 3 { if len(parts) < 3 {
@ -3415,10 +3347,7 @@ func (al *AgentLoop) applyExplicitSkillCommand(
return true, false, "" return true, false, ""
} }
func (al *AgentLoop) buildCommandsRuntime( func (al *AgentLoop) buildCommandsRuntime(agent *AgentInstance, opts *processOptions) *commands.Runtime {
agent *AgentInstance,
opts *processOptions,
) *commands.Runtime {
registry := al.GetRegistry() registry := al.GetRegistry()
cfg := al.GetConfig() cfg := al.GetConfig()
rt := &commands.Runtime{ rt := &commands.Runtime{
@ -3462,10 +3391,7 @@ func (al *AgentLoop) buildCommandsRuntime(
rt.ListSkillNames = agent.ContextBuilder.ListSkillNames rt.ListSkillNames = agent.ContextBuilder.ListSkillNames
} }
rt.GetModelInfo = func() (string, string) { rt.GetModelInfo = func() (string, string) {
return agent.Model, resolvedCandidateProvider( return agent.Model, resolvedCandidateProvider(agent.Candidates, cfg.Agents.Defaults.Provider)
agent.Candidates,
cfg.Agents.Defaults.Provider,
)
} }
rt.SwitchModel = func(value string) (string, error) { rt.SwitchModel = func(value string) (string, error) {
value = strings.TrimSpace(value) value = strings.TrimSpace(value)
@ -3479,12 +3405,7 @@ func (al *AgentLoop) buildCommandsRuntime(
return "", fmt.Errorf("failed to initialize model %q: %w", value, err) return "", fmt.Errorf("failed to initialize model %q: %w", value, err)
} }
nextCandidates := resolveModelCandidates( nextCandidates := resolveModelCandidates(cfg, cfg.Agents.Defaults.Provider, modelCfg.Model, agent.Fallbacks)
cfg,
cfg.Agents.Defaults.Provider,
modelCfg.Model,
agent.Fallbacks,
)
if len(nextCandidates) == 0 { if len(nextCandidates) == 0 {
return "", fmt.Errorf("model %q did not resolve to any provider candidates", value) return "", fmt.Errorf("model %q did not resolve to any provider candidates", value)
} }

View file

@ -65,11 +65,7 @@ func (al *AgentLoop) ensureMCPInitialized(ctx context.Context) error {
} }
if al.cfg.Tools.MCP.Servers == nil || len(al.cfg.Tools.MCP.Servers) == 0 { if al.cfg.Tools.MCP.Servers == nil || len(al.cfg.Tools.MCP.Servers) == 0 {
logger.WarnCF( logger.WarnCF("agent", "MCP is enabled but no servers are configured, skipping MCP initialization", nil)
"agent",
"MCP is enabled but no servers are configured, skipping MCP initialization",
nil,
)
return nil return nil
} }
@ -80,11 +76,7 @@ func (al *AgentLoop) ensureMCPInitialized(ctx context.Context) error {
} }
} }
if !findValidServer { if !findValidServer {
logger.WarnCF( logger.WarnCF("agent", "MCP is enabled but no valid servers are configured, skipping MCP initialization", nil)
"agent",
"MCP is enabled but no valid servers are configured, skipping MCP initialization",
nil,
)
return nil return nil
} }
@ -201,14 +193,10 @@ func (al *AgentLoop) ensureMCPInitialized(ctx context.Context) error {
} }
if useRegex { if useRegex {
agent.Tools.Register( agent.Tools.Register(tools.NewRegexSearchTool(agent.Tools, ttl, maxSearchResults))
tools.NewRegexSearchTool(agent.Tools, ttl, maxSearchResults),
)
} }
if useBM25 { if useBM25 {
agent.Tools.Register( agent.Tools.Register(tools.NewBM25SearchTool(agent.Tools, ttl, maxSearchResults))
tools.NewBM25SearchTool(agent.Tools, ttl, maxSearchResults),
)
} }
} }
} }

View file

@ -25,11 +25,7 @@ import (
// Non-image files (documents, audio, video) have their local path injected // Non-image files (documents, audio, video) have their local path injected
// into Content so the agent can access them via file tools like read_file. // into Content so the agent can access them via file tools like read_file.
// Returns a new slice; original messages are not mutated. // Returns a new slice; original messages are not mutated.
func resolveMediaRefs( func resolveMediaRefs(messages []providers.Message, store media.MediaStore, maxSize int) []providers.Message {
messages []providers.Message,
store media.MediaStore,
maxSize int,
) []providers.Message {
if store == nil { if store == nil {
return messages return messages
} }

View file

@ -591,9 +591,7 @@ 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( al.SetChannelManager(newStartedTestChannelManager(t, msgBus, store, "telegram", telegramChannel))
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 {
@ -615,10 +613,7 @@ func TestProcessMessage_MediaToolHandledSkipsFollowUpLLMAndFinalText(t *testing.
t.Fatalf("processMessage() error = %v", err) t.Fatalf("processMessage() error = %v", err)
} }
if response != "" { if response != "" {
t.Fatalf( t.Fatalf("expected no final response when media tool already handled delivery, got %q", response)
"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)
@ -631,20 +626,13 @@ func TestProcessMessage_MediaToolHandledSkipsFollowUpLLMAndFinalText(t *testing.
} }
if len(telegramChannel.sentMedia) != 1 { if len(telegramChannel.sentMedia) != 1 {
t.Fatalf( t.Fatalf("expected exactly 1 synchronously sent media message, got %d", len(telegramChannel.sentMedia))
"expected exactly 1 synchronously sent media message, got %d",
len(telegramChannel.sentMedia),
)
} }
if telegramChannel.sentMedia[0].Channel != "telegram" || if telegramChannel.sentMedia[0].Channel != "telegram" || telegramChannel.sentMedia[0].ChatID != "chat1" {
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( t.Fatalf("expected exactly 1 sent media part, got %d", len(telegramChannel.sentMedia[0].Parts))
"expected exactly 1 sent media part, got %d",
len(telegramChannel.sentMedia[0].Parts),
)
} }
select { select {
@ -672,8 +660,7 @@ func TestProcessMessage_MediaToolHandledSkipsFollowUpLLMAndFinalText(t *testing.
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" || if last.Role != "assistant" || last.Content != "Requested output delivered via tool attachment." {
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)
} }
} }
@ -698,9 +685,7 @@ 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( al.SetChannelManager(newStartedTestChannelManager(t, msgBus, store, "telegram", telegramChannel))
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 {
@ -729,10 +714,7 @@ 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( t.Fatalf("expected exactly 1 synchronously sent media message, got %d", len(telegramChannel.sentMedia))
"expected exactly 1 synchronously sent media message, got %d",
len(telegramChannel.sentMedia),
)
} }
} }
@ -751,9 +733,7 @@ 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( al.SetChannelManager(newStartedTestChannelManager(t, msgBus, store, "telegram", telegramChannel))
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 {
@ -786,20 +766,13 @@ func TestProcessMessage_MediaArtifactCanBeForwardedBySendFile(t *testing.T) {
} }
if len(telegramChannel.sentMedia) != 1 { if len(telegramChannel.sentMedia) != 1 {
t.Fatalf( t.Fatalf("expected exactly 1 synchronously sent media message, got %d", len(telegramChannel.sentMedia))
"expected exactly 1 synchronously sent media message, got %d",
len(telegramChannel.sentMedia),
)
} }
if telegramChannel.sentMedia[0].Channel != "telegram" || if telegramChannel.sentMedia[0].Channel != "telegram" || telegramChannel.sentMedia[0].ChatID != "chat1" {
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( t.Fatalf("expected exactly 1 sent media part, got %d", len(telegramChannel.sentMedia[0].Parts))
"expected exactly 1 sent media part, got %d",
len(telegramChannel.sentMedia[0].Parts),
)
} }
select { select {
@ -1210,10 +1183,7 @@ func (m *handledMediaWithSteeringTool) Parameters() map[string]any {
} }
} }
func (m *handledMediaWithSteeringTool) Execute( func (m *handledMediaWithSteeringTool) Execute(ctx context.Context, args map[string]any) *tools.ToolResult {
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)
} }
@ -1366,11 +1336,7 @@ func newStrictChatCompletionTestServer(
})) }))
} }
func (h testHelper) executeAndGetResponse( func (h testHelper) executeAndGetResponse(tb testing.TB, ctx context.Context, msg bus.InboundMessage) string {
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()
@ -1501,10 +1467,7 @@ 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( t.Fatalf("LLM should be called exactly once after /foo passthrough, calls=%d", provider.calls)
"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{
@ -1654,10 +1617,7 @@ func TestProcessMessage_SwitchModelRejectsUnknownAlias(t *testing.T) {
} }
if provider.calls != 0 { if provider.calls != 0 {
t.Fatalf( t.Fatalf("LLM should not be called for rejected /switch and /show, calls=%d", provider.calls)
"LLM should not be called for rejected /switch and /show, calls=%d",
provider.calls,
)
} }
} }
@ -1675,13 +1635,7 @@ func TestProcessMessage_SwitchModelRoutesSubsequentRequestsToSelectedProvider(t
remoteCalls := 0 remoteCalls := 0
remoteModel := "" remoteModel := ""
remoteServer := newChatCompletionTestServer( remoteServer := newChatCompletionTestServer(t, "remote", "remote reply", &remoteCalls, &remoteModel)
t,
"remote",
"remote reply",
&remoteCalls,
&remoteModel,
)
defer remoteServer.Close() defer remoteServer.Close()
cfg := &config.Config{ cfg := &config.Config{
@ -2004,9 +1958,7 @@ 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( contextErr := fmt.Errorf("InvalidParameter: Total tokens of image and text exceed max message tokens")
"InvalidParameter: Total tokens of image and text exceed max message tokens",
)
provider := &failFirstMockProvider{ provider := &failFirstMockProvider{
failures: 1, failures: 1,
failError: contextErr, failError: contextErr,
@ -2087,13 +2039,7 @@ 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( response, err := al.ProcessDirectWithChannel(context.Background(), "hello", "empty-response", "test", "chat1")
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)
} }
@ -2125,13 +2071,7 @@ 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( response, err := al.ProcessDirectWithChannel(context.Background(), "hello", "tool-limit", "test", "chat1")
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)
} }
@ -2449,9 +2389,7 @@ func TestHandleReasoning(t *testing.T) {
break break
} }
if msg.Content == "should timeout" { if msg.Content == "should timeout" {
t.Fatal( t.Fatal("expected reasoning message to be dropped when bus is full, but it was published")
"expected reasoning message to be dropped when bus is full, but it was published",
)
} }
} }
} }
@ -2545,12 +2483,7 @@ 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( response, err := al.ProcessHeartbeat(context.Background(), "check heartbeat tasks", "telegram", "chat-1")
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)
} }
@ -3035,14 +2968,8 @@ 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( agent.Sessions.AddFullMessage(sessionKey, providers.Message{Role: "user", Content: "heavy message"})
sessionKey, agent.Sessions.AddFullMessage(sessionKey, providers.Message{Role: "assistant", Content: "response"})
providers.Message{Role: "user", Content: "heavy message"},
)
agent.Sessions.AddFullMessage(
sessionKey,
providers.Message{Role: "assistant", Content: "response"},
)
} }
response, err := al.processMessage(context.Background(), bus.InboundMessage{ response, err := al.processMessage(context.Background(), bus.InboundMessage{

View file

@ -26,8 +26,7 @@ func buildModelListResolver(cfg *config.Config) func(raw string) (string, bool)
return "", false return "", false
} }
if mc, err := cfg.GetModelConfig(raw); err == nil && mc != nil && if mc, err := cfg.GetModelConfig(raw); err == nil && mc != nil && strings.TrimSpace(mc.Model) != "" {
strings.TrimSpace(mc.Model) != "" {
return ensureProtocol(mc.Model), true return ensureProtocol(mc.Model), true
} }
@ -79,10 +78,7 @@ func resolvedCandidateProvider(candidates []providers.FallbackCandidate, fallbac
return fallback return fallback
} }
func resolvedModelConfig( func resolvedModelConfig(cfg *config.Config, modelName, workspace string) (*config.ModelConfig, error) {
cfg *config.Config,
modelName, workspace string,
) (*config.ModelConfig, error) {
if cfg == nil { if cfg == nil {
return nil, fmt.Errorf("config is nil") return nil, fmt.Errorf("config is nil")
} }

View file

@ -325,10 +325,7 @@ func (al *AgentLoop) agentForSession(sessionKey string) *AgentInstance {
// user has since enqueued steering messages. // user has since enqueued steering messages.
// //
// If no steering messages are pending, it returns an empty string. // If no steering messages are pending, it returns an empty string.
func (al *AgentLoop) Continue( func (al *AgentLoop) Continue(ctx context.Context, sessionKey, channel, chatID string) (string, error) {
ctx context.Context,
sessionKey, channel, chatID string,
) (string, error) {
if active := al.GetActiveTurn(); active != nil { if active := al.GetActiveTurn(); active != nil {
return "", fmt.Errorf("turn %s is still active", active.TurnID) return "", fmt.Errorf("turn %s is still active", active.TurnID)
} }

View file

@ -896,10 +896,7 @@ func TestAgentLoop_Run_AutoContinuesLateSteeringMessage(t *testing.T) {
defer cancelNoExtra() defer cancelNoExtra()
select { select {
case out2 := <-msgBus.OutboundChan(): case out2 := <-msgBus.OutboundChan():
t.Fatalf( t.Fatalf("expected stale direct response to be suppressed, got extra outbound %q", out2.Content)
"expected stale direct response to be suppressed, got extra outbound %q",
out2.Content,
)
case <-noExtraCtx.Done(): case <-noExtraCtx.Done():
} }
@ -1047,11 +1044,7 @@ func TestAgentLoop_Continue_PreservesSteeringMedia(t *testing.T) {
if err = os.WriteFile(pngPath, pngHeader, 0o644); err != nil { if err = os.WriteFile(pngPath, pngHeader, 0o644); err != nil {
t.Fatalf("WriteFile failed: %v", err) t.Fatalf("WriteFile failed: %v", err)
} }
ref, err := store.Store( ref, err := store.Store(pngPath, media.MediaMeta{Filename: "steer.png", ContentType: "image/png"}, "test")
pngPath,
media.MediaMeta{Filename: "steer.png", ContentType: "image/png"},
"test",
)
if err != nil { if err != nil {
t.Fatalf("Store failed: %v", err) t.Fatalf("Store failed: %v", err)
} }
@ -1243,10 +1236,7 @@ func TestAgentLoop_InterruptGraceful_UsesTerminalNoToolCall(t *testing.T) {
t.Fatalf("expected 2 provider calls, got %d", calls) t.Fatalf("expected 2 provider calls, got %d", calls)
} }
if terminalToolsCount != 0 { if terminalToolsCount != 0 {
t.Fatalf( t.Fatalf("expected graceful terminal call to disable tools, got %d tool defs", terminalToolsCount)
"expected graceful terminal call to disable tools, got %d tool defs",
terminalToolsCount,
)
} }
foundHint := false foundHint := false
@ -1257,8 +1247,7 @@ func TestAgentLoop_InterruptGraceful_UsesTerminalNoToolCall(t *testing.T) {
if msg.Role == "user" && msg.Content == expectedHint { if msg.Role == "user" && msg.Content == expectedHint {
foundHint = true foundHint = true
} }
if msg.Role == "tool" && msg.ToolCallID == "call_2" && if msg.Role == "tool" && msg.ToolCallID == "call_2" && msg.Content == "Skipped due to graceful interrupt." {
msg.Content == "Skipped due to graceful interrupt." {
foundSkipped = true foundSkipped = true
} }
} }
@ -1550,8 +1539,7 @@ func TestAgentLoop_Steering_SkippedToolsHaveErrorResults(t *testing.T) {
foundSkipped := false foundSkipped := false
for _, m := range msgs { for _, m := range msgs {
if m.Role == "tool" && m.ToolCallID == "call_2" && if m.Role == "tool" && m.ToolCallID == "call_2" && m.Content == "Skipped due to queued user message." {
m.Content == "Skipped due to queued user message." {
foundSkipped = true foundSkipped = true
break break
} }
@ -1559,13 +1547,7 @@ func TestAgentLoop_Steering_SkippedToolsHaveErrorResults(t *testing.T) {
if !foundSkipped { if !foundSkipped {
// Log what we actually got // Log what we actually got
for i, m := range msgs { for i, m := range msgs {
t.Logf( t.Logf("msg[%d]: role=%s toolCallID=%s content=%s", i, m.Role, m.ToolCallID, truncate(m.Content, 80))
"msg[%d]: role=%s toolCallID=%s content=%s",
i,
m.Role,
m.ToolCallID,
truncate(m.Content, 80),
)
} }
t.Fatal("expected skipped tool result for call_2") t.Fatal("expected skipped tool result for call_2")
} }

View file

@ -505,12 +505,7 @@ func spawnSubTurn(
// Event emissions: // Event emissions:
// - SubTurnResultDeliveredEvent: successful delivery to channel // - SubTurnResultDeliveredEvent: successful delivery to channel
// - SubTurnOrphanResultEvent: delivery failed (parent finished or channel full) // - SubTurnOrphanResultEvent: delivery failed (parent finished or channel full)
func deliverSubTurnResult( func deliverSubTurnResult(al *AgentLoop, parentTS *turnState, childID string, result *tools.ToolResult) {
al *AgentLoop,
parentTS *turnState,
childID string,
result *tools.ToolResult,
) {
// Let GC clean up the pendingResults channel; parent Finish will no longer close it. // Let GC clean up the pendingResults channel; parent Finish will no longer close it.
// We use defer/recover to catch any unlikely channel panics if it were ever closed. // We use defer/recover to catch any unlikely channel panics if it were ever closed.
defer func() { defer func() {
@ -521,14 +516,9 @@ func deliverSubTurnResult(
"recover": r, "recover": r,
}) })
if result != nil && al != nil { if result != nil && al != nil {
al.emitEvent( al.emitEvent(EventKindSubTurnOrphan,
EventKindSubTurnOrphan,
parentTS.eventMeta("deliverSubTurnResult", "subturn.orphan"), parentTS.eventMeta("deliverSubTurnResult", "subturn.orphan"),
SubTurnOrphanPayload{ SubTurnOrphanPayload{ParentTurnID: parentTS.turnID, ChildTurnID: childID, Reason: "panic"},
ParentTurnID: parentTS.turnID,
ChildTurnID: childID,
Reason: "panic",
},
) )
} }
} }
@ -541,14 +531,9 @@ func deliverSubTurnResult(
// If parent turn has already finished, treat this as an orphan result // If parent turn has already finished, treat this as an orphan result
if isFinished || resultChan == nil { if isFinished || resultChan == nil {
if result != nil && al != nil { if result != nil && al != nil {
al.emitEvent( al.emitEvent(EventKindSubTurnOrphan,
EventKindSubTurnOrphan,
parentTS.eventMeta("deliverSubTurnResult", "subturn.orphan"), parentTS.eventMeta("deliverSubTurnResult", "subturn.orphan"),
SubTurnOrphanPayload{ SubTurnOrphanPayload{ParentTurnID: parentTS.turnID, ChildTurnID: childID, Reason: "parent_finished"},
ParentTurnID: parentTS.turnID,
ChildTurnID: childID,
Reason: "parent_finished",
},
) )
} }
return return

View file

@ -571,8 +571,7 @@ func TestHardAbortSessionRollback(t *testing.T) {
} }
// Verify the content matches the initial state // Verify the content matches the initial state
if finalHistory[0].Content != "initial message 1" || if finalHistory[0].Content != "initial message 1" || finalHistory[1].Content != "initial response 1" {
finalHistory[1].Content != "initial response 1" {
t.Error("history content does not match initial state after rollback") t.Error("history content does not match initial state after rollback")
} }
} }
@ -1291,12 +1290,7 @@ func TestDeliverSubTurnResult_RaceWithFinish(t *testing.T) {
finalOrphan := orphanCount finalOrphan := orphanCount
mu.Unlock() mu.Unlock()
t.Logf( t.Logf("Delivered: %d, Orphan: %d, Total: %d", finalDelivered, finalOrphan, finalDelivered+finalOrphan)
"Delivered: %d, Orphan: %d, Total: %d",
finalDelivered,
finalOrphan,
finalDelivered+finalOrphan,
)
// With the new drainPendingResults behavior, the total events may be >= numResults // With the new drainPendingResults behavior, the total events may be >= numResults
// because Finish() drains remaining results from the channel and emits them as orphans. // because Finish() drains remaining results from the channel and emits them as orphans.

View file

@ -65,11 +65,7 @@ func DefaultConfig() *Config {
Enabled: true, Enabled: true,
Text: FlexibleStringSlice{"Thinking... 💭"}, Text: FlexibleStringSlice{"Thinking... 💭"},
}, },
Streaming: StreamingConfig{ Streaming: StreamingConfig{Enabled: true, ThrottleSeconds: 3, MinGrowthChars: 200},
Enabled: true,
ThrottleSeconds: 3,
MinGrowthChars: 200,
},
UseMarkdownV2: false, UseMarkdownV2: false,
}, },
Feishu: FeishuConfig{ Feishu: FeishuConfig{

View file

@ -335,8 +335,7 @@ func v0ConvertProvidersToModelList(cfg *configV0) []modelConfigV0 {
providerNames: []string{"github_copilot", "copilot"}, providerNames: []string{"github_copilot", "copilot"},
protocol: "github-copilot", protocol: "github-copilot",
buildConfig: func(p providersConfigV0) (modelConfigV0, bool) { buildConfig: func(p providersConfigV0) (modelConfigV0, bool) {
if p.GitHubCopilot.APIKey == "" && p.GitHubCopilot.APIBase == "" && if p.GitHubCopilot.APIKey == "" && p.GitHubCopilot.APIBase == "" && p.GitHubCopilot.ConnectMode == "" {
p.GitHubCopilot.ConnectMode == "" {
return modelConfigV0{}, false return modelConfigV0{}, false
} }
return modelConfigV0{ return modelConfigV0{

View file

@ -72,11 +72,7 @@ func TestMigration_Integration_LegacyConfigWithoutWorkspace(t *testing.T) {
// CRITICAL: Verify that user's settings are preserved // CRITICAL: Verify that user's settings are preserved
// This was the bug - these settings were lost when Workspace was empty // This was the bug - these settings were lost when Workspace was empty
if cfg.Agents.Defaults.Provider != "openai" { if cfg.Agents.Defaults.Provider != "openai" {
t.Errorf( t.Errorf("Provider = %q, want %q (user's setting should be preserved)", cfg.Agents.Defaults.Provider, "openai")
"Provider = %q, want %q (user's setting should be preserved)",
cfg.Agents.Defaults.Provider,
"openai",
)
} }
// Old "model" field is migrated to "model_name" field // Old "model" field is migrated to "model_name" field
if cfg.Agents.Defaults.ModelName != "gpt-4o" { if cfg.Agents.Defaults.ModelName != "gpt-4o" {
@ -303,11 +299,7 @@ func TestMigration_Integration_PreservesAllAgentsFields(t *testing.T) {
t.Errorf("Agent.ID = %q, want %q", cfg.Agents.List[0].ID, "special-agent") t.Errorf("Agent.ID = %q, want %q", cfg.Agents.List[0].ID, "special-agent")
} }
if cfg.Agents.List[0].Workspace != "/special/workspace" { if cfg.Agents.List[0].Workspace != "/special/workspace" {
t.Errorf( t.Errorf("Agent.Workspace = %q, want %q", cfg.Agents.List[0].Workspace, "/special/workspace")
"Agent.Workspace = %q, want %q",
cfg.Agents.List[0].Workspace,
"/special/workspace",
)
} }
// Workspace should have default since it was empty in legacy config // Workspace should have default since it was empty in legacy config
@ -370,10 +362,7 @@ func TestMigration_Integration_ChannelsConfigMigrated(t *testing.T) {
// OneBot: group_trigger_prefix should be migrated to group_trigger.prefixes // OneBot: group_trigger_prefix should be migrated to group_trigger.prefixes
if len(cfg.Channels.OneBot.GroupTrigger.Prefixes) != 2 { if len(cfg.Channels.OneBot.GroupTrigger.Prefixes) != 2 {
t.Errorf( t.Errorf("len(OneBot.GroupTrigger.Prefixes) = %d, want 2", len(cfg.Channels.OneBot.GroupTrigger.Prefixes))
"len(OneBot.GroupTrigger.Prefixes) = %d, want 2",
len(cfg.Channels.OneBot.GroupTrigger.Prefixes),
)
} else { } else {
if cfg.Channels.OneBot.GroupTrigger.Prefixes[0] != "/" { if cfg.Channels.OneBot.GroupTrigger.Prefixes[0] != "/" {
t.Errorf("Prefixes[0] = %q, want %q", cfg.Channels.OneBot.GroupTrigger.Prefixes[0], "/") t.Errorf("Prefixes[0] = %q, want %q", cfg.Channels.OneBot.GroupTrigger.Prefixes[0], "/")
@ -454,25 +443,13 @@ func TestMigration_Integration_RoundTrip_SerializeAndLoad(t *testing.T) {
// Verify configs are identical // Verify configs are identical
if cfg2.Agents.Defaults.Provider != cfg1.Agents.Defaults.Provider { if cfg2.Agents.Defaults.Provider != cfg1.Agents.Defaults.Provider {
t.Errorf( t.Errorf("Provider changed from %q to %q", cfg1.Agents.Defaults.Provider, cfg2.Agents.Defaults.Provider)
"Provider changed from %q to %q",
cfg1.Agents.Defaults.Provider,
cfg2.Agents.Defaults.Provider,
)
} }
if cfg2.Agents.Defaults.ModelName != cfg1.Agents.Defaults.ModelName { if cfg2.Agents.Defaults.ModelName != cfg1.Agents.Defaults.ModelName {
t.Errorf( t.Errorf("ModelName changed from %q to %q", cfg1.Agents.Defaults.ModelName, cfg2.Agents.Defaults.ModelName)
"ModelName changed from %q to %q",
cfg1.Agents.Defaults.ModelName,
cfg2.Agents.Defaults.ModelName,
)
} }
if cfg2.Agents.Defaults.MaxTokens != cfg1.Agents.Defaults.MaxTokens { if cfg2.Agents.Defaults.MaxTokens != cfg1.Agents.Defaults.MaxTokens {
t.Errorf( t.Errorf("MaxTokens changed from %d to %d", cfg1.Agents.Defaults.MaxTokens, cfg2.Agents.Defaults.MaxTokens)
"MaxTokens changed from %d to %d",
cfg1.Agents.Defaults.MaxTokens,
cfg2.Agents.Defaults.MaxTokens,
)
} }
} }
@ -580,11 +557,7 @@ func TestMigration_Integration_ModelNameField(t *testing.T) {
// GetModelName() should return model_name, not model (deprecated) // GetModelName() should return model_name, not model (deprecated)
if cfg.Agents.Defaults.GetModelName() != "deepseek-reasoner" { if cfg.Agents.Defaults.GetModelName() != "deepseek-reasoner" {
t.Errorf( t.Errorf("GetModelName() = %q, want %q", cfg.Agents.Defaults.GetModelName(), "deepseek-reasoner")
"GetModelName() = %q, want %q",
cfg.Agents.Defaults.GetModelName(),
"deepseek-reasoner",
)
} }
if len(cfg.Agents.Defaults.ModelFallbacks) != 1 { if len(cfg.Agents.Defaults.ModelFallbacks) != 1 {

View file

@ -91,11 +91,9 @@ func TestConvertProvidersToModelList_LiteLLM(t *testing.T) {
func TestConvertProvidersToModelList_Multiple(t *testing.T) { func TestConvertProvidersToModelList_Multiple(t *testing.T) {
cfg := &configV0{ cfg := &configV0{
Providers: providersConfigV0{ Providers: providersConfigV0{
OpenAI: openAIProviderConfigV0{ OpenAI: openAIProviderConfigV0{providerConfigV0: providerConfigV0{APIKey: "openai-key"}},
providerConfigV0: providerConfigV0{APIKey: "openai-key"}, Groq: providerConfigV0{APIKey: "groq-key"},
}, Zhipu: providerConfigV0{APIKey: "zhipu-key"},
Groq: providerConfigV0{APIKey: "groq-key"},
Zhipu: providerConfigV0{APIKey: "zhipu-key"},
}, },
} }
@ -144,13 +142,8 @@ func TestConvertProvidersToModelList_AllProviders(t *testing.T) {
// Other providers have no configuration, so they won't be converted. // Other providers have no configuration, so they won't be converted.
cfg := &configV0{ cfg := &configV0{
Providers: providersConfigV0{ Providers: providersConfigV0{
OpenAI: openAIProviderConfigV0{ OpenAI: openAIProviderConfigV0{providerConfigV0: providerConfigV0{APIKey: "key1"}},
providerConfigV0: providerConfigV0{APIKey: "key1"}, LiteLLM: providerConfigV0{APIKey: "key-litellm", APIBase: "http://localhost:4000/v1"},
},
LiteLLM: providerConfigV0{
APIKey: "key-litellm",
APIBase: "http://localhost:4000/v1",
},
Anthropic: providerConfigV0{APIKey: "key2"}, Anthropic: providerConfigV0{APIKey: "key2"},
OpenRouter: providerConfigV0{APIKey: "key3"}, OpenRouter: providerConfigV0{APIKey: "key3"},
Groq: providerConfigV0{APIKey: "key4"}, Groq: providerConfigV0{APIKey: "key4"},
@ -268,11 +261,7 @@ func TestConvertProvidersToModelList_PreservesUserModel_DeepSeek(t *testing.T) {
// Should use user's model, not default // Should use user's model, not default
if result[0].Model != "deepseek/deepseek-reasoner" { if result[0].Model != "deepseek/deepseek-reasoner" {
t.Errorf( t.Errorf("Model = %q, want %q (user's configured model)", result[0].Model, "deepseek/deepseek-reasoner")
"Model = %q, want %q (user's configured model)",
result[0].Model,
"deepseek/deepseek-reasoner",
)
} }
} }
@ -382,9 +371,7 @@ func TestConvertProvidersToModelList_MultipleProviders_PreservesUserModel(t *tes
}, },
}, },
Providers: providersConfigV0{ Providers: providersConfigV0{
OpenAI: openAIProviderConfigV0{ OpenAI: openAIProviderConfigV0{providerConfigV0: providerConfigV0{APIKey: "sk-openai"}},
providerConfigV0: providerConfigV0{APIKey: "sk-openai"},
},
DeepSeek: providerConfigV0{APIKey: "sk-deepseek"}, DeepSeek: providerConfigV0{APIKey: "sk-deepseek"},
}, },
} }
@ -404,11 +391,7 @@ func TestConvertProvidersToModelList_MultipleProviders_PreservesUserModel(t *tes
} }
case "deepseek": case "deepseek":
if mc.Model != "deepseek/deepseek-reasoner" { if mc.Model != "deepseek/deepseek-reasoner" {
t.Errorf( t.Errorf("DeepSeek Model = %q, want %q (user's)", mc.Model, "deepseek/deepseek-reasoner")
"DeepSeek Model = %q, want %q (user's)",
mc.Model,
"deepseek/deepseek-reasoner",
)
} }
} }
} }
@ -506,11 +489,7 @@ func TestConvertProvidersToModelList_NoProviderField_SingleProvider(t *testing.T
// ModelName should be the user's model value for backward compatibility // ModelName should be the user's model value for backward compatibility
if result[0].ModelName != "glm-4.7" { if result[0].ModelName != "glm-4.7" {
t.Errorf( t.Errorf("ModelName = %q, want %q (user's model for backward compatibility)", result[0].ModelName, "glm-4.7")
"ModelName = %q, want %q (user's model for backward compatibility)",
result[0].ModelName,
"glm-4.7",
)
} }
// Model should use the user's model with protocol prefix // Model should use the user's model with protocol prefix
@ -531,10 +510,8 @@ func TestConvertProvidersToModelList_NoProviderField_MultipleProviders(t *testin
}, },
}, },
Providers: providersConfigV0{ Providers: providersConfigV0{
OpenAI: openAIProviderConfigV0{ OpenAI: openAIProviderConfigV0{providerConfigV0: providerConfigV0{APIKey: "openai-key"}},
providerConfigV0: providerConfigV0{APIKey: "openai-key"}, Zhipu: providerConfigV0{APIKey: "zhipu-key"},
},
Zhipu: providerConfigV0{APIKey: "zhipu-key"},
}, },
} }
@ -594,11 +571,7 @@ func TestBuildModelWithProtocol_NoPrefix(t *testing.T) {
func TestBuildModelWithProtocol_AlreadyHasPrefix(t *testing.T) { func TestBuildModelWithProtocol_AlreadyHasPrefix(t *testing.T) {
result := buildModelWithProtocol("openrouter", "openrouter/auto") result := buildModelWithProtocol("openrouter", "openrouter/auto")
if result != "openrouter/auto" { if result != "openrouter/auto" {
t.Errorf( t.Errorf("buildModelWithProtocol(openrouter, openrouter/auto) = %q, want %q", result, "openrouter/auto")
"buildModelWithProtocol(openrouter, openrouter/auto) = %q, want %q",
result,
"openrouter/auto",
)
} }
} }
@ -640,10 +613,6 @@ func TestConvertProvidersToModelList_LegacyModelWithProtocolPrefix(t *testing.T)
// Model should NOT have duplicated prefix // Model should NOT have duplicated prefix
if result[0].Model != "openrouter/auto" { if result[0].Model != "openrouter/auto" {
t.Errorf( t.Errorf("Model = %q, want %q (should not duplicate prefix)", result[0].Model, "openrouter/auto")
"Model = %q, want %q (should not duplicate prefix)",
result[0].Model,
"openrouter/auto",
)
} }
} }

View file

@ -17,11 +17,7 @@ func TestGetModelConfig_Found(t *testing.T) {
Version: CurrentVersion, Version: CurrentVersion,
ModelList: []*ModelConfig{ ModelList: []*ModelConfig{
{ModelName: "test-model", Model: "openai/gpt-4o", APIKeys: SimpleSecureStrings("key1")}, {ModelName: "test-model", Model: "openai/gpt-4o", APIKeys: SimpleSecureStrings("key1")},
{ {ModelName: "other-model", Model: "anthropic/claude", APIKeys: SimpleSecureStrings("key2")},
ModelName: "other-model",
Model: "anthropic/claude",
APIKeys: SimpleSecureStrings("key2"),
},
}, },
} }
@ -118,16 +114,8 @@ func TestGetModelConfig_RoundRobinStartsFromFirstMatch(t *testing.T) {
func TestGetModelConfig_Concurrent(t *testing.T) { func TestGetModelConfig_Concurrent(t *testing.T) {
cfg := &Config{ cfg := &Config{
ModelList: []*ModelConfig{ ModelList: []*ModelConfig{
{ {ModelName: "concurrent-model", Model: "openai/gpt-4o-1", APIKeys: SimpleSecureStrings("key1")},
ModelName: "concurrent-model", {ModelName: "concurrent-model", Model: "openai/gpt-4o-2", APIKeys: SimpleSecureStrings("key2")},
Model: "openai/gpt-4o-1",
APIKeys: SimpleSecureStrings("key1"),
},
{
ModelName: "concurrent-model",
Model: "openai/gpt-4o-2",
APIKeys: SimpleSecureStrings("key2"),
},
}, },
} }
@ -302,11 +290,7 @@ func TestConfig_ValidateModelList(t *testing.T) {
} }
if err != nil && tt.errMsg != "" { if err != nil && tt.errMsg != "" {
if !strings.Contains(err.Error(), tt.errMsg) { if !strings.Contains(err.Error(), tt.errMsg) {
t.Errorf( t.Errorf("ValidateModelList() error = %v, want error containing %q", err, tt.errMsg)
"ValidateModelList() error = %v, want error containing %q",
err,
tt.errMsg,
)
} }
} }
}) })

View file

@ -117,10 +117,7 @@ func TestExpandMultiKeyModels_WithExistingFallbacks(t *testing.T) {
ModelName: "gpt-4", ModelName: "gpt-4",
Model: "openai/gpt-4o", Model: "openai/gpt-4o",
} }
modelCfg.APIKeys = SimpleSecureStrings( modelCfg.APIKeys = SimpleSecureStrings("key0", "key1") // Use internal field for multi-key testing
"key0",
"key1",
) // Use internal field for multi-key testing
modelCfg.Fallbacks = []string{"claude-3"} modelCfg.Fallbacks = []string{"claude-3"}
models := []*ModelConfig{modelCfg} models := []*ModelConfig{modelCfg}
@ -199,10 +196,7 @@ func TestExpandMultiKeyModels_PreservesOtherFields(t *testing.T) {
RequestTimeout: 30, RequestTimeout: 30,
ThinkingLevel: "high", ThinkingLevel: "high",
} }
modelCfg.APIKeys = SimpleSecureStrings( modelCfg.APIKeys = SimpleSecureStrings("key0", "key1") // Use internal field for multi-key testing
"key0",
"key1",
) // Use internal field for multi-key testing
models := []*ModelConfig{modelCfg} models := []*ModelConfig{modelCfg}
result := expandMultiKeyModels(models) result := expandMultiKeyModels(models)

View file

@ -304,13 +304,11 @@ func (s *SecureString) UnmarshalJSON(value []byte) error {
func (s SecureString) MarshalYAML() (any, error) { func (s SecureString) MarshalYAML() (any, error) {
// Preserve raw value if it is already a reference (enc:// or file://) // Preserve raw value if it is already a reference (enc:// or file://)
if strings.HasPrefix(s.raw, credential.EncScheme) || if strings.HasPrefix(s.raw, credential.EncScheme) || strings.HasPrefix(s.raw, credential.FileScheme) {
strings.HasPrefix(s.raw, credential.FileScheme) {
return s.raw, nil return s.raw, nil
} }
// If resolved is a reference format (e.g. set via Set), copy back to raw // If resolved is a reference format (e.g. set via Set), copy back to raw
if strings.HasPrefix(s.resolved, credential.EncScheme) || if strings.HasPrefix(s.resolved, credential.EncScheme) || strings.HasPrefix(s.resolved, credential.FileScheme) {
strings.HasPrefix(s.resolved, credential.FileScheme) {
s.raw = s.resolved s.raw = s.resolved
return s.raw, nil return s.raw, nil
} }

View file

@ -35,10 +35,7 @@ func TestJSONUnmarshalPrivateFields(t *testing.T) {
t.Errorf("PublicField = %q, want 'pub'", s.PublicField) t.Errorf("PublicField = %q, want 'pub'", s.PublicField)
} }
if s.privateField != "" { if s.privateField != "" {
t.Errorf( t.Errorf("privateField = %q, want empty because unexported fields are ignored", s.privateField)
"privateField = %q, want empty because unexported fields are ignored",
s.privateField,
)
} }
} }
@ -355,21 +352,13 @@ skills:
// Verify Channel tokens via Key() methods // Verify Channel tokens via Key() methods
// Telegram // Telegram
assert.Equal( assert.Equal(t, "123456789:ABCdefGHIjklMNOpqrsTUVwxyz", cfg.Channels.Telegram.Token.String())
t,
"123456789:ABCdefGHIjklMNOpqrsTUVwxyz",
cfg.Channels.Telegram.Token.String(),
)
t.Logf("Telegram Token(): %s", cfg.Channels.Telegram.Token.String()) t.Logf("Telegram Token(): %s", cfg.Channels.Telegram.Token.String())
// Feishu // Feishu
assert.Equal(t, "feishu_test_app_secret", cfg.Channels.Feishu.AppSecret.String()) assert.Equal(t, "feishu_test_app_secret", cfg.Channels.Feishu.AppSecret.String())
assert.Equal(t, "feishu_test_encrypt_key", cfg.Channels.Feishu.EncryptKey.String()) assert.Equal(t, "feishu_test_encrypt_key", cfg.Channels.Feishu.EncryptKey.String())
assert.Equal( assert.Equal(t, "feishu_test_verification_token", cfg.Channels.Feishu.VerificationToken.String())
t,
"feishu_test_verification_token",
cfg.Channels.Feishu.VerificationToken.String(),
)
t.Logf("Feishu AppSecret(): %s", cfg.Channels.Feishu.AppSecret.String()) t.Logf("Feishu AppSecret(): %s", cfg.Channels.Feishu.AppSecret.String())
t.Logf("Feishu EncryptKey(): %s", cfg.Channels.Feishu.EncryptKey.String()) t.Logf("Feishu EncryptKey(): %s", cfg.Channels.Feishu.EncryptKey.String())
t.Logf("Feishu VerificationToken(): %s", cfg.Channels.Feishu.VerificationToken.String()) t.Logf("Feishu VerificationToken(): %s", cfg.Channels.Feishu.VerificationToken.String())
@ -394,11 +383,7 @@ skills:
// LINE // LINE
assert.Equal(t, "line_test_channel_secret", cfg.Channels.LINE.ChannelSecret.String()) assert.Equal(t, "line_test_channel_secret", cfg.Channels.LINE.ChannelSecret.String())
assert.Equal( assert.Equal(t, "line_test_channel_access_token", cfg.Channels.LINE.ChannelAccessToken.String())
t,
"line_test_channel_access_token",
cfg.Channels.LINE.ChannelAccessToken.String(),
)
t.Logf("LINE ChannelSecret(): %s", cfg.Channels.LINE.ChannelSecret.String()) t.Logf("LINE ChannelSecret(): %s", cfg.Channels.LINE.ChannelSecret.String())
t.Logf("LINE ChannelAccessToken(): %s", cfg.Channels.LINE.ChannelAccessToken.String()) t.Logf("LINE ChannelAccessToken(): %s", cfg.Channels.LINE.ChannelAccessToken.String())
@ -446,11 +431,7 @@ skills:
assert.Equal(t, "ghp-github-from-file-abc123", cfg.Tools.Skills.Github.Token.String()) assert.Equal(t, "ghp-github-from-file-abc123", cfg.Tools.Skills.Github.Token.String())
t.Logf("Github Token(): %s", cfg.Tools.Skills.Github.Token.String()) t.Logf("Github Token(): %s", cfg.Tools.Skills.Github.Token.String())
assert.Equal( assert.Equal(t, "clawhub-auth-token-from-file", cfg.Tools.Skills.Registries.ClawHub.AuthToken.String())
t,
"clawhub-auth-token-from-file",
cfg.Tools.Skills.Registries.ClawHub.AuthToken.String(),
)
t.Logf("ClawHub AuthToken(): %s", cfg.Tools.Skills.Registries.ClawHub.AuthToken.String()) t.Logf("ClawHub AuthToken(): %s", cfg.Tools.Skills.Registries.ClawHub.AuthToken.String())
t.Log("All security keys are successfully accessible via their respective Key() methods") t.Log("All security keys are successfully accessible via their respective Key() methods")

View file

@ -15,10 +15,7 @@ import (
// JobExecutor is the interface for executing cron jobs through the agent // JobExecutor is the interface for executing cron jobs through the agent
type JobExecutor interface { type JobExecutor interface {
ProcessDirectWithChannel( ProcessDirectWithChannel(ctx context.Context, content, sessionKey, channel, chatID string) (string, error)
ctx context.Context,
content, sessionKey, channel, chatID string,
) (string, error)
// PublishResponseIfNeeded sends response to the outbound bus only when the // PublishResponseIfNeeded sends response to the outbound bus only when the
// agent did not already deliver content through the message tool in this round. // agent did not already deliver content through the message tool in this round.
PublishResponseIfNeeded(ctx context.Context, channel, chatID, response string) PublishResponseIfNeeded(ctx context.Context, channel, chatID, response string)
@ -37,13 +34,8 @@ type CronTool struct {
// NewCronTool creates a new CronTool // NewCronTool creates a new CronTool
// execTimeout: 0 means no timeout, >0 sets the timeout duration // execTimeout: 0 means no timeout, >0 sets the timeout duration
func NewCronTool( func NewCronTool(
cronService *cron.CronService, cronService *cron.CronService, executor JobExecutor, msgBus *bus.MessageBus, workspace string, restrict bool,
executor JobExecutor, execTimeout time.Duration, config *config.Config,
msgBus *bus.MessageBus,
workspace string,
restrict bool,
execTimeout time.Duration,
config *config.Config,
) (*CronTool, error) { ) (*CronTool, error) {
allowCommand := true allowCommand := true
execEnabled := true execEnabled := true
@ -164,9 +156,7 @@ func (t *CronTool) addJob(ctx context.Context, args map[string]any) *ToolResult
chatID := ToolChatID(ctx) chatID := ToolChatID(ctx)
if channel == "" || chatID == "" { if channel == "" || chatID == "" {
return ErrorResult( return ErrorResult("no session context (channel/chat_id not set). Use this tool in an active conversation.")
"no session context (channel/chat_id not set). Use this tool in an active conversation.",
)
} }
message, ok := args["message"].(string) message, ok := args["message"].(string)
@ -218,9 +208,7 @@ func (t *CronTool) addJob(ctx context.Context, args map[string]any) *ToolResult
// Validate type parameter (server-side whitelist, not just LLM schema hint) // Validate type parameter (server-side whitelist, not just LLM schema hint)
msgType, _ := args["type"].(string) msgType, _ := args["type"].(string)
if msgType != "" && msgType != "message" && msgType != "directive" { if msgType != "" && msgType != "message" && msgType != "directive" {
return ErrorResult( return ErrorResult(fmt.Sprintf("invalid type %q, must be 'message' or 'directive'", msgType))
fmt.Sprintf("invalid type %q, must be 'message' or 'directive'", msgType),
)
} }
// GHSA-pv8c-p6jf-3fpp: command scheduling requires internal channel. When // GHSA-pv8c-p6jf-3fpp: command scheduling requires internal channel. When

View file

@ -49,11 +49,7 @@ func (s *stubJobExecutor) PublishResponseIfNeeded(
s.publishedChatID = chatID s.publishedChatID = chatID
} }
func newTestCronToolWithExecutorAndConfig( func newTestCronToolWithExecutorAndConfig(t *testing.T, executor JobExecutor, cfg *config.Config) *CronTool {
t *testing.T,
executor JobExecutor,
cfg *config.Config,
) *CronTool {
t.Helper() t.Helper()
storePath := filepath.Join(t.TempDir(), "cron.json") storePath := filepath.Join(t.TempDir(), "cron.json")
cronService := cron.NewCronService(storePath, nil) cronService := cron.NewCronService(storePath, nil)
@ -106,10 +102,7 @@ func TestCronTool_CommandDoesNotRequireConfirmByDefault(t *testing.T) {
}) })
if result.IsError { if result.IsError {
t.Fatalf( t.Fatalf("expected command scheduling without confirm to succeed by default, got: %s", result.ForLLM)
"expected command scheduling without confirm to succeed by default, got: %s",
result.ForLLM,
)
} }
if !strings.Contains(result.ForLLM, "Cron job added") { if !strings.Contains(result.ForLLM, "Cron job added") {
t.Errorf("expected 'Cron job added', got: %s", result.ForLLM) t.Errorf("expected 'Cron job added', got: %s", result.ForLLM)
@ -197,10 +190,7 @@ func TestCronTool_CommandAllowedFromInternalChannel(t *testing.T) {
}) })
if result.IsError { if result.IsError {
t.Fatalf( t.Fatalf("expected command scheduling to succeed from internal channel, got: %s", result.ForLLM)
"expected command scheduling to succeed from internal channel, got: %s",
result.ForLLM,
)
} }
if !strings.Contains(result.ForLLM, "Cron job added") { if !strings.Contains(result.ForLLM, "Cron job added") {
t.Errorf("expected 'Cron job added', got: %s", result.ForLLM) t.Errorf("expected 'Cron job added', got: %s", result.ForLLM)
@ -235,10 +225,7 @@ func TestCronTool_NonCommandJobAllowedFromRemoteChannel(t *testing.T) {
}) })
if result.IsError { if result.IsError {
t.Fatalf( t.Fatalf("expected non-command reminder to succeed from remote channel, got: %s", result.ForLLM)
"expected non-command reminder to succeed from remote channel, got: %s",
result.ForLLM,
)
} }
} }
@ -310,11 +297,7 @@ func TestCronTool_ExecuteJobPublishesAgentResponse(t *testing.T) {
t.Fatalf("sessionKey = %q, want cron-job-1", executor.lastKey) t.Fatalf("sessionKey = %q, want cron-job-1", executor.lastKey)
} }
if executor.lastChan != "telegram" || executor.lastChatID != "chat-1" { if executor.lastChan != "telegram" || executor.lastChatID != "chat-1" {
t.Fatalf( t.Fatalf("executor target = %s/%s, want telegram/chat-1", executor.lastChan, executor.lastChatID)
"executor target = %s/%s, want telegram/chat-1",
executor.lastChan,
executor.lastChatID,
)
} }
if executor.lastPrompt != "send me a poem" { if executor.lastPrompt != "send me a poem" {
t.Fatalf("prompt = %q, want original message", executor.lastPrompt) t.Fatalf("prompt = %q, want original message", executor.lastPrompt)
@ -323,11 +306,7 @@ func TestCronTool_ExecuteJobPublishesAgentResponse(t *testing.T) {
t.Fatalf("published response = %q, want generated reply", executor.publishedResp) t.Fatalf("published response = %q, want generated reply", executor.publishedResp)
} }
if executor.publishedChan != "telegram" || executor.publishedChatID != "chat-1" { if executor.publishedChan != "telegram" || executor.publishedChatID != "chat-1" {
t.Fatalf( t.Fatalf("published target = %s/%s, want telegram/chat-1", executor.publishedChan, executor.publishedChatID)
"published target = %s/%s, want telegram/chat-1",
executor.publishedChan,
executor.publishedChatID,
)
} }
} }
@ -363,10 +342,7 @@ func TestCronTool_ExecuteJobSkipsWhenMessageToolAlreadySent(t *testing.T) {
} }
if executor.publishedResp != "" { if executor.publishedResp != "" {
t.Fatalf( t.Fatalf("expected no published response when message tool already sent, got: %q", executor.publishedResp)
"expected no published response when message tool already sent, got: %q",
executor.publishedResp,
)
} }
} }
@ -410,9 +386,7 @@ func TestCronTool_ExecuteJobDirectiveWithDeliverRoutesToAgent(t *testing.T) {
} }
if executor.lastPrompt == "" { if executor.lastPrompt == "" {
t.Fatal( t.Fatal("expected agent to be called for directive+deliver, but ProcessDirectWithChannel was not invoked")
"expected agent to be called for directive+deliver, but ProcessDirectWithChannel was not invoked",
)
} }
if executor.publishedResp != "agent processed" { if executor.publishedResp != "agent processed" {
t.Fatalf("published response = %q, want %q", executor.publishedResp, "agent processed") t.Fatalf("published response = %q, want %q", executor.publishedResp, "agent processed")

View file

@ -16,11 +16,7 @@ type EditFileTool struct {
} }
// NewEditFileTool creates a new EditFileTool with optional directory restriction. // NewEditFileTool creates a new EditFileTool with optional directory restriction.
func NewEditFileTool( func NewEditFileTool(workspace string, restrict bool, allowPaths ...[]*regexp.Regexp) *EditFileTool {
workspace string,
restrict bool,
allowPaths ...[]*regexp.Regexp,
) *EditFileTool {
var patterns []*regexp.Regexp var patterns []*regexp.Regexp
if len(allowPaths) > 0 { if len(allowPaths) > 0 {
patterns = allowPaths[0] patterns = allowPaths[0]
@ -83,11 +79,7 @@ type AppendFileTool struct {
fs fileSystem fs fileSystem
} }
func NewAppendFileTool( func NewAppendFileTool(workspace string, restrict bool, allowPaths ...[]*regexp.Regexp) *AppendFileTool {
workspace string,
restrict bool,
allowPaths ...[]*regexp.Regexp,
) *AppendFileTool {
var patterns []*regexp.Regexp var patterns []*regexp.Regexp
if len(allowPaths) > 0 { if len(allowPaths) > 0 {
patterns = allowPaths[0] patterns = allowPaths[0]
@ -174,10 +166,7 @@ func replaceEditContent(content []byte, oldText, newText string) ([]byte, error)
count := strings.Count(contentStr, oldText) count := strings.Count(contentStr, oldText)
if count > 1 { if count > 1 {
return nil, fmt.Errorf( return nil, fmt.Errorf("old_text appears %d times. Please provide more context to make it unique", count)
"old_text appears %d times. Please provide more context to make it unique",
count,
)
} }
newContent := strings.Replace(contentStr, oldText, newText, 1) newContent := strings.Replace(contentStr, oldText, newText, 1)

View file

@ -76,8 +76,7 @@ func TestEditTool_EditFile_NotFound(t *testing.T) {
} }
// Should mention file not found // Should mention file not found
if !strings.Contains(result.ForLLM, "not found") && if !strings.Contains(result.ForLLM, "not found") && !strings.Contains(result.ForUser, "not found") {
!strings.Contains(result.ForUser, "not found") {
t.Errorf("Expected 'file not found' message, got ForLLM: %s", result.ForLLM) t.Errorf("Expected 'file not found' message, got ForLLM: %s", result.ForLLM)
} }
} }
@ -104,8 +103,7 @@ func TestEditTool_EditFile_OldTextNotFound(t *testing.T) {
} }
// Should mention old_text not found // Should mention old_text not found
if !strings.Contains(result.ForLLM, "not found") && if !strings.Contains(result.ForLLM, "not found") && !strings.Contains(result.ForUser, "not found") {
!strings.Contains(result.ForUser, "not found") {
t.Errorf("Expected 'not found' message, got ForLLM: %s", result.ForLLM) t.Errorf("Expected 'not found' message, got ForLLM: %s", result.ForLLM)
} }
} }

View file

@ -20,11 +20,7 @@ import (
const MaxReadFileSize = 64 * 1024 // 64KB limit to avoid context overflow const MaxReadFileSize = 64 * 1024 // 64KB limit to avoid context overflow
func validatePathWithAllowPaths( func validatePathWithAllowPaths(path, workspace string, restrict bool, patterns []*regexp.Regexp) (string, error) {
path, workspace string,
restrict bool,
patterns []*regexp.Regexp,
) (string, error) {
if workspace == "" { if workspace == "" {
return path, fmt.Errorf("workspace is not defined") return path, fmt.Errorf("workspace is not defined")
} }
@ -487,11 +483,7 @@ type WriteFileTool struct {
fs fileSystem fs fileSystem
} }
func NewWriteFileTool( func NewWriteFileTool(workspace string, restrict bool, allowPaths ...[]*regexp.Regexp) *WriteFileTool {
workspace string,
restrict bool,
allowPaths ...[]*regexp.Regexp,
) *WriteFileTool {
var patterns []*regexp.Regexp var patterns []*regexp.Regexp
if len(allowPaths) > 0 { if len(allowPaths) > 0 {
patterns = allowPaths[0] patterns = allowPaths[0]
@ -544,9 +536,7 @@ func (t *WriteFileTool) Execute(ctx context.Context, args map[string]any) *ToolR
if !overwrite { if !overwrite {
if _, err := t.fs.Open(path); err == nil { if _, err := t.fs.Open(path); err == nil {
return ErrorResult( return ErrorResult(fmt.Sprintf("file: %s already exists. Set overwrite=true to replace.", path))
fmt.Sprintf("file: %s already exists. Set overwrite=true to replace.", path),
)
} }
} }

View file

@ -59,13 +59,8 @@ func TestFilesystemTool_ReadFile_NotFound(t *testing.T) {
} }
// Should contain error message // Should contain error message
if !strings.Contains(result.ForLLM, "failed to open file") && if !strings.Contains(result.ForLLM, "failed to open file") && !strings.Contains(result.ForUser, "failed to read") {
!strings.Contains(result.ForUser, "failed to read") { t.Errorf("Expected error message, got ForLLM: %s, ForUser: %s", result.ForLLM, result.ForUser)
t.Errorf(
"Expected error message, got ForLLM: %s, ForUser: %s",
result.ForLLM,
result.ForUser,
)
} }
} }
@ -83,8 +78,7 @@ func TestFilesystemTool_ReadFile_MissingPath(t *testing.T) {
} }
// Should mention required parameter // Should mention required parameter
if !strings.Contains(result.ForLLM, "path is required") && if !strings.Contains(result.ForLLM, "path is required") && !strings.Contains(result.ForUser, "path is required") {
!strings.Contains(result.ForUser, "path is required") {
t.Errorf("Expected 'path is required' message, got ForLLM: %s", result.ForLLM) t.Errorf("Expected 'path is required' message, got ForLLM: %s", result.ForLLM)
} }
} }
@ -303,12 +297,7 @@ func TestFilesystemTool_WriteFile_OverwriteSandboxed(t *testing.T) {
"content": "replaced in sandbox", "content": "replaced in sandbox",
"overwrite": true, "overwrite": true,
}) })
assert.False( assert.False(t, result.IsError, "expected success in sandbox mode with overwrite=true, got: %s", result.ForLLM)
t,
result.IsError,
"expected success in sandbox mode with overwrite=true, got: %s",
result.ForLLM,
)
data, err := os.ReadFile(filepath.Join(workspace, testFile)) data, err := os.ReadFile(filepath.Join(workspace, testFile))
assert.NoError(t, err) assert.NoError(t, err)
@ -336,8 +325,7 @@ func TestFilesystemTool_ListDir_Success(t *testing.T) {
} }
// Should list files and directories // Should list files and directories
if !strings.Contains(result.ForLLM, "file1.txt") || if !strings.Contains(result.ForLLM, "file1.txt") || !strings.Contains(result.ForLLM, "file2.txt") {
!strings.Contains(result.ForLLM, "file2.txt") {
t.Errorf("Expected files in listing, got: %s", result.ForLLM) t.Errorf("Expected files in listing, got: %s", result.ForLLM)
} }
if !strings.Contains(result.ForLLM, "subdir") { if !strings.Contains(result.ForLLM, "subdir") {
@ -361,13 +349,8 @@ func TestFilesystemTool_ListDir_NotFound(t *testing.T) {
} }
// Should contain error message // Should contain error message
if !strings.Contains(result.ForLLM, "failed to read") && if !strings.Contains(result.ForLLM, "failed to read") && !strings.Contains(result.ForUser, "failed to read") {
!strings.Contains(result.ForUser, "failed to read") { t.Errorf("Expected error message, got ForLLM: %s, ForUser: %s", result.ForLLM, result.ForUser)
t.Errorf(
"Expected error message, got ForLLM: %s, ForUser: %s",
result.ForLLM,
result.ForUser,
)
} }
} }
@ -414,8 +397,7 @@ func TestFilesystemTool_ReadFile_RejectsSymlinkEscape(t *testing.T) {
// os.Root might return different errors depending on platform/implementation // os.Root might return different errors depending on platform/implementation
// but it definitely should error. // but it definitely should error.
// Our wrapper returns "access denied or file not found" // Our wrapper returns "access denied or file not found"
if !strings.Contains(result.ForLLM, "access denied") && if !strings.Contains(result.ForLLM, "access denied") && !strings.Contains(result.ForLLM, "file not found") &&
!strings.Contains(result.ForLLM, "file not found") &&
!strings.Contains(result.ForLLM, "no such file") { !strings.Contains(result.ForLLM, "no such file") {
t.Fatalf("expected symlink escape error, got: %s", result.ForLLM) t.Fatalf("expected symlink escape error, got: %s", result.ForLLM)
} }
@ -434,20 +416,10 @@ func TestFilesystemTool_EmptyWorkspace_AccessDenied(t *testing.T) {
}) })
// We EXPECT IsError=true (access blocked due to empty workspace) // We EXPECT IsError=true (access blocked due to empty workspace)
assert.True( assert.True(t, result.IsError, "Security Regression: Empty workspace allowed access! content: %s", result.ForLLM)
t,
result.IsError,
"Security Regression: Empty workspace allowed access! content: %s",
result.ForLLM,
)
// Verify it failed for the right reason // Verify it failed for the right reason
assert.Contains( assert.Contains(t, result.ForLLM, "workspace is not defined", "Expected 'workspace is not defined' error")
t,
result.ForLLM,
"workspace is not defined",
"Expected 'workspace is not defined' error",
)
} }
// TestRootMkdirAll verifies that root.MkdirAll (used by atomicWriteFileInRoot) handles all cases: // TestRootMkdirAll verifies that root.MkdirAll (used by atomicWriteFileInRoot) handles all cases:
@ -681,10 +653,7 @@ func TestWhitelistFs_BlocksSymlinkEscapeInAllowedDir(t *testing.T) {
patterns := []*regexp.Regexp{regexp.MustCompile(`^` + regexp.QuoteMeta(allowedDir))} patterns := []*regexp.Regexp{regexp.MustCompile(`^` + regexp.QuoteMeta(allowedDir))}
tool := NewReadFileTool(workspace, true, MaxReadFileSize, patterns) tool := NewReadFileTool(workspace, true, MaxReadFileSize, patterns)
result := tool.Execute( result := tool.Execute(context.Background(), map[string]any{"path": filepath.Join(linkPath, "secret.txt")})
context.Background(),
map[string]any{"path": filepath.Join(linkPath, "secret.txt")},
)
if !result.IsError { if !result.IsError {
t.Fatalf("expected symlink escape from allowed dir to be blocked, got: %s", result.ForLLM) t.Fatalf("expected symlink escape from allowed dir to be blocked, got: %s", result.ForLLM)
} }

View file

@ -65,9 +65,7 @@ func (t *I2CTool) Parameters() map[string]any {
func (t *I2CTool) Execute(ctx context.Context, args map[string]any) *ToolResult { func (t *I2CTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
if runtime.GOOS != "linux" { if runtime.GOOS != "linux" {
return ErrorResult( return ErrorResult("I2C is only supported on Linux. This tool requires /dev/i2c-* device files.")
"I2C is only supported on Linux. This tool requires /dev/i2c-* device files.",
)
} }
action, ok := args["action"].(string) action, ok := args["action"].(string)
@ -85,9 +83,7 @@ func (t *I2CTool) Execute(ctx context.Context, args map[string]any) *ToolResult
case "write": case "write":
return t.writeDevice(args) return t.writeDevice(args)
default: default:
return ErrorResult( return ErrorResult(fmt.Sprintf("unknown action: %s (valid: detect, scan, read, write)", action))
fmt.Sprintf("unknown action: %s (valid: detect, scan, read, write)", action),
)
} }
} }

View file

@ -55,12 +55,7 @@ func smbusProbe(fd int, addr int, hasQuick bool) bool {
size: i2cSmbusQuick, size: i2cSmbusQuick,
data: nil, data: nil,
} }
_, _, errno := syscall.Syscall( _, _, errno := syscall.Syscall(syscall.SYS_IOCTL, uintptr(fd), i2cSmbus, uintptr(unsafe.Pointer(&args)))
syscall.SYS_IOCTL,
uintptr(fd),
i2cSmbus,
uintptr(unsafe.Pointer(&args)),
)
return errno == 0 return errno == 0
} }
@ -72,12 +67,7 @@ func smbusProbe(fd int, addr int, hasQuick bool) bool {
size: i2cSmbusByte, size: i2cSmbusByte,
data: &data, data: &data,
} }
_, _, errno := syscall.Syscall( _, _, errno := syscall.Syscall(syscall.SYS_IOCTL, uintptr(fd), i2cSmbus, uintptr(unsafe.Pointer(&args)))
syscall.SYS_IOCTL,
uintptr(fd),
i2cSmbus,
uintptr(unsafe.Pointer(&args)),
)
return errno == 0 return errno == 0
} }
@ -93,29 +83,16 @@ func (t *I2CTool) scan(args map[string]any) *ToolResult {
devPath := fmt.Sprintf("/dev/i2c-%s", bus) devPath := fmt.Sprintf("/dev/i2c-%s", bus)
fd, err := syscall.Open(devPath, syscall.O_RDWR, 0) fd, err := syscall.Open(devPath, syscall.O_RDWR, 0)
if err != nil { if err != nil {
return ErrorResult( return ErrorResult(fmt.Sprintf("failed to open %s: %v (check permissions and i2c-dev module)", devPath, err))
fmt.Sprintf(
"failed to open %s: %v (check permissions and i2c-dev module)",
devPath,
err,
),
)
} }
defer syscall.Close(fd) defer syscall.Close(fd)
// Query adapter capabilities to determine available probe methods. // Query adapter capabilities to determine available probe methods.
// I2C_FUNCS writes an unsigned long, which is word-sized on Linux. // I2C_FUNCS writes an unsigned long, which is word-sized on Linux.
var funcs uintptr var funcs uintptr
_, _, errno := syscall.Syscall( _, _, errno := syscall.Syscall(syscall.SYS_IOCTL, uintptr(fd), i2cFuncs, uintptr(unsafe.Pointer(&funcs)))
syscall.SYS_IOCTL,
uintptr(fd),
i2cFuncs,
uintptr(unsafe.Pointer(&funcs)),
)
if errno != 0 { if errno != 0 {
return ErrorResult( return ErrorResult(fmt.Sprintf("failed to query I2C adapter capabilities on %s: %v", devPath, errno))
fmt.Sprintf("failed to query I2C adapter capabilities on %s: %v", devPath, errno),
)
} }
hasQuick := funcs&i2cFuncSmbusQuick != 0 hasQuick := funcs&i2cFuncSmbusQuick != 0
@ -123,10 +100,7 @@ func (t *I2CTool) scan(args map[string]any) *ToolResult {
if !hasQuick && !hasReadByte { if !hasQuick && !hasReadByte {
return ErrorResult( return ErrorResult(
fmt.Sprintf( fmt.Sprintf("I2C adapter %s supports neither SMBus Quick nor Read Byte — cannot probe safely", devPath),
"I2C adapter %s supports neither SMBus Quick nor Read Byte — cannot probe safely",
devPath,
),
) )
} }
@ -158,9 +132,7 @@ func (t *I2CTool) scan(args map[string]any) *ToolResult {
} }
if len(found) == 0 { if len(found) == 0 {
return SilentResult( return SilentResult(fmt.Sprintf("No devices found on %s. Check wiring and pull-up resistors.", devPath))
fmt.Sprintf("No devices found on %s. Check wiring and pull-up resistors.", devPath),
)
} }
result, _ := json.MarshalIndent(map[string]any{ result, _ := json.MarshalIndent(map[string]any{

View file

@ -314,10 +314,7 @@ func (t *MCPTool) normalizeResultContent(ctx context.Context, content []mcp.Cont
return result return result
} }
func (t *MCPTool) storeEmbeddedResource( func (t *MCPTool) storeEmbeddedResource(ctx context.Context, content *mcp.EmbeddedResource) (string, string) {
ctx context.Context,
content *mcp.EmbeddedResource,
) (string, string) {
if content == nil || content.Resource == nil { if content == nil || content.Resource == nil {
return "", "[MCP returned an embedded resource without data.]" return "", "[MCP returned an embedded resource without data.]"
} }
@ -377,39 +374,23 @@ func (t *MCPTool) storeBinaryContent(
dir := media.TempDir() dir := media.TempDir()
if err := os.MkdirAll(dir, 0o700); err != nil { if err := os.MkdirAll(dir, 0o700); err != nil {
return "", fmt.Sprintf( return "", fmt.Sprintf("[MCP returned %s content (%s) but it could not be stored.]", kind, mimeType)
"[MCP returned %s content (%s) but it could not be stored.]",
kind,
mimeType,
)
} }
ext := extensionForMIMEType(mimeType) ext := extensionForMIMEType(mimeType)
tmpFile, err := os.CreateTemp(dir, "mcp-*"+ext) tmpFile, err := os.CreateTemp(dir, "mcp-*"+ext)
if err != nil { if err != nil {
return "", fmt.Sprintf( return "", fmt.Sprintf("[MCP returned %s content (%s) but it could not be stored.]", kind, mimeType)
"[MCP returned %s content (%s) but it could not be stored.]",
kind,
mimeType,
)
} }
tmpPath := tmpFile.Name() tmpPath := tmpFile.Name()
if _, err = tmpFile.Write(data); err != nil { if _, err = tmpFile.Write(data); err != nil {
_ = tmpFile.Close() _ = tmpFile.Close()
_ = os.Remove(tmpPath) _ = os.Remove(tmpPath)
return "", fmt.Sprintf( return "", fmt.Sprintf("[MCP returned %s content (%s) but it could not be stored.]", kind, mimeType)
"[MCP returned %s content (%s) but it could not be stored.]",
kind,
mimeType,
)
} }
if err = tmpFile.Close(); err != nil { if err = tmpFile.Close(); err != nil {
_ = os.Remove(tmpPath) _ = os.Remove(tmpPath)
return "", fmt.Sprintf( return "", fmt.Sprintf("[MCP returned %s content (%s) but it could not be stored.]", kind, mimeType)
"[MCP returned %s content (%s) but it could not be stored.]",
kind,
mimeType,
)
} }
scope := fmt.Sprintf( scope := fmt.Sprintf(
@ -489,10 +470,7 @@ func summarizeEmbeddedResource(content *mcp.EmbeddedResource) string {
normalizedMIMEType(resource.MIMEType), normalizedMIMEType(resource.MIMEType),
) )
} }
return fmt.Sprintf( return fmt.Sprintf("[MCP returned embedded resource (%s).]", normalizedMIMEType(resource.MIMEType))
"[MCP returned embedded resource (%s).]",
normalizedMIMEType(resource.MIMEType),
)
} }
func annotationsAllowUser(annotations *mcp.Annotations) bool { func annotationsAllowUser(annotations *mcp.Annotations) bool {

View file

@ -571,10 +571,7 @@ func TestMCPTool_Execute_EmbeddedResourceBlobStoredAsMedia(t *testing.T) {
result := mcpTool.Execute(WithToolContext(context.Background(), "telegram", "chat-42"), nil) result := mcpTool.Execute(WithToolContext(context.Background(), "telegram", "chat-42"), nil)
if len(result.Media) != 1 { if len(result.Media) != 1 {
t.Fatalf( t.Fatalf("expected embedded resource blob to be stored as media, got %d refs", len(result.Media))
"expected embedded resource blob to be stored as media, got %d refs",
len(result.Media),
)
} }
path, _, err := store.ResolveWithMeta(result.Media[0]) path, _, err := store.ResolveWithMeta(result.Media[0])
if err != nil { if err != nil {

View file

@ -43,10 +43,7 @@ func TestMessageTool_Execute_Success(t *testing.T) {
// - ForLLM contains send status description // - ForLLM contains send status description
if result.ForLLM != "Message sent to test-channel:test-chat-id" { if result.ForLLM != "Message sent to test-channel:test-chat-id" {
t.Errorf( t.Errorf("Expected ForLLM 'Message sent to test-channel:test-chat-id', got '%s'", result.ForLLM)
"Expected ForLLM 'Message sent to test-channel:test-chat-id', got '%s'",
result.ForLLM,
)
} }
// - ForUser is empty (user already received message directly) // - ForUser is empty (user already received message directly)
@ -91,10 +88,7 @@ func TestMessageTool_Execute_WithCustomChannel(t *testing.T) {
t.Error("Expected Silent=true") t.Error("Expected Silent=true")
} }
if result.ForLLM != "Message sent to custom-channel:custom-chat-id" { if result.ForLLM != "Message sent to custom-channel:custom-chat-id" {
t.Errorf( t.Errorf("Expected ForLLM 'Message sent to custom-channel:custom-chat-id', got '%s'", result.ForLLM)
"Expected ForLLM 'Message sent to custom-channel:custom-chat-id', got '%s'",
result.ForLLM,
)
} }
} }

View file

@ -215,43 +215,28 @@ func storeInlineDataURL(
payload = strings.NewReplacer("\n", "", "\r", "", "\t", "", " ", "").Replace(payload) payload = strings.NewReplacer("\n", "", "\r", "", "\t", "", " ", "").Replace(payload)
decoded, err := base64.StdEncoding.DecodeString(payload) decoded, err := base64.StdEncoding.DecodeString(payload)
if err != nil { if err != nil {
return "", fmt.Sprintf( return "", fmt.Sprintf("[Tool returned inline media content (%s) that could not be decoded.]", mimeType)
"[Tool returned inline media content (%s) that could not be decoded.]",
mimeType,
)
} }
dir := media.TempDir() dir := media.TempDir()
if err = os.MkdirAll(dir, 0o700); err != nil { if err = os.MkdirAll(dir, 0o700); err != nil {
return "", fmt.Sprintf( return "", fmt.Sprintf("[Tool returned inline media content (%s) but it could not be stored.]", mimeType)
"[Tool returned inline media content (%s) but it could not be stored.]",
mimeType,
)
} }
ext := extensionForMIMEType(mimeType) ext := extensionForMIMEType(mimeType)
tmpFile, err := os.CreateTemp(dir, "tool-inline-*"+ext) tmpFile, err := os.CreateTemp(dir, "tool-inline-*"+ext)
if err != nil { if err != nil {
return "", fmt.Sprintf( return "", fmt.Sprintf("[Tool returned inline media content (%s) but it could not be stored.]", mimeType)
"[Tool returned inline media content (%s) but it could not be stored.]",
mimeType,
)
} }
tmpPath := tmpFile.Name() tmpPath := tmpFile.Name()
if _, err = tmpFile.Write(decoded); err != nil { if _, err = tmpFile.Write(decoded); err != nil {
tmpFile.Close() tmpFile.Close()
_ = os.Remove(tmpPath) _ = os.Remove(tmpPath)
return "", fmt.Sprintf( return "", fmt.Sprintf("[Tool returned inline media content (%s) but it could not be stored.]", mimeType)
"[Tool returned inline media content (%s) but it could not be stored.]",
mimeType,
)
} }
if err = tmpFile.Close(); err != nil { if err = tmpFile.Close(); err != nil {
_ = os.Remove(tmpPath) _ = os.Remove(tmpPath)
return "", fmt.Sprintf( return "", fmt.Sprintf("[Tool returned inline media content (%s) but it could not be stored.]", mimeType)
"[Tool returned inline media content (%s) but it could not be stored.]",
mimeType,
)
} }
filename := sanitizeIdentifierComponent(toolName) + ext filename := sanitizeIdentifierComponent(toolName) + ext
@ -270,10 +255,7 @@ func storeInlineDataURL(
}, scope) }, scope)
if err != nil { if err != nil {
_ = os.Remove(tmpPath) _ = os.Remove(tmpPath)
return "", fmt.Sprintf( return "", fmt.Sprintf("[Tool returned inline media content (%s) but it could not be registered.]", mimeType)
"[Tool returned inline media content (%s) but it could not be registered.]",
mimeType,
)
} }
return ref, fmt.Sprintf(inlineMediaStoredMessage, mimeType) return ref, fmt.Sprintf(inlineMediaStoredMessage, mimeType)

View file

@ -80,10 +80,7 @@ func (tr *ToolResult) ContentForLLM() string {
} }
} }
if len(tr.ArtifactTags) > 0 { if len(tr.ArtifactTags) > 0 {
artifactNote := "Local artifact paths: " + strings.Join( artifactNote := "Local artifact paths: " + strings.Join(tr.ArtifactTags, " ") + "\n" + artifactPathsLLMNote
tr.ArtifactTags,
" ",
) + "\n" + artifactPathsLLMNote
if content == "" { if content == "" {
content = artifactNote content = artifactNote
} else if !strings.Contains(content, artifactNote) { } else if !strings.Contains(content, artifactNote) {

View file

@ -142,11 +142,7 @@ func TestToolResultJSONSerialization(t *testing.T) {
t.Errorf("ForLLM mismatch: got '%s', want '%s'", decoded.ForLLM, tt.result.ForLLM) t.Errorf("ForLLM mismatch: got '%s', want '%s'", decoded.ForLLM, tt.result.ForLLM)
} }
if decoded.ForUser != tt.result.ForUser { if decoded.ForUser != tt.result.ForUser {
t.Errorf( t.Errorf("ForUser mismatch: got '%s', want '%s'", decoded.ForUser, tt.result.ForUser)
"ForUser mismatch: got '%s', want '%s'",
decoded.ForUser,
tt.result.ForUser,
)
} }
if decoded.Silent != tt.result.Silent { if decoded.Silent != tt.result.Silent {
t.Errorf("Silent mismatch: got %v, want %v", decoded.Silent, tt.result.Silent) t.Errorf("Silent mismatch: got %v, want %v", decoded.Silent, tt.result.Silent)

View file

@ -56,38 +56,19 @@ func (t *RegexSearchTool) Execute(ctx context.Context, args map[string]any) *Too
} }
if len(pattern) > MaxRegexPatternLength { if len(pattern) > MaxRegexPatternLength {
logger.WarnCF( logger.WarnCF("discovery", "Regex pattern rejected (too long)", map[string]any{"len": len(pattern)})
"discovery", return ErrorResult(fmt.Sprintf("Pattern too long: max %d characters allowed", MaxRegexPatternLength))
"Regex pattern rejected (too long)",
map[string]any{"len": len(pattern)},
)
return ErrorResult(
fmt.Sprintf("Pattern too long: max %d characters allowed", MaxRegexPatternLength),
)
} }
logger.DebugCF("discovery", "Regex search", map[string]any{"pattern": pattern}) logger.DebugCF("discovery", "Regex search", map[string]any{"pattern": pattern})
res, err := t.registry.SearchRegex(pattern, t.maxSearchResults) res, err := t.registry.SearchRegex(pattern, t.maxSearchResults)
if err != nil { if err != nil {
logger.WarnCF( logger.WarnCF("discovery", "Invalid regex pattern", map[string]any{"pattern": pattern, "error": err.Error()})
"discovery", return ErrorResult(fmt.Sprintf("Invalid regex pattern syntax: %v. Please fix your regex and try again.", err))
"Invalid regex pattern",
map[string]any{"pattern": pattern, "error": err.Error()},
)
return ErrorResult(
fmt.Sprintf(
"Invalid regex pattern syntax: %v. Please fix your regex and try again.",
err,
),
)
} }
logger.InfoCF( logger.InfoCF("discovery", "Regex search completed", map[string]any{"pattern": pattern, "results": len(res)})
"discovery",
"Regex search completed",
map[string]any{"pattern": pattern, "results": len(res)},
)
return formatDiscoveryResponse(t.registry, res, t.ttl) return formatDiscoveryResponse(t.registry, res, t.ttl)
} }
@ -157,11 +138,7 @@ func (t *BM25SearchTool) Execute(ctx context.Context, args map[string]any) *Tool
} }
} }
logger.InfoCF( logger.InfoCF("discovery", "BM25 search completed", map[string]any{"query": query, "results": len(results)})
"discovery",
"BM25 search completed",
map[string]any{"query": query, "results": len(results)},
)
return formatDiscoveryResponse(t.registry, results, t.ttl) return formatDiscoveryResponse(t.registry, results, t.ttl)
} }
@ -173,10 +150,7 @@ type ToolSearchResult struct {
Description string `json:"description"` Description string `json:"description"`
} }
func (r *ToolRegistry) SearchRegex( func (r *ToolRegistry) SearchRegex(pattern string, maxSearchResults int) ([]ToolSearchResult, error) {
pattern string,
maxSearchResults int,
) ([]ToolSearchResult, error) {
if maxSearchResults <= 0 { if maxSearchResults <= 0 {
return nil, nil return nil, nil
} }
@ -214,11 +188,7 @@ func (r *ToolRegistry) SearchRegex(
return results, nil return results, nil
} }
func formatDiscoveryResponse( func formatDiscoveryResponse(registry *ToolRegistry, results []ToolSearchResult, ttl int) *ToolResult {
registry *ToolRegistry,
results []ToolSearchResult,
ttl int,
) *ToolResult {
if len(results) == 0 { if len(results) == 0 {
return SilentResult("No tools found matching the query.") return SilentResult("No tools found matching the query.")
} }
@ -304,11 +274,7 @@ func (t *BM25SearchTool) getOrBuildEngine() *bm25CachedEngine {
cached := &bm25CachedEngine{engine: buildBM25Engine(docs)} cached := &bm25CachedEngine{engine: buildBM25Engine(docs)}
t.cachedEngine = cached t.cachedEngine = cached
t.cacheVersion = snap.Version t.cacheVersion = snap.Version
logger.DebugCF( logger.DebugCF("discovery", "BM25 engine rebuilt", map[string]any{"docs": len(docs), "version": snap.Version})
"discovery",
"BM25 engine rebuilt",
map[string]any{"docs": len(docs), "version": snap.Version},
)
return cached return cached
} }

View file

@ -93,10 +93,7 @@ func TestRegexSearchTool_Execute(t *testing.T) {
reg.mu.RLock() reg.mu.RLock()
defer reg.mu.RUnlock() defer reg.mu.RUnlock()
if reg.tools["mcp_read_file"].TTL != 5 { if reg.tools["mcp_read_file"].TTL != 5 {
t.Errorf( t.Errorf("Expected TTL of 'mcp_read_file' to be promoted to 5, got %d", reg.tools["mcp_read_file"].TTL)
"Expected TTL of 'mcp_read_file' to be promoted to 5, got %d",
reg.tools["mcp_read_file"].TTL,
)
} }
if reg.tools["mcp_fetch_net"].TTL != 0 { if reg.tools["mcp_fetch_net"].TTL != 0 {
t.Errorf("Expected 'mcp_fetch_net' to NOT be promoted (TTL=0)") t.Errorf("Expected 'mcp_fetch_net' to NOT be promoted (TTL=0)")

View file

@ -142,10 +142,7 @@ func (t *SendFileTool) Execute(ctx context.Context, args map[string]any) *ToolRe
return ErrorResult(fmt.Sprintf("failed to register media: %v", err)) return ErrorResult(fmt.Sprintf("failed to register media: %v", err))
} }
return MediaResult( return MediaResult(fmt.Sprintf("File %q sent to user", filename), []string{ref}).WithResponseHandled()
fmt.Sprintf("File %q sent to user", filename),
[]string{ref},
).WithResponseHandled()
} }
// detectMediaType determines the MIME type of a file. // detectMediaType determines the MIME type of a file.

View file

@ -79,11 +79,7 @@ func TestSendFileTool_FileTooLarge(t *testing.T) {
func TestSendFileTool_DefaultMaxSize(t *testing.T) { func TestSendFileTool_DefaultMaxSize(t *testing.T) {
tool := NewSendFileTool("/tmp", false, 0, nil) tool := NewSendFileTool("/tmp", false, 0, nil)
if tool.maxFileSize != config.DefaultMaxMediaSize { if tool.maxFileSize != config.DefaultMaxMediaSize {
t.Errorf( t.Errorf("expected default max size %d, got %d", config.DefaultMaxMediaSize, tool.maxFileSize)
"expected default max size %d, got %d",
config.DefaultMaxMediaSize,
tool.maxFileSize,
)
} }
} }
@ -166,11 +162,7 @@ func TestSendFileTool_AllowsWhitelistedMediaTempPath(t *testing.T) {
t.Cleanup(func() { _ = os.Remove(testPath) }) t.Cleanup(func() { _ = os.Remove(testPath) })
pattern := regexp.MustCompile( pattern := regexp.MustCompile(
"^" + regexp.QuoteMeta( "^" + regexp.QuoteMeta(filepath.Clean(mediaDir)) + "(?:" + regexp.QuoteMeta(string(os.PathSeparator)) + "|$)",
filepath.Clean(mediaDir),
) + "(?:" + regexp.QuoteMeta(
string(os.PathSeparator),
) + "|$)",
) )
store := media.NewFileMediaStore() store := media.NewFileMediaStore()

View file

@ -113,11 +113,7 @@ var (
} }
) )
func NewExecTool( func NewExecTool(workingDir string, restrict bool, allowPaths ...[]*regexp.Regexp) (*ExecTool, error) {
workingDir string,
restrict bool,
allowPaths ...[]*regexp.Regexp,
) (*ExecTool, error) {
return NewExecToolWithConfig(workingDir, restrict, nil, allowPaths...) return NewExecToolWithConfig(workingDir, restrict, nil, allowPaths...)
} }
@ -197,16 +193,8 @@ func (t *ExecTool) Parameters() map[string]any {
"type": "object", "type": "object",
"properties": map[string]any{ "properties": map[string]any{
"action": map[string]any{ "action": map[string]any{
"type": "string", "type": "string",
"enum": []string{ "enum": []string{"run", "list", "poll", "read", "write", "kill", "send-keys"},
"run",
"list",
"poll",
"read",
"write",
"kill",
"send-keys",
},
"description": "Action: run (execute command), list (show sessions), poll (check status), read (get output), write (send input), kill (terminate), send-keys (send keys to PTY)", "description": "Action: run (execute command), list (show sessions), poll (check status), read (get output), write (send input), kill (terminate), send-keys (send keys to PTY)",
}, },
"command": map[string]any{ "command": map[string]any{
@ -312,12 +300,7 @@ func (t *ExecTool) executeRun(ctx context.Context, args map[string]any) *ToolRes
cwd := t.workingDir cwd := t.workingDir
if wd, ok := args["cwd"].(string); ok && wd != "" { if wd, ok := args["cwd"].(string); ok && wd != "" {
if t.restrictToWorkspace && t.workingDir != "" { if t.restrictToWorkspace && t.workingDir != "" {
resolvedWD, err := validatePathWithAllowPaths( resolvedWD, err := validatePathWithAllowPaths(wd, t.workingDir, true, t.allowedPathPatterns)
wd,
t.workingDir,
true,
t.allowedPathPatterns,
)
if err != nil { if err != nil {
return ErrorResult("Command blocked by safety guard (" + err.Error() + ")") return ErrorResult("Command blocked by safety guard (" + err.Error() + ")")
} }
@ -343,9 +326,7 @@ func (t *ExecTool) executeRun(ctx context.Context, args map[string]any) *ToolRes
if t.restrictToWorkspace && t.workingDir != "" && cwd != t.workingDir { if t.restrictToWorkspace && t.workingDir != "" && cwd != t.workingDir {
resolved, err := filepath.EvalSymlinks(cwd) resolved, err := filepath.EvalSymlinks(cwd)
if err != nil { if err != nil {
return ErrorResult( return ErrorResult(fmt.Sprintf("Command blocked by safety guard (path resolution failed: %v)", err))
fmt.Sprintf("Command blocked by safety guard (path resolution failed: %v)", err),
)
} }
if isAllowedPath(resolved, t.allowedPathPatterns) { if isAllowedPath(resolved, t.allowedPathPatterns) {
cwd = resolved cwd = resolved
@ -383,14 +364,7 @@ func (t *ExecTool) runSync(ctx context.Context, command, cwd string) *ToolResult
var cmd *exec.Cmd var cmd *exec.Cmd
if runtime.GOOS == "windows" { if runtime.GOOS == "windows" {
cmd = exec.CommandContext( cmd = exec.CommandContext(cmdCtx, "powershell", "-NoProfile", "-NonInteractive", "-Command", command)
cmdCtx,
"powershell",
"-NoProfile",
"-NonInteractive",
"-Command",
command,
)
} else { } else {
cmd = exec.CommandContext(cmdCtx, "sh", "-c", command) cmd = exec.CommandContext(cmdCtx, "sh", "-c", command)
} }
@ -468,10 +442,7 @@ func (t *ExecTool) runSync(ctx context.Context, command, cwd string) *ToolResult
maxLen := 10000 maxLen := 10000
if len(output) > maxLen { if len(output) > maxLen {
output = output[:maxLen] + fmt.Sprintf( output = output[:maxLen] + fmt.Sprintf("\n... (truncated, %d more chars)", len(output)-maxLen)
"\n... (truncated, %d more chars)",
len(output)-maxLen,
)
} }
if err != nil { if err != nil {
@ -489,11 +460,7 @@ func (t *ExecTool) runSync(ctx context.Context, command, cwd string) *ToolResult
} }
} }
func (t *ExecTool) runBackground( func (t *ExecTool) runBackground(ctx context.Context, command, cwd string, ptyEnabled bool) *ToolResult {
ctx context.Context,
command, cwd string,
ptyEnabled bool,
) *ToolResult {
sessionID := generateSessionID() sessionID := generateSessionID()
session := &ProcessSession{ session := &ProcessSession{
ID: sessionID, ID: sessionID,
@ -586,8 +553,7 @@ func (t *ExecTool) runBackground(
n, err := session.ptyMaster.Read(buf) n, err := session.ptyMaster.Read(buf)
if n > 0 { if n > 0 {
raw := string(buf[:n]) raw := string(buf[:n])
if mode := detectPtyKeyMode(raw); mode != PtyKeyModeNotFound && if mode := detectPtyKeyMode(raw); mode != PtyKeyModeNotFound && mode != session.GetPtyKeyMode() {
mode != session.GetPtyKeyMode() {
session.SetPtyKeyMode(mode) session.SetPtyKeyMode(mode)
} }
@ -768,16 +734,12 @@ func (t *ExecTool) executeWrite(args map[string]any) *ToolResult {
} }
if session.IsDone() { if session.IsDone() {
return ErrorResult( return ErrorResult(fmt.Sprintf("process already exited with code %d", session.GetExitCode()))
fmt.Sprintf("process already exited with code %d", session.GetExitCode()),
)
} }
if err := session.Write(data); err != nil { if err := session.Write(data); err != nil {
if errors.Is(err, ErrSessionDone) { if errors.Is(err, ErrSessionDone) {
return ErrorResult( return ErrorResult(fmt.Sprintf("process already exited with code %d", session.GetExitCode()))
fmt.Sprintf("process already exited with code %d", session.GetExitCode()),
)
} }
return ErrorResult(fmt.Sprintf("failed to write to session: %v", err)) return ErrorResult(fmt.Sprintf("failed to write to session: %v", err))
} }
@ -808,9 +770,7 @@ func (t *ExecTool) executeKill(args map[string]any) *ToolResult {
} }
if session.IsDone() { if session.IsDone() {
return ErrorResult( return ErrorResult(fmt.Sprintf("process already exited with code %d", session.GetExitCode()))
fmt.Sprintf("process already exited with code %d", session.GetExitCode()),
)
} }
if err := session.Kill(); err != nil { if err := session.Kill(); err != nil {
@ -1032,16 +992,12 @@ func (t *ExecTool) executeSendKeys(args map[string]any) *ToolResult {
} }
if session.IsDone() { if session.IsDone() {
return ErrorResult( return ErrorResult(fmt.Sprintf("process already exited with code %d", session.GetExitCode()))
fmt.Sprintf("process already exited with code %d", session.GetExitCode()),
)
} }
if err := session.Write(data); err != nil { if err := session.Write(data); err != nil {
if errors.Is(err, ErrSessionDone) { if errors.Is(err, ErrSessionDone) {
return ErrorResult( return ErrorResult(fmt.Sprintf("process already exited with code %d", session.GetExitCode()))
fmt.Sprintf("process already exited with code %d", session.GetExitCode()),
)
} }
return ErrorResult(fmt.Sprintf("failed to send keys: %v", err)) return ErrorResult(fmt.Sprintf("failed to send keys: %v", err))
} }

View file

@ -100,13 +100,8 @@ func TestShellTool_Timeout(t *testing.T) {
} }
// Should mention timeout // Should mention timeout
if !strings.Contains(result.ForLLM, "timed out") && if !strings.Contains(result.ForLLM, "timed out") && !strings.Contains(result.ForUser, "timed out") {
!strings.Contains(result.ForUser, "timed out") { t.Errorf("Expected timeout message, got ForLLM: %s, ForUser: %s", result.ForLLM, result.ForUser)
t.Errorf(
"Expected timeout message, got ForLLM: %s, ForUser: %s",
result.ForLLM,
result.ForUser,
)
} }
} }
@ -161,11 +156,7 @@ func TestShellTool_DangerousCommand(t *testing.T) {
} }
if !strings.Contains(result.ForLLM, "blocked") && !strings.Contains(result.ForUser, "blocked") { if !strings.Contains(result.ForLLM, "blocked") && !strings.Contains(result.ForUser, "blocked") {
t.Errorf( t.Errorf("Expected 'blocked' message, got ForLLM: %s, ForUser: %s", result.ForLLM, result.ForUser)
"Expected 'blocked' message, got ForLLM: %s, ForUser: %s",
result.ForLLM,
result.ForUser,
)
} }
} }
@ -186,11 +177,7 @@ func TestShellTool_DangerousCommand_KillBlocked(t *testing.T) {
t.Errorf("Expected kill command to be blocked") t.Errorf("Expected kill command to be blocked")
} }
if !strings.Contains(result.ForLLM, "blocked") && !strings.Contains(result.ForUser, "blocked") { if !strings.Contains(result.ForLLM, "blocked") && !strings.Contains(result.ForUser, "blocked") {
t.Errorf( t.Errorf("Expected blocked message, got ForLLM: %s, ForUser: %s", result.ForLLM, result.ForUser)
"Expected blocked message, got ForLLM: %s, ForUser: %s",
result.ForLLM,
result.ForUser,
)
} }
} }
@ -282,10 +269,7 @@ func TestShellTool_WorkingDir_OutsideWorkspace(t *testing.T) {
}) })
if !result.IsError { if !result.IsError {
t.Fatalf( t.Fatalf("expected working_dir outside workspace to be blocked, got output: %s", result.ForLLM)
"expected working_dir outside workspace to be blocked, got output: %s",
result.ForLLM,
)
} }
if !strings.Contains(result.ForLLM, "blocked") { if !strings.Contains(result.ForLLM, "blocked") {
t.Errorf("expected 'blocked' in error, got: %s", result.ForLLM) t.Errorf("expected 'blocked' in error, got: %s", result.ForLLM)
@ -460,10 +444,7 @@ func TestShellTool_DevNullAllowed(t *testing.T) {
} }
for _, cmd := range commands { for _, cmd := range commands {
result := tool.Execute( result := tool.Execute(context.Background(), map[string]any{"action": "run", "command": cmd})
context.Background(),
map[string]any{"action": "run", "command": cmd},
)
if result.IsError && strings.Contains(result.ForLLM, "blocked") { if result.IsError && strings.Contains(result.ForLLM, "blocked") {
t.Errorf("command should not be blocked: %s\n error: %s", cmd, result.ForLLM) t.Errorf("command should not be blocked: %s\n error: %s", cmd, result.ForLLM)
} }
@ -492,10 +473,7 @@ func TestShellTool_BlockDevices(t *testing.T) {
} }
for _, cmd := range blocked { for _, cmd := range blocked {
result := tool.Execute( result := tool.Execute(context.Background(), map[string]any{"action": "run", "command": cmd})
context.Background(),
map[string]any{"action": "run", "command": cmd},
)
if !result.IsError { if !result.IsError {
t.Errorf("expected block device write to be blocked: %s", cmd) t.Errorf("expected block device write to be blocked: %s", cmd)
} }
@ -519,16 +497,9 @@ func TestShellTool_SafePathsInWorkspaceRestriction(t *testing.T) {
} }
for _, cmd := range commands { for _, cmd := range commands {
result := tool.Execute( result := tool.Execute(context.Background(), map[string]any{"action": "run", "command": cmd})
context.Background(),
map[string]any{"action": "run", "command": cmd},
)
if result.IsError && strings.Contains(result.ForLLM, "path outside working dir") { if result.IsError && strings.Contains(result.ForLLM, "path outside working dir") {
t.Errorf( t.Errorf("safe path should not be blocked by workspace check: %s\n error: %s", cmd, result.ForLLM)
"safe path should not be blocked by workspace check: %s\n error: %s",
cmd,
result.ForLLM,
)
} }
} }
} }
@ -620,10 +591,7 @@ func TestShellTool_CustomAllowPatterns(t *testing.T) {
"command": "git push origin main", "command": "git push origin main",
}) })
if result.IsError && strings.Contains(result.ForLLM, "blocked") { if result.IsError && strings.Contains(result.ForLLM, "blocked") {
t.Errorf( t.Errorf("custom allow pattern should exempt 'git push origin main', got: %s", result.ForLLM)
"custom allow pattern should exempt 'git push origin main', got: %s",
result.ForLLM,
)
} }
// "git push upstream main" should still be blocked (does not match allow pattern). // "git push upstream main" should still be blocked (does not match allow pattern).
@ -661,11 +629,7 @@ func TestShellTool_URLsNotBlocked(t *testing.T) {
result := tool.Execute(ctx, map[string]any{"action": "run", "command": cmd}) result := tool.Execute(ctx, map[string]any{"action": "run", "command": cmd})
cancel() cancel()
if result.IsError && strings.Contains(result.ForLLM, "path outside working dir") { if result.IsError && strings.Contains(result.ForLLM, "path outside working dir") {
t.Errorf( t.Errorf("command with URL should not be blocked by workspace check: %s\n error: %s", cmd, result.ForLLM)
"command with URL should not be blocked by workspace check: %s\n error: %s",
cmd,
result.ForLLM,
)
} }
} }
} }
@ -688,10 +652,7 @@ func TestShellTool_FileURISandboxing(t *testing.T) {
} }
for _, cmd := range blockedCommands { for _, cmd := range blockedCommands {
result := tool.Execute( result := tool.Execute(context.Background(), map[string]any{"action": "run", "command": cmd})
context.Background(),
map[string]any{"action": "run", "command": cmd},
)
if !result.IsError || !strings.Contains(result.ForLLM, "path outside working dir") { if !result.IsError || !strings.Contains(result.ForLLM, "path outside working dir") {
t.Errorf("file:// URI outside workspace should be blocked: %s", cmd) t.Errorf("file:// URI outside workspace should be blocked: %s", cmd)
} }
@ -709,16 +670,9 @@ func TestShellTool_FileURISandboxing(t *testing.T) {
} }
for _, cmd := range allowedCommands { for _, cmd := range allowedCommands {
result := tool.Execute( result := tool.Execute(context.Background(), map[string]any{"action": "run", "command": cmd})
context.Background(),
map[string]any{"action": "run", "command": cmd},
)
if result.IsError && strings.Contains(result.ForLLM, "path outside working dir") { if result.IsError && strings.Contains(result.ForLLM, "path outside working dir") {
t.Errorf( t.Errorf("file:// URI inside workspace should be allowed: %s\n error: %s", cmd, result.ForLLM)
"file:// URI inside workspace should be allowed: %s\n error: %s",
cmd,
result.ForLLM,
)
} }
} }
} }
@ -742,10 +696,7 @@ func TestShellTool_URLBypassPrevented(t *testing.T) {
} }
for _, cmd := range blockedCommands { for _, cmd := range blockedCommands {
result := tool.Execute( result := tool.Execute(context.Background(), map[string]any{"action": "run", "command": cmd})
context.Background(),
map[string]any{"action": "run", "command": cmd},
)
if !result.IsError || !strings.Contains(result.ForLLM, "path outside working dir") { if !result.IsError || !strings.Contains(result.ForLLM, "path outside working dir") {
t.Errorf("bypass attempt should be blocked: %q\n got: %s", cmd, result.ForLLM) t.Errorf("bypass attempt should be blocked: %q\n got: %s", cmd, result.ForLLM)
} }
@ -1270,9 +1221,7 @@ func TestShellTool_PTY_ProcessGroupKill(t *testing.T) {
// The binary is created in /tmp/test_pgroup.c and compiled as part of test setup. // The binary is created in /tmp/test_pgroup.c and compiled as part of test setup.
testBinary := "/tmp/test_pgroup" testBinary := "/tmp/test_pgroup"
if _, err := os.Stat(testBinary); os.IsNotExist(err) { if _, err := os.Stat(testBinary); os.IsNotExist(err) {
t.Skip( t.Skip("Test binary /tmp/test_pgroup not found - run: gcc -o /tmp/test_pgroup /tmp/test_pgroup.c")
"Test binary /tmp/test_pgroup not found - run: gcc -o /tmp/test_pgroup /tmp/test_pgroup.c",
)
} }
tool, err := NewExecTool("", false) tool, err := NewExecTool("", false)
@ -1606,16 +1555,8 @@ func TestDetectPtyKeyMode(t *testing.T) {
{"rmkx only", "\x1b[?1l\x1b>", PtyKeyModeCSI}, {"rmkx only", "\x1b[?1l\x1b>", PtyKeyModeCSI},
{"both smkx first", "\x1b[?1h\x1b=...\x1b[?1l\x1b>", PtyKeyModeCSI}, {"both smkx first", "\x1b[?1h\x1b=...\x1b[?1l\x1b>", PtyKeyModeCSI},
{"both rmkx first", "\x1b[?1l\x1b>...\x1b[?1h\x1b=", PtyKeyModeSS3}, {"both rmkx first", "\x1b[?1l\x1b>...\x1b[?1h\x1b=", PtyKeyModeSS3},
{ {"multiple toggles smkx last", "\x1b[?1h\x1b=...\x1b[?1l\x1b>...\x1b[?1h\x1b=", PtyKeyModeSS3},
"multiple toggles smkx last", {"multiple toggles rmkx last", "\x1b[?1l\x1b>...\x1b[?1h\x1b=...\x1b[?1l\x1b>", PtyKeyModeCSI},
"\x1b[?1h\x1b=...\x1b[?1l\x1b>...\x1b[?1h\x1b=",
PtyKeyModeSS3,
},
{
"multiple toggles rmkx last",
"\x1b[?1l\x1b>...\x1b[?1h\x1b=...\x1b[?1l\x1b>",
PtyKeyModeCSI,
},
{"partial smkx", "\x1b[?1h", PtyKeyModeSS3}, {"partial smkx", "\x1b[?1h", PtyKeyModeSS3},
{"partial rmkx", "\x1b[?1l", PtyKeyModeCSI}, {"partial rmkx", "\x1b[?1l", PtyKeyModeCSI},
} }

View file

@ -96,11 +96,7 @@ func (t *InstallSkillTool) Execute(ctx context.Context, args map[string]any) *To
if !force { if !force {
if _, err := os.Stat(targetDir); err == nil { if _, err := os.Stat(targetDir); err == nil {
return ErrorResult( return ErrorResult(
fmt.Sprintf( fmt.Sprintf("skill %q already installed at %s. Use force=true to reinstall.", slug, targetDir),
"skill %q already installed at %s. Use force=true to reinstall.",
slug,
targetDir,
),
) )
} }
} else { } else {
@ -146,9 +142,7 @@ func (t *InstallSkillTool) Execute(ctx context.Context, args map[string]any) *To
"error": rmErr.Error(), "error": rmErr.Error(),
}) })
} }
return ErrorResult( return ErrorResult(fmt.Sprintf("skill %q is flagged as malicious and cannot be installed", slug))
fmt.Sprintf("skill %q is flagged as malicious and cannot be installed", slug),
)
} }
// Write origin metadata. // Write origin metadata.
@ -168,10 +162,7 @@ func (t *InstallSkillTool) Execute(ctx context.Context, args map[string]any) *To
// Build result with moderation warning if suspicious. // Build result with moderation warning if suspicious.
var output string var output string
if result.IsSuspicious { if result.IsSuspicious {
output = fmt.Sprintf( output = fmt.Sprintf("⚠️ Warning: skill %q is flagged as suspicious (may contain risky patterns).\n\n", slug)
"⚠️ Warning: skill %q is flagged as suspicious (may contain risky patterns).\n\n",
slug,
)
} }
output += fmt.Sprintf("Successfully installed skill %q v%s from %s registry.\nLocation: %s\n", output += fmt.Sprintf("Successfully installed skill %q v%s from %s registry.\nLocation: %s\n",
slug, result.Version, registry.Name(), targetDir) slug, result.Version, registry.Name(), targetDir)

View file

@ -17,10 +17,7 @@ type FindSkillsTool struct {
// NewFindSkillsTool creates a new FindSkillsTool. // NewFindSkillsTool creates a new FindSkillsTool.
// registryMgr is the shared registry manager (built from config in createToolRegistry). // registryMgr is the shared registry manager (built from config in createToolRegistry).
// cache is the search cache for deduplicating similar queries. // cache is the search cache for deduplicating similar queries.
func NewFindSkillsTool( func NewFindSkillsTool(registryMgr *skills.RegistryManager, cache *skills.SearchCache) *FindSkillsTool {
registryMgr *skills.RegistryManager,
cache *skills.SearchCache,
) *FindSkillsTool {
return &FindSkillsTool{ return &FindSkillsTool{
registryMgr: registryMgr, registryMgr: registryMgr,
cache: cache, cache: cache,

View file

@ -77,12 +77,10 @@ func (t *SpawnStatusTool) Execute(ctx context.Context, args map[string]any) *Too
} }
// Restrict lookup to tasks that belong to this conversation. // Restrict lookup to tasks that belong to this conversation.
if callerChannel != "" && taskCopy.OriginChannel != "" && if callerChannel != "" && taskCopy.OriginChannel != "" && taskCopy.OriginChannel != callerChannel {
taskCopy.OriginChannel != callerChannel {
return ErrorResult(fmt.Sprintf("No subagent found with task ID: %s", taskID)) return ErrorResult(fmt.Sprintf("No subagent found with task ID: %s", taskID))
} }
if callerChatID != "" && taskCopy.OriginChatID != "" && if callerChatID != "" && taskCopy.OriginChatID != "" && taskCopy.OriginChatID != callerChatID {
taskCopy.OriginChatID != callerChatID {
return ErrorResult(fmt.Sprintf("No subagent found with task ID: %s", taskID)) return ErrorResult(fmt.Sprintf("No subagent found with task ID: %s", taskID))
} }

View file

@ -195,12 +195,7 @@ func TestSpawnStatusTool_TaskID_NonString(t *testing.T) {
for _, badVal := range []any{42, 3.14, true, map[string]any{"x": 1}, []string{"a"}} { for _, badVal := range []any{42, 3.14, true, map[string]any{"x": 1}, []string{"a"}} {
result := tool.Execute(context.Background(), map[string]any{"task_id": badVal}) result := tool.Execute(context.Background(), map[string]any{"task_id": badVal})
if !result.IsError { if !result.IsError {
t.Errorf( t.Errorf("Expected error for task_id=%T(%v), got success: %s", badVal, badVal, result.ForLLM)
"Expected error for task_id=%T(%v), got success: %s",
badVal,
badVal,
result.ForLLM,
)
} }
if !strings.Contains(result.ForLLM, "task_id must be a string") { if !strings.Contains(result.ForLLM, "task_id must be a string") {
t.Errorf("Expected type-error message, got: %s", result.ForLLM) t.Errorf("Expected type-error message, got: %s", result.ForLLM)
@ -324,10 +319,7 @@ func TestSpawnStatusTool_SortByCreatedTimestamp(t *testing.T) {
t.Fatalf("Both task IDs should appear in output:\n%s", result.ForLLM) t.Fatalf("Both task IDs should appear in output:\n%s", result.ForLLM)
} }
if pos2 > pos10 { if pos2 > pos10 {
t.Errorf( t.Errorf("Expected subagent-2 (created first) to appear before subagent-10, but got:\n%s", result.ForLLM)
"Expected subagent-2 (created first) to appear before subagent-10, but got:\n%s",
result.ForLLM,
)
} }
} }

View file

@ -69,9 +69,7 @@ func (t *SPITool) Parameters() map[string]any {
func (t *SPITool) Execute(ctx context.Context, args map[string]any) *ToolResult { func (t *SPITool) Execute(ctx context.Context, args map[string]any) *ToolResult {
if runtime.GOOS != "linux" { if runtime.GOOS != "linux" {
return ErrorResult( return ErrorResult("SPI is only supported on Linux. This tool requires /dev/spidev* device files.")
"SPI is only supported on Linux. This tool requires /dev/spidev* device files.",
)
} }
action, ok := args["action"].(string) action, ok := args["action"].(string)
@ -126,9 +124,7 @@ func (t *SPITool) list() *ToolResult {
// parseSPIArgs extracts and validates common SPI parameters // parseSPIArgs extracts and validates common SPI parameters
// //
//nolint:unused // Used by spi_linux.go //nolint:unused // Used by spi_linux.go
func parseSPIArgs( func parseSPIArgs(args map[string]any) (device string, speed uint32, mode uint8, bits uint8, errMsg string) {
args map[string]any,
) (device string, speed uint32, mode uint8, bits uint8, errMsg string) {
dev, ok := args["device"].(string) dev, ok := args["device"].(string)
if !ok || dev == "" { if !ok || dev == "" {
return "", 0, 0, 0, "device is required (e.g. \"2.0\" for /dev/spidev2.0)" return "", 0, 0, 0, "device is required (e.g. \"2.0\" for /dev/spidev2.0)"

View file

@ -38,46 +38,25 @@ type spiTransfer struct {
func configureSPI(devPath string, mode uint8, bits uint8, speed uint32) (int, *ToolResult) { func configureSPI(devPath string, mode uint8, bits uint8, speed uint32) (int, *ToolResult) {
fd, err := syscall.Open(devPath, syscall.O_RDWR, 0) fd, err := syscall.Open(devPath, syscall.O_RDWR, 0)
if err != nil { if err != nil {
return -1, ErrorResult( return -1, ErrorResult(fmt.Sprintf("failed to open %s: %v (check permissions and spidev module)", devPath, err))
fmt.Sprintf(
"failed to open %s: %v (check permissions and spidev module)",
devPath,
err,
),
)
} }
// Set SPI mode // Set SPI mode
_, _, errno := syscall.Syscall( _, _, errno := syscall.Syscall(syscall.SYS_IOCTL, uintptr(fd), spiIocWrMode, uintptr(unsafe.Pointer(&mode)))
syscall.SYS_IOCTL,
uintptr(fd),
spiIocWrMode,
uintptr(unsafe.Pointer(&mode)),
)
if errno != 0 { if errno != 0 {
syscall.Close(fd) syscall.Close(fd)
return -1, ErrorResult(fmt.Sprintf("failed to set SPI mode %d: %v", mode, errno)) return -1, ErrorResult(fmt.Sprintf("failed to set SPI mode %d: %v", mode, errno))
} }
// Set bits per word // Set bits per word
_, _, errno = syscall.Syscall( _, _, errno = syscall.Syscall(syscall.SYS_IOCTL, uintptr(fd), spiIocWrBitsPerWord, uintptr(unsafe.Pointer(&bits)))
syscall.SYS_IOCTL,
uintptr(fd),
spiIocWrBitsPerWord,
uintptr(unsafe.Pointer(&bits)),
)
if errno != 0 { if errno != 0 {
syscall.Close(fd) syscall.Close(fd)
return -1, ErrorResult(fmt.Sprintf("failed to set bits per word %d: %v", bits, errno)) return -1, ErrorResult(fmt.Sprintf("failed to set bits per word %d: %v", bits, errno))
} }
// Set max speed // Set max speed
_, _, errno = syscall.Syscall( _, _, errno = syscall.Syscall(syscall.SYS_IOCTL, uintptr(fd), spiIocWrMaxSpeedHz, uintptr(unsafe.Pointer(&speed)))
syscall.SYS_IOCTL,
uintptr(fd),
spiIocWrMaxSpeedHz,
uintptr(unsafe.Pointer(&speed)),
)
if errno != 0 { if errno != 0 {
syscall.Close(fd) syscall.Close(fd)
return -1, ErrorResult(fmt.Sprintf("failed to set SPI speed %d Hz: %v", speed, errno)) return -1, ErrorResult(fmt.Sprintf("failed to set SPI speed %d Hz: %v", speed, errno))
@ -138,12 +117,7 @@ func (t *SPITool) transfer(args map[string]any) *ToolResult {
bitsPerWord: bits, bitsPerWord: bits,
} }
_, _, errno := syscall.Syscall( _, _, errno := syscall.Syscall(syscall.SYS_IOCTL, uintptr(fd), spiIocMessage1, uintptr(unsafe.Pointer(&xfer)))
syscall.SYS_IOCTL,
uintptr(fd),
spiIocMessage1,
uintptr(unsafe.Pointer(&xfer)),
)
runtime.KeepAlive(txBuf) runtime.KeepAlive(txBuf)
runtime.KeepAlive(rxBuf) runtime.KeepAlive(rxBuf)
if errno != 0 { if errno != 0 {
@ -200,12 +174,7 @@ func (t *SPITool) readDevice(args map[string]any) *ToolResult {
bitsPerWord: bits, bitsPerWord: bits,
} }
_, _, errno := syscall.Syscall( _, _, errno := syscall.Syscall(syscall.SYS_IOCTL, uintptr(fd), spiIocMessage1, uintptr(unsafe.Pointer(&xfer)))
syscall.SYS_IOCTL,
uintptr(fd),
spiIocMessage1,
uintptr(unsafe.Pointer(&xfer)),
)
runtime.KeepAlive(txBuf) runtime.KeepAlive(txBuf)
runtime.KeepAlive(rxBuf) runtime.KeepAlive(rxBuf)
if errno != 0 { if errno != 0 {

View file

@ -316,11 +316,7 @@ func TestSubagentTool_ForUserTruncation(t *testing.T) {
// ForUser should be truncated to 500 chars + "..." // ForUser should be truncated to 500 chars + "..."
maxUserLen := 500 maxUserLen := 500
if len(result.ForUser) > maxUserLen+3 { // +3 for "..." if len(result.ForUser) > maxUserLen+3 { // +3 for "..."
t.Errorf( t.Errorf("ForUser should be truncated to ~%d chars, got: %d", maxUserLen, len(result.ForUser))
"ForUser should be truncated to ~%d chars, got: %d",
maxUserLen,
len(result.ForUser),
)
} }
// ForLLM should have full content // ForLLM should have full content

View file

@ -64,13 +64,7 @@ func RunToolLoop(
llmOpts = map[string]any{} llmOpts = map[string]any{}
} }
// 3. Call LLM // 3. Call LLM
response, err := config.Provider.Chat( response, err := config.Provider.Chat(ctx, messages, providerToolDefs, config.Model, llmOpts)
ctx,
messages,
providerToolDefs,
config.Model,
llmOpts,
)
if err != nil { if err != nil {
logger.ErrorCF("toolloop", "LLM call failed", logger.ErrorCF("toolloop", "LLM call failed",
map[string]any{ map[string]any{
@ -154,14 +148,7 @@ func RunToolLoop(
var toolResult *ToolResult var toolResult *ToolResult
if config.Tools != nil { if config.Tools != nil {
toolResult = config.Tools.ExecuteWithContext( toolResult = config.Tools.ExecuteWithContext(ctx, tc.Name, tc.Arguments, channel, chatID, nil)
ctx,
tc.Name,
tc.Arguments,
channel,
chatID,
nil,
)
} else { } else {
toolResult = ErrorResult("No tools available") toolResult = ErrorResult("No tools available")
} }

View file

@ -151,10 +151,7 @@ func TestValidateToolArgs(t *testing.T) {
schema: map[string]any{ schema: map[string]any{
"type": "object", "type": "object",
"properties": map[string]any{ "properties": map[string]any{
"color": map[string]any{ "color": map[string]any{"type": "string", "enum": []any{"red", "green", "blue"}},
"type": "string",
"enum": []any{"red", "green", "blue"},
},
}, },
}, },
args: map[string]any{"color": "red"}, args: map[string]any{"color": "red"},
@ -164,10 +161,7 @@ func TestValidateToolArgs(t *testing.T) {
schema: map[string]any{ schema: map[string]any{
"type": "object", "type": "object",
"properties": map[string]any{ "properties": map[string]any{
"color": map[string]any{ "color": map[string]any{"type": "string", "enum": []any{"red", "green", "blue"}},
"type": "string",
"enum": []any{"red", "green", "blue"},
},
}, },
}, },
args: map[string]any{"color": "yellow"}, args: map[string]any{"color": "yellow"},
@ -178,10 +172,7 @@ func TestValidateToolArgs(t *testing.T) {
schema: map[string]any{ schema: map[string]any{
"type": "object", "type": "object",
"properties": map[string]any{ "properties": map[string]any{
"color": map[string]any{ "color": map[string]any{"type": "string", "enum": []string{"red", "green", "blue"}},
"type": "string",
"enum": []string{"red", "green", "blue"},
},
}, },
}, },
args: map[string]any{"color": "green"}, args: map[string]any{"color": "green"},
@ -191,10 +182,7 @@ func TestValidateToolArgs(t *testing.T) {
schema: map[string]any{ schema: map[string]any{
"type": "object", "type": "object",
"properties": map[string]any{ "properties": map[string]any{
"color": map[string]any{ "color": map[string]any{"type": "string", "enum": []string{"red", "green", "blue"}},
"type": "string",
"enum": []string{"red", "green", "blue"},
},
}, },
}, },
args: map[string]any{"color": "yellow"}, args: map[string]any{"color": "yellow"},
@ -354,11 +342,7 @@ func TestValidateToolArgs_RegistryIntegration(t *testing.T) {
} }
// Extra property — should fail with validation error // Extra property — should fail with validation error
result = r.Execute( result = r.Execute(context.Background(), "read_file", map[string]any{"path": "/x", "__inject": true})
context.Background(),
"read_file",
map[string]any{"path": "/x", "__inject": true},
)
if !result.IsError { if !result.IsError {
t.Error("expected validation error for extra property") t.Error("expected validation error for extra property")
} }

View file

@ -54,8 +54,7 @@ func TestWebTool_WebFetch_Success(t *testing.T) {
} }
// ForUser should contain summary // ForUser should contain summary
if !strings.Contains(result.ForUser, "bytes") && if !strings.Contains(result.ForUser, "bytes") && !strings.Contains(result.ForUser, "extractor") {
!strings.Contains(result.ForUser, "extractor") {
t.Errorf("Expected ForUser to contain summary, got: %s", result.ForUser) t.Errorf("Expected ForUser to contain summary, got: %s", result.ForUser)
} }
} }
@ -76,11 +75,7 @@ func TestWebTool_WebFetch_JSON(t *testing.T) {
tool, err := NewWebFetchTool(50000, format, testFetchLimit) tool, err := NewWebFetchTool(50000, format, testFetchLimit)
if err != nil { if err != nil {
logger.ErrorCF( logger.ErrorCF("agent", "Failed to create web fetch tool", map[string]any{"error": err.Error()})
"agent",
"Failed to create web fetch tool",
map[string]any{"error": err.Error()},
)
} }
ctx := context.Background() ctx := context.Background()
@ -105,11 +100,7 @@ func TestWebTool_WebFetch_JSON(t *testing.T) {
func TestWebTool_WebFetch_InvalidURL(t *testing.T) { func TestWebTool_WebFetch_InvalidURL(t *testing.T) {
tool, err := NewWebFetchTool(50000, format, testFetchLimit) tool, err := NewWebFetchTool(50000, format, testFetchLimit)
if err != nil { if err != nil {
logger.ErrorCF( logger.ErrorCF("agent", "Failed to create web fetch tool", map[string]any{"error": err.Error()})
"agent",
"Failed to create web fetch tool",
map[string]any{"error": err.Error()},
)
} }
ctx := context.Background() ctx := context.Background()
@ -134,11 +125,7 @@ func TestWebTool_WebFetch_InvalidURL(t *testing.T) {
func TestWebTool_WebFetch_UnsupportedScheme(t *testing.T) { func TestWebTool_WebFetch_UnsupportedScheme(t *testing.T) {
tool, err := NewWebFetchTool(50000, format, testFetchLimit) tool, err := NewWebFetchTool(50000, format, testFetchLimit)
if err != nil { if err != nil {
logger.ErrorCF( logger.ErrorCF("agent", "Failed to create web fetch tool", map[string]any{"error": err.Error()})
"agent",
"Failed to create web fetch tool",
map[string]any{"error": err.Error()},
)
} }
ctx := context.Background() ctx := context.Background()
@ -154,8 +141,7 @@ func TestWebTool_WebFetch_UnsupportedScheme(t *testing.T) {
} }
// Should mention only http/https allowed // Should mention only http/https allowed
if !strings.Contains(result.ForLLM, "http/https") && if !strings.Contains(result.ForLLM, "http/https") && !strings.Contains(result.ForUser, "http/https") {
!strings.Contains(result.ForUser, "http/https") {
t.Errorf("Expected scheme error message, got ForLLM: %s", result.ForLLM) t.Errorf("Expected scheme error message, got ForLLM: %s", result.ForLLM)
} }
} }
@ -164,11 +150,7 @@ func TestWebTool_WebFetch_UnsupportedScheme(t *testing.T) {
func TestWebTool_WebFetch_MissingURL(t *testing.T) { func TestWebTool_WebFetch_MissingURL(t *testing.T) {
tool, err := NewWebFetchTool(50000, format, testFetchLimit) tool, err := NewWebFetchTool(50000, format, testFetchLimit)
if err != nil { if err != nil {
logger.ErrorCF( logger.ErrorCF("agent", "Failed to create web fetch tool", map[string]any{"error": err.Error()})
"agent",
"Failed to create web fetch tool",
map[string]any{"error": err.Error()},
)
} }
ctx := context.Background() ctx := context.Background()
@ -182,8 +164,7 @@ func TestWebTool_WebFetch_MissingURL(t *testing.T) {
} }
// Should mention URL is required // Should mention URL is required
if !strings.Contains(result.ForLLM, "url is required") && if !strings.Contains(result.ForLLM, "url is required") && !strings.Contains(result.ForUser, "url is required") {
!strings.Contains(result.ForUser, "url is required") {
t.Errorf("Expected 'url is required' message, got ForLLM: %s", result.ForLLM) t.Errorf("Expected 'url is required' message, got ForLLM: %s", result.ForLLM)
} }
} }
@ -203,11 +184,7 @@ func TestWebTool_WebFetch_Truncation(t *testing.T) {
tool, err := NewWebFetchTool(1000, format, testFetchLimit) // Limit to 1000 chars tool, err := NewWebFetchTool(1000, format, testFetchLimit) // Limit to 1000 chars
if err != nil { if err != nil {
logger.ErrorCF( logger.ErrorCF("agent", "Failed to create web fetch tool", map[string]any{"error": err.Error()})
"agent",
"Failed to create web fetch tool",
map[string]any{"error": err.Error()},
)
} }
ctx := context.Background() ctx := context.Background()
@ -239,10 +216,7 @@ func TestWebTool_WebFetch_Truncation(t *testing.T) {
// Text should end with the truncation notice // Text should end with the truncation notice
if text, ok := resultMap["text"].(string); ok { if text, ok := resultMap["text"].(string); ok {
if !strings.HasSuffix(text, "[Content truncated due to size limit]") { if !strings.HasSuffix(text, "[Content truncated due to size limit]") {
t.Errorf( t.Errorf("Expected text to end with truncation notice, got: %q", text[max(0, len(text)-60):])
"Expected text to end with truncation notice, got: %q",
text[max(0, len(text)-60):],
)
} }
} }
} }
@ -289,13 +263,11 @@ func TestWebTool_WebFetch_TruncationNotice(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
server := httptest.NewServer( server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", tt.contentType)
w.Header().Set("Content-Type", tt.contentType) w.WriteHeader(http.StatusOK)
w.WriteHeader(http.StatusOK) w.Write([]byte(tt.body))
w.Write([]byte(tt.body)) }))
}),
)
defer server.Close() defer server.Close()
tool, err := NewWebFetchTool(maxChars, tt.format, testFetchLimit) tool, err := NewWebFetchTool(maxChars, tt.format, testFetchLimit)
@ -319,11 +291,7 @@ func TestWebTool_WebFetch_TruncationNotice(t *testing.T) {
} }
if !strings.HasSuffix(text, truncationNotice) { if !strings.HasSuffix(text, truncationNotice) {
t.Errorf( t.Errorf("expected text to end with %q, got suffix: %q", truncationNotice, text[max(0, len(text)-60):])
"expected text to end with %q, got suffix: %q",
truncationNotice,
text[max(0, len(text)-60):],
)
} }
if truncated, ok := resultMap["truncated"].(bool); !ok || !truncated { if truncated, ok := resultMap["truncated"].(bool); !ok || !truncated {
@ -392,11 +360,7 @@ func TestWebFetchTool_PayloadTooLarge(t *testing.T) {
// Initialize the tool // Initialize the tool
tool, err := NewWebFetchTool(50000, format, testFetchLimit) tool, err := NewWebFetchTool(50000, format, testFetchLimit)
if err != nil { if err != nil {
logger.ErrorCF( logger.ErrorCF("agent", "Failed to create web fetch tool", map[string]any{"error": err.Error()})
"agent",
"Failed to create web fetch tool",
map[string]any{"error": err.Error()},
)
} }
// Prepare the arguments pointing to the URL of our local mock server // Prepare the arguments pointing to the URL of our local mock server
@ -416,8 +380,7 @@ func TestWebFetchTool_PayloadTooLarge(t *testing.T) {
// Search for the exact error string we set earlier in the Execute method // Search for the exact error string we set earlier in the Execute method
expectedErrorMsg := fmt.Sprintf("size exceeded %d bytes limit", testFetchLimit) expectedErrorMsg := fmt.Sprintf("size exceeded %d bytes limit", testFetchLimit)
if !strings.Contains(result.ForLLM, expectedErrorMsg) && if !strings.Contains(result.ForLLM, expectedErrorMsg) && !strings.Contains(result.ForUser, expectedErrorMsg) {
!strings.Contains(result.ForUser, expectedErrorMsg) {
t.Errorf("test failed: expected error %q, but got: %+v", expectedErrorMsg, result) t.Errorf("test failed: expected error %q, but got: %+v", expectedErrorMsg, result)
} }
} }
@ -570,11 +533,7 @@ func TestWebTool_WebFetch_HTMLExtraction(t *testing.T) {
tool, err := NewWebFetchTool(50000, format, testFetchLimit) tool, err := NewWebFetchTool(50000, format, testFetchLimit)
if err != nil { if err != nil {
logger.ErrorCF( logger.ErrorCF("agent", "Failed to create web fetch tool", map[string]any{"error": err.Error()})
"agent",
"Failed to create web fetch tool",
map[string]any{"error": err.Error()},
)
} }
ctx := context.Background() ctx := context.Background()
@ -759,13 +718,7 @@ func TestWebTool_WebFetch_PrivateHostAllowedByCIDRWhitelist(t *testing.T) {
defer server.Close() defer server.Close()
host, _ := serverHostAndPort(t, server.URL) host, _ := serverHostAndPort(t, server.URL)
tool, err := NewWebFetchToolWithConfig( tool, err := NewWebFetchToolWithConfig(50000, "", format, testFetchLimit, []string{singleHostCIDR(t, host)})
50000,
"",
format,
testFetchLimit,
[]string{singleHostCIDR(t, host)},
)
if err != nil { if err != nil {
t.Fatalf("Failed to create web fetch tool: %v", err) t.Fatalf("Failed to create web fetch tool: %v", err)
} }
@ -800,10 +753,7 @@ func TestWebTool_WebFetch_PrivateHostAllowedForTests(t *testing.T) {
}) })
if result.IsError { if result.IsError {
t.Errorf( t.Errorf("expected success when private host access is allowed in tests, got %q", result.ForLLM)
"expected success when private host access is allowed in tests, got %q",
result.ForLLM,
)
} }
} }
@ -1023,11 +973,7 @@ func TestIsPrivateOrRestrictedIP_Table(t *testing.T) {
func TestWebTool_WebFetch_MissingDomain(t *testing.T) { func TestWebTool_WebFetch_MissingDomain(t *testing.T) {
tool, err := NewWebFetchTool(50000, format, testFetchLimit) tool, err := NewWebFetchTool(50000, format, testFetchLimit)
if err != nil { if err != nil {
logger.ErrorCF( logger.ErrorCF("agent", "Failed to create web fetch tool", map[string]any{"error": err.Error()})
"agent",
"Failed to create web fetch tool",
map[string]any{"error": err.Error()},
)
} }
ctx := context.Background() ctx := context.Background()
@ -1049,19 +995,9 @@ func TestWebTool_WebFetch_MissingDomain(t *testing.T) {
} }
func TestNewWebFetchToolWithProxy(t *testing.T) { func TestNewWebFetchToolWithProxy(t *testing.T) {
tool, err := NewWebFetchToolWithProxy( tool, err := NewWebFetchToolWithProxy(1024, "http://127.0.0.1:7890", format, testFetchLimit, nil)
1024,
"http://127.0.0.1:7890",
format,
testFetchLimit,
nil,
)
if err != nil { if err != nil {
logger.ErrorCF( logger.ErrorCF("agent", "Failed to create web fetch tool", map[string]any{"error": err.Error()})
"agent",
"Failed to create web fetch tool",
map[string]any{"error": err.Error()},
)
} else if tool.maxChars != 1024 { } else if tool.maxChars != 1024 {
t.Fatalf("maxChars = %d, want %d", tool.maxChars, 1024) t.Fatalf("maxChars = %d, want %d", tool.maxChars, 1024)
} }
@ -1072,11 +1008,7 @@ func TestNewWebFetchToolWithProxy(t *testing.T) {
tool, err = NewWebFetchToolWithProxy(0, "http://127.0.0.1:7890", format, testFetchLimit, nil) tool, err = NewWebFetchToolWithProxy(0, "http://127.0.0.1:7890", format, testFetchLimit, nil)
if err != nil { if err != nil {
logger.ErrorCF( logger.ErrorCF("agent", "Failed to create web fetch tool", map[string]any{"error": err.Error()})
"agent",
"Failed to create web fetch tool",
map[string]any{"error": err.Error()},
)
} }
if tool.maxChars != 50000 { if tool.maxChars != 50000 {
@ -1085,13 +1017,7 @@ func TestNewWebFetchToolWithProxy(t *testing.T) {
} }
func TestNewWebFetchToolWithConfig_InvalidPrivateHostWhitelist(t *testing.T) { func TestNewWebFetchToolWithConfig_InvalidPrivateHostWhitelist(t *testing.T) {
_, err := NewWebFetchToolWithConfig( _, err := NewWebFetchToolWithConfig(1024, "", format, testFetchLimit, []string{"not-an-ip-or-cidr"})
1024,
"",
format,
testFetchLimit,
[]string{"not-an-ip-or-cidr"},
)
if err == nil { if err == nil {
t.Fatal("expected invalid whitelist entry to fail") t.Fatal("expected invalid whitelist entry to fail")
} }
@ -1247,11 +1173,7 @@ func TestWebTool_TavilySearch_RangeMapping(t *testing.T) {
w.WriteHeader(http.StatusOK) w.WriteHeader(http.StatusOK)
json.NewEncoder(w).Encode(map[string]any{ json.NewEncoder(w).Encode(map[string]any{
"results": []map[string]any{ "results": []map[string]any{
{ {"title": "Recent result", "url": "https://example.com/recent", "content": "snippet"},
"title": "Recent result",
"url": "https://example.com/recent",
"content": "snippet",
},
}, },
}) })
})) }))
@ -1381,10 +1303,7 @@ func TestWebFetchTool_CloudflareChallenge_RetryFailsToo(t *testing.T) {
// Should not be an error — the retry response is used as-is (403 is a valid HTTP response) // Should not be an error — the retry response is used as-is (403 is a valid HTTP response)
if result.IsError { if result.IsError {
t.Fatalf( t.Fatalf("expected non-error result even when retry is also blocked, got: %s", result.ForLLM)
"expected non-error result even when retry is also blocked, got: %s",
result.ForLLM,
)
} }
// Status in the JSON result should reflect the 403 // Status in the JSON result should reflect the 403
if !strings.Contains(result.ForLLM, "403") { if !strings.Contains(result.ForLLM, "403") {
@ -1549,10 +1468,7 @@ func TestWebTool_GLMSearch_Success(t *testing.T) {
t.Errorf("Expected Content-Type application/json, got %s", r.Header.Get("Content-Type")) t.Errorf("Expected Content-Type application/json, got %s", r.Header.Get("Content-Type"))
} }
if r.Header.Get("Authorization") != "Bearer test-glm-key" { if r.Header.Get("Authorization") != "Bearer test-glm-key" {
t.Errorf( t.Errorf("Expected Authorization Bearer test-glm-key, got %s", r.Header.Get("Authorization"))
"Expected Authorization Bearer test-glm-key, got %s",
r.Header.Get("Authorization"),
)
} }
var payload map[string]any var payload map[string]any
@ -1618,21 +1534,14 @@ func TestWebTool_GLMSearch_RangeMapping(t *testing.T) {
t.Fatalf("failed to decode payload: %v", err) t.Fatalf("failed to decode payload: %v", err)
} }
if payload["search_recency_filter"] != "oneMonth" { if payload["search_recency_filter"] != "oneMonth" {
t.Fatalf( t.Fatalf("expected search_recency_filter=oneMonth, got %v", payload["search_recency_filter"])
"expected search_recency_filter=oneMonth, got %v",
payload["search_recency_filter"],
)
} }
w.Header().Set("Content-Type", "application/json") w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK) w.WriteHeader(http.StatusOK)
json.NewEncoder(w).Encode(map[string]any{ json.NewEncoder(w).Encode(map[string]any{
"search_result": []map[string]any{ "search_result": []map[string]any{
{ {"title": "Recent GLM Result", "content": "snippet", "link": "https://example.com/glm-range"},
"title": "Recent GLM Result",
"content": "snippet",
"link": "https://example.com/glm-range",
},
}, },
}) })
})) }))
@ -1664,21 +1573,14 @@ func TestWebTool_BaiduSearch_RangeMapping(t *testing.T) {
t.Fatalf("failed to decode payload: %v", err) t.Fatalf("failed to decode payload: %v", err)
} }
if payload["search_recency_filter"] != "week" { if payload["search_recency_filter"] != "week" {
t.Fatalf( t.Fatalf("expected search_recency_filter=week for day fallback, got %v", payload["search_recency_filter"])
"expected search_recency_filter=week for day fallback, got %v",
payload["search_recency_filter"],
)
} }
w.Header().Set("Content-Type", "application/json") w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK) w.WriteHeader(http.StatusOK)
json.NewEncoder(w).Encode(map[string]any{ json.NewEncoder(w).Encode(map[string]any{
"references": []map[string]any{ "references": []map[string]any{
{ {"title": "Recent Baidu Result", "url": "https://example.com/baidu", "content": "snippet"},
"title": "Recent Baidu Result",
"url": "https://example.com/baidu",
"content": "snippet",
},
}, },
}) })
})) }))