fix(hooks): deep-clone typed payloads and add regressions

This commit is contained in:
xj 2026-02-25 22:39:19 -08:00
parent d343d5e5f0
commit 70f812d358
3 changed files with 184 additions and 7 deletions

View file

@ -3,7 +3,9 @@ package agent
import ( import (
"context" "context"
"os" "os"
"strings"
"testing" "testing"
"time"
"github.com/sipeed/picoclaw/pkg/bus" "github.com/sipeed/picoclaw/pkg/bus"
"github.com/sipeed/picoclaw/pkg/config" "github.com/sipeed/picoclaw/pkg/config"
@ -189,3 +191,88 @@ func TestBeforeToolCallHooksCannotLeaveToolArgsNil(t *testing.T) {
t.Fatal("expected tool args to be reinitialized to non-nil map") t.Fatal("expected tool args to be reinitialized to non-nil map")
} }
} }
func TestSetHooksNilRestoresDirectMessageCallback(t *testing.T) {
tmpDir, err := os.MkdirTemp("", "agent-plugin-test-*")
if err != nil {
t.Fatalf("Failed to create temp dir: %v", err)
}
defer os.RemoveAll(tmpDir)
cfg := &config.Config{
Agents: config.AgentsConfig{
Defaults: config.AgentDefaults{
Workspace: tmpDir,
Model: "test-model",
MaxTokens: 4096,
MaxToolIterations: 10,
},
},
}
msgBus := bus.NewMessageBus()
al := NewAgentLoop(cfg, msgBus, &mockProvider{})
agent := al.registry.GetDefaultAgent()
if agent == nil {
t.Fatal("expected default agent")
}
tool, ok := agent.Tools.Get("message")
if !ok {
t.Fatal("expected message tool")
}
mt, ok := tool.(*tools.MessageTool)
if !ok {
t.Fatal("expected message tool type")
}
reg := hooks.NewHookRegistry()
reg.OnMessageSending("block-all", 0, func(_ context.Context, e *hooks.MessageSendingEvent) error {
e.Cancel = true
e.CancelReason = "blocked-by-hook"
return nil
})
if err := al.SetHooks(reg); err != nil {
t.Fatalf("SetHooks(reg): %v", err)
}
blocked := mt.Execute(context.Background(), map[string]any{
"content": "first",
"channel": "cli",
"chat_id": "direct",
})
if !blocked.IsError {
t.Fatal("expected message tool call to fail while hooks are active")
}
if blocked.Err == nil || !strings.Contains(blocked.Err.Error(), "blocked-by-hook") {
t.Fatalf("expected hook cancel reason in error, got %#v", blocked.Err)
}
ctxNoMsg, cancelNoMsg := context.WithTimeout(context.Background(), 20*time.Millisecond)
defer cancelNoMsg()
if _, got := msgBus.SubscribeOutbound(ctxNoMsg); got {
t.Fatal("did not expect outbound message while hook cancellation is active")
}
if err := al.SetHooks(nil); err != nil {
t.Fatalf("SetHooks(nil): %v", err)
}
delivered := mt.Execute(context.Background(), map[string]any{
"content": "second",
"channel": "cli",
"chat_id": "direct",
})
if delivered.IsError {
t.Fatalf("expected message tool to succeed after SetHooks(nil), got %#v", delivered)
}
ctxMsg, cancelMsg := context.WithTimeout(context.Background(), time.Second)
defer cancelMsg()
msg, got := msgBus.SubscribeOutbound(ctxMsg)
if !got {
t.Fatal("expected outbound message after SetHooks(nil)")
}
if msg.Content != "second" || msg.Channel != "cli" || msg.ChatID != "direct" {
t.Fatalf("unexpected outbound message: %#v", msg)
}
}

View file

@ -9,6 +9,7 @@ package hooks
import ( import (
"context" "context"
"fmt" "fmt"
"reflect"
"sync" "sync"
"github.com/sipeed/picoclaw/pkg/logger" "github.com/sipeed/picoclaw/pkg/logger"
@ -148,13 +149,59 @@ func cloneMapStringAny(src map[string]any) map[string]any {
} }
func cloneAny(v any) any { func cloneAny(v any) any {
switch tv := v.(type) { if v == nil {
case map[string]any: return nil
return cloneMapStringAny(tv) }
case []any: cloned := cloneReflectValue(reflect.ValueOf(v))
out := make([]any, len(tv)) if !cloned.IsValid() {
for i := range tv { return nil
out[i] = cloneAny(tv[i]) }
return cloned.Interface()
}
func cloneReflectValue(v reflect.Value) reflect.Value {
if !v.IsValid() {
return v
}
switch v.Kind() {
case reflect.Pointer:
if v.IsNil() {
return reflect.Zero(v.Type())
}
out := reflect.New(v.Type().Elem())
out.Elem().Set(cloneReflectValue(v.Elem()))
return out
case reflect.Interface:
if v.IsNil() {
return reflect.Zero(v.Type())
}
out := reflect.New(v.Type()).Elem()
out.Set(cloneReflectValue(v.Elem()))
return out
case reflect.Map:
if v.IsNil() {
return reflect.Zero(v.Type())
}
out := reflect.MakeMapWithSize(v.Type(), v.Len())
iter := v.MapRange()
for iter.Next() {
out.SetMapIndex(iter.Key(), cloneReflectValue(iter.Value()))
}
return out
case reflect.Slice:
if v.IsNil() {
return reflect.Zero(v.Type())
}
out := reflect.MakeSlice(v.Type(), v.Len(), v.Len())
for i := range v.Len() {
out.Index(i).Set(cloneReflectValue(v.Index(i)))
}
return out
case reflect.Array:
out := reflect.New(v.Type()).Elem()
for i := range v.Len() {
out.Index(i).Set(cloneReflectValue(v.Index(i)))
} }
return out return out
default: default:

View file

@ -14,6 +14,7 @@ import (
"testing" "testing"
"time" "time"
"github.com/sipeed/picoclaw/pkg/providers"
"github.com/sipeed/picoclaw/pkg/tools" "github.com/sipeed/picoclaw/pkg/tools"
) )
@ -183,6 +184,48 @@ func TestVoidHooksReceiveIsolatedAfterToolCallEvents(t *testing.T) {
} }
} }
func TestVoidHooksReceiveIsolatedLLMInputToolSchema(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
r.OnLLMInput("mutator", 0, func(_ context.Context, e *LLMInputEvent) error {
required, ok := e.Tools[0].Function.Parameters["required"].([]string)
if !ok {
t.Fatal("required should be []string")
}
required[0] = "mutated"
e.Tools[0].Function.Parameters["required"] = append(required, "extra")
return nil
})
event := &LLMInputEvent{
AgentID: "a1",
Model: "m1",
Tools: []providers.ToolDefinition{
{
Type: "function",
Function: providers.ToolFunctionDefinition{
Name: "message",
Parameters: map[string]any{
"type": "object",
"required": []string{"content"},
},
},
},
},
}
r.TriggerLLMInput(ctx, event)
required, ok := event.Tools[0].Function.Parameters["required"].([]string)
if !ok {
t.Fatal("required should remain []string")
}
if len(required) != 1 || required[0] != "content" {
t.Fatalf("expected required to remain unchanged, got %#v", required)
}
}
func TestModifyingHookPriority(t *testing.T) { func TestModifyingHookPriority(t *testing.T) {
r := NewHookRegistry() r := NewHookRegistry()
ctx := context.Background() ctx := context.Background()