Files
goclaw/internal/tools/registry_test.go
T
Plateau Nguyen db98e477d4 fix(gateway): re-apply tool rate limiter after system_configs overlay (#1111)
setupToolRegistry creates the rate limiter from cfg.Tools.RateLimitPerHour
during early bootstrap (gateway.go line 137). System_configs DB overlay
runs ~50 lines later via cfg.ApplySystemConfigs (line 191), so any DB
override of tools.rate_limit_per_hour was silently lost - the limiter
object was already initialised from the JSON5 default.

Symptom in production: editing tools.rate_limit_per_hour via system_configs
table or the config HTTP API had no effect on running gateways. Operators
had to inject a config.json file to change the value, defeating the
DB-as-source-of-truth pattern that other tunables rely on.

Re-apply the limiter after ApplySystemConfigs runs. Safe ordering: server
has not started yet, no in-flight tool calls, and SetRateLimiter is a
plain field assignment with no shutdown cost on the discarded limiter.
nil case (rate_limit_per_hour <= 0) also handled so DB writes can disable
the limiter without restart.

Tests: new TestRegistry_SetRateLimiter_ReplacesPriorLimiter covers both
the replace-with-higher-limit path and the nil-disables path. All
existing rate limiter tests still pass.

Note: the same ordering pattern likely affects other config consumed
inside setupToolRegistry (tools.scrub_credentials, MCP server wiring).
Out of scope for this hotfix - they need a deeper restructure.
2026-06-21 16:58:15 +07:00

416 lines
12 KiB
Go

package tools
import (
"context"
"strings"
"testing"
)
// mockTool is a minimal tool for testing the registry.
type mockTool struct {
name string
execFn func(ctx context.Context, args map[string]any) *Result
}
func (m *mockTool) Name() string { return m.name }
func (m *mockTool) Description() string { return "mock tool" }
func (m *mockTool) Parameters() map[string]any {
return map[string]any{"type": "object", "properties": map[string]any{}}
}
func (m *mockTool) Execute(ctx context.Context, args map[string]any) *Result {
if m.execFn != nil {
return m.execFn(ctx, args)
}
return NewResult("ok")
}
func TestRegistry_RegisterAndGet(t *testing.T) {
reg := NewRegistry()
tool := &mockTool{name: "test_tool"}
reg.Register(tool)
got, ok := reg.Get("test_tool")
if !ok {
t.Fatal("tool not found")
}
if got.Name() != "test_tool" {
t.Errorf("expected test_tool, got %s", got.Name())
}
}
func TestRegistry_GetUnknown(t *testing.T) {
reg := NewRegistry()
_, ok := reg.Get("nonexistent")
if ok {
t.Error("expected tool not found")
}
}
func TestRegistry_Unregister(t *testing.T) {
reg := NewRegistry()
reg.Register(&mockTool{name: "t1"})
reg.Unregister("t1")
if _, ok := reg.Get("t1"); ok {
t.Error("tool should be unregistered")
}
}
func TestRegistry_Count(t *testing.T) {
reg := NewRegistry()
reg.Register(&mockTool{name: "t1"})
reg.Register(&mockTool{name: "t2"})
if reg.Count() != 2 {
t.Errorf("expected 2, got %d", reg.Count())
}
}
func TestRegistry_ExecuteUnknownTool(t *testing.T) {
reg := NewRegistry()
result := reg.Execute(context.Background(), "missing", nil)
if !result.IsError {
t.Error("expected error result for unknown tool")
}
}
func TestRegistry_ExecuteWithContext_InjectsContextValues(t *testing.T) {
reg := NewRegistry()
var gotChannel, gotChatID, gotPeerKind, gotSandboxKey string
var gotAsyncCB AsyncCallback
reg.Register(&mockTool{
name: "ctx_tool",
execFn: func(ctx context.Context, args map[string]any) *Result {
gotChannel = ToolChannelFromCtx(ctx)
gotChatID = ToolChatIDFromCtx(ctx)
gotPeerKind = ToolPeerKindFromCtx(ctx)
gotSandboxKey = ToolSandboxKeyFromCtx(ctx)
gotAsyncCB = ToolAsyncCBFromCtx(ctx)
return NewResult("done")
},
})
called := false
cb := AsyncCallback(func(ctx context.Context, result *Result) { called = true })
reg.ExecuteWithContext(context.Background(), "ctx_tool", nil,
"telegram", "chat-1", "group", "sess-1", cb)
if gotChannel != "telegram" {
t.Errorf("channel: expected telegram, got %q", gotChannel)
}
if gotChatID != "chat-1" {
t.Errorf("chatID: expected chat-1, got %q", gotChatID)
}
if gotPeerKind != "group" {
t.Errorf("peerKind: expected group, got %q", gotPeerKind)
}
if gotSandboxKey != "sess-1" {
t.Errorf("sandboxKey: expected sess-1, got %q", gotSandboxKey)
}
if gotAsyncCB == nil {
t.Error("asyncCB should not be nil")
}
gotAsyncCB(context.Background(), nil)
if !called {
t.Error("asyncCB was not properly propagated")
}
}
func TestRegistry_ExecuteWithContext_ScrubsCredentials(t *testing.T) {
reg := NewRegistry()
reg.Register(&mockTool{
name: "leaky_tool",
execFn: func(ctx context.Context, args map[string]any) *Result {
return &Result{
ForLLM: "key is sk-abcdefghijklmnopqrstuvwxyz1234567890",
ForUser: "token: ghp_ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghij",
}
},
})
result := reg.Execute(context.Background(), "leaky_tool", nil)
if result.ForLLM == "key is sk-abcdefghijklmnopqrstuvwxyz1234567890" {
t.Error("ForLLM should have credentials scrubbed")
}
if result.ForUser == "token: ghp_ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghij" {
t.Error("ForUser should have credentials scrubbed")
}
}
func TestRegistry_ExecuteWithContext_RateLimiting(t *testing.T) {
reg := NewRegistry()
reg.SetRateLimiter(NewToolRateLimiter(2))
reg.Register(&mockTool{name: "rl_tool"})
// First 2 calls allowed
for i := range 2 {
result := reg.ExecuteWithContext(context.Background(), "rl_tool", nil,
"", "", "", "session-1", nil)
if result.IsError {
t.Errorf("call %d should succeed: %s", i, result.ForLLM)
}
}
// 3rd call blocked
result := reg.ExecuteWithContext(context.Background(), "rl_tool", nil,
"", "", "", "session-1", nil)
if !result.IsError {
t.Error("3rd call should be rate-limited")
}
// Different session key allowed
result = reg.ExecuteWithContext(context.Background(), "rl_tool", nil,
"", "", "", "session-2", nil)
if result.IsError {
t.Error("different session should be allowed")
}
}
// SetRateLimiter must be re-callable so that startup code can swap the limiter
// after DB-overlaid config is applied (cmd/gateway.go re-applies once
// system_configs has overlaid the JSON5 default).
func TestRegistry_SetRateLimiter_ReplacesPriorLimiter(t *testing.T) {
reg := NewRegistry()
reg.SetRateLimiter(NewToolRateLimiter(1)) // simulate JSON5 default
reg.Register(&mockTool{name: "tool"})
// Replace with higher limit (simulates DB overlay = 5)
reg.SetRateLimiter(NewToolRateLimiter(5))
for i := range 5 {
result := reg.ExecuteWithContext(context.Background(), "tool", nil,
"", "", "", "session-replace", nil)
if result.IsError {
t.Errorf("call %d should succeed under new 5/h limit: %s", i, result.ForLLM)
}
}
result := reg.ExecuteWithContext(context.Background(), "tool", nil,
"", "", "", "session-replace", nil)
if !result.IsError {
t.Error("6th call should hit the 5/h limit")
}
// Disable rate limiting via nil — verifies the gateway path that disables
// the limiter when cfg.Tools.RateLimitPerHour <= 0.
reg.SetRateLimiter(nil)
for i := range 10 {
result := reg.ExecuteWithContext(context.Background(), "tool", nil,
"", "", "", "session-replace", nil)
if result.IsError {
t.Errorf("call %d after nil limiter should be unbounded", i)
}
}
}
func TestRegistry_ExecuteWithContext_NoRateLimitWithoutSessionKey(t *testing.T) {
reg := NewRegistry()
reg.SetRateLimiter(NewToolRateLimiter(1))
reg.Register(&mockTool{name: "tool"})
// Without sessionKey, rate limiting is skipped
for i := range 5 {
result := reg.ExecuteWithContext(context.Background(), "tool", nil,
"", "", "", "", nil)
if result.IsError {
t.Errorf("call %d should succeed (no sessionKey): %s", i, result.ForLLM)
}
}
}
func TestRegistry_ExecuteWithContext_EmptyContextValues(t *testing.T) {
reg := NewRegistry()
var gotChannel, gotSandboxKey string
reg.Register(&mockTool{
name: "empty_ctx",
execFn: func(ctx context.Context, args map[string]any) *Result {
gotChannel = ToolChannelFromCtx(ctx)
gotSandboxKey = ToolSandboxKeyFromCtx(ctx)
return NewResult("ok")
},
})
// Empty strings should NOT be injected into context
reg.ExecuteWithContext(context.Background(), "empty_ctx", nil,
"", "", "", "", nil)
if gotChannel != "" {
t.Errorf("empty channel should not be injected, got %q", gotChannel)
}
if gotSandboxKey != "" {
t.Errorf("empty sandboxKey should not be injected, got %q", gotSandboxKey)
}
}
// --- Panic recovery tests ---
func TestRegistry_ExecuteWithContext_PanicRecovery(t *testing.T) {
reg := NewRegistry()
reg.Register(&mockTool{
name: "panicking_tool",
execFn: func(ctx context.Context, args map[string]any) *Result {
panic("unexpected nil pointer")
},
})
// Without panic recovery this would crash the test process.
result := reg.ExecuteWithContext(context.Background(), "panicking_tool", nil,
"", "", "", "", nil)
if !result.IsError {
t.Fatal("expected error result from panicking tool")
}
if !strings.Contains(result.ForLLM, "panicked") {
t.Errorf("error should mention panic, got: %s", result.ForLLM)
}
if !strings.Contains(result.ForLLM, "panicking_tool") {
t.Errorf("error should mention tool name, got: %s", result.ForLLM)
}
}
func TestRegistry_Execute_PanicRecovery(t *testing.T) {
reg := NewRegistry()
reg.Register(&mockTool{
name: "panic_via_execute",
execFn: func(ctx context.Context, args map[string]any) *Result {
var s []int
_ = s[10] // index out of range panic
return NewResult("unreachable")
},
})
result := reg.Execute(context.Background(), "panic_via_execute", nil)
if !result.IsError {
t.Fatal("expected error result from panicking tool")
}
}
func TestRegistry_ExecuteWithContext_PanicRecovery_Concurrent(t *testing.T) {
reg := NewRegistry()
reg.Register(&mockTool{
name: "concurrent_panic",
execFn: func(ctx context.Context, args map[string]any) *Result {
panic("boom")
},
})
const goroutines = 20
results := make(chan *Result, goroutines)
for range goroutines {
go func() {
results <- reg.ExecuteWithContext(context.Background(), "concurrent_panic", nil,
"", "", "", "session-1", nil)
}()
}
for range goroutines {
r := <-results
if !r.IsError {
t.Error("expected error result from panicking tool in concurrent execution")
}
}
}
// --- TryActivateDeferred / SetDeferredActivator tests ---
func TestRegistry_TryActivateDeferred_NoActivator(t *testing.T) {
reg := NewRegistry()
// No activator set — must return false without panicking.
if reg.TryActivateDeferred("any_tool") {
t.Error("expected false when no activator is set")
}
}
func TestRegistry_TryActivateDeferred_ActivatorCalledWithCorrectName(t *testing.T) {
reg := NewRegistry()
var called string
reg.SetDeferredActivator(func(name string) bool {
called = name
return false
})
reg.TryActivateDeferred("mcp_foo__bar")
if called != "mcp_foo__bar" {
t.Errorf("activator called with %q, want %q", called, "mcp_foo__bar")
}
}
func TestRegistry_TryActivateDeferred_ReturnsTrueWhenActivated(t *testing.T) {
reg := NewRegistry()
reg.SetDeferredActivator(func(name string) bool {
// Simulate activating: register the tool in the registry
if name == "mcp_svc__get_data" {
reg.Register(&mockTool{name: name})
return true
}
return false
})
if !reg.TryActivateDeferred("mcp_svc__get_data") {
t.Error("expected true for activatable tool")
}
if _, ok := reg.Get("mcp_svc__get_data"); !ok {
t.Error("tool should be in registry after activation")
}
}
func TestRegistry_TryActivateDeferred_ReturnsFalseForUnknown(t *testing.T) {
reg := NewRegistry()
reg.SetDeferredActivator(func(name string) bool { return false })
if reg.TryActivateDeferred("nonexistent_tool") {
t.Error("expected false for unknown tool")
}
if _, ok := reg.Get("nonexistent_tool"); ok {
t.Error("tool should not appear in registry")
}
}
func TestRegistry_SetDeferredActivator_OverwritesPrevious(t *testing.T) {
reg := NewRegistry()
calls := 0
reg.SetDeferredActivator(func(name string) bool { calls++; return false })
reg.SetDeferredActivator(func(name string) bool { calls += 10; return false })
reg.TryActivateDeferred("any")
if calls != 10 {
t.Errorf("expected only the second activator to run (calls=10), got %d", calls)
}
}
func TestRegistry_TryActivateDeferred_Concurrent(t *testing.T) {
// Verify no data race when many goroutines call TryActivateDeferred simultaneously.
reg := NewRegistry()
reg.SetDeferredActivator(func(name string) bool {
reg.Register(&mockTool{name: name})
return true
})
const goroutines = 50
done := make(chan struct{}, goroutines)
for i := range goroutines {
toolName := "mcp_server__tool"
if i%2 == 0 {
toolName = "mcp_other__tool"
}
go func(n string) {
reg.TryActivateDeferred(n)
done <- struct{}{}
}(toolName)
}
for range goroutines {
<-done
}
}
func TestRegistry_TryActivateDeferred_NilActivatorAfterSet(t *testing.T) {
reg := NewRegistry()
reg.SetDeferredActivator(func(name string) bool { return true })
// Overwrite with nil — should behave as if no activator.
reg.SetDeferredActivator(nil)
if reg.TryActivateDeferred("any") {
t.Error("expected false after setting nil activator")
}
}