diff --git a/pkg/hooks/hooks.go b/pkg/hooks/hooks.go index f1b9959fa..2b9e22133 100644 --- a/pkg/hooks/hooks.go +++ b/pkg/hooks/hooks.go @@ -12,6 +12,7 @@ import ( "sync" "github.com/sipeed/picoclaw/pkg/logger" + "github.com/sipeed/picoclaw/pkg/providers" ) // HookHandler is the callback signature for all hooks. @@ -124,8 +125,146 @@ func (r *HookRegistry) OnSessionEnd(name string, priority int, handler HookHandl // Trigger methods — void hooks +func cloneMapStringString(src map[string]string) map[string]string { + if src == nil { + return nil + } + dst := make(map[string]string, len(src)) + for k, v := range src { + dst[k] = v + } + return dst +} + +func cloneMapStringAny(src map[string]any) map[string]any { + if src == nil { + return nil + } + dst := make(map[string]any, len(src)) + for k, v := range src { + dst[k] = cloneAny(v) + } + return dst +} + +func cloneAny(v any) any { + switch tv := v.(type) { + case map[string]any: + return cloneMapStringAny(tv) + case []any: + out := make([]any, len(tv)) + for i := range tv { + out[i] = cloneAny(tv[i]) + } + return out + default: + return v + } +} + +func cloneToolCall(tc providers.ToolCall) providers.ToolCall { + out := tc + out.Arguments = cloneMapStringAny(tc.Arguments) + if tc.Function != nil { + f := *tc.Function + out.Function = &f + } + if tc.ExtraContent != nil { + ec := *tc.ExtraContent + if tc.ExtraContent.Google != nil { + g := *tc.ExtraContent.Google + ec.Google = &g + } + out.ExtraContent = &ec + } + return out +} + +func cloneMessage(msg providers.Message) providers.Message { + out := msg + if msg.ToolCalls != nil { + out.ToolCalls = make([]providers.ToolCall, len(msg.ToolCalls)) + for i := range msg.ToolCalls { + out.ToolCalls[i] = cloneToolCall(msg.ToolCalls[i]) + } + } + if msg.SystemParts != nil { + out.SystemParts = make([]providers.ContentBlock, len(msg.SystemParts)) + for i := range msg.SystemParts { + part := msg.SystemParts[i] + if part.CacheControl != nil { + cc := *part.CacheControl + part.CacheControl = &cc + } + out.SystemParts[i] = part + } + } + return out +} + +func cloneToolDefinition(td providers.ToolDefinition) providers.ToolDefinition { + out := td + out.Function = td.Function + out.Function.Parameters = cloneMapStringAny(td.Function.Parameters) + return out +} + +func cloneVoidEvent[T any](event *T) *T { + if event == nil { + return nil + } + + switch e := any(event).(type) { + case *MessageReceivedEvent: + c := *e + if e.Media != nil { + c.Media = append([]string(nil), e.Media...) + } + c.Metadata = cloneMapStringString(e.Metadata) + return any(&c).(*T) + case *AfterToolCallEvent: + c := *e + c.Args = cloneMapStringAny(e.Args) + if e.Result != nil { + r := *e.Result + c.Result = &r + } + return any(&c).(*T) + case *LLMInputEvent: + c := *e + if e.Messages != nil { + c.Messages = make([]providers.Message, len(e.Messages)) + for i := range e.Messages { + c.Messages[i] = cloneMessage(e.Messages[i]) + } + } + if e.Tools != nil { + c.Tools = make([]providers.ToolDefinition, len(e.Tools)) + for i := range e.Tools { + c.Tools[i] = cloneToolDefinition(e.Tools[i]) + } + } + return any(&c).(*T) + case *LLMOutputEvent: + c := *e + if e.ToolCalls != nil { + c.ToolCalls = make([]providers.ToolCall, len(e.ToolCalls)) + for i := range e.ToolCalls { + c.ToolCalls[i] = cloneToolCall(e.ToolCalls[i]) + } + } + return any(&c).(*T) + case *SessionEvent: + c := *e + return any(&c).(*T) + default: + c := *event + return &c + } +} + // triggerVoid runs all handlers concurrently and waits for completion. -// Handlers MUST NOT mutate the event — it is shared across goroutines. +// Each handler receives a cloned event to avoid shared-state mutation races. // Errors are logged but do not propagate to the caller. func triggerVoid[T any](ctx context.Context, hooks []HookRegistration[T], event *T, hookName string) { if len(hooks) == 0 { @@ -136,6 +275,7 @@ func triggerVoid[T any](ctx context.Context, hooks []HookRegistration[T], event wg.Add(1) go func(reg HookRegistration[T]) { defer wg.Done() + eventCopy := cloneVoidEvent(event) defer func() { if r := recover(); r != nil { logger.ErrorCF("hooks", "Hook panic", @@ -146,7 +286,7 @@ func triggerVoid[T any](ctx context.Context, hooks []HookRegistration[T], event }) } }() - if err := reg.Handler(ctx, event); err != nil { + if err := reg.Handler(ctx, eventCopy); err != nil { logger.WarnCF("hooks", "Hook error", map[string]any{ "hook": hookName, @@ -204,7 +344,7 @@ func triggerModifying[T any]( } // TriggerMessageReceived fires all message_received handlers concurrently. -// Handlers must not mutate the event. +// Handler mutations are isolated per hook invocation and are not propagated. func (r *HookRegistry) TriggerMessageReceived(ctx context.Context, event *MessageReceivedEvent) { r.mu.RLock() hooks := r.messageReceived @@ -231,7 +371,7 @@ func (r *HookRegistry) TriggerBeforeToolCall(ctx context.Context, event *BeforeT } // TriggerAfterToolCall fires all after_tool_call handlers concurrently. -// Handlers must not mutate the event. +// Handler mutations are isolated per hook invocation and are not propagated. func (r *HookRegistry) TriggerAfterToolCall(ctx context.Context, event *AfterToolCallEvent) { r.mu.RLock() hooks := r.afterToolCall @@ -240,7 +380,7 @@ func (r *HookRegistry) TriggerAfterToolCall(ctx context.Context, event *AfterToo } // TriggerLLMInput fires all llm_input handlers concurrently. -// Handlers must not mutate the event. +// Handler mutations are isolated per hook invocation and are not propagated. func (r *HookRegistry) TriggerLLMInput(ctx context.Context, event *LLMInputEvent) { r.mu.RLock() hooks := r.llmInput @@ -249,7 +389,7 @@ func (r *HookRegistry) TriggerLLMInput(ctx context.Context, event *LLMInputEvent } // TriggerLLMOutput fires all llm_output handlers concurrently. -// Handlers must not mutate the event. +// Handler mutations are isolated per hook invocation and are not propagated. func (r *HookRegistry) TriggerLLMOutput(ctx context.Context, event *LLMOutputEvent) { r.mu.RLock() hooks := r.llmOutput @@ -258,7 +398,7 @@ func (r *HookRegistry) TriggerLLMOutput(ctx context.Context, event *LLMOutputEve } // TriggerSessionStart fires all session_start handlers concurrently. -// Handlers must not mutate the event. +// Handler mutations are isolated per hook invocation and are not propagated. func (r *HookRegistry) TriggerSessionStart(ctx context.Context, event *SessionEvent) { r.mu.RLock() hooks := r.sessionStart @@ -267,7 +407,7 @@ func (r *HookRegistry) TriggerSessionStart(ctx context.Context, event *SessionEv } // TriggerSessionEnd fires all session_end handlers concurrently. -// Handlers must not mutate the event. +// Handler mutations are isolated per hook invocation and are not propagated. func (r *HookRegistry) TriggerSessionEnd(ctx context.Context, event *SessionEvent) { r.mu.RLock() hooks := r.sessionEnd diff --git a/pkg/hooks/hooks_test.go b/pkg/hooks/hooks_test.go index 1f9aee18d..a293aed79 100644 --- a/pkg/hooks/hooks_test.go +++ b/pkg/hooks/hooks_test.go @@ -13,6 +13,8 @@ import ( "sync/atomic" "testing" "time" + + "github.com/sipeed/picoclaw/pkg/tools" ) func TestNewHookRegistry(t *testing.T) { @@ -96,6 +98,91 @@ func TestVoidHooksConcurrent(t *testing.T) { } } +func TestVoidHooksReceiveIsolatedMessageReceivedEvents(t *testing.T) { + r := NewHookRegistry() + ctx := context.Background() + + r.OnMessageReceived("mutator-a", 0, func(_ context.Context, e *MessageReceivedEvent) error { + e.Content = "changed-a" + e.Media[0] = "changed-media-a" + e.Metadata["k"] = "changed-a" + e.Metadata["new-a"] = "x" + return nil + }) + r.OnMessageReceived("mutator-b", 1, func(_ context.Context, e *MessageReceivedEvent) error { + e.Content = "changed-b" + e.Media = append(e.Media, "extra") + e.Metadata["k"] = "changed-b" + e.Metadata["new-b"] = "y" + return nil + }) + + event := &MessageReceivedEvent{ + Content: "original", + Media: []string{"m1"}, + Metadata: map[string]string{"k": "v"}, + } + r.TriggerMessageReceived(ctx, event) + + if event.Content != "original" { + t.Fatalf("expected original content to remain unchanged, got %q", event.Content) + } + if len(event.Media) != 1 || event.Media[0] != "m1" { + t.Fatalf("expected original media to remain unchanged, got %#v", event.Media) + } + if got := event.Metadata["k"]; got != "v" { + t.Fatalf("expected metadata[k] to remain v, got %q", got) + } + if _, ok := event.Metadata["new-a"]; ok { + t.Fatal("unexpected mutation leaked from hook mutator-a") + } + if _, ok := event.Metadata["new-b"]; ok { + t.Fatal("unexpected mutation leaked from hook mutator-b") + } +} + +func TestVoidHooksReceiveIsolatedAfterToolCallEvents(t *testing.T) { + r := NewHookRegistry() + ctx := context.Background() + + r.OnAfterToolCall("mutator-a", 0, func(_ context.Context, e *AfterToolCallEvent) error { + e.Args["k"] = "changed-a" + e.Result.ForLLM = "mutated-a" + return nil + }) + r.OnAfterToolCall("mutator-b", 1, func(_ context.Context, e *AfterToolCallEvent) error { + e.Args["k"] = "changed-b" + e.Args["new"] = "v" + e.Result.ForUser = "mutated-b" + return nil + }) + + event := &AfterToolCallEvent{ + ToolName: "shell", + Args: map[string]any{"k": "original"}, + Result: &tools.ToolResult{ + ForLLM: "for-llm", + ForUser: "for-user", + }, + } + + // Use a local copy so we can compare immutable expectations. + r.TriggerAfterToolCall(ctx, event) + + if got := event.Args["k"]; got != "original" { + t.Fatalf("expected args[k] to remain original, got %#v", got) + } + if _, ok := event.Args["new"]; ok { + t.Fatal("unexpected args mutation leaked from hook") + } + if event.Result.ForLLM != "for-llm" { + t.Fatalf("expected original result.ForLLM to remain unchanged, got %q", event.Result.ForLLM) + } + if event.Result.ForUser != "for-user" { + t.Fatalf("expected original result.ForUser to remain unchanged, got %q", event.Result.ForUser) + } +} + func TestModifyingHookPriority(t *testing.T) { r := NewHookRegistry() ctx := context.Background()