feat(plugin): add config-driven selection core and resolver

This commit is contained in:
xj 2026-02-28 17:38:46 -08:00
parent 9c9524f934
commit bc6d6b1200
14 changed files with 2681 additions and 0 deletions

View file

@ -0,0 +1,64 @@
package pluginruntime
import (
"fmt"
"github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/pkg/plugin"
"github.com/sipeed/picoclaw/pkg/plugin/builtin"
)
type Summary struct {
Enabled []string
Disabled []string
UnknownEnabled []string
UnknownDisabled []string
Warnings []string
}
func ResolveConfiguredPlugins(cfg *config.Config) ([]plugin.Plugin, Summary, error) {
if cfg == nil {
return nil, Summary{}, fmt.Errorf("config is nil")
}
resolved, err := plugin.ResolveSelection(
builtin.Names(),
plugin.SelectionInput{
DefaultEnabled: cfg.Plugins.DefaultEnabled,
Enabled: cfg.Plugins.Enabled,
Disabled: cfg.Plugins.Disabled,
},
)
summary := Summary{
Enabled: resolved.EnabledNames,
Disabled: resolved.DisabledNames,
UnknownEnabled: resolved.UnknownEnabled,
UnknownDisabled: resolved.UnknownDisabled,
Warnings: resolved.Warnings,
}
if err != nil {
return nil, summary, err
}
catalog := builtin.Catalog()
normalizedCatalog := make(map[string]builtin.Factory, len(catalog))
for name, factory := range catalog {
normalizedCatalog[plugin.NormalizePluginName(name)] = factory
}
instances := make([]plugin.Plugin, 0, len(resolved.EnabledNames))
for _, name := range resolved.EnabledNames {
factory, ok := normalizedCatalog[name]
if !ok {
return nil, summary, fmt.Errorf("builtin plugin %q has no factory", name)
}
instance := factory()
if instance == nil {
return nil, summary, fmt.Errorf("builtin plugin %q factory returned nil", name)
}
instances = append(instances, instance)
}
return instances, summary, nil
}

View file

@ -0,0 +1,106 @@
package pluginruntime
import (
"slices"
"strings"
"testing"
"github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/pkg/plugin"
"github.com/sipeed/picoclaw/pkg/plugin/builtin"
)
func TestResolveConfiguredPlugins_UnknownEnabledReturnsError(t *testing.T) {
cfg := config.DefaultConfig()
cfg.Plugins = config.PluginsConfig{
DefaultEnabled: false,
Enabled: []string{"missing-plugin"},
Disabled: []string{},
}
instances, summary, err := ResolveConfiguredPlugins(cfg)
if err == nil {
t.Fatal("expected error for unknown enabled plugin")
}
if !strings.Contains(err.Error(), "missing-plugin") {
t.Fatalf("expected error to mention missing plugin, got %v", err)
}
if len(instances) != 0 {
t.Fatalf("expected no instances on error, got %d", len(instances))
}
if !slices.Equal(summary.UnknownEnabled, []string{"missing-plugin"}) {
t.Fatalf("UnknownEnabled mismatch: got %v", summary.UnknownEnabled)
}
}
func TestResolveConfiguredPlugins_ReturnsDeterministicInstances(t *testing.T) {
available := builtin.Names()
if len(available) == 0 {
t.Fatal("expected at least one builtin plugin")
}
enabled := slices.Clone(available)
slices.Reverse(enabled)
cfg := config.DefaultConfig()
cfg.Plugins = config.PluginsConfig{
DefaultEnabled: false,
Enabled: enabled,
Disabled: []string{},
}
instances, summary, err := ResolveConfiguredPlugins(cfg)
if err != nil {
t.Fatalf("ResolveConfiguredPlugins() error = %v", err)
}
gotNames := pluginNames(instances)
if !slices.Equal(gotNames, available) {
t.Fatalf("plugin names mismatch: got %v, want %v", gotNames, available)
}
if !slices.Equal(summary.Enabled, available) {
t.Fatalf("summary enabled mismatch: got %v, want %v", summary.Enabled, available)
}
}
func TestResolveConfiguredPlugins_UnknownDisabledWarns(t *testing.T) {
cfg := config.DefaultConfig()
cfg.Plugins = config.PluginsConfig{
DefaultEnabled: true,
Enabled: []string{},
Disabled: []string{"missing-plugin"},
}
instances, summary, err := ResolveConfiguredPlugins(cfg)
if err != nil {
t.Fatalf("ResolveConfiguredPlugins() error = %v", err)
}
expectedEnabled := builtin.Names()
if !slices.Equal(pluginNames(instances), expectedEnabled) {
t.Fatalf("plugin names mismatch: got %v, want %v", pluginNames(instances), expectedEnabled)
}
if !slices.Equal(summary.UnknownDisabled, []string{"missing-plugin"}) {
t.Fatalf("UnknownDisabled mismatch: got %v", summary.UnknownDisabled)
}
if !hasWarningSubstring(summary.Warnings, `unknown disabled plugin "missing-plugin" ignored`) {
t.Fatalf("expected unknown disabled warning, got %v", summary.Warnings)
}
}
func pluginNames(instances []plugin.Plugin) []string {
names := make([]string, 0, len(instances))
for _, instance := range instances {
names = append(names, instance.Name())
}
return names
}
func hasWarningSubstring(warnings []string, sub string) bool {
for _, warning := range warnings {
if strings.Contains(warning, sub) {
return true
}
}
return false
}

View file

@ -53,6 +53,7 @@ type Config struct {
Session SessionConfig `json:"session,omitempty"` Session SessionConfig `json:"session,omitempty"`
Channels ChannelsConfig `json:"channels"` Channels ChannelsConfig `json:"channels"`
Providers ProvidersConfig `json:"providers,omitempty"` Providers ProvidersConfig `json:"providers,omitempty"`
Plugins PluginsConfig `json:"plugins,omitempty"`
ModelList []ModelConfig `json:"model_list"` // New model-centric provider configuration ModelList []ModelConfig `json:"model_list"` // New model-centric provider configuration
Gateway GatewayConfig `json:"gateway"` Gateway GatewayConfig `json:"gateway"`
Tools ToolsConfig `json:"tools"` Tools ToolsConfig `json:"tools"`
@ -167,6 +168,12 @@ type SessionConfig struct {
IdentityLinks map[string][]string `json:"identity_links,omitempty"` IdentityLinks map[string][]string `json:"identity_links,omitempty"`
} }
type PluginsConfig struct {
DefaultEnabled bool `json:"default_enabled"`
Enabled []string `json:"enabled,omitempty"`
Disabled []string `json:"disabled,omitempty"`
}
type AgentDefaults struct { type AgentDefaults struct {
Workspace string `json:"workspace" env:"PICOCLAW_AGENTS_DEFAULTS_WORKSPACE"` Workspace string `json:"workspace" env:"PICOCLAW_AGENTS_DEFAULTS_WORKSPACE"`
RestrictToWorkspace bool `json:"restrict_to_workspace" env:"PICOCLAW_AGENTS_DEFAULTS_RESTRICT_TO_WORKSPACE"` RestrictToWorkspace bool `json:"restrict_to_workspace" env:"PICOCLAW_AGENTS_DEFAULTS_RESTRICT_TO_WORKSPACE"`

View file

@ -243,6 +243,51 @@ func TestDefaultConfig_Temperature(t *testing.T) {
} }
} }
func TestDefaultConfig_PluginsDefaults(t *testing.T) {
cfg := DefaultConfig()
if !cfg.Plugins.DefaultEnabled {
t.Error("Plugins.DefaultEnabled should be true by default")
}
if cfg.Plugins.Enabled == nil {
t.Error("Plugins.Enabled should be initialized to an empty slice")
}
if len(cfg.Plugins.Enabled) != 0 {
t.Errorf("Plugins.Enabled len = %d, want 0", len(cfg.Plugins.Enabled))
}
if cfg.Plugins.Disabled == nil {
t.Error("Plugins.Disabled should be initialized to an empty slice")
}
if len(cfg.Plugins.Disabled) != 0 {
t.Errorf("Plugins.Disabled len = %d, want 0", len(cfg.Plugins.Disabled))
}
}
func TestConfig_PluginsJSONUnmarshal(t *testing.T) {
jsonData := `{
"plugins": {
"default_enabled": false,
"enabled": ["plugin-a", "plugin-b"],
"disabled": ["plugin-c"]
}
}`
cfg := DefaultConfig()
if err := json.Unmarshal([]byte(jsonData), cfg); err != nil {
t.Fatalf("unmarshal: %v", err)
}
if cfg.Plugins.DefaultEnabled {
t.Error("Plugins.DefaultEnabled = true, want false")
}
if len(cfg.Plugins.Enabled) != 2 || cfg.Plugins.Enabled[0] != "plugin-a" || cfg.Plugins.Enabled[1] != "plugin-b" {
t.Errorf("Plugins.Enabled = %v, want [plugin-a plugin-b]", cfg.Plugins.Enabled)
}
if len(cfg.Plugins.Disabled) != 1 || cfg.Plugins.Disabled[0] != "plugin-c" {
t.Errorf("Plugins.Disabled = %v, want [plugin-c]", cfg.Plugins.Disabled)
}
}
// TestDefaultConfig_Gateway verifies gateway defaults // TestDefaultConfig_Gateway verifies gateway defaults
func TestDefaultConfig_Gateway(t *testing.T) { func TestDefaultConfig_Gateway(t *testing.T) {
cfg := DefaultConfig() cfg := DefaultConfig()

View file

@ -134,6 +134,11 @@ func DefaultConfig() *Config {
Providers: ProvidersConfig{ Providers: ProvidersConfig{
OpenAI: OpenAIProviderConfig{WebSearch: true}, OpenAI: OpenAIProviderConfig{WebSearch: true},
}, },
Plugins: PluginsConfig{
DefaultEnabled: true,
Enabled: []string{},
Disabled: []string{},
},
ModelList: []ModelConfig{ ModelList: []ModelConfig{
// ============================================ // ============================================
// Add your API key to the model you want to use // Add your API key to the model you want to use

499
pkg/hooks/hooks.go Normal file
View file

@ -0,0 +1,499 @@
// PicoClaw - Ultra-lightweight personal AI agent
// Inspired by and based on nanobot: https://github.com/HKUDS/nanobot
// License: MIT
//
// Copyright (c) 2026 PicoClaw contributors
package hooks
import (
"context"
"fmt"
"reflect"
"sync"
"time"
"github.com/sipeed/picoclaw/pkg/logger"
"github.com/sipeed/picoclaw/pkg/providers"
)
const voidHookWaitBudget = 50 * time.Millisecond
// HookHandler is the callback signature for all hooks.
type HookHandler[T any] func(ctx context.Context, event *T) error
// HookRegistration tracks a handler with its priority and name.
type HookRegistration[T any] struct {
Handler HookHandler[T]
Priority int // Lower = runs first
Name string
}
// HookRegistry manages all lifecycle hooks.
type HookRegistry struct {
messageReceived []HookRegistration[MessageReceivedEvent]
messageSending []HookRegistration[MessageSendingEvent]
beforeToolCall []HookRegistration[BeforeToolCallEvent]
afterToolCall []HookRegistration[AfterToolCallEvent]
llmInput []HookRegistration[LLMInputEvent]
llmOutput []HookRegistration[LLMOutputEvent]
sessionStart []HookRegistration[SessionEvent]
sessionEnd []HookRegistration[SessionEvent]
mu sync.RWMutex
}
// NewHookRegistry creates an empty hook registry.
func NewHookRegistry() *HookRegistry {
return &HookRegistry{}
}
// insertSorted inserts a registration into a new slice sorted by priority.
// Always allocates a new backing array so concurrent readers of the old slice are safe.
func insertSorted[T any](slice []HookRegistration[T], reg HookRegistration[T]) []HookRegistration[T] {
i := 0
for i < len(slice) && slice[i].Priority <= reg.Priority {
i++
}
result := make([]HookRegistration[T], len(slice)+1)
copy(result, slice[:i])
result[i] = reg
copy(result[i+1:], slice[i:])
return result
}
// Registration methods
func (r *HookRegistry) OnMessageReceived(name string, priority int, handler HookHandler[MessageReceivedEvent]) {
r.mu.Lock()
defer r.mu.Unlock()
r.messageReceived = insertSorted(r.messageReceived, HookRegistration[MessageReceivedEvent]{
Handler: handler, Priority: priority, Name: name,
})
}
func (r *HookRegistry) OnMessageSending(name string, priority int, handler HookHandler[MessageSendingEvent]) {
r.mu.Lock()
defer r.mu.Unlock()
r.messageSending = insertSorted(r.messageSending, HookRegistration[MessageSendingEvent]{
Handler: handler, Priority: priority, Name: name,
})
}
func (r *HookRegistry) OnBeforeToolCall(name string, priority int, handler HookHandler[BeforeToolCallEvent]) {
r.mu.Lock()
defer r.mu.Unlock()
r.beforeToolCall = insertSorted(r.beforeToolCall, HookRegistration[BeforeToolCallEvent]{
Handler: handler, Priority: priority, Name: name,
})
}
func (r *HookRegistry) OnAfterToolCall(name string, priority int, handler HookHandler[AfterToolCallEvent]) {
r.mu.Lock()
defer r.mu.Unlock()
r.afterToolCall = insertSorted(r.afterToolCall, HookRegistration[AfterToolCallEvent]{
Handler: handler, Priority: priority, Name: name,
})
}
func (r *HookRegistry) OnLLMInput(name string, priority int, handler HookHandler[LLMInputEvent]) {
r.mu.Lock()
defer r.mu.Unlock()
r.llmInput = insertSorted(r.llmInput, HookRegistration[LLMInputEvent]{
Handler: handler, Priority: priority, Name: name,
})
}
func (r *HookRegistry) OnLLMOutput(name string, priority int, handler HookHandler[LLMOutputEvent]) {
r.mu.Lock()
defer r.mu.Unlock()
r.llmOutput = insertSorted(r.llmOutput, HookRegistration[LLMOutputEvent]{
Handler: handler, Priority: priority, Name: name,
})
}
func (r *HookRegistry) OnSessionStart(name string, priority int, handler HookHandler[SessionEvent]) {
r.mu.Lock()
defer r.mu.Unlock()
r.sessionStart = insertSorted(r.sessionStart, HookRegistration[SessionEvent]{
Handler: handler, Priority: priority, Name: name,
})
}
func (r *HookRegistry) OnSessionEnd(name string, priority int, handler HookHandler[SessionEvent]) {
r.mu.Lock()
defer r.mu.Unlock()
r.sessionEnd = insertSorted(r.sessionEnd, HookRegistration[SessionEvent]{
Handler: handler, Priority: priority, Name: name,
})
}
// 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 {
if v == nil {
return nil
}
cloned := cloneReflectValue(reflect.ValueOf(v))
if !cloned.IsValid() {
return nil
}
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
case reflect.Struct:
out := reflect.New(v.Type()).Elem()
for i := range v.NumField() {
field := out.Field(i)
if !field.CanSet() {
// Preserve original value for structs with non-settable fields.
return v
}
field.Set(cloneReflectValue(v.Field(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.
// It waits up to a small budget to collect immediate completions, then
// continues fail-open to avoid blocking the core agent pipeline.
// 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 {
return
}
var wg sync.WaitGroup
for _, h := range hooks {
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",
map[string]any{
"hook": hookName,
"handler": reg.Name,
"panic": fmt.Sprintf("%v", r),
})
}
}()
if err := reg.Handler(ctx, eventCopy); err != nil {
logger.WarnCF("hooks", "Hook error",
map[string]any{
"hook": hookName,
"handler": reg.Name,
"error": err.Error(),
})
}
}(h)
}
done := make(chan struct{})
go func() {
wg.Wait()
close(done)
}()
select {
case <-done:
case <-ctx.Done():
logger.WarnCF("hooks", "Void hook dispatch interrupted by context",
map[string]any{
"hook": hookName,
})
case <-time.After(voidHookWaitBudget):
logger.WarnCF("hooks", "Void hook dispatch exceeded wait budget; continuing",
map[string]any{
"hook": hookName,
"wait_budget_ms": voidHookWaitBudget.Milliseconds(),
})
}
}
// triggerModifying runs handlers sequentially by priority, stopping if Cancel is set.
// The cancelCheck function inspects the event to determine if Cancel was set.
func triggerModifying[T any](
ctx context.Context,
hooks []HookRegistration[T],
event *T,
hookName string,
cancelCheck func(*T) bool,
) {
if len(hooks) == 0 {
return
}
for _, h := range hooks {
func() {
defer func() {
if r := recover(); r != nil {
logger.ErrorCF("hooks", "Hook panic",
map[string]any{
"hook": hookName,
"handler": h.Name,
"panic": fmt.Sprintf("%v", r),
})
}
}()
if err := h.Handler(ctx, event); err != nil {
logger.WarnCF("hooks", "Hook error",
map[string]any{
"hook": hookName,
"handler": h.Name,
"error": err.Error(),
})
}
}()
if cancelCheck(event) {
logger.InfoCF("hooks", "Hook canceled operation",
map[string]any{
"hook": hookName,
"handler": h.Name,
})
return
}
}
}
// TriggerMessageReceived fires all message_received handlers concurrently.
// 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
r.mu.RUnlock()
triggerVoid(ctx, hooks, event, "message_received")
}
func (r *HookRegistry) TriggerMessageSending(ctx context.Context, event *MessageSendingEvent) {
r.mu.RLock()
hooks := r.messageSending
r.mu.RUnlock()
triggerModifying(ctx, hooks, event, "message_sending", func(e *MessageSendingEvent) bool {
return e.Cancel
})
}
func (r *HookRegistry) TriggerBeforeToolCall(ctx context.Context, event *BeforeToolCallEvent) {
r.mu.RLock()
hooks := r.beforeToolCall
r.mu.RUnlock()
triggerModifying(ctx, hooks, event, "before_tool_call", func(e *BeforeToolCallEvent) bool {
return e.Cancel
})
}
// TriggerAfterToolCall fires all after_tool_call handlers concurrently.
// 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
r.mu.RUnlock()
triggerVoid(ctx, hooks, event, "after_tool_call")
}
// TriggerLLMInput fires all llm_input handlers concurrently.
// 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
r.mu.RUnlock()
triggerVoid(ctx, hooks, event, "llm_input")
}
// TriggerLLMOutput fires all llm_output handlers concurrently.
// 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
r.mu.RUnlock()
triggerVoid(ctx, hooks, event, "llm_output")
}
// TriggerSessionStart fires all session_start handlers concurrently.
// 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
r.mu.RUnlock()
triggerVoid(ctx, hooks, event, "session_start")
}
// TriggerSessionEnd fires all session_end handlers concurrently.
// 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
r.mu.RUnlock()
triggerVoid(ctx, hooks, event, "session_end")
}

657
pkg/hooks/hooks_test.go Normal file
View file

@ -0,0 +1,657 @@
// PicoClaw - Ultra-lightweight personal AI agent
// Inspired by and based on nanobot: https://github.com/HKUDS/nanobot
// License: MIT
//
// Copyright (c) 2026 PicoClaw contributors
package hooks
import (
"context"
"fmt"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/sipeed/picoclaw/pkg/providers"
"github.com/sipeed/picoclaw/pkg/tools"
)
func TestNewHookRegistry(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
// Triggering all hooks on an empty registry should not panic.
r.TriggerMessageReceived(ctx, &MessageReceivedEvent{Content: "hello"})
r.TriggerMessageSending(ctx, &MessageSendingEvent{Content: "hello"})
r.TriggerBeforeToolCall(ctx, &BeforeToolCallEvent{ToolName: "t"})
r.TriggerAfterToolCall(ctx, &AfterToolCallEvent{ToolName: "t"})
r.TriggerLLMInput(ctx, &LLMInputEvent{AgentID: "a"})
r.TriggerLLMOutput(ctx, &LLMOutputEvent{AgentID: "a"})
r.TriggerSessionStart(ctx, &SessionEvent{AgentID: "a"})
r.TriggerSessionEnd(ctx, &SessionEvent{AgentID: "a"})
}
func TestVoidHookExecution(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
var called atomic.Bool
r.OnMessageReceived("test", 0, func(_ context.Context, e *MessageReceivedEvent) error {
called.Store(true)
if e.Content != "ping" {
t.Errorf("Expected content 'ping', got '%s'", e.Content)
}
return nil
})
r.TriggerMessageReceived(ctx, &MessageReceivedEvent{Content: "ping"})
if !called.Load() {
t.Error("Expected handler to be called")
}
}
func TestVoidHooksConcurrent(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
var count atomic.Int32
started := make(chan struct{}, 5)
release := make(chan struct{})
done := make(chan struct{})
for i := range 5 {
r.OnMessageReceived("hook-"+string(rune('A'+i)), i, func(_ context.Context, _ *MessageReceivedEvent) error {
started <- struct{}{}
<-release
count.Add(1)
return nil
})
}
go func() {
r.TriggerMessageReceived(ctx, &MessageReceivedEvent{Content: "test"})
close(done)
}()
// All 5 handlers must reach the barrier concurrently.
for i := range 5 {
select {
case <-started:
case <-time.After(1 * time.Second):
t.Fatalf("timeout waiting for handler %d to start", i+1)
}
}
// Release all handlers.
close(release)
select {
case <-done:
case <-time.After(1 * time.Second):
t.Fatal("timeout waiting for handlers to complete")
}
if count.Load() != 5 {
t.Errorf("Expected 5 handlers called, got %d", count.Load())
}
}
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 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 TestVoidHooksReceiveIsolatedStructValuesInMap(t *testing.T) {
type schemaSpec struct {
Required []string
Meta map[string]string
}
r := NewHookRegistry()
ctx := context.Background()
r.OnLLMInput("struct-mutator", 0, func(_ context.Context, e *LLMInputEvent) error {
spec, ok := e.Tools[0].Function.Parameters["schema"].(schemaSpec)
if !ok {
t.Fatal("schema should be schemaSpec")
}
spec.Required[0] = "mutated"
spec.Meta["k"] = "changed"
e.Tools[0].Function.Parameters["schema"] = spec
return nil
})
event := &LLMInputEvent{
AgentID: "a1",
Model: "m1",
Tools: []providers.ToolDefinition{
{
Type: "function",
Function: providers.ToolFunctionDefinition{
Name: "message",
Parameters: map[string]any{
"schema": schemaSpec{
Required: []string{"content"},
Meta: map[string]string{"k": "v"},
},
},
},
},
},
}
r.TriggerLLMInput(ctx, event)
spec, ok := event.Tools[0].Function.Parameters["schema"].(schemaSpec)
if !ok {
t.Fatal("schema should remain schemaSpec")
}
if len(spec.Required) != 1 || spec.Required[0] != "content" {
t.Fatalf("expected required to remain unchanged, got %#v", spec.Required)
}
if got := spec.Meta["k"]; got != "v" {
t.Fatalf("expected meta[k] to remain v, got %q", got)
}
}
func TestVoidHooksFailOpenOnSlowHandler(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
started := make(chan struct{})
release := make(chan struct{})
done := make(chan struct{})
r.OnLLMInput("slow", 0, func(_ context.Context, _ *LLMInputEvent) error {
close(started)
<-release
close(done)
return nil
})
begin := time.Now()
r.TriggerLLMInput(ctx, &LLMInputEvent{AgentID: "a1"})
elapsed := time.Since(begin)
if elapsed > voidHookWaitBudget*3 {
t.Fatalf("expected fail-open dispatch within budget, got %s", elapsed)
}
select {
case <-started:
case <-time.After(1 * time.Second):
t.Fatal("timeout waiting for slow handler to start")
}
close(release)
select {
case <-done:
case <-time.After(1 * time.Second):
t.Fatal("timeout waiting for slow handler to finish after release")
}
}
func TestModifyingHookPriority(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
var mu sync.Mutex
var order []string
// Register in reverse priority order to verify sorting.
r.OnMessageSending("third", 30, func(_ context.Context, _ *MessageSendingEvent) error {
mu.Lock()
order = append(order, "third")
mu.Unlock()
return nil
})
r.OnMessageSending("first", 10, func(_ context.Context, _ *MessageSendingEvent) error {
mu.Lock()
order = append(order, "first")
mu.Unlock()
return nil
})
r.OnMessageSending("second", 20, func(_ context.Context, _ *MessageSendingEvent) error {
mu.Lock()
order = append(order, "second")
mu.Unlock()
return nil
})
r.TriggerMessageSending(ctx, &MessageSendingEvent{Content: "hi"})
if len(order) != 3 {
t.Fatalf("Expected 3 handlers, got %d", len(order))
}
if order[0] != "first" || order[1] != "second" || order[2] != "third" {
t.Errorf("Expected [first second third], got %v", order)
}
}
func TestModifyingHookCancel(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
var secondCalled bool
r.OnMessageSending("canceler", 10, func(_ context.Context, e *MessageSendingEvent) error {
e.Cancel = true
e.CancelReason = "blocked"
return nil
})
r.OnMessageSending("after-cancel", 20, func(_ context.Context, _ *MessageSendingEvent) error {
secondCalled = true
return nil
})
event := &MessageSendingEvent{Content: "hi"}
r.TriggerMessageSending(ctx, event)
if !event.Cancel {
t.Error("Expected Cancel to be true")
}
if secondCalled {
t.Error("Expected second handler NOT to be called after cancel")
}
}
func TestBeforeToolCallModification(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
r.OnBeforeToolCall("modifier", 10, func(_ context.Context, e *BeforeToolCallEvent) error {
e.Args["injected"] = "value"
return nil
})
event := &BeforeToolCallEvent{
ToolName: "search",
Args: map[string]any{"query": "test"},
}
r.TriggerBeforeToolCall(ctx, event)
if event.Args["injected"] != "value" {
t.Error("Expected injected arg to persist")
}
if event.Args["query"] != "test" {
t.Error("Expected original arg to remain")
}
}
func TestMessageSendingFilter(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
r.OnMessageSending("rewriter", 10, func(_ context.Context, e *MessageSendingEvent) error {
e.Content = "[filtered] " + e.Content
return nil
})
event := &MessageSendingEvent{Content: "hello world"}
r.TriggerMessageSending(ctx, event)
if event.Content != "[filtered] hello world" {
t.Errorf("Expected '[filtered] hello world', got '%s'", event.Content)
}
}
func TestZeroCostWhenEmpty(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
// This is primarily a safety/smoke test — no panics, no allocations of note.
for range 100 {
r.TriggerMessageReceived(ctx, &MessageReceivedEvent{})
r.TriggerMessageSending(ctx, &MessageSendingEvent{})
r.TriggerBeforeToolCall(ctx, &BeforeToolCallEvent{})
r.TriggerAfterToolCall(ctx, &AfterToolCallEvent{})
r.TriggerLLMInput(ctx, &LLMInputEvent{})
r.TriggerLLMOutput(ctx, &LLMOutputEvent{})
r.TriggerSessionStart(ctx, &SessionEvent{})
r.TriggerSessionEnd(ctx, &SessionEvent{})
}
}
func TestLLMInputOutput(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
var inputCalled, outputCalled atomic.Bool
r.OnLLMInput("input-hook", 0, func(_ context.Context, e *LLMInputEvent) error {
if e.Model != "gpt-4" {
t.Errorf("Expected model 'gpt-4', got '%s'", e.Model)
}
inputCalled.Store(true)
return nil
})
r.OnLLMOutput("output-hook", 0, func(_ context.Context, e *LLMOutputEvent) error {
if e.Content != "response" {
t.Errorf("Expected content 'response', got '%s'", e.Content)
}
outputCalled.Store(true)
return nil
})
r.TriggerLLMInput(ctx, &LLMInputEvent{AgentID: "a1", Model: "gpt-4", Iteration: 1})
r.TriggerLLMOutput(ctx, &LLMOutputEvent{AgentID: "a1", Model: "gpt-4", Content: "response", Iteration: 1})
if !inputCalled.Load() {
t.Error("Expected LLM input hook to be called")
}
if !outputCalled.Load() {
t.Error("Expected LLM output hook to be called")
}
}
func TestSessionStartEnd(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
var startCalled, endCalled atomic.Bool
r.OnSessionStart("start-hook", 0, func(_ context.Context, e *SessionEvent) error {
if e.SessionKey != "sess-1" {
t.Errorf("Expected session key 'sess-1', got '%s'", e.SessionKey)
}
startCalled.Store(true)
return nil
})
r.OnSessionEnd("end-hook", 0, func(_ context.Context, e *SessionEvent) error {
if e.SessionKey != "sess-1" {
t.Errorf("Expected session key 'sess-1', got '%s'", e.SessionKey)
}
endCalled.Store(true)
return nil
})
event := &SessionEvent{AgentID: "a1", SessionKey: "sess-1", Channel: "test", ChatID: "c1"}
r.TriggerSessionStart(ctx, event)
r.TriggerSessionEnd(ctx, event)
if !startCalled.Load() {
t.Error("Expected session start hook to be called")
}
if !endCalled.Load() {
t.Error("Expected session end hook to be called")
}
}
func TestConcurrentRegistrationAndTrigger(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
var wg sync.WaitGroup
// Goroutines registering hooks.
for i := range 10 {
wg.Add(1)
go func(idx int) {
defer wg.Done()
r.OnMessageReceived(
fmt.Sprintf("reg-hook-%d", idx),
idx,
func(_ context.Context, _ *MessageReceivedEvent) error {
return nil
},
)
}(i)
}
// Goroutines triggering hooks concurrently.
for range 10 {
wg.Add(1)
go func() {
defer wg.Done()
r.TriggerMessageReceived(ctx, &MessageReceivedEvent{Content: "race"})
}()
}
wg.Wait()
}
func TestInsertSorted(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
var order []int
// Register with priorities: 50, 10, 30, 20, 40
priorities := []int{50, 10, 30, 20, 40}
for _, p := range priorities {
r.OnBeforeToolCall(fmt.Sprintf("p-%d", p), p, func(_ context.Context, _ *BeforeToolCallEvent) error {
order = append(order, p)
return nil
})
}
r.TriggerBeforeToolCall(ctx, &BeforeToolCallEvent{ToolName: "test", Args: map[string]any{}})
expected := []int{10, 20, 30, 40, 50}
if len(order) != len(expected) {
t.Fatalf("Expected %d handlers, got %d", len(expected), len(order))
}
for i, v := range expected {
if order[i] != v {
t.Errorf("Position %d: expected priority %d, got %d", i, v, order[i])
}
}
}
func TestAfterToolCallExecution(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
var called bool
var capturedName string
r.OnAfterToolCall("logger", 0, func(_ context.Context, event *AfterToolCallEvent) error {
called = true
capturedName = event.ToolName
return nil
})
r.TriggerAfterToolCall(ctx, &AfterToolCallEvent{
ToolName: "shell",
Args: map[string]any{"cmd": "ls"},
Channel: "telegram",
ChatID: "123",
})
if !called {
t.Error("Expected after_tool_call handler to be called")
}
if capturedName != "shell" {
t.Errorf("Expected ToolName 'shell', got '%s'", capturedName)
}
}
func TestHandlerErrorsSwallowed(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
// Test void hooks: error in one handler doesn't prevent others from running
var secondCalled bool
r.OnMessageReceived("erroring", 10, func(_ context.Context, _ *MessageReceivedEvent) error {
return fmt.Errorf("handler error")
})
r.OnMessageReceived("observer", 20, func(_ context.Context, _ *MessageReceivedEvent) error {
secondCalled = true
return nil
})
r.TriggerMessageReceived(ctx, &MessageReceivedEvent{Content: "test"})
if !secondCalled {
t.Error("Expected second void handler to run despite first handler's error")
}
// Test modifying hooks: error doesn't stop chain (only Cancel does)
var modifySecondCalled bool
r.OnMessageSending("erroring", 10, func(_ context.Context, _ *MessageSendingEvent) error {
return fmt.Errorf("handler error")
})
r.OnMessageSending("modifier", 20, func(_ context.Context, _ *MessageSendingEvent) error {
modifySecondCalled = true
return nil
})
r.TriggerMessageSending(ctx, &MessageSendingEvent{Content: "test"})
if !modifySecondCalled {
t.Error("Expected second modifying handler to run despite first handler's error")
}
}
func TestPanicRecovery(t *testing.T) {
r := NewHookRegistry()
ctx := context.Background()
// Void hook: panic in one handler shouldn't crash, other handlers should still run
var safeHandlerCalled bool
r.OnLLMInput("panicker", 10, func(_ context.Context, _ *LLMInputEvent) error {
panic("boom")
})
r.OnLLMInput("safe", 10, func(_ context.Context, _ *LLMInputEvent) error {
safeHandlerCalled = true
return nil
})
// Should not panic
r.TriggerLLMInput(ctx, &LLMInputEvent{AgentID: "test"})
if !safeHandlerCalled {
t.Error("Expected safe handler to run despite panicking sibling")
}
// Modifying hook: panic in handler shouldn't crash
r.OnBeforeToolCall("panicker", 10, func(_ context.Context, _ *BeforeToolCallEvent) error {
panic("boom")
})
// Should not panic
r.TriggerBeforeToolCall(ctx, &BeforeToolCallEvent{ToolName: "test"})
}

82
pkg/hooks/types.go Normal file
View file

@ -0,0 +1,82 @@
// PicoClaw - Ultra-lightweight personal AI agent
// Inspired by and based on nanobot: https://github.com/HKUDS/nanobot
// License: MIT
//
// Copyright (c) 2026 PicoClaw contributors
package hooks
import (
"time"
"github.com/sipeed/picoclaw/pkg/providers"
"github.com/sipeed/picoclaw/pkg/tools"
)
// MessageReceivedEvent is fired when an inbound message is consumed from the bus.
type MessageReceivedEvent struct {
Channel string
SenderID string
ChatID string
Content string
Media []string
Metadata map[string]string
}
// MessageSendingEvent is fired before an outbound message is published.
// Handlers can modify Content or set Cancel to block delivery.
type MessageSendingEvent struct {
Channel string
ChatID string
Content string // Modifiable
Cancel bool
CancelReason string
}
// BeforeToolCallEvent is fired before a tool is executed.
// Handlers can modify Args, or set Cancel to block execution.
type BeforeToolCallEvent struct {
ToolName string
Args map[string]any // Modifiable; guaranteed non-nil when triggered via AgentLoop.
Channel string
ChatID string
Cancel bool
CancelReason string // Message returned to LLM when canceled
}
// AfterToolCallEvent is fired after a tool completes execution.
type AfterToolCallEvent struct {
ToolName string
Args map[string]any
Channel string
ChatID string
Duration time.Duration
Result *tools.ToolResult
}
// LLMInputEvent is fired before the LLM provider is called.
type LLMInputEvent struct {
AgentID string
Model string
Messages []providers.Message
Tools []providers.ToolDefinition
Iteration int
}
// LLMOutputEvent is fired after the LLM provider responds.
type LLMOutputEvent struct {
AgentID string
Model string
Content string
ToolCalls []providers.ToolCall
Iteration int
Duration time.Duration
}
// SessionEvent is fired at session start and end.
type SessionEvent struct {
AgentID string
SessionKey string
Channel string
ChatID string
}

View file

@ -0,0 +1,31 @@
package builtin
import (
"sort"
"github.com/sipeed/picoclaw/pkg/plugin"
"github.com/sipeed/picoclaw/pkg/plugin/demoplugin"
)
// Factory creates one builtin plugin instance.
type Factory func() plugin.Plugin
// Catalog returns compile-time builtin plugin factories by name.
func Catalog() map[string]Factory {
return map[string]Factory{
"policy-demo": func() plugin.Plugin {
return demoplugin.NewPolicyDemoPlugin(demoplugin.PolicyDemoConfig{})
},
}
}
// Names returns sorted builtin plugin names.
func Names() []string {
catalog := Catalog()
names := make([]string, 0, len(catalog))
for name := range catalog {
names = append(names, name)
}
sort.Strings(names)
return names
}

View file

@ -0,0 +1,32 @@
package builtin
import (
"slices"
"testing"
)
func TestCatalogContainsPolicyDemo(t *testing.T) {
catalog := Catalog()
factory, ok := catalog["policy-demo"]
if !ok {
t.Fatalf("Catalog() missing %q plugin", "policy-demo")
}
if factory == nil {
t.Fatalf("Catalog()[%q] factory is nil", "policy-demo")
}
if got := factory(); got == nil {
t.Fatalf("Catalog()[%q]() returned nil plugin", "policy-demo")
}
}
func TestNamesSorted(t *testing.T) {
first := Names()
second := Names()
if !slices.IsSorted(first) {
t.Fatalf("Names() is not sorted: %v", first)
}
if !slices.Equal(first, second) {
t.Fatalf("Names() is not deterministic across calls: %v vs %v", first, second)
}
}

View file

@ -0,0 +1,315 @@
package demoplugin
import (
"context"
"fmt"
"strings"
"sync"
"time"
"github.com/sipeed/picoclaw/pkg/hooks"
"github.com/sipeed/picoclaw/pkg/plugin"
)
// PolicyDemoConfig controls the demo plugin behavior.
type PolicyDemoConfig struct {
BlockedTools []string
RedactPrefixes []string
ChannelToolAllowlist map[string][]string
DenyOutboundPatterns []string
MaxToolTimeoutSecond int
}
// PolicyDemoStats provides basic evidence that hook paths were executed.
type PolicyDemoStats struct {
BeforeToolCalls int
BlockedToolCalls int
MessageSends int
RedactedMessages int
BlockedMessages int
SessionStarts int
SessionEnds int
AfterToolCalls int
TotalToolDuration time.Duration
}
// PolicyDemoPlugin demonstrates why plugins are needed: it enforces runtime policy
// at tool-call and outbound-message lifecycle points and collects audit metrics.
type PolicyDemoPlugin struct {
blockedTools map[string]struct{}
prefixes []string
channelAllowlist map[string]map[string]struct{}
denyPatterns []string
maxTimeout int
mu sync.Mutex
stats PolicyDemoStats
}
func NewPolicyDemoPlugin(cfg PolicyDemoConfig) *PolicyDemoPlugin {
blocked := make(map[string]struct{}, len(cfg.BlockedTools))
for _, t := range cfg.BlockedTools {
t = normalizeLower(t)
if t == "" {
continue
}
blocked[t] = struct{}{}
}
prefixes := make([]string, 0, len(cfg.RedactPrefixes))
for _, p := range cfg.RedactPrefixes {
p = strings.TrimSpace(p)
if p == "" {
continue
}
prefixes = append(prefixes, p)
}
allowlist := make(map[string]map[string]struct{}, len(cfg.ChannelToolAllowlist))
for channel, tools := range cfg.ChannelToolAllowlist {
channel = normalizeLower(channel)
if channel == "" {
continue
}
toolSet := make(map[string]struct{}, len(tools))
for _, t := range tools {
t = normalizeLower(t)
if t == "" {
continue
}
toolSet[t] = struct{}{}
}
allowlist[channel] = toolSet
}
patterns := make([]string, 0, len(cfg.DenyOutboundPatterns))
for _, p := range cfg.DenyOutboundPatterns {
p = strings.TrimSpace(p)
if p == "" {
continue
}
patterns = append(patterns, p)
}
maxTimeout := cfg.MaxToolTimeoutSecond
if maxTimeout < 0 {
maxTimeout = 0
}
return &PolicyDemoPlugin{
blockedTools: blocked,
prefixes: prefixes,
channelAllowlist: allowlist,
denyPatterns: patterns,
maxTimeout: maxTimeout,
}
}
func (p *PolicyDemoPlugin) Name() string {
return "policy-demo"
}
func (p *PolicyDemoPlugin) APIVersion() string {
return plugin.APIVersion
}
func (p *PolicyDemoPlugin) Snapshot() PolicyDemoStats {
p.mu.Lock()
defer p.mu.Unlock()
return p.stats
}
func (p *PolicyDemoPlugin) Register(r *hooks.HookRegistry) error {
r.OnBeforeToolCall("policy-demo-tool-policy", 100, func(_ context.Context, e *hooks.BeforeToolCallEvent) error {
tool := normalizeLower(e.ToolName)
p.incBeforeToolCalls()
if _, blocked := p.blockedTools[tool]; blocked {
e.Cancel = true
e.CancelReason = "blocked by policy-demo plugin"
p.incBlockedToolCalls()
return nil
}
channel := normalizeLower(e.Channel)
if allow, ok := p.channelAllowlist[channel]; ok {
if _, allowed := allow[tool]; !allowed {
e.Cancel = true
e.CancelReason = fmt.Sprintf("tool %q is not allowed on channel %q", e.ToolName, e.Channel)
p.incBlockedToolCalls()
return nil
}
}
if p.maxTimeout > 0 {
clampArgNumber(e.Args, "timeout", p.maxTimeout)
clampArgNumber(e.Args, "timeout_seconds", p.maxTimeout)
}
return nil
})
r.OnMessageSending("policy-demo-redact-and-guard", 50, func(_ context.Context, e *hooks.MessageSendingEvent) error {
p.incMessageSends()
for _, pattern := range p.denyPatterns {
if strings.Contains(e.Content, pattern) {
e.Cancel = true
e.CancelReason = "blocked by policy-demo outbound guard"
p.incBlockedMessages()
return nil
}
}
content := e.Content
redacted := false
for _, prefix := range p.prefixes {
next := strings.ReplaceAll(content, prefix, "[redacted]-")
if next != content {
redacted = true
}
content = next
}
e.Content = content
if redacted {
p.incRedactedMessages()
}
return nil
})
r.OnSessionStart("policy-demo-session-start-audit", 0, func(_ context.Context, _ *hooks.SessionEvent) error {
p.incSessionStarts()
return nil
})
r.OnSessionEnd("policy-demo-session-end-audit", 0, func(_ context.Context, _ *hooks.SessionEvent) error {
p.incSessionEnds()
return nil
})
r.OnAfterToolCall("policy-demo-after-tool-audit", 0, func(_ context.Context, e *hooks.AfterToolCallEvent) error {
p.incAfterToolCall(e.Duration)
return nil
})
return nil
}
func normalizeLower(s string) string {
return strings.ToLower(strings.TrimSpace(s))
}
func clampArgNumber(args map[string]any, key string, limit int) {
if args == nil || limit <= 0 {
return
}
v, ok := args[key]
if !ok {
return
}
n, ok := toInt(v)
if !ok {
return
}
if n > limit {
args[key] = limit
}
}
func toInt(v any) (int, bool) {
maxInt := int(^uint(0) >> 1)
maxIntU64 := uint64(maxInt)
maxInt64 := int64(maxInt)
minInt64 := -maxInt64 - 1
switch n := v.(type) {
case int:
return n, true
case int8:
return int(n), true
case int16:
return int(n), true
case int32:
return int(n), true
case int64:
if n < minInt64 || n > maxInt64 {
return 0, false
}
return int(n), true
case uint:
if uint64(n) > maxIntU64 {
return 0, false
}
return int(n), true
case uint8:
return int(n), true
case uint16:
return int(n), true
case uint32:
if uint64(n) > maxIntU64 {
return 0, false
}
return int(n), true
case uint64:
if n > maxIntU64 {
return 0, false
}
return int(n), true
case float32:
// Truncation is intentional for timeout normalization.
return int(n), true
case float64:
// Truncation is intentional for timeout normalization.
return int(n), true
default:
return 0, false
}
}
func (p *PolicyDemoPlugin) incBeforeToolCalls() {
p.mu.Lock()
defer p.mu.Unlock()
p.stats.BeforeToolCalls++
}
func (p *PolicyDemoPlugin) incBlockedToolCalls() {
p.mu.Lock()
defer p.mu.Unlock()
p.stats.BlockedToolCalls++
}
func (p *PolicyDemoPlugin) incMessageSends() {
p.mu.Lock()
defer p.mu.Unlock()
p.stats.MessageSends++
}
func (p *PolicyDemoPlugin) incRedactedMessages() {
p.mu.Lock()
defer p.mu.Unlock()
p.stats.RedactedMessages++
}
func (p *PolicyDemoPlugin) incBlockedMessages() {
p.mu.Lock()
defer p.mu.Unlock()
p.stats.BlockedMessages++
}
func (p *PolicyDemoPlugin) incSessionStarts() {
p.mu.Lock()
defer p.mu.Unlock()
p.stats.SessionStarts++
}
func (p *PolicyDemoPlugin) incSessionEnds() {
p.mu.Lock()
defer p.mu.Unlock()
p.stats.SessionEnds++
}
func (p *PolicyDemoPlugin) incAfterToolCall(d time.Duration) {
p.mu.Lock()
defer p.mu.Unlock()
p.stats.AfterToolCalls++
p.stats.TotalToolDuration += d
}

View file

@ -0,0 +1,189 @@
package demoplugin
import (
"context"
"strconv"
"testing"
"time"
"github.com/sipeed/picoclaw/pkg/hooks"
"github.com/sipeed/picoclaw/pkg/plugin"
)
func TestPolicyDemoPluginBlocksConfiguredTool(t *testing.T) {
pm := plugin.NewManager()
p := NewPolicyDemoPlugin(PolicyDemoConfig{
BlockedTools: []string{"shell"},
})
if err := pm.Register(p); err != nil {
t.Fatalf("register plugin: %v", err)
}
e := &hooks.BeforeToolCallEvent{ToolName: "shell", Args: map[string]any{}, Channel: "cli"}
pm.HookRegistry().TriggerBeforeToolCall(context.Background(), e)
if !e.Cancel {
t.Fatal("expected tool call to be canceled")
}
if e.CancelReason == "" {
t.Fatal("expected cancel reason")
}
stats := p.Snapshot()
if stats.BeforeToolCalls != 1 || stats.BlockedToolCalls != 1 {
t.Fatalf("unexpected stats: %+v", stats)
}
}
func TestPolicyDemoPluginRedactsOutboundContent(t *testing.T) {
pm := plugin.NewManager()
p := NewPolicyDemoPlugin(PolicyDemoConfig{
RedactPrefixes: []string{"sk-"},
})
if err := pm.Register(p); err != nil {
t.Fatalf("register plugin: %v", err)
}
e := &hooks.MessageSendingEvent{Content: "token=sk-abc123"}
pm.HookRegistry().TriggerMessageSending(context.Background(), e)
if e.Cancel {
t.Fatal("did not expect cancellation")
}
if e.Content != "token=[redacted]-abc123" {
t.Fatalf("unexpected redaction result: %q", e.Content)
}
stats := p.Snapshot()
if stats.MessageSends != 1 || stats.RedactedMessages != 1 {
t.Fatalf("unexpected stats: %+v", stats)
}
}
func TestPolicyDemoPluginChannelAllowlist(t *testing.T) {
pm := plugin.NewManager()
p := NewPolicyDemoPlugin(PolicyDemoConfig{
ChannelToolAllowlist: map[string][]string{
"telegram": {"web_search"},
},
})
if err := pm.Register(p); err != nil {
t.Fatalf("register plugin: %v", err)
}
blocked := &hooks.BeforeToolCallEvent{ToolName: "shell", Args: map[string]any{}, Channel: "telegram"}
pm.HookRegistry().TriggerBeforeToolCall(context.Background(), blocked)
if !blocked.Cancel {
t.Fatal("expected tool to be blocked by channel allowlist")
}
allowed := &hooks.BeforeToolCallEvent{ToolName: "web_search", Args: map[string]any{}, Channel: "telegram"}
pm.HookRegistry().TriggerBeforeToolCall(context.Background(), allowed)
if allowed.Cancel {
t.Fatalf("did not expect allowlisted tool to be blocked: %s", allowed.CancelReason)
}
}
func TestPolicyDemoPluginOutboundGuard(t *testing.T) {
pm := plugin.NewManager()
p := NewPolicyDemoPlugin(PolicyDemoConfig{
DenyOutboundPatterns: []string{"4111-1111-1111-1111", "@corp.internal"},
})
if err := pm.Register(p); err != nil {
t.Fatalf("register plugin: %v", err)
}
e := &hooks.MessageSendingEvent{Content: "card=4111-1111-1111-1111"}
pm.HookRegistry().TriggerMessageSending(context.Background(), e)
if !e.Cancel {
t.Fatal("expected outbound message to be blocked")
}
if e.CancelReason == "" {
t.Fatal("expected block reason")
}
stats := p.Snapshot()
if stats.BlockedMessages != 1 {
t.Fatalf("expected blocked message count to be 1, got %+v", stats)
}
}
func TestPolicyDemoPluginNormalizesTimeoutArg(t *testing.T) {
pm := plugin.NewManager()
p := NewPolicyDemoPlugin(PolicyDemoConfig{MaxToolTimeoutSecond: 30})
if err := pm.Register(p); err != nil {
t.Fatalf("register plugin: %v", err)
}
e := &hooks.BeforeToolCallEvent{
ToolName: "web_fetch",
Channel: "cli",
Args: map[string]any{
"timeout": 120,
"timeout_seconds": 90.0,
},
}
pm.HookRegistry().TriggerBeforeToolCall(context.Background(), e)
if got, ok := e.Args["timeout"].(int); !ok || got != 30 {
t.Fatalf("expected timeout to be clamped to 30, got %#v", e.Args["timeout"])
}
if got, ok := e.Args["timeout_seconds"].(int); !ok || got != 30 {
t.Fatalf("expected timeout_seconds to be clamped to 30, got %#v", e.Args["timeout_seconds"])
}
}
func TestPolicyDemoPluginAuditHooks(t *testing.T) {
pm := plugin.NewManager()
p := NewPolicyDemoPlugin(PolicyDemoConfig{})
if err := pm.Register(p); err != nil {
t.Fatalf("register plugin: %v", err)
}
pm.HookRegistry().TriggerSessionStart(context.Background(), &hooks.SessionEvent{AgentID: "a1", SessionKey: "s1"})
pm.HookRegistry().TriggerAfterToolCall(
context.Background(),
&hooks.AfterToolCallEvent{
ToolName: "web_search",
Duration: 45 * time.Millisecond,
},
)
pm.HookRegistry().TriggerSessionEnd(context.Background(), &hooks.SessionEvent{AgentID: "a1", SessionKey: "s1"})
stats := p.Snapshot()
if stats.SessionStarts != 1 || stats.SessionEnds != 1 {
t.Fatalf("unexpected session stats: %+v", stats)
}
if stats.AfterToolCalls != 1 || stats.TotalToolDuration != 45*time.Millisecond {
t.Fatalf("unexpected after_tool_call stats: %+v", stats)
}
}
func TestPolicyDemoPluginNoConfigNoEffect(t *testing.T) {
pm := plugin.NewManager()
p := NewPolicyDemoPlugin(PolicyDemoConfig{})
if err := pm.Register(p); err != nil {
t.Fatalf("register plugin: %v", err)
}
toolEvent := &hooks.BeforeToolCallEvent{ToolName: "shell", Args: map[string]any{}, Channel: "telegram"}
pm.HookRegistry().TriggerBeforeToolCall(context.Background(), toolEvent)
if toolEvent.Cancel {
t.Fatal("did not expect cancellation with empty config")
}
msgEvent := &hooks.MessageSendingEvent{Content: "token=sk-abc123"}
pm.HookRegistry().TriggerMessageSending(context.Background(), msgEvent)
if msgEvent.Content != "token=sk-abc123" {
t.Fatalf("did not expect content rewrite, got %q", msgEvent.Content)
}
}
func TestToIntRejectsInt64OverflowOn32Bit(t *testing.T) {
if strconv.IntSize != 32 {
t.Skip("overflow scenario is specific to 32-bit int")
}
if _, ok := toInt(int64(1 << 40)); ok {
t.Fatal("expected overflow conversion to fail on 32-bit int")
}
}

275
pkg/plugin/manager.go Normal file
View file

@ -0,0 +1,275 @@
// PicoClaw - Ultra-lightweight personal AI agent
// Inspired by and based on nanobot: https://github.com/HKUDS/nanobot
// License: MIT
//
// Copyright (c) 2026 PicoClaw contributors
package plugin
import (
"errors"
"fmt"
"sort"
"slices"
"strings"
"sync"
"github.com/sipeed/picoclaw/pkg/hooks"
)
// APIVersion identifies the compile-time plugin contract version.
const APIVersion = "v1alpha1"
// SelectionInput controls plugin enable/disable resolution.
type SelectionInput struct {
DefaultEnabled bool
Enabled []string
Disabled []string
}
// SelectionResult is the normalized output of plugin enable/disable resolution.
type SelectionResult struct {
EnabledNames []string
DisabledNames []string
UnknownEnabled []string
UnknownDisabled []string
Warnings []string
}
// Plugin is the Phase-1 compile-time contract for PicoClaw extensions.
type Plugin interface {
Name() string
APIVersion() string
Register(registry *hooks.HookRegistry) error
}
// PluginInfo describes plugin metadata for introspection APIs.
type PluginInfo struct {
Name string `json:"name"`
APIVersion string `json:"api_version"`
Status string `json:"status"`
}
// PluginDescriptor optionally provides richer plugin metadata.
type PluginDescriptor interface {
Info() PluginInfo
}
// NormalizePluginName normalizes plugin names for deterministic matching.
func NormalizePluginName(name string) string {
return strings.ToLower(strings.TrimSpace(name))
}
// ResolveSelection resolves final enabled/disabled plugin names deterministically.
func ResolveSelection(available []string, in SelectionInput) (SelectionResult, error) {
result := SelectionResult{}
availableSet := make(map[string]struct{}, len(available))
for _, name := range available {
normalized := NormalizePluginName(name)
if normalized == "" {
continue
}
availableSet[normalized] = struct{}{}
}
enabledSet := make(map[string]struct{}, len(in.Enabled))
for _, name := range in.Enabled {
normalized := NormalizePluginName(name)
if _, exists := enabledSet[normalized]; exists {
result.Warnings = append(result.Warnings, fmt.Sprintf("duplicate enabled plugin %q ignored", normalized))
continue
}
enabledSet[normalized] = struct{}{}
}
disabledSet := make(map[string]struct{}, len(in.Disabled))
for _, name := range in.Disabled {
normalized := NormalizePluginName(name)
if _, exists := disabledSet[normalized]; exists {
result.Warnings = append(result.Warnings, fmt.Sprintf("duplicate disabled plugin %q ignored", normalized))
continue
}
disabledSet[normalized] = struct{}{}
}
for name := range enabledSet {
if _, ok := availableSet[name]; !ok {
result.UnknownEnabled = append(result.UnknownEnabled, name)
}
}
sort.Strings(result.UnknownEnabled)
for name := range disabledSet {
if _, ok := availableSet[name]; !ok {
result.UnknownDisabled = append(result.UnknownDisabled, name)
}
}
sort.Strings(result.UnknownDisabled)
for _, name := range result.UnknownDisabled {
result.Warnings = append(result.Warnings, fmt.Sprintf("unknown disabled plugin %q ignored", name))
}
resolvedEnabled := make(map[string]struct{}, len(availableSet))
if len(enabledSet) > 0 {
for name := range enabledSet {
if _, ok := availableSet[name]; !ok {
continue
}
if _, disabled := disabledSet[name]; disabled {
continue
}
resolvedEnabled[name] = struct{}{}
}
} else if in.DefaultEnabled {
for name := range availableSet {
if _, disabled := disabledSet[name]; disabled {
continue
}
resolvedEnabled[name] = struct{}{}
}
}
for name := range resolvedEnabled {
result.EnabledNames = append(result.EnabledNames, name)
}
sort.Strings(result.EnabledNames)
for name := range availableSet {
if _, enabled := resolvedEnabled[name]; enabled {
continue
}
result.DisabledNames = append(result.DisabledNames, name)
}
sort.Strings(result.DisabledNames)
if len(result.UnknownEnabled) > 0 {
return result, fmt.Errorf("unknown enabled plugins: %s", strings.Join(result.UnknownEnabled, ", "))
}
return result, nil
}
// Manager owns a shared hook registry and loaded plugin metadata.
type Manager struct {
mu sync.RWMutex
registry *hooks.HookRegistry
names []string
plugins []Plugin
seen map[string]struct{}
}
// NewManager creates an empty plugin manager with a fresh hook registry.
func NewManager() *Manager {
return &Manager{
registry: hooks.NewHookRegistry(),
seen: make(map[string]struct{}),
}
}
// HookRegistry returns the shared registry where plugins register hooks.
func (m *Manager) HookRegistry() *hooks.HookRegistry {
return m.registry
}
// Names returns loaded plugin names in registration order.
func (m *Manager) Names() []string {
m.mu.RLock()
defer m.mu.RUnlock()
return slices.Clone(m.names)
}
// DescribeAll returns plugin metadata in registration order.
func (m *Manager) DescribeAll() []PluginInfo {
m.mu.RLock()
defer m.mu.RUnlock()
infos := make([]PluginInfo, 0, len(m.plugins))
for i, p := range m.plugins {
fallbackName := ""
if i < len(m.names) {
fallbackName = m.names[i]
}
infos = append(infos, normalizePluginInfo(p, fallbackName))
}
return infos
}
// DescribeEnabled returns metadata for currently enabled plugins.
func (m *Manager) DescribeEnabled() []PluginInfo {
return m.DescribeAll()
}
// Register loads one plugin into the shared hook registry.
func (m *Manager) Register(p Plugin) error {
if p == nil {
return errors.New("plugin is nil")
}
name := strings.TrimSpace(p.Name())
if name == "" {
return errors.New("plugin name is required")
}
if got := strings.TrimSpace(p.APIVersion()); got != APIVersion {
if got == "" {
got = "<empty>"
}
return fmt.Errorf(
"plugin %q api version mismatch: got %s, want %s",
name,
got,
APIVersion,
)
}
m.mu.Lock()
defer m.mu.Unlock()
if _, exists := m.seen[name]; exists {
return fmt.Errorf("plugin %q already registered", name)
}
if err := p.Register(m.registry); err != nil {
return fmt.Errorf("register plugin %q: %w", name, err)
}
m.seen[name] = struct{}{}
m.names = append(m.names, name)
m.plugins = append(m.plugins, p)
return nil
}
// RegisterAll loads plugins sequentially.
func (m *Manager) RegisterAll(plugins ...Plugin) error {
for _, p := range plugins {
if err := m.Register(p); err != nil {
return err
}
}
return nil
}
func normalizePluginInfo(p Plugin, fallbackName string) PluginInfo {
info := PluginInfo{
Name: strings.TrimSpace(fallbackName),
APIVersion: strings.TrimSpace(p.APIVersion()),
Status: "enabled",
}
if descriptor, ok := p.(PluginDescriptor); ok {
described := descriptor.Info()
if name := strings.TrimSpace(described.Name); name != "" {
info.Name = name
}
if version := strings.TrimSpace(described.APIVersion); version != "" {
info.APIVersion = version
}
if status := strings.TrimSpace(described.Status); status != "" {
info.Status = status
}
}
if info.Name == "" {
info.Name = strings.TrimSpace(p.Name())
}
if info.APIVersion == "" {
info.APIVersion = APIVersion
}
if info.Status == "" {
info.Status = "enabled"
}
return info
}

374
pkg/plugin/manager_test.go Normal file
View file

@ -0,0 +1,374 @@
package plugin
import (
"context"
"errors"
"slices"
"strings"
"testing"
"github.com/sipeed/picoclaw/pkg/hooks"
)
type testPlugin struct {
name string
apiVersion string
registerFn func(*hooks.HookRegistry) error
}
func (p testPlugin) Name() string {
return p.name
}
func (p testPlugin) Register(r *hooks.HookRegistry) error {
if p.registerFn != nil {
return p.registerFn(r)
}
return nil
}
func (p testPlugin) APIVersion() string {
if p.apiVersion == "" {
return APIVersion
}
return p.apiVersion
}
type descriptorTestPlugin struct {
testPlugin
info PluginInfo
}
func (p descriptorTestPlugin) Info() PluginInfo {
return p.info
}
func TestNewManager(t *testing.T) {
m := NewManager()
if m == nil {
t.Fatal("expected manager")
}
if m.HookRegistry() == nil {
t.Fatal("expected non-nil hook registry")
}
if len(m.Names()) != 0 {
t.Fatalf("expected empty names, got %v", m.Names())
}
}
func TestRegisterPluginAndTriggerHook(t *testing.T) {
m := NewManager()
called := false
p := testPlugin{
name: "audit",
registerFn: func(r *hooks.HookRegistry) error {
r.OnSessionStart("audit-session", 0, func(_ context.Context, _ *hooks.SessionEvent) error {
called = true
return nil
})
return nil
},
}
if err := m.Register(p); err != nil {
t.Fatalf("Register() error = %v", err)
}
if got := m.Names(); len(got) != 1 || got[0] != "audit" {
t.Fatalf("unexpected names: %v", got)
}
m.HookRegistry().TriggerSessionStart(context.Background(), &hooks.SessionEvent{
AgentID: "a1",
SessionKey: "s1",
})
if !called {
t.Fatal("expected plugin hook to be called")
}
}
func TestRegisterRejectsNilPlugin(t *testing.T) {
m := NewManager()
if err := m.Register(nil); err == nil {
t.Fatal("expected error for nil plugin")
}
}
func TestRegisterRejectsEmptyName(t *testing.T) {
m := NewManager()
if err := m.Register(testPlugin{}); err == nil {
t.Fatal("expected error for empty name")
}
}
func TestRegisterRejectsDuplicateName(t *testing.T) {
m := NewManager()
p := testPlugin{name: "dup"}
if err := m.Register(p); err != nil {
t.Fatalf("unexpected first register error: %v", err)
}
if err := m.Register(p); err == nil {
t.Fatal("expected duplicate name error")
}
}
func TestRegisterPropagatesPluginError(t *testing.T) {
m := NewManager()
want := errors.New("register failed")
p := testPlugin{
name: "bad",
registerFn: func(_ *hooks.HookRegistry) error {
return want
},
}
err := m.Register(p)
if err == nil {
t.Fatal("expected error")
}
if !errors.Is(err, want) {
t.Fatalf("expected wrapped error %v, got %v", want, err)
}
}
func TestRegisterRejectsPluginVersionMismatch(t *testing.T) {
m := NewManager()
p := testPlugin{
name: "old-plugin",
apiVersion: "v0",
}
err := m.Register(p)
if err == nil {
t.Fatal("expected version mismatch error")
}
}
func TestDescribeAll_UsesDescriptorWhenImplemented(t *testing.T) {
m := NewManager()
p := descriptorTestPlugin{
testPlugin: testPlugin{name: "descriptor"},
info: PluginInfo{
Name: " descriptor-visible ",
APIVersion: " custom-v1 ",
Status: " active ",
},
}
if err := m.Register(p); err != nil {
t.Fatalf("Register() error = %v", err)
}
got := m.DescribeAll()
want := []PluginInfo{
{
Name: "descriptor-visible",
APIVersion: "custom-v1",
Status: "active",
},
}
if !slices.Equal(got, want) {
t.Fatalf("DescribeAll() mismatch: got %v, want %v", got, want)
}
}
func TestDescribeAll_FallsBackForPlainPlugin(t *testing.T) {
m := NewManager()
p := testPlugin{name: "plain"}
if err := m.Register(p); err != nil {
t.Fatalf("Register() error = %v", err)
}
got := m.DescribeAll()
want := []PluginInfo{
{
Name: "plain",
APIVersion: APIVersion,
Status: "enabled",
},
}
if !slices.Equal(got, want) {
t.Fatalf("DescribeAll() mismatch: got %v, want %v", got, want)
}
}
func TestDescribeEnabled_MatchesDescribeAllForNow(t *testing.T) {
m := NewManager()
plain := testPlugin{name: "plain"}
described := descriptorTestPlugin{
testPlugin: testPlugin{name: "described"},
info: PluginInfo{
Name: " described-visible ",
},
}
if err := m.RegisterAll(plain, described); err != nil {
t.Fatalf("RegisterAll() error = %v", err)
}
all := m.DescribeAll()
enabled := m.DescribeEnabled()
if !slices.Equal(enabled, all) {
t.Fatalf("DescribeEnabled() mismatch: got %v, want %v", enabled, all)
}
wantAll := []PluginInfo{
{
Name: "plain",
APIVersion: APIVersion,
Status: "enabled",
},
{
Name: "described-visible",
APIVersion: APIVersion,
Status: "enabled",
},
}
if !slices.Equal(all, wantAll) {
t.Fatalf("DescribeAll() order/content mismatch: got %v, want %v", all, wantAll)
}
}
func TestResolveSelection_DefaultEnabled(t *testing.T) {
result, err := ResolveSelection(
[]string{"beta", "alpha", "gamma"},
SelectionInput{
DefaultEnabled: true,
Disabled: []string{"beta"},
},
)
if err != nil {
t.Fatalf("ResolveSelection() error = %v", err)
}
if !slices.Equal(result.EnabledNames, []string{"alpha", "gamma"}) {
t.Fatalf("EnabledNames mismatch: got %v", result.EnabledNames)
}
if !slices.Equal(result.DisabledNames, []string{"beta"}) {
t.Fatalf("DisabledNames mismatch: got %v", result.DisabledNames)
}
}
func TestResolveSelection_EnabledListOnly(t *testing.T) {
result, err := ResolveSelection(
[]string{"a", "b", "c"},
SelectionInput{
DefaultEnabled: true,
Enabled: []string{"c", "a"},
},
)
if err != nil {
t.Fatalf("ResolveSelection() error = %v", err)
}
if !slices.Equal(result.EnabledNames, []string{"a", "c"}) {
t.Fatalf("EnabledNames mismatch: got %v", result.EnabledNames)
}
if !slices.Equal(result.DisabledNames, []string{"b"}) {
t.Fatalf("DisabledNames mismatch: got %v", result.DisabledNames)
}
}
func TestResolveSelection_DisabledWinsOverlap(t *testing.T) {
result, err := ResolveSelection(
[]string{"a", "b", "c"},
SelectionInput{
Enabled: []string{"a", "b"},
Disabled: []string{"b"},
},
)
if err != nil {
t.Fatalf("ResolveSelection() error = %v", err)
}
if !slices.Equal(result.EnabledNames, []string{"a"}) {
t.Fatalf("EnabledNames mismatch: got %v", result.EnabledNames)
}
if !slices.Equal(result.DisabledNames, []string{"b", "c"}) {
t.Fatalf("DisabledNames mismatch: got %v", result.DisabledNames)
}
}
func TestResolveSelection_UnknownEnabledFails(t *testing.T) {
result, err := ResolveSelection(
[]string{"a"},
SelectionInput{
Enabled: []string{"missing"},
},
)
if err == nil {
t.Fatal("expected error for unknown enabled plugin")
}
if !strings.Contains(err.Error(), "missing") {
t.Fatalf("expected error to mention unknown plugin, got %v", err)
}
if !slices.Equal(result.UnknownEnabled, []string{"missing"}) {
t.Fatalf("UnknownEnabled mismatch: got %v", result.UnknownEnabled)
}
}
func TestResolveSelection_UnknownDisabledWarns(t *testing.T) {
result, err := ResolveSelection(
[]string{"a"},
SelectionInput{
DefaultEnabled: true,
Disabled: []string{"missing"},
},
)
if err != nil {
t.Fatalf("ResolveSelection() error = %v", err)
}
if !slices.Equal(result.EnabledNames, []string{"a"}) {
t.Fatalf("EnabledNames mismatch: got %v", result.EnabledNames)
}
if len(result.DisabledNames) != 0 {
t.Fatalf("DisabledNames mismatch: got %v", result.DisabledNames)
}
if !slices.Equal(result.UnknownDisabled, []string{"missing"}) {
t.Fatalf("UnknownDisabled mismatch: got %v", result.UnknownDisabled)
}
if !hasWarningSubstring(result.Warnings, `unknown disabled plugin "missing" ignored`) {
t.Fatalf("expected unknown disabled warning, got %v", result.Warnings)
}
}
func TestResolveSelection_NormalizationAndDedupe(t *testing.T) {
result, err := ResolveSelection(
[]string{" Alpha ", "beta", "gamma"},
SelectionInput{
Enabled: []string{"ALPHA", " alpha ", "BETA", "beta"},
Disabled: []string{" beta", "BETA", "missing", " MISSING "},
},
)
if err != nil {
t.Fatalf("ResolveSelection() error = %v", err)
}
if !slices.Equal(result.EnabledNames, []string{"alpha"}) {
t.Fatalf("EnabledNames mismatch: got %v", result.EnabledNames)
}
if !slices.Equal(result.DisabledNames, []string{"beta", "gamma"}) {
t.Fatalf("DisabledNames mismatch: got %v", result.DisabledNames)
}
if !slices.Equal(result.UnknownDisabled, []string{"missing"}) {
t.Fatalf("UnknownDisabled mismatch: got %v", result.UnknownDisabled)
}
if !hasWarningSubstring(result.Warnings, `duplicate enabled plugin "alpha" ignored`) {
t.Fatalf("expected duplicate enabled warning for alpha, got %v", result.Warnings)
}
if !hasWarningSubstring(result.Warnings, `duplicate enabled plugin "beta" ignored`) {
t.Fatalf("expected duplicate enabled warning for beta, got %v", result.Warnings)
}
if !hasWarningSubstring(result.Warnings, `duplicate disabled plugin "beta" ignored`) {
t.Fatalf("expected duplicate disabled warning for beta, got %v", result.Warnings)
}
if !hasWarningSubstring(result.Warnings, `duplicate disabled plugin "missing" ignored`) {
t.Fatalf("expected duplicate disabled warning for missing, got %v", result.Warnings)
}
if !hasWarningSubstring(result.Warnings, `unknown disabled plugin "missing" ignored`) {
t.Fatalf("expected unknown disabled warning, got %v", result.Warnings)
}
}
func hasWarningSubstring(warnings []string, sub string) bool {
for _, warning := range warnings {
if strings.Contains(warning, sub) {
return true
}
}
return false
}