fix(hooks): isolate void-hook events per handler

This commit is contained in:
xj 2026-02-25 22:28:07 -08:00
parent 3132756623
commit d343d5e5f0
2 changed files with 235 additions and 8 deletions

View file

@ -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

View file

@ -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()