fix(bus,cron): panic recovery + deadlock fix + unit tests for 9 packages

Bug fixes:
- bus: add panic recovery in Broadcast() — panicking subscriber no longer
  crashes entire event bus goroutine
- cron: fix deadlock in RunJob() — recordRun called while holding mutex,
  extracted recordRunLocked for callers already holding lock
- cron: add run log recording to executeJobByID — automatic scheduler
  was not populating run log, only manual RunJob did

New tests (~170 cases across 9 packages):
- crypto: roundtrip, key derivation (hex/b64/raw), nonce uniqueness,
  backward compat, wrong key, corruption
- permissions: role hierarchy, RoleFromScopes, CanAccess, scope precedence
- bus: pub/sub delivery, buffer full, panic recovery, concurrent safety
- providers/retry: IsRetryableError, backoff, jitter, Retry-After, context
  cancellation, hook callback
- sessions: idempotency, concurrent writes, defensive copy, save/load
  roundtrip, metadata accumulation
- scheduler: draining, drop policies, adaptive throttle, stale completion,
  lane concurrency, debounce, interrupt mode
- config: JSON5 parsing, env overrides, FlexibleStringSlice, owner IDs
- cron: schedule validation, computeNextRun, CRUD, job execution, failure
  tracking, persistence roundtrip
- agent: history sanitization edge cases — all-tool history, partial results,
  cross-turn dedup, orphaned tools, 1000-msg performance
This commit is contained in:
viettranx committed 2026-03-28 18:15:46 +07:00
1 parent a7800604b6
commit 90b7396d74
11 files changed
+2915 -4

No files matched your search

+264
View File
@@ -0,0 +1,264 @@
package agent
import (
"strings"
"testing"
"github.com/nextlevelbuilder/goclaw/internal/providers"
)
// --- sanitizeHistory: all-tool-message history ---
// After aggressive truncation, history could be ALL tool messages.
// Should return nil (no useful messages), not crash.
func TestSanitizeHistory_AllToolMessages(t *testing.T) {
msgs := []providers.Message{
{Role: "tool", Content: "result1", ToolCallID: "tc1"},
{Role: "tool", Content: "result2", ToolCallID: "tc2"},
{Role: "tool", Content: "result3", ToolCallID: "tc3"},
}
result, dropped := sanitizeHistory(msgs)
if result != nil {
t.Fatalf("all-tool history should return nil, got %d messages", len(result))
}
if dropped != 3 {
t.Fatalf("expected 3 dropped, got %d", dropped)
}
}
// --- sanitizeHistory: single user message (no tools) ---
func TestSanitizeHistory_SingleUserMessage(t *testing.T) {
msgs := []providers.Message{
{Role: "user", Content: "hello"},
}
result, dropped := sanitizeHistory(msgs)
if len(result) != 1 {
t.Fatalf("expected 1 message, got %d", len(result))
}
if dropped != 0 {
t.Fatalf("expected 0 dropped, got %d", dropped)
}
}
// --- sanitizeHistory: assistant with tool_calls but ALL results missing ---
// Every tool_call should get a synthesized "[Tool result missing]" placeholder.
func TestSanitizeHistory_AllToolResultsMissing(t *testing.T) {
msgs := []providers.Message{
{Role: "user", Content: "do something"},
{Role: "assistant", Content: "", ToolCalls: []providers.ToolCall{
{ID: "tc1", Name: "read_file", Arguments: map[string]any{"path": "a.go"}},
{ID: "tc2", Name: "read_file", Arguments: map[string]any{"path": "b.go"}},
{ID: "tc3", Name: "read_file", Arguments: map[string]any{"path": "c.go"}},
}},
// No tool results at all — all missing
{Role: "user", Content: "next message"},
}
result, dropped := sanitizeHistory(msgs)
// Should have: user + assistant + 3 synthesized tool results + user
if len(result) != 6 {
t.Fatalf("expected 6 messages (user + assistant + 3 synth + user), got %d", len(result))
}
if dropped != 3 {
t.Fatalf("expected 3 synthesized (counted as dropped), got %d", dropped)
}
// Verify synthesized messages
for i := 2; i <= 4; i++ {
if result[i].Role != "tool" {
t.Fatalf("message %d: expected role 'tool', got %q", i, result[i].Role)
}
if !strings.Contains(result[i].Content, "missing") {
t.Fatalf("message %d: expected 'missing' placeholder, got %q", i, result[i].Content)
}
}
}
// --- sanitizeHistory: interleaved tool results from wrong assistant ---
// tool_result IDs that don't match the preceding assistant's tool_calls
// should be dropped.
func TestSanitizeHistory_ToolResultsFromWrongAssistant(t *testing.T) {
msgs := []providers.Message{
{Role: "user", Content: "first"},
{Role: "assistant", Content: "", ToolCalls: []providers.ToolCall{
{ID: "tc_a", Name: "read_file", Arguments: map[string]any{}},
}},
// Tool result with wrong ID (from a different assistant turn)
{Role: "tool", Content: "wrong result", ToolCallID: "tc_z"},
// Correct result
// (missing — will be synthesized)
{Role: "user", Content: "next"},
}
result, dropped := sanitizeHistory(msgs)
// tc_z dropped, tc_a synthesized
if dropped != 2 {
t.Fatalf("expected 2 dropped (1 mismatched + 1 synthesized), got %d", dropped)
}
// Should have: user + assistant + synth(tc_a) + user
if len(result) != 4 {
t.Fatalf("expected 4 messages, got %d", len(result))
}
if result[2].ToolCallID != "tc_a" {
t.Fatalf("synthesized result should have tc_a, got %q", result[2].ToolCallID)
}
}
// --- sanitizeHistory: multiple tool_calls with partial results ---
// 3 tool_calls but only 1 result provided → 2 synthesized
func TestSanitizeHistory_PartialToolResults(t *testing.T) {
msgs := []providers.Message{
{Role: "user", Content: "go"},
{Role: "assistant", Content: "", ToolCalls: []providers.ToolCall{
{ID: "tc1", Name: "read_file", Arguments: map[string]any{"path": "a"}},
{ID: "tc2", Name: "write_file", Arguments: map[string]any{"path": "b"}},
{ID: "tc3", Name: "exec", Arguments: map[string]any{"cmd": "ls"}},
}},
{Role: "tool", Content: "content of a", ToolCallID: "tc1"},
// tc2 and tc3 missing
{Role: "user", Content: "done"},
}
result, dropped := sanitizeHistory(msgs)
// user + assistant + tc1_result + tc2_synth + tc3_synth + user
if len(result) != 6 {
t.Fatalf("expected 6 messages, got %d", len(result))
}
if dropped != 2 {
t.Fatalf("expected 2 synthesized, got %d", dropped)
}
// Verify order: real result first, then synthesized
if result[2].ToolCallID != "tc1" || strings.Contains(result[2].Content, "missing") {
t.Fatal("tc1 should be the real result, not synthesized")
}
if result[3].ToolCallID != "tc2" || !strings.Contains(result[3].Content, "missing") {
t.Fatal("tc2 should be synthesized")
}
if result[4].ToolCallID != "tc3" || !strings.Contains(result[4].Content, "missing") {
t.Fatal("tc3 should be synthesized")
}
}
// --- sanitizeHistory: tool message between two user messages (orphaned mid-history) ---
func TestSanitizeHistory_OrphanedToolBetweenUsers(t *testing.T) {
msgs := []providers.Message{
{Role: "user", Content: "first"},
{Role: "assistant", Content: "response"},
{Role: "tool", Content: "orphaned", ToolCallID: "tc_orphan"}, // no preceding tool_calls
{Role: "user", Content: "second"},
}
result, dropped := sanitizeHistory(msgs)
if dropped != 1 {
t.Fatalf("expected 1 dropped orphan, got %d", dropped)
}
// user + assistant + user (orphan dropped)
if len(result) != 3 {
t.Fatalf("expected 3 messages, got %d", len(result))
}
}
// --- sanitizeHistory: dedup with 2 identical IDs across 2 turns ---
// Turn 1: ID "dup" → kept as-is (first occurrence)
// Turn 2: ID "dup" → rewritten to "dup_dedup_0", tool result matched via idQueue
func TestSanitizeHistory_DedupAcrossTwoTurns(t *testing.T) {
msgs := []providers.Message{
// Turn 1
{Role: "user", Content: "1"},
{Role: "assistant", ToolCalls: []providers.ToolCall{{ID: "dup", Name: "read_file", Arguments: map[string]any{}}}},
{Role: "tool", Content: "r1", ToolCallID: "dup"},
// Turn 2 — same ID
{Role: "user", Content: "2"},
{Role: "assistant", ToolCalls: []providers.ToolCall{{ID: "dup", Name: "read_file", Arguments: map[string]any{}}}},
{Role: "tool", Content: "r2", ToolCallID: "dup"},
}
result, dropped := sanitizeHistory(msgs)
// All 6 messages should be preserved (dedup rewrites turn 2's ID)
if len(result) != 6 {
t.Fatalf("expected 6 messages preserved via dedup, got %d", len(result))
}
if dropped != 0 {
t.Fatalf("expected 0 dropped, got %d", dropped)
}
// Turn 1 assistant should keep original ID
if result[1].ToolCalls[0].ID != "dup" {
t.Fatalf("turn 1 ID should be 'dup', got %q", result[1].ToolCalls[0].ID)
}
// Turn 2 assistant should have rewritten ID
if result[4].ToolCalls[0].ID == "dup" {
t.Fatal("turn 2 ID should be rewritten, still 'dup'")
}
// Turn 2 tool result should match the rewritten ID
if result[5].ToolCallID != result[4].ToolCalls[0].ID {
t.Fatalf("turn 2 tool result ID %q should match assistant ID %q",
result[5].ToolCallID, result[4].ToolCalls[0].ID)
}
}
// --- sanitizeHistory: performance with large history ---
// Verify sanitization doesn't degrade catastrophically with many messages.
func TestSanitizeHistory_LargeHistory_Performance(t *testing.T) {
// Build 1000-message history with proper tool pairing
msgs := make([]providers.Message, 0, 1000)
for i := 0; i < 250; i++ {
tcID := "tc_" + strings.Repeat("x", 5) + "_" + string(rune('a'+i%26)) + string(rune('0'+i%10))
msgs = append(msgs,
providers.Message{Role: "user", Content: "question " + tcID},
providers.Message{Role: "assistant", ToolCalls: []providers.ToolCall{
{ID: tcID, Name: "read_file", Arguments: map[string]any{"path": "file.go"}},
}},
providers.Message{Role: "tool", Content: "result", ToolCallID: tcID},
providers.Message{Role: "assistant", Content: "answer"},
)
}
result, dropped := sanitizeHistory(msgs)
if dropped != 0 {
t.Fatalf("well-formed history should have 0 drops, got %d", dropped)
}
if len(result) != 1000 {
t.Fatalf("expected 1000 messages, got %d", len(result))
}
}
// --- sanitizeHistory: empty ToolCallID in tool result ---
func TestSanitizeHistory_EmptyToolCallID(t *testing.T) {
msgs := []providers.Message{
{Role: "user", Content: "go"},
{Role: "assistant", ToolCalls: []providers.ToolCall{
{ID: "tc1", Name: "read_file", Arguments: map[string]any{}},
}},
{Role: "tool", Content: "result", ToolCallID: ""}, // empty ID — won't match
}
result, dropped := sanitizeHistory(msgs)
// Empty ID tool result should be dropped (mismatched)
// tc1 should get a synthesized result
if dropped != 2 {
t.Fatalf("expected 2 dropped (1 mismatched + 1 synthesized), got %d", dropped)
}
// user + assistant + synth(tc1)
if len(result) != 3 {
t.Fatalf("expected 3 messages, got %d", len(result))
}
}
+17 -2
View File
@@ -2,6 +2,8 @@ package bus
import (
"context"
"fmt"
"log/slog"
"sync"
)
@@ -113,11 +115,24 @@ func (mb *MessageBus) Unsubscribe(id string) {
}
// Broadcast sends an event to all subscribers (non-blocking per subscriber).
// Panicking handlers are caught and logged to prevent one bad subscriber
// from crashing the entire event bus.
func (mb *MessageBus) Broadcast(event Event) {
mb.subMu.RLock()
defer mb.subMu.RUnlock()
for _, handler := range mb.subscribers {
handler(event) // handlers should be non-blocking
for id, handler := range mb.subscribers {
func() {
defer func() {
if r := recover(); r != nil {
slog.Error("bus: subscriber panicked",
"subscriber", id,
"event", event.Name,
"panic", fmt.Sprint(r),
)
}
}()
handler(event)
}()
}
}
+263
View File
@@ -0,0 +1,263 @@
package bus
import (
"context"
"sync"
"sync/atomic"
"testing"
"time"
)
// --- Pub/Sub delivery ---
func TestPublishInbound_ConsumeInbound(t *testing.T) {
mb := New()
defer mb.Close()
msg := InboundMessage{Channel: "telegram", Content: "hello"}
mb.PublishInbound(msg)
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
got, ok := mb.ConsumeInbound(ctx)
if !ok {
t.Fatal("expected to consume message")
}
if got.Content != "hello" {
t.Fatalf("content mismatch: got %q, want %q", got.Content, "hello")
}
}
func TestPublishOutbound_SubscribeOutbound(t *testing.T) {
mb := New()
defer mb.Close()
msg := OutboundMessage{Channel: "discord", Content: "world"}
mb.PublishOutbound(msg)
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
got, ok := mb.SubscribeOutbound(ctx)
if !ok {
t.Fatal("expected to receive message")
}
if got.Content != "world" {
t.Fatalf("content mismatch: got %q, want %q", got.Content, "world")
}
}
// --- Context cancellation ---
func TestConsumeInbound_ContextCancelled(t *testing.T) {
mb := New()
defer mb.Close()
ctx, cancel := context.WithCancel(context.Background())
cancel() // cancel immediately
_, ok := mb.ConsumeInbound(ctx)
if ok {
t.Fatal("expected false on cancelled context")
}
}
// --- TryPublish non-blocking ---
func TestTryPublishInbound_BufferFull(t *testing.T) {
mb := &MessageBus{
inbound: make(chan InboundMessage, 1), // tiny buffer
outbound: make(chan OutboundMessage, 1),
handlers: make(map[string]MessageHandler),
subscribers: make(map[string]EventHandler),
}
// First message fits
if !mb.TryPublishInbound(InboundMessage{Content: "1"}) {
t.Fatal("first message should succeed")
}
// Second message should be dropped (buffer full)
if mb.TryPublishInbound(InboundMessage{Content: "2"}) {
t.Fatal("second message should fail (buffer full)")
}
}
func TestTryPublishOutbound_BufferFull(t *testing.T) {
mb := &MessageBus{
inbound: make(chan InboundMessage, 1),
outbound: make(chan OutboundMessage, 1),
handlers: make(map[string]MessageHandler),
subscribers: make(map[string]EventHandler),
}
if !mb.TryPublishOutbound(OutboundMessage{Content: "1"}) {
t.Fatal("first message should succeed")
}
if mb.TryPublishOutbound(OutboundMessage{Content: "2"}) {
t.Fatal("second message should fail (buffer full)")
}
}
// --- Broadcast delivery ---
func TestBroadcast_DeliveredToAllSubscribers(t *testing.T) {
mb := New()
defer mb.Close()
var count atomic.Int32
mb.Subscribe("sub1", func(e Event) { count.Add(1) })
mb.Subscribe("sub2", func(e Event) { count.Add(1) })
mb.Subscribe("sub3", func(e Event) { count.Add(1) })
mb.Broadcast(Event{Name: "test"})
if got := count.Load(); got != 3 {
t.Fatalf("expected 3 deliveries, got %d", got)
}
}
// --- Broadcast panic recovery: panicking handler must NOT crash other subscribers ---
func TestBroadcast_PanickingHandler_DoesNotCrashBus(t *testing.T) {
mb := New()
defer mb.Close()
var delivered atomic.Int32
mb.Subscribe("panicker", func(e Event) {
panic("subscriber exploded")
})
mb.Subscribe("normal", func(e Event) {
delivered.Add(1)
})
// This must not panic — the bus should catch the panicking handler
mb.Broadcast(Event{Name: "test"})
// The normal handler may or may not be called depending on iteration order,
// but the important thing is we didn't crash.
// Broadcast a second time to verify bus is still functional.
mb.Broadcast(Event{Name: "test2"})
// After two broadcasts, normal handler should have been called at least once
if got := delivered.Load(); got == 0 {
t.Fatal("normal handler should have been called at least once after two broadcasts")
}
}
// --- Subscribe / Unsubscribe ---
func TestUnsubscribe_StopsDelivery(t *testing.T) {
mb := New()
defer mb.Close()
var count atomic.Int32
mb.Subscribe("temp", func(e Event) { count.Add(1) })
mb.Broadcast(Event{Name: "before"})
if count.Load() != 1 {
t.Fatal("expected delivery before unsubscribe")
}
mb.Unsubscribe("temp")
mb.Broadcast(Event{Name: "after"})
if count.Load() != 1 {
t.Fatal("expected no delivery after unsubscribe")
}
}
func TestSubscribe_OverwritesPrevious(t *testing.T) {
mb := New()
defer mb.Close()
var first, second atomic.Int32
mb.Subscribe("id1", func(e Event) { first.Add(1) })
mb.Subscribe("id1", func(e Event) { second.Add(1) }) // overwrite
mb.Broadcast(Event{Name: "test"})
if first.Load() != 0 {
t.Fatal("first handler should have been replaced")
}
if second.Load() != 1 {
t.Fatal("second handler should have been called")
}
}
// --- Handler registration ---
func TestRegisterHandler_GetHandler(t *testing.T) {
mb := New()
defer mb.Close()
called := false
mb.RegisterHandler("telegram", func(msg InboundMessage) error {
called = true
return nil
})
handler, ok := mb.GetHandler("telegram")
if !ok {
t.Fatal("expected handler to be registered")
}
_ = handler(InboundMessage{})
if !called {
t.Fatal("expected handler to be called")
}
_, ok = mb.GetHandler("nonexistent")
if ok {
t.Fatal("expected no handler for unregistered channel")
}
}
// --- Concurrent safety ---
func TestBroadcast_ConcurrentSubscribeUnsubscribe(t *testing.T) {
mb := New()
defer mb.Close()
var wg sync.WaitGroup
done := make(chan struct{})
// Broadcast in a goroutine
wg.Add(1)
go func() {
defer wg.Done()
for {
select {
case <-done:
return
default:
mb.Broadcast(Event{Name: "concurrent"})
}
}
}()
// Subscribe/unsubscribe rapidly
for i := 0; i < 100; i++ {
mb.Subscribe("rapid", func(e Event) {})
mb.Unsubscribe("rapid")
}
close(done)
wg.Wait()
// No panic = success
}
func TestPublishInbound_ConcurrentProducers(t *testing.T) {
mb := New()
defer mb.Close()
const n = 100
var wg sync.WaitGroup
wg.Add(n)
for i := 0; i < n; i++ {
go func() {
defer wg.Done()
mb.TryPublishInbound(InboundMessage{Content: "msg"})
}()
}
wg.Wait()
// No panic = success
}
+201
View File
@@ -0,0 +1,201 @@
package config
import (
"encoding/json"
"os"
"path/filepath"
"testing"
)
// --- Default ---
func TestDefault_SensibleDefaults(t *testing.T) {
cfg := Default()
if cfg.Gateway.Port != 18790 {
t.Fatalf("default port: got %d, want 18790", cfg.Gateway.Port)
}
if cfg.Gateway.RateLimitRPM != 20 {
t.Fatalf("default rate limit: got %d, want 20", cfg.Gateway.RateLimitRPM)
}
if cfg.Agents.Defaults.Provider != "anthropic" {
t.Fatalf("default provider: got %q", cfg.Agents.Defaults.Provider)
}
if cfg.Agents.Defaults.MaxToolIterations != DefaultMaxIterations {
t.Fatalf("default max iterations: got %d", cfg.Agents.Defaults.MaxToolIterations)
}
if cfg.Tools.Web.DuckDuckGo.MaxResults != 5 {
t.Fatalf("default ddg max results: got %d", cfg.Tools.Web.DuckDuckGo.MaxResults)
}
}
// --- Load with missing file → uses defaults ---
func TestLoad_MissingFile_UsesDefaults(t *testing.T) {
cfg, err := Load("/nonexistent/path/config.json")
if err != nil {
t.Fatalf("missing file should not error: %v", err)
}
if cfg.Gateway.Port != 18790 {
t.Fatalf("expected default port, got %d", cfg.Gateway.Port)
}
}
// --- Load with valid JSON5 ---
func TestLoad_ValidJSON5(t *testing.T) {
dir := t.TempDir()
cfgPath := filepath.Join(dir, "config.json5")
// JSON5: comments and trailing commas allowed
content := `{
// custom port
"gateway": {
"port": 9999,
"rate_limit_rpm": 100,
},
}`
os.WriteFile(cfgPath, []byte(content), 0644)
cfg, err := Load(cfgPath)
if err != nil {
t.Fatalf("load error: %v", err)
}
if cfg.Gateway.Port != 9999 {
t.Fatalf("port: got %d, want 9999", cfg.Gateway.Port)
}
if cfg.Gateway.RateLimitRPM != 100 {
t.Fatalf("rate limit: got %d, want 100", cfg.Gateway.RateLimitRPM)
}
// Unset fields should retain defaults
if cfg.Agents.Defaults.Provider != "anthropic" {
t.Fatalf("default provider should be preserved: got %q", cfg.Agents.Defaults.Provider)
}
}
// --- Load with invalid JSON5 → error ---
func TestLoad_InvalidJSON5(t *testing.T) {
dir := t.TempDir()
cfgPath := filepath.Join(dir, "config.json5")
os.WriteFile(cfgPath, []byte(`{invalid json!!!`), 0644)
_, err := Load(cfgPath)
if err == nil {
t.Fatal("expected error for invalid JSON5")
}
}
// --- Env var override precedence ---
func TestLoad_EnvVarOverrides(t *testing.T) {
dir := t.TempDir()
cfgPath := filepath.Join(dir, "config.json5")
os.WriteFile(cfgPath, []byte(`{"gateway":{"port":8080}}`), 0644)
// Env override should win
t.Setenv("GOCLAW_PORT", "7777")
cfg, err := Load(cfgPath)
if err != nil {
t.Fatalf("load error: %v", err)
}
if cfg.Gateway.Port != 7777 {
t.Fatalf("env override: got port %d, want 7777", cfg.Gateway.Port)
}
}
func TestLoad_EnvVarOverrides_InvalidPort(t *testing.T) {
t.Setenv("GOCLAW_PORT", "not-a-number")
cfg, err := Load("/nonexistent/path")
if err != nil {
t.Fatalf("load error: %v", err)
}
// Invalid port should keep default
if cfg.Gateway.Port != 18790 {
t.Fatalf("invalid port env should keep default: got %d", cfg.Gateway.Port)
}
}
// --- Env var for API keys ---
func TestLoad_EnvVarAPIKeys(t *testing.T) {
t.Setenv("GOCLAW_ANTHROPIC_API_KEY", "sk-test-key")
cfg, err := Load("/nonexistent/path")
if err != nil {
t.Fatalf("load error: %v", err)
}
if cfg.Providers.Anthropic.APIKey != "sk-test-key" {
t.Fatalf("anthropic key: got %q", cfg.Providers.Anthropic.APIKey)
}
}
// --- FlexibleStringSlice ---
func TestFlexibleStringSlice_StringArray(t *testing.T) {
var f FlexibleStringSlice
err := json.Unmarshal([]byte(`["a","b","c"]`), &f)
if err != nil {
t.Fatalf("unmarshal error: %v", err)
}
if len(f) != 3 || f[0] != "a" || f[1] != "b" || f[2] != "c" {
t.Fatalf("got %v", f)
}
}
func TestFlexibleStringSlice_MixedArray(t *testing.T) {
var f FlexibleStringSlice
// Numbers and strings mixed
err := json.Unmarshal([]byte(`["user1", 12345, "user2"]`), &f)
if err != nil {
t.Fatalf("unmarshal error: %v", err)
}
if len(f) != 3 || f[0] != "user1" || f[1] != "12345" || f[2] != "user2" {
t.Fatalf("got %v", f)
}
}
func TestFlexibleStringSlice_EmptyArray(t *testing.T) {
var f FlexibleStringSlice
err := json.Unmarshal([]byte(`[]`), &f)
if err != nil {
t.Fatalf("unmarshal error: %v", err)
}
if len(f) != 0 {
t.Fatalf("expected empty, got %v", f)
}
}
// --- Owner IDs parsing ---
func TestLoad_OwnerIDsParsing(t *testing.T) {
t.Setenv("GOCLAW_OWNER_IDS", " alice , bob , charlie ")
cfg, err := Load("/nonexistent/path")
if err != nil {
t.Fatalf("load error: %v", err)
}
if len(cfg.Gateway.OwnerIDs) != 3 {
t.Fatalf("expected 3 owner IDs, got %d: %v", len(cfg.Gateway.OwnerIDs), cfg.Gateway.OwnerIDs)
}
if cfg.Gateway.OwnerIDs[0] != "alice" || cfg.Gateway.OwnerIDs[1] != "bob" || cfg.Gateway.OwnerIDs[2] != "charlie" {
t.Fatalf("owner IDs not trimmed: %v", cfg.Gateway.OwnerIDs)
}
}
func TestLoad_OwnerIDsEmpty(t *testing.T) {
t.Setenv("GOCLAW_OWNER_IDS", "")
cfg, err := Load("/nonexistent/path")
if err != nil {
t.Fatalf("load error: %v", err)
}
// Empty string should not produce [""]
for _, id := range cfg.Gateway.OwnerIDs {
if id == "" {
t.Fatal("empty owner ID should not be included")
}
}
}
+9 -2
View File
@@ -80,8 +80,8 @@ func (cs *Service) RunJob(jobID string, force bool) (bool, string, error) {
break
}
// Record run log
cs.recordRun(jobID, err, result)
// Record run log (already holding cs.mu via defer above)
cs.recordRunLocked(jobID, err, result)
if err != nil {
return true, "", err
@@ -111,7 +111,11 @@ func (cs *Service) GetRunLog(jobID string, limit int) []RunLogEntry {
func (cs *Service) recordRun(jobID string, err error, resultText string) {
cs.mu.Lock()
defer cs.mu.Unlock()
cs.recordRunLocked(jobID, err, resultText)
}
// recordRunLocked appends a run log entry. Must be called with cs.mu held.
func (cs *Service) recordRunLocked(jobID string, err error, resultText string) {
entry := RunLogEntry{
Ts: nowMS(),
JobID: jobID,
@@ -252,6 +256,9 @@ func (cs *Service) executeJobByID(jobID string) {
break
}
// Record run log (already holding cs.mu via defer above)
cs.recordRunLocked(jobID, err, result)
cs.saveUnsafe()
}
+357
View File
@@ -0,0 +1,357 @@
package cron
import (
"fmt"
"os"
"path/filepath"
"sync/atomic"
"testing"
"time"
)
// --- Schedule validation ---
func TestValidateSchedule(t *testing.T) {
cs := NewService("", nil)
tests := []struct {
name string
sched Schedule
wantErr bool
}{
{"at_valid", Schedule{Kind: "at", AtMS: ptrInt64(time.Now().Add(time.Hour).UnixMilli())}, false},
{"at_missing_timestamp", Schedule{Kind: "at"}, true},
{"every_valid", Schedule{Kind: "every", EveryMS: ptrInt64(5000)}, false},
{"every_zero_interval", Schedule{Kind: "every", EveryMS: ptrInt64(0)}, true},
{"every_negative_interval", Schedule{Kind: "every", EveryMS: ptrInt64(-1)}, true},
{"every_nil_interval", Schedule{Kind: "every"}, true},
{"cron_valid", Schedule{Kind: "cron", Expr: "*/5 * * * *"}, false},
{"cron_empty_expr", Schedule{Kind: "cron", Expr: ""}, true},
{"cron_invalid_expr", Schedule{Kind: "cron", Expr: "bad cron"}, true},
{"cron_valid_with_tz", Schedule{Kind: "cron", Expr: "0 9 * * *", TZ: "Asia/Saigon"}, false},
{"cron_invalid_tz", Schedule{Kind: "cron", Expr: "0 9 * * *", TZ: "Invalid/Zone"}, true},
{"unknown_kind", Schedule{Kind: "invalid"}, true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
err := cs.validateSchedule(&tt.sched)
if (err != nil) != tt.wantErr {
t.Fatalf("validateSchedule() error = %v, wantErr = %v", err, tt.wantErr)
}
})
}
}
// --- computeNextRun ---
func TestComputeNextRun(t *testing.T) {
cs := NewService("", nil)
now := time.Now().UnixMilli()
t.Run("at_future", func(t *testing.T) {
future := now + 60000
sched := Schedule{Kind: "at", AtMS: &future}
next := cs.computeNextRun(&sched, now)
if next == nil || *next != future {
t.Fatalf("expected %d, got %v", future, next)
}
})
t.Run("at_past", func(t *testing.T) {
past := now - 60000
sched := Schedule{Kind: "at", AtMS: &past}
next := cs.computeNextRun(&sched, now)
if next != nil {
t.Fatalf("past at-schedule should return nil, got %d", *next)
}
})
t.Run("every_5s", func(t *testing.T) {
interval := int64(5000)
sched := Schedule{Kind: "every", EveryMS: &interval}
next := cs.computeNextRun(&sched, now)
if next == nil {
t.Fatal("expected non-nil next")
}
expected := now + 5000
if *next != expected {
t.Fatalf("expected %d, got %d", expected, *next)
}
})
t.Run("every_nil_interval", func(t *testing.T) {
sched := Schedule{Kind: "every"}
next := cs.computeNextRun(&sched, now)
if next != nil {
t.Fatal("nil interval should return nil")
}
})
t.Run("cron_every_minute", func(t *testing.T) {
sched := Schedule{Kind: "cron", Expr: "* * * * *"}
next := cs.computeNextRun(&sched, now)
if next == nil {
t.Fatal("expected non-nil next for every-minute cron")
}
// Should be within next 60 seconds
diff := *next - now
if diff < 0 || diff > 61000 {
t.Fatalf("next run should be within 61s, got diff=%dms", diff)
}
})
t.Run("cron_empty_expr", func(t *testing.T) {
sched := Schedule{Kind: "cron", Expr: ""}
next := cs.computeNextRun(&sched, now)
if next != nil {
t.Fatal("empty expr should return nil")
}
})
t.Run("unknown_kind", func(t *testing.T) {
sched := Schedule{Kind: "wut"}
next := cs.computeNextRun(&sched, now)
if next != nil {
t.Fatal("unknown kind should return nil")
}
})
}
// --- AddJob / RemoveJob / ListJobs / EnableJob ---
func TestService_CRUD(t *testing.T) {
dir := t.TempDir()
storePath := filepath.Join(dir, "cron.json")
cs := NewService(storePath, nil)
// Add a job
interval := int64(60000)
job, err := cs.AddJob("test-job", Schedule{Kind: "every", EveryMS: &interval}, "hello", false, "", "", "agent-1")
if err != nil {
t.Fatalf("AddJob error: %v", err)
}
if job.ID == "" {
t.Fatal("job ID should not be empty")
}
if job.Name != "test-job" {
t.Fatalf("job name: got %q", job.Name)
}
if !job.Enabled {
t.Fatal("new job should be enabled")
}
if job.State.NextRunAtMS == nil {
t.Fatal("new job should have NextRunAtMS set")
}
// List jobs
jobs := cs.ListJobs(false)
if len(jobs) != 1 {
t.Fatalf("expected 1 job, got %d", len(jobs))
}
// Disable job
if err := cs.EnableJob(job.ID, false); err != nil {
t.Fatalf("EnableJob error: %v", err)
}
jobs = cs.ListJobs(false) // excludes disabled
if len(jobs) != 0 {
t.Fatalf("expected 0 enabled jobs, got %d", len(jobs))
}
jobs = cs.ListJobs(true) // includes disabled
if len(jobs) != 1 {
t.Fatalf("expected 1 total job, got %d", len(jobs))
}
// Re-enable
if err := cs.EnableJob(job.ID, true); err != nil {
t.Fatalf("EnableJob error: %v", err)
}
// Remove
if err := cs.RemoveJob(job.ID); err != nil {
t.Fatalf("RemoveJob error: %v", err)
}
jobs = cs.ListJobs(true)
if len(jobs) != 0 {
t.Fatalf("expected 0 jobs after remove, got %d", len(jobs))
}
// Verify persisted
if _, err := os.Stat(storePath); os.IsNotExist(err) {
t.Fatal("store file should exist")
}
}
func TestService_AddJob_InvalidSchedule(t *testing.T) {
cs := NewService("", nil)
_, err := cs.AddJob("bad", Schedule{Kind: "unknown"}, "msg", false, "", "", "")
if err == nil {
t.Fatal("expected error for invalid schedule")
}
}
func TestService_RemoveJob_NotFound(t *testing.T) {
cs := NewService("", nil)
err := cs.RemoveJob("nonexistent")
if err == nil {
t.Fatal("expected error for nonexistent job")
}
}
func TestService_EnableJob_NotFound(t *testing.T) {
cs := NewService("", nil)
err := cs.EnableJob("nonexistent", true)
if err == nil {
t.Fatal("expected error for nonexistent job")
}
}
// --- At-schedule sets DeleteAfterRun ---
func TestService_AddJob_AtSchedule_DeleteAfterRun(t *testing.T) {
cs := NewService("", nil)
future := time.Now().Add(time.Hour).UnixMilli()
job, err := cs.AddJob("one-shot", Schedule{Kind: "at", AtMS: &future}, "run once", false, "", "", "")
if err != nil {
t.Fatalf("AddJob error: %v", err)
}
if !job.DeleteAfterRun {
t.Fatal("at-schedule should set DeleteAfterRun=true")
}
}
// --- Job execution callback ---
func TestService_StartStop_JobExecution(t *testing.T) {
dir := t.TempDir()
storePath := filepath.Join(dir, "cron.json")
var execCount atomic.Int32
handler := func(job *Job) (string, error) {
execCount.Add(1)
return "done", nil
}
cs := NewService(storePath, handler)
// Add a fast-interval job (every 100ms) — but runLoop ticks every 1s
interval := int64(100)
_, err := cs.AddJob("fast", Schedule{Kind: "every", EveryMS: &interval}, "tick", false, "", "", "")
if err != nil {
t.Fatalf("AddJob error: %v", err)
}
if err := cs.Start(); err != nil {
t.Fatalf("Start error: %v", err)
}
// runLoop ticks every 1s, wait enough for at least 1 tick
time.Sleep(1500 * time.Millisecond)
cs.Stop()
count := execCount.Load()
if count == 0 {
t.Fatal("expected at least 1 job execution")
}
}
// --- Handler not set → no panic ---
func TestService_NilHandler_NoPanic(t *testing.T) {
dir := t.TempDir()
storePath := filepath.Join(dir, "cron.json")
cs := NewService(storePath, nil) // no handler
interval := int64(100)
cs.AddJob("no-handler", Schedule{Kind: "every", EveryMS: &interval}, "tick", false, "", "", "")
cs.Start()
time.Sleep(1500 * time.Millisecond) // wait for at least 1 tick
cs.Stop() // should not panic
}
// --- Job failure with retry ---
func TestService_JobFailure_Updates_LastError(t *testing.T) {
dir := t.TempDir()
storePath := filepath.Join(dir, "cron.json")
handler := func(job *Job) (string, error) {
return "", fmt.Errorf("intentional failure")
}
cs := NewService(storePath, handler)
cs.SetRetryConfig(RetryConfig{MaxRetries: 0}) // no retry
interval := int64(100)
job, _ := cs.AddJob("failing", Schedule{Kind: "every", EveryMS: &interval}, "fail", false, "", "", "")
cs.Start()
time.Sleep(1500 * time.Millisecond) // wait for at least 1 tick
cs.Stop()
// Check last error
found, ok := cs.GetJob(job.ID)
if !ok {
t.Fatal("job should exist")
}
if found.State.LastStatus != "error" {
t.Fatalf("expected last status 'error', got %q", found.State.LastStatus)
}
if found.State.LastError == "" {
t.Fatal("expected non-empty last error")
}
}
// --- Persistence: save and reload ---
func TestService_Persistence_Roundtrip(t *testing.T) {
dir := t.TempDir()
storePath := filepath.Join(dir, "cron.json")
cs1 := NewService(storePath, nil)
interval := int64(60000)
cs1.AddJob("persist-test", Schedule{Kind: "every", EveryMS: &interval}, "msg", false, "", "", "agent-1")
// New service should load the persisted job
cs2 := NewService(storePath, nil)
cs2.Start()
defer cs2.Stop()
jobs := cs2.ListJobs(true)
if len(jobs) != 1 {
t.Fatalf("expected 1 persisted job, got %d", len(jobs))
}
if jobs[0].Name != "persist-test" {
t.Fatalf("job name mismatch: got %q", jobs[0].Name)
}
}
// --- Run log ---
func TestService_RunLog_PopulatedByAutoExecution(t *testing.T) {
dir := t.TempDir()
cs := NewService(filepath.Join(dir, "cron.json"), func(job *Job) (string, error) {
return "ok", nil
})
interval := int64(100)
job, _ := cs.AddJob("logger", Schedule{Kind: "every", EveryMS: &interval}, "tick", false, "", "", "")
cs.Start()
time.Sleep(1500 * time.Millisecond)
cs.Stop()
log := cs.GetRunLog(job.ID, 50)
if len(log) == 0 {
t.Fatal("expected at least 1 run log entry from automatic execution")
}
if log[0].Status != "ok" {
t.Fatalf("expected status 'ok', got %q", log[0].Status)
}
}
// --- helpers ---
func ptrInt64(v int64) *int64 { return &v }
+261
View File
@@ -0,0 +1,261 @@
package crypto
import (
"encoding/base64"
"encoding/hex"
"strings"
"testing"
)
// testKey32 is a deterministic 32-byte raw key for tests.
const testKey32 = "01234567890123456789012345678901" // exactly 32 bytes
func testKeyHex() string {
return hex.EncodeToString([]byte(testKey32)) // 64 hex chars
}
func testKeyBase64() string {
return base64.StdEncoding.EncodeToString([]byte(testKey32)) // 44 chars ending with =
}
// --- DeriveKey ---
func TestDeriveKey_Hex64(t *testing.T) {
k, err := DeriveKey(testKeyHex())
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(k) != 32 {
t.Fatalf("expected 32 bytes, got %d", len(k))
}
}
func TestDeriveKey_Base64_44(t *testing.T) {
k, err := DeriveKey(testKeyBase64())
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(k) != 32 {
t.Fatalf("expected 32 bytes, got %d", len(k))
}
}
func TestDeriveKey_Raw32(t *testing.T) {
k, err := DeriveKey(testKey32)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if string(k) != testKey32 {
t.Fatalf("raw key mismatch")
}
}
func TestDeriveKey_InvalidLength(t *testing.T) {
tests := []struct {
name string
key string
}{
{"too_short", "abc"},
{"16_bytes", "0123456789abcdef"},
{"48_bytes", strings.Repeat("a", 48)},
{"empty", ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
_, err := DeriveKey(tt.key)
if err == nil {
t.Fatalf("expected error for key len %d, got nil", len(tt.key))
}
})
}
}
func TestDeriveKey_Hex64_InvalidHex(t *testing.T) {
// 64 chars but not valid hex → should fall through to raw 32 check, then fail
key := strings.Repeat("zz", 32) // 64 chars, invalid hex
_, err := DeriveKey(key)
if err == nil {
t.Fatal("expected error for invalid hex, got nil")
}
}
// --- Encrypt with invalid key ---
func TestEncrypt_InvalidKey(t *testing.T) {
_, err := Encrypt("secret", "too-short-key")
if err == nil {
t.Fatal("expected error for invalid key length")
}
}
// --- Encrypt / Decrypt roundtrip ---
func TestEncryptDecrypt_Roundtrip(t *testing.T) {
tests := []struct {
name string
plaintext string
}{
{"simple_ascii", "my-secret-api-key-12345"},
{"unicode", "你好世界🌍emoji日本語"},
{"special_chars", `key="value"&foo=bar<>'"!@#$%`},
{"long_string", strings.Repeat("abcdefghij", 1000)},
{"single_char", "x"},
{"whitespace", " \t\n "},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
enc, err := Encrypt(tt.plaintext, testKey32)
if err != nil {
t.Fatalf("encrypt failed: %v", err)
}
if !strings.HasPrefix(enc, prefix) {
t.Fatalf("encrypted value missing prefix, got: %s", enc[:20])
}
dec, err := Decrypt(enc, testKey32)
if err != nil {
t.Fatalf("decrypt failed: %v", err)
}
if dec != tt.plaintext {
t.Fatalf("roundtrip mismatch: got %q, want %q", dec, tt.plaintext)
}
})
}
}
func TestEncryptDecrypt_EmptyPlaintext(t *testing.T) {
enc, err := Encrypt("", testKey32)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if enc != "" {
t.Fatalf("expected empty string, got %q", enc)
}
}
func TestEncryptDecrypt_EmptyKey(t *testing.T) {
// Empty key → plaintext returned unchanged for both encrypt and decrypt
enc, err := Encrypt("secret", "")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if enc != "secret" {
t.Fatalf("expected plaintext passthrough, got %q", enc)
}
dec, err := Decrypt("aes-gcm:someciphertext", "")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if dec != "aes-gcm:someciphertext" {
t.Fatalf("expected ciphertext passthrough, got %q", dec)
}
}
// --- Nonce uniqueness: same plaintext + key → different ciphertexts ---
func TestEncrypt_NonceUniqueness(t *testing.T) {
const plaintext = "identical-plaintext"
enc1, _ := Encrypt(plaintext, testKey32)
enc2, _ := Encrypt(plaintext, testKey32)
if enc1 == enc2 {
t.Fatal("two encryptions of same plaintext must produce different ciphertext (unique nonce)")
}
}
// --- Backward compatibility: decrypt unencrypted string returns as-is ---
func TestDecrypt_BackwardCompat_PlainText(t *testing.T) {
plain := "sk-ant-not-encrypted-just-raw"
dec, err := Decrypt(plain, testKey32)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if dec != plain {
t.Fatalf("expected passthrough for unencrypted value, got %q", dec)
}
}
// --- Error cases ---
func TestDecrypt_CorruptedCiphertext(t *testing.T) {
// Valid prefix but garbage base64 → should return as-is (line 68: treat as plain text)
garbage := prefix + "!!!not-base64!!!"
dec, err := Decrypt(garbage, testKey32)
if err != nil {
t.Fatalf("unexpected error for garbage base64: %v", err)
}
if dec != garbage {
t.Fatalf("expected passthrough for invalid base64, got %q", dec)
}
}
func TestDecrypt_WrongKey(t *testing.T) {
enc, _ := Encrypt("secret", testKey32)
otherKey := "98765432109876543210987654321098" // different 32-byte key
_, err := Decrypt(enc, otherKey)
if err == nil {
t.Fatal("expected error for wrong key, got nil")
}
if !strings.Contains(err.Error(), "decrypt failed") {
t.Fatalf("expected 'decrypt failed' error, got: %v", err)
}
}
func TestDecrypt_TooShortCiphertext(t *testing.T) {
// Valid prefix + valid base64 but too short for nonce → return as-is
short := prefix + base64.StdEncoding.EncodeToString([]byte("x"))
dec, err := Decrypt(short, testKey32)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if dec != short {
t.Fatalf("expected passthrough for too-short ciphertext, got %q", dec)
}
}
// --- IsEncrypted ---
func TestIsEncrypted(t *testing.T) {
if !IsEncrypted("aes-gcm:abc123") {
t.Fatal("expected true for prefixed value")
}
if IsEncrypted("plain-text-value") {
t.Fatal("expected false for unprefixed value")
}
if IsEncrypted("") {
t.Fatal("expected false for empty string")
}
}
// --- Cross-key-format roundtrip: encrypt with one format, decrypt with another ---
func TestEncryptDecrypt_CrossKeyFormats(t *testing.T) {
// All three key formats represent the same 32-byte key
keyRaw := testKey32
keyHex := testKeyHex()
keyB64 := testKeyBase64()
plaintext := "cross-format-test-value"
// Encrypt with raw, decrypt with hex
enc, _ := Encrypt(plaintext, keyRaw)
dec, err := Decrypt(enc, keyHex)
if err != nil {
t.Fatalf("hex decrypt failed: %v", err)
}
if dec != plaintext {
t.Fatalf("cross-format mismatch: raw→hex")
}
// Encrypt with hex, decrypt with base64
enc, _ = Encrypt(plaintext, keyHex)
dec, err = Decrypt(enc, keyB64)
if err != nil {
t.Fatalf("base64 decrypt failed: %v", err)
}
if dec != plaintext {
t.Fatalf("cross-format mismatch: hex→base64")
}
}
+280
View File
@@ -0,0 +1,280 @@
package permissions
import (
"testing"
"github.com/nextlevelbuilder/goclaw/pkg/protocol"
)
// --- Role hierarchy ---
func TestRoleLevel_Ordering(t *testing.T) {
// Owner > Admin > Operator > Viewer > unknown
levels := []struct {
role Role
level int
}{
{RoleOwner, 4},
{RoleAdmin, 3},
{RoleOperator, 2},
{RoleViewer, 1},
{Role("unknown"), 0},
{Role(""), 0},
}
for _, tt := range levels {
t.Run(string(tt.role), func(t *testing.T) {
got := roleLevel(tt.role)
if got != tt.level {
t.Fatalf("roleLevel(%q) = %d, want %d", tt.role, got, tt.level)
}
})
}
}
func TestHasMinRole(t *testing.T) {
tests := []struct {
name string
role Role
required Role
want bool
}{
{"owner_meets_admin", RoleOwner, RoleAdmin, true},
{"admin_meets_admin", RoleAdmin, RoleAdmin, true},
{"operator_fails_admin", RoleOperator, RoleAdmin, false},
{"viewer_fails_operator", RoleViewer, RoleOperator, false},
{"operator_meets_viewer", RoleOperator, RoleViewer, true},
{"admin_meets_viewer", RoleAdmin, RoleViewer, true},
{"viewer_meets_viewer", RoleViewer, RoleViewer, true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := HasMinRole(tt.role, tt.required)
if got != tt.want {
t.Fatalf("HasMinRole(%q, %q) = %v, want %v", tt.role, tt.required, got, tt.want)
}
})
}
}
// --- RoleFromScopes ---
func TestRoleFromScopes(t *testing.T) {
tests := []struct {
name string
scopes []Scope
want Role
}{
{"admin_scope", []Scope{ScopeAdmin}, RoleAdmin},
{"admin_overrides_read", []Scope{ScopeRead, ScopeAdmin}, RoleAdmin},
{"write_is_operator", []Scope{ScopeWrite}, RoleOperator},
{"approvals_is_operator", []Scope{ScopeApprovals}, RoleOperator},
{"pairing_is_operator", []Scope{ScopePairing}, RoleOperator},
{"read_is_viewer", []Scope{ScopeRead}, RoleViewer},
{"empty_scopes", []Scope{}, RoleViewer},
{"nil_scopes", nil, RoleViewer},
{"provision_only_is_viewer", []Scope{ScopeProvision}, RoleViewer},
{"read_and_write", []Scope{ScopeRead, ScopeWrite}, RoleOperator},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := RoleFromScopes(tt.scopes)
if got != tt.want {
t.Fatalf("RoleFromScopes(%v) = %q, want %q", tt.scopes, got, tt.want)
}
})
}
}
// --- CanAccess: role-based method access ---
func TestCanAccess_AdminMethods(t *testing.T) {
pe := NewPolicyEngine(nil)
adminMethods := []string{
protocol.MethodConfigApply,
protocol.MethodAgentsCreate,
protocol.MethodAgentsDelete,
protocol.MethodAPIKeysCreate,
protocol.MethodTeamsCreate,
}
for _, method := range adminMethods {
t.Run(method, func(t *testing.T) {
if !pe.CanAccess(RoleAdmin, method) {
t.Fatalf("admin should access %s", method)
}
if !pe.CanAccess(RoleOwner, method) {
t.Fatalf("owner should access %s", method)
}
if pe.CanAccess(RoleOperator, method) {
t.Fatalf("operator should NOT access %s", method)
}
if pe.CanAccess(RoleViewer, method) {
t.Fatalf("viewer should NOT access %s", method)
}
})
}
}
func TestCanAccess_WriteMethods(t *testing.T) {
pe := NewPolicyEngine(nil)
writeMethods := []string{
protocol.MethodChatSend,
protocol.MethodSessionsDelete,
protocol.MethodCronCreate,
}
for _, method := range writeMethods {
t.Run(method, func(t *testing.T) {
if !pe.CanAccess(RoleOperator, method) {
t.Fatalf("operator should access %s", method)
}
if !pe.CanAccess(RoleAdmin, method) {
t.Fatalf("admin should access %s", method)
}
if pe.CanAccess(RoleViewer, method) {
t.Fatalf("viewer should NOT access write method %s", method)
}
})
}
}
func TestCanAccess_ReadMethods_AnyRole(t *testing.T) {
pe := NewPolicyEngine(nil)
// A method not in admin or write lists → defaults to viewer
readMethod := "sessions.list" // not in admin/write lists
for _, role := range []Role{RoleViewer, RoleOperator, RoleAdmin, RoleOwner} {
if !pe.CanAccess(role, readMethod) {
t.Fatalf("%s should access read method %s", role, readMethod)
}
}
}
func TestCanAccess_UnknownMethod_DefaultsToViewer(t *testing.T) {
pe := NewPolicyEngine(nil)
// Unknown method defaults to viewer role requirement
if !pe.CanAccess(RoleViewer, "totally.unknown.method") {
t.Fatal("viewer should access unknown method (defaults to viewer)")
}
}
// --- CanAccessWithScopes: scope-based method access ---
func TestCanAccessWithScopes(t *testing.T) {
pe := NewPolicyEngine(nil)
tests := []struct {
name string
scopes []Scope
method string
want bool
}{
// Admin method requires ScopeAdmin
{"admin_scope_for_admin_method", []Scope{ScopeAdmin}, protocol.MethodAgentsCreate, true},
{"read_scope_for_admin_method", []Scope{ScopeRead}, protocol.MethodAgentsCreate, false},
{"write_scope_for_admin_method", []Scope{ScopeWrite}, protocol.MethodAgentsCreate, false},
// Approvals method requires ScopeApprovals or ScopeAdmin
{"approvals_scope_for_approvals", []Scope{ScopeApprovals}, "approvals.list", true},
{"admin_scope_for_approvals", []Scope{ScopeAdmin}, "approvals.list", true},
{"read_scope_for_approvals", []Scope{ScopeRead}, "approvals.list", false},
// Write method requires ScopeWrite or ScopeAdmin
{"write_scope_for_chat", []Scope{ScopeWrite}, protocol.MethodChatSend, true},
{"admin_scope_for_chat", []Scope{ScopeAdmin}, protocol.MethodChatSend, true},
{"read_scope_for_chat", []Scope{ScopeRead}, protocol.MethodChatSend, false},
// Read method allows ScopeRead, ScopeWrite, or ScopeAdmin
{"read_scope_for_read_method", []Scope{ScopeRead}, "sessions.list", true},
{"write_scope_for_read_method", []Scope{ScopeWrite}, "sessions.list", true},
{"admin_scope_for_read_method", []Scope{ScopeAdmin}, "sessions.list", true},
// Empty scopes
{"empty_scopes", []Scope{}, protocol.MethodAgentsCreate, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := pe.CanAccessWithScopes(tt.scopes, tt.method)
if got != tt.want {
t.Fatalf("CanAccessWithScopes(%v, %q) = %v, want %v", tt.scopes, tt.method, got, tt.want)
}
})
}
}
// --- IsOwner ---
func TestIsOwner(t *testing.T) {
pe := NewPolicyEngine([]string{"alice", "bob"})
if !pe.IsOwner("alice") {
t.Fatal("alice should be owner")
}
if !pe.IsOwner("bob") {
t.Fatal("bob should be owner")
}
if pe.IsOwner("charlie") {
t.Fatal("charlie should not be owner")
}
if pe.IsOwner("") {
t.Fatal("empty string should not be owner")
}
}
func TestIsOwner_EmptyList(t *testing.T) {
pe := NewPolicyEngine(nil)
if pe.IsOwner("anyone") {
t.Fatal("no one should be owner with empty list")
}
}
// --- ValidScope ---
func TestValidScope(t *testing.T) {
for scope := range AllScopes {
if !ValidScope(string(scope)) {
t.Fatalf("expected %q to be valid", scope)
}
}
if ValidScope("nonexistent.scope") {
t.Fatal("expected invalid scope to be rejected")
}
if ValidScope("") {
t.Fatal("expected empty scope to be rejected")
}
}
// --- MethodScopes: verify pairing/approvals special routes ---
func TestMethodScopes_PairingMethod(t *testing.T) {
scopes := MethodScopes("pairing.request")
if len(scopes) != 2 {
t.Fatalf("expected 2 scopes for pairing method, got %d", len(scopes))
}
// Should require ScopePairing or ScopeAdmin
hasPairing, hasAdmin := false, false
for _, s := range scopes {
if s == ScopePairing {
hasPairing = true
}
if s == ScopeAdmin {
hasAdmin = true
}
}
if !hasPairing || !hasAdmin {
t.Fatalf("pairing method should require [pairing, admin], got %v", scopes)
}
}
func TestMethodScopes_ApprovalMethod(t *testing.T) {
scopes := MethodScopes("approvals.list")
hasScopeApprovals, hasAdmin := false, false
for _, s := range scopes {
if s == ScopeApprovals {
hasScopeApprovals = true
}
if s == ScopeAdmin {
hasAdmin = true
}
}
if !hasScopeApprovals || !hasAdmin {
t.Fatalf("approvals method should require [approvals, admin], got %v", scopes)
}
}
+330
View File
@@ -0,0 +1,330 @@
package providers
import (
"context"
"errors"
"fmt"
"sync/atomic"
"testing"
"time"
)
// --- IsRetryableError ---
func TestIsRetryableError(t *testing.T) {
tests := []struct {
name string
err error
want bool
}{
{"nil", nil, false},
{"http_429_rate_limit", &HTTPError{Status: 429}, true},
{"http_500_server_error", &HTTPError{Status: 500}, true},
{"http_502_bad_gateway", &HTTPError{Status: 502}, true},
{"http_503_unavailable", &HTTPError{Status: 503}, true},
{"http_504_timeout", &HTTPError{Status: 504}, true},
{"http_400_bad_request", &HTTPError{Status: 400}, false},
{"http_401_unauthorized", &HTTPError{Status: 401}, false},
{"http_403_forbidden", &HTTPError{Status: 403}, false},
{"http_404_not_found", &HTTPError{Status: 404}, false},
{"connection_reset", errors.New("connection reset by peer"), true},
{"broken_pipe", errors.New("write: broken pipe"), true},
{"eof", errors.New("unexpected EOF"), true},
{"timeout_string", errors.New("i/o timeout"), true},
{"generic_error", errors.New("something went wrong"), false},
{"wrapped_retryable", fmt.Errorf("provider: %w", &HTTPError{Status: 429}), true},
{"wrapped_non_retryable", fmt.Errorf("provider: %w", &HTTPError{Status: 400}), false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := IsRetryableError(tt.err)
if got != tt.want {
t.Fatalf("IsRetryableError(%v) = %v, want %v", tt.err, got, tt.want)
}
})
}
}
// --- computeDelay ---
func TestComputeDelay_ExponentialBackoff(t *testing.T) {
cfg := RetryConfig{
MinDelay: 100 * time.Millisecond,
MaxDelay: 10 * time.Second,
Jitter: 0, // no jitter for deterministic test
}
err := &HTTPError{Status: 500}
// attempt 1: 100ms * 2^0 = 100ms
d1 := computeDelay(cfg, 1, err)
if d1 != 100*time.Millisecond {
t.Fatalf("attempt 1: got %v, want 100ms", d1)
}
// attempt 2: 100ms * 2^1 = 200ms
d2 := computeDelay(cfg, 2, err)
if d2 != 200*time.Millisecond {
t.Fatalf("attempt 2: got %v, want 200ms", d2)
}
// attempt 3: 100ms * 2^2 = 400ms
d3 := computeDelay(cfg, 3, err)
if d3 != 400*time.Millisecond {
t.Fatalf("attempt 3: got %v, want 400ms", d3)
}
// attempt 4: 100ms * 2^3 = 800ms
d4 := computeDelay(cfg, 4, err)
if d4 != 800*time.Millisecond {
t.Fatalf("attempt 4: got %v, want 800ms", d4)
}
}
func TestComputeDelay_CappedAtMaxDelay(t *testing.T) {
cfg := RetryConfig{
MinDelay: 1 * time.Second,
MaxDelay: 5 * time.Second,
Jitter: 0,
}
err := &HTTPError{Status: 500}
// attempt 10: 1s * 2^9 = 512s → capped at 5s
d := computeDelay(cfg, 10, err)
if d != 5*time.Second {
t.Fatalf("attempt 10: got %v, want 5s (capped)", d)
}
}
func TestComputeDelay_JitterRange(t *testing.T) {
cfg := RetryConfig{
MinDelay: 1 * time.Second,
MaxDelay: 30 * time.Second,
Jitter: 0.25, // ±25%
}
err := &HTTPError{Status: 500}
// attempt 1: base = 1s, jitter ±25% → [750ms, 1250ms]
min := 750 * time.Millisecond
max := 1250 * time.Millisecond
for i := 0; i < 100; i++ {
d := computeDelay(cfg, 1, err)
if d < min || d > max {
t.Fatalf("jitter out of range: got %v, want [%v, %v]", d, min, max)
}
}
}
func TestComputeDelay_NeverNegative(t *testing.T) {
cfg := RetryConfig{
MinDelay: 10 * time.Millisecond,
MaxDelay: 100 * time.Millisecond,
Jitter: 0.9, // extreme jitter
}
err := &HTTPError{Status: 500}
for i := 0; i < 200; i++ {
d := computeDelay(cfg, 1, err)
if d < 0 {
t.Fatalf("negative delay: %v", d)
}
}
}
func TestComputeDelay_RetryAfterOverride(t *testing.T) {
cfg := RetryConfig{
MinDelay: 100 * time.Millisecond,
MaxDelay: 30 * time.Second,
Jitter: 0.1,
}
// HTTPError with RetryAfter should override computed delay
err := &HTTPError{Status: 429, RetryAfter: 42 * time.Second}
d := computeDelay(cfg, 1, err)
if d != 42*time.Second {
t.Fatalf("expected Retry-After override: got %v, want 42s", d)
}
}
// --- ParseRetryAfter ---
func TestParseRetryAfter(t *testing.T) {
tests := []struct {
name string
value string
want time.Duration
}{
{"empty", "", 0},
{"integer_seconds", "30", 30 * time.Second},
{"zero", "0", 0},
{"negative_int", "-5", -5 * time.Second}, // strconv.Atoi succeeds → returns negative duration (caller should clamp)
{"non_numeric", "abc", 0}, // neither int nor date
{"float", "1.5", 0}, // not a valid int
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := ParseRetryAfter(tt.value)
if got != tt.want {
t.Fatalf("ParseRetryAfter(%q) = %v, want %v", tt.value, got, tt.want)
}
})
}
}
func TestParseRetryAfter_RFC1123(t *testing.T) {
// Use a future date so time.Until() returns positive
future := time.Now().Add(60 * time.Second).UTC().Format(time.RFC1123)
got := ParseRetryAfter(future)
// Should be roughly 60s (allow ±5s tolerance for test execution time)
if got < 55*time.Second || got > 65*time.Second {
t.Fatalf("RFC1123 parse: got %v, want ~60s", got)
}
}
func TestParseRetryAfter_PastDate(t *testing.T) {
past := time.Now().Add(-60 * time.Second).UTC().Format(time.RFC1123)
got := ParseRetryAfter(past)
if got != 0 {
t.Fatalf("past date should return 0, got %v", got)
}
}
// --- RetryDo ---
func TestRetryDo_SuccessOnFirstAttempt(t *testing.T) {
cfg := RetryConfig{Attempts: 3, MinDelay: time.Millisecond}
var calls int
result, err := RetryDo(context.Background(), cfg, func() (string, error) {
calls++
return "ok", nil
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if result != "ok" {
t.Fatalf("got %q, want %q", result, "ok")
}
if calls != 1 {
t.Fatalf("expected 1 call, got %d", calls)
}
}
func TestRetryDo_SuccessAfterRetries(t *testing.T) {
cfg := RetryConfig{Attempts: 3, MinDelay: time.Millisecond, MaxDelay: 10 * time.Millisecond}
var calls int
result, err := RetryDo(context.Background(), cfg, func() (string, error) {
calls++
if calls < 3 {
return "", &HTTPError{Status: 500, Body: "server error"}
}
return "recovered", nil
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if result != "recovered" {
t.Fatalf("got %q, want %q", result, "recovered")
}
if calls != 3 {
t.Fatalf("expected 3 calls, got %d", calls)
}
}
func TestRetryDo_NonRetryableError_NoRetry(t *testing.T) {
cfg := RetryConfig{Attempts: 5, MinDelay: time.Millisecond}
var calls int
_, err := RetryDo(context.Background(), cfg, func() (string, error) {
calls++
return "", &HTTPError{Status: 400, Body: "bad request"}
})
if err == nil {
t.Fatal("expected error")
}
if calls != 1 {
t.Fatalf("non-retryable error should not retry: got %d calls, want 1", calls)
}
}
func TestRetryDo_MaxAttemptsExhausted(t *testing.T) {
cfg := RetryConfig{Attempts: 3, MinDelay: time.Millisecond, MaxDelay: 10 * time.Millisecond}
var calls int
_, err := RetryDo(context.Background(), cfg, func() (string, error) {
calls++
return "", &HTTPError{Status: 503, Body: "unavailable"}
})
if err == nil {
t.Fatal("expected error after exhausting retries")
}
if calls != 3 {
t.Fatalf("expected 3 attempts, got %d", calls)
}
}
func TestRetryDo_ContextCancellation(t *testing.T) {
cfg := RetryConfig{Attempts: 10, MinDelay: 5 * time.Second, MaxDelay: 5 * time.Second}
ctx, cancel := context.WithCancel(context.Background())
go func() {
time.Sleep(50 * time.Millisecond)
cancel()
}()
start := time.Now()
_, err := RetryDo(ctx, cfg, func() (string, error) {
return "", &HTTPError{Status: 500}
})
elapsed := time.Since(start)
if !errors.Is(err, context.Canceled) {
t.Fatalf("expected context.Canceled, got: %v", err)
}
// Should have been cancelled quickly, not waited for 5s backoff
if elapsed > 2*time.Second {
t.Fatalf("context cancellation took too long: %v", elapsed)
}
}
func TestRetryDo_HookCalledOnRetry(t *testing.T) {
cfg := RetryConfig{Attempts: 3, MinDelay: time.Millisecond, MaxDelay: 10 * time.Millisecond}
var hookCalls atomic.Int32
ctx := WithRetryHook(context.Background(), func(attempt, maxAttempts int, err error) {
hookCalls.Add(1)
})
RetryDo(ctx, cfg, func() (string, error) {
return "", &HTTPError{Status: 500}
})
// Hook called before each retry (not the first attempt, not the last failure)
// With 3 attempts: attempt 1 fails → hook → attempt 2 fails → hook → attempt 3 fails → done
if got := hookCalls.Load(); got != 2 {
t.Fatalf("expected 2 hook calls, got %d", got)
}
}
func TestRetryDo_ZeroAttempts_DefaultsToOne(t *testing.T) {
cfg := RetryConfig{Attempts: 0}
var calls int
_, err := RetryDo(context.Background(), cfg, func() (string, error) {
calls++
return "", &HTTPError{Status: 500}
})
if err == nil {
t.Fatal("expected error")
}
if calls != 1 {
t.Fatalf("zero attempts should default to 1: got %d calls", calls)
}
}
// --- HTTPError ---
func TestHTTPError_ErrorString(t *testing.T) {
err := &HTTPError{Status: 429, Body: "rate limited"}
got := err.Error()
want := "HTTP 429: rate limited"
if got != want {
t.Fatalf("got %q, want %q", got, want)
}
}
+466
View File
@@ -0,0 +1,466 @@
package scheduler
import (
"context"
"errors"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/nextlevelbuilder/goclaw/internal/agent"
)
// mockRunFn creates a run function that completes after a delay.
func mockRunFn(delay time.Duration) RunFunc {
return func(ctx context.Context, req agent.RunRequest) (*agent.RunResult, error) {
select {
case <-ctx.Done():
return nil, ctx.Err()
case <-time.After(delay):
return &agent.RunResult{Content: "done:" + req.RunID}, nil
}
}
}
// fastRunFn completes immediately.
func fastRunFn() RunFunc {
return func(ctx context.Context, req agent.RunRequest) (*agent.RunResult, error) {
return &agent.RunResult{Content: "ok"}, nil
}
}
// --- Draining mode rejects new requests ---
func TestScheduler_MarkDraining(t *testing.T) {
sched := NewScheduler(nil, DefaultQueueConfig(), fastRunFn())
defer sched.Stop()
sched.MarkDraining()
ch := sched.Schedule(context.Background(), LaneMain, agent.RunRequest{
SessionKey: "agent:a1:s1",
RunID: "run-1",
})
outcome := <-ch
if !errors.Is(outcome.Err, ErrGatewayDraining) {
t.Fatalf("expected ErrGatewayDraining, got: %v", outcome.Err)
}
}
// --- DropNew policy: full queue rejects incoming ---
func TestSessionQueue_DropNewPolicy(t *testing.T) {
cfg := QueueConfig{
Mode: QueueModeQueue,
Cap: 2,
Drop: DropNew,
DebounceMs: 0,
MaxConcurrent: 1,
}
// Use a slow run to keep the queue occupied
blockCh := make(chan struct{})
runFn := func(ctx context.Context, req agent.RunRequest) (*agent.RunResult, error) {
<-blockCh
return &agent.RunResult{}, nil
}
laneMgr := NewLaneManager([]LaneConfig{{Name: LaneMain, Concurrency: 10}})
sq := NewSessionQueue("test-session", LaneMain, cfg, laneMgr, runFn)
// Enqueue 3 requests: 1 active + 2 in queue = at capacity
ctx := context.Background()
sq.Enqueue(ctx, agent.RunRequest{RunID: "r1", SessionKey: "s"})
time.Sleep(10 * time.Millisecond) // let r1 start
sq.Enqueue(ctx, agent.RunRequest{RunID: "r2", SessionKey: "s"})
sq.Enqueue(ctx, agent.RunRequest{RunID: "r3", SessionKey: "s"})
// Queue is full (cap=2). Next one should be rejected.
ch := sq.Enqueue(ctx, agent.RunRequest{RunID: "r4", SessionKey: "s"})
outcome := <-ch
if !errors.Is(outcome.Err, ErrQueueFull) {
t.Fatalf("expected ErrQueueFull, got: %v", outcome.Err)
}
close(blockCh) // unblock
}
// --- DropOld policy: full queue drops oldest ---
func TestSessionQueue_DropOldPolicy(t *testing.T) {
cfg := QueueConfig{
Mode: QueueModeQueue,
Cap: 2,
Drop: DropOld,
DebounceMs: 0,
MaxConcurrent: 1,
}
blockCh := make(chan struct{})
runFn := func(ctx context.Context, req agent.RunRequest) (*agent.RunResult, error) {
<-blockCh
return &agent.RunResult{}, nil
}
laneMgr := NewLaneManager([]LaneConfig{{Name: LaneMain, Concurrency: 10}})
sq := NewSessionQueue("test-session", LaneMain, cfg, laneMgr, runFn)
ctx := context.Background()
sq.Enqueue(ctx, agent.RunRequest{RunID: "r1", SessionKey: "s"})
time.Sleep(10 * time.Millisecond)
ch2 := sq.Enqueue(ctx, agent.RunRequest{RunID: "r2", SessionKey: "s"})
sq.Enqueue(ctx, agent.RunRequest{RunID: "r3", SessionKey: "s"})
// Queue is full. Adding r4 should drop r2 (oldest queued).
sq.Enqueue(ctx, agent.RunRequest{RunID: "r4", SessionKey: "s"})
// r2 should have been dropped
outcome := <-ch2
if !errors.Is(outcome.Err, ErrQueueDropped) {
t.Fatalf("expected ErrQueueDropped for r2, got: %v", outcome.Err)
}
close(blockCh)
}
// --- Adaptive throttle: reduces concurrency near 60% context usage ---
func TestSessionQueue_AdaptiveThrottle(t *testing.T) {
cfg := QueueConfig{
Mode: QueueModeQueue,
Cap: 10,
Drop: DropOld,
DebounceMs: 0,
MaxConcurrent: 5,
}
laneMgr := NewLaneManager([]LaneConfig{{Name: LaneMain, Concurrency: 20}})
sq := NewSessionQueue("test-session", LaneMain, cfg, laneMgr, fastRunFn())
// Without token estimate → effectiveMaxConcurrent = 5
sq.mu.Lock()
got := sq.effectiveMaxConcurrent()
sq.mu.Unlock()
if got != 5 {
t.Fatalf("without estimate: got %d, want 5", got)
}
// With token estimate under threshold (50%) → still 5
sq.tokenEstimateFn = func(key string) (int, int) {
return 50000, 100000 // 50%
}
sq.mu.Lock()
got = sq.effectiveMaxConcurrent()
sq.mu.Unlock()
if got != 5 {
t.Fatalf("under threshold: got %d, want 5", got)
}
// At 60% threshold → drops to 1
sq.tokenEstimateFn = func(key string) (int, int) {
return 60000, 100000 // 60%
}
sq.mu.Lock()
got = sq.effectiveMaxConcurrent()
sq.mu.Unlock()
if got != 1 {
t.Fatalf("at threshold: got %d, want 1", got)
}
// Above threshold → still 1
sq.tokenEstimateFn = func(key string) (int, int) {
return 80000, 100000 // 80%
}
sq.mu.Lock()
got = sq.effectiveMaxConcurrent()
sq.mu.Unlock()
if got != 1 {
t.Fatalf("above threshold: got %d, want 1", got)
}
}
// --- Generation-based stale completion filtering ---
func TestSessionQueue_StaleCompletion_AfterReset(t *testing.T) {
var runCount atomic.Int32
blockCh := make(chan struct{})
runFn := func(ctx context.Context, req agent.RunRequest) (*agent.RunResult, error) {
runCount.Add(1)
select {
case <-ctx.Done():
return nil, ctx.Err()
case <-blockCh:
return &agent.RunResult{Content: "from-" + req.RunID}, nil
}
}
cfg := QueueConfig{Mode: QueueModeQueue, Cap: 10, DebounceMs: 0, MaxConcurrent: 1}
laneMgr := NewLaneManager([]LaneConfig{{Name: LaneMain, Concurrency: 10}})
sq := NewSessionQueue("test", LaneMain, cfg, laneMgr, runFn)
// Start run r1
ch1 := sq.Enqueue(context.Background(), agent.RunRequest{RunID: "r1", SessionKey: "test"})
time.Sleep(20 * time.Millisecond) // let r1 start
// Reset bumps generation, cancels r1
sq.Reset()
// r1's context is cancelled → it should complete with error
close(blockCh)
outcome := <-ch1
// r1 was either cancelled or completed, but the key test is the generation check
_ = outcome
// After reset, new runs should work normally
sq2ch := sq.Enqueue(context.Background(), agent.RunRequest{RunID: "r2", SessionKey: "test"})
outcome2 := <-sq2ch
if outcome2.Err != nil {
t.Fatalf("post-reset run should succeed: %v", outcome2.Err)
}
}
// --- CancelAll sets abort cutoff for stale messages ---
func TestSessionQueue_CancelAll_StaleMessages(t *testing.T) {
blockCh := make(chan struct{})
runFn := func(ctx context.Context, req agent.RunRequest) (*agent.RunResult, error) {
select {
case <-ctx.Done():
return nil, ctx.Err()
case <-blockCh:
return &agent.RunResult{}, nil
}
}
cfg := QueueConfig{Mode: QueueModeQueue, Cap: 10, DebounceMs: 0, MaxConcurrent: 1}
laneMgr := NewLaneManager([]LaneConfig{{Name: LaneMain, Concurrency: 10}})
sq := NewSessionQueue("test", LaneMain, cfg, laneMgr, runFn)
ctx := context.Background()
sq.Enqueue(ctx, agent.RunRequest{RunID: "r1", SessionKey: "test"})
time.Sleep(10 * time.Millisecond)
// Queue r2 before cancellation
ch2 := sq.Enqueue(ctx, agent.RunRequest{RunID: "r2", SessionKey: "test"})
// CancelAll → sets abort cutoff → r2 was queued before cutoff → stale
sq.CancelAll()
outcome := <-ch2
// r2 was in queue when CancelAll drained it → drainQueue sends context.Canceled
if !errors.Is(outcome.Err, context.Canceled) {
t.Fatalf("expected context.Canceled for drained message, got: %v", outcome.Err)
}
close(blockCh)
}
// --- Lane concurrency enforcement ---
func TestLane_ConcurrencyEnforcement(t *testing.T) {
lane := NewLane("test", 2)
defer lane.Stop()
var maxConcurrent atomic.Int32
var current atomic.Int32
var wg sync.WaitGroup
for i := 0; i < 10; i++ {
wg.Add(1)
err := lane.Submit(context.Background(), func() {
defer wg.Done()
cur := current.Add(1)
// Track max concurrent
for {
old := maxConcurrent.Load()
if cur <= old || maxConcurrent.CompareAndSwap(old, cur) {
break
}
}
time.Sleep(10 * time.Millisecond)
current.Add(-1)
})
if err != nil {
wg.Done()
}
}
wg.Wait()
if got := maxConcurrent.Load(); got > 2 {
t.Fatalf("max concurrent exceeded limit: got %d, want ≤2", got)
}
}
// --- Lane Submit with cancelled context ---
func TestLane_Submit_CancelledContext(t *testing.T) {
lane := NewLane("test", 1)
defer lane.Stop()
// Fill the single slot
blockCh := make(chan struct{})
lane.Submit(context.Background(), func() {
<-blockCh
})
// Submit with already-cancelled context
ctx, cancel := context.WithCancel(context.Background())
cancel()
err := lane.Submit(ctx, func() {
t.Fatal("should not execute")
})
if err == nil {
t.Fatal("expected error for cancelled context")
}
close(blockCh)
}
// --- LaneManager fallback to main ---
func TestLaneManager_FallbackToMain(t *testing.T) {
lm := NewLaneManager([]LaneConfig{
{Name: LaneMain, Concurrency: 5},
})
// Known lane
if lm.Get(LaneMain) == nil {
t.Fatal("main lane should exist")
}
// Unknown lane → falls back to main
fallback := lm.Get("nonexistent")
main := lm.Get(LaneMain)
if fallback != main {
t.Fatal("unknown lane should fall back to main")
}
}
// --- HasActiveSessionsForAgent ---
func TestScheduler_HasActiveSessionsForAgent(t *testing.T) {
blockCh := make(chan struct{})
runFn := func(ctx context.Context, req agent.RunRequest) (*agent.RunResult, error) {
<-blockCh
return &agent.RunResult{}, nil
}
cfg := DefaultQueueConfig()
cfg.DebounceMs = 0
sched := NewScheduler(nil, cfg, runFn)
sched.Schedule(context.Background(), LaneMain, agent.RunRequest{
SessionKey: "agent:agent-123:scope1",
RunID: "run-1",
})
time.Sleep(20 * time.Millisecond) // let run start
if !sched.HasActiveSessionsForAgent("agent-123") {
t.Fatal("expected active sessions for agent-123")
}
if sched.HasActiveSessionsForAgent("other-agent") {
t.Fatal("expected no active sessions for other-agent")
}
close(blockCh)
time.Sleep(20 * time.Millisecond)
if sched.HasActiveSessionsForAgent("agent-123") {
t.Fatal("expected no active sessions after completion")
}
}
// --- Debounce collapsing ---
func TestSessionQueue_Debounce_CollapsesRapidMessages(t *testing.T) {
var runCount atomic.Int32
runFn := func(ctx context.Context, req agent.RunRequest) (*agent.RunResult, error) {
runCount.Add(1)
return &agent.RunResult{}, nil
}
cfg := QueueConfig{
Mode: QueueModeQueue,
Cap: 10,
Drop: DropOld,
DebounceMs: 200, // 200ms debounce
MaxConcurrent: 1,
}
laneMgr := NewLaneManager([]LaneConfig{{Name: LaneMain, Concurrency: 10}})
sq := NewSessionQueue("test", LaneMain, cfg, laneMgr, runFn)
ctx := context.Background()
// Send 5 rapid messages within debounce window
var channels []<-chan RunOutcome
for i := 0; i < 5; i++ {
ch := sq.Enqueue(ctx, agent.RunRequest{
RunID: "r" + string(rune('0'+i)),
SessionKey: "test",
})
channels = append(channels, ch)
}
// Wait for debounce + execution of all queued messages
for i, ch := range channels {
select {
case <-ch:
case <-time.After(3 * time.Second):
t.Fatalf("message %d timed out", i)
}
}
// Debounce collapses the initial scheduleNext calls into one timer fire,
// but all 5 messages still execute sequentially. The key behavior is that
// execution doesn't start until after debounce delay (200ms), not immediately.
if got := runCount.Load(); got != 5 {
t.Fatalf("all 5 messages should eventually run, got %d", got)
}
}
// --- Interrupt mode ---
func TestSessionQueue_InterruptMode(t *testing.T) {
var runIDs []string
var mu sync.Mutex
runFn := func(ctx context.Context, req agent.RunRequest) (*agent.RunResult, error) {
mu.Lock()
runIDs = append(runIDs, req.RunID)
mu.Unlock()
select {
case <-ctx.Done():
return nil, ctx.Err()
case <-time.After(100 * time.Millisecond):
return &agent.RunResult{}, nil
}
}
cfg := QueueConfig{
Mode: QueueModeInterrupt,
Cap: 10,
DebounceMs: 0,
MaxConcurrent: 1,
}
laneMgr := NewLaneManager([]LaneConfig{{Name: LaneMain, Concurrency: 10}})
sq := NewSessionQueue("test", LaneMain, cfg, laneMgr, runFn)
ctx := context.Background()
sq.Enqueue(ctx, agent.RunRequest{RunID: "r1", SessionKey: "test"})
time.Sleep(20 * time.Millisecond)
// Interrupt: should cancel r1 and start r2
ch2 := sq.Enqueue(ctx, agent.RunRequest{RunID: "r2", SessionKey: "test"})
outcome := <-ch2
if outcome.Err != nil {
t.Fatalf("r2 should complete successfully: %v", outcome.Err)
}
}
+467
View File
@@ -0,0 +1,467 @@
package sessions
import (
"context"
"os"
"sync"
"testing"
"github.com/nextlevelbuilder/goclaw/internal/providers"
)
// --- SessionKey ---
func TestSessionKey(t *testing.T) {
got := SessionKey("agent-123", "telegram:direct:chat456")
want := "agent:agent-123:telegram:direct:chat456"
if got != want {
t.Fatalf("SessionKey = %q, want %q", got, want)
}
}
// --- GetOrCreate idempotency ---
func TestGetOrCreate_Idempotent(t *testing.T) {
m := NewManager("")
ctx := context.Background()
s1 := m.GetOrCreate(ctx, "agent:a1:scope1")
s2 := m.GetOrCreate(ctx, "agent:a1:scope1")
if s1 != s2 {
t.Fatal("GetOrCreate should return same session pointer for same key")
}
}
func TestGetOrCreate_DifferentKeys(t *testing.T) {
m := NewManager("")
ctx := context.Background()
s1 := m.GetOrCreate(ctx, "agent:a1:scope1")
s2 := m.GetOrCreate(ctx, "agent:a2:scope2")
if s1 == s2 {
t.Fatal("different keys should return different sessions")
}
}
// --- AddMessage ---
func TestAddMessage_CreatesSession(t *testing.T) {
m := NewManager("")
ctx := context.Background()
m.AddMessage(ctx, "agent:a1:s1", providers.Message{Role: "user", Content: "hello"})
history := m.GetHistory(ctx, "agent:a1:s1")
if len(history) != 1 {
t.Fatalf("expected 1 message, got %d", len(history))
}
if history[0].Content != "hello" {
t.Fatalf("content mismatch: got %q", history[0].Content)
}
}
func TestAddMessage_AppendsToExisting(t *testing.T) {
m := NewManager("")
ctx := context.Background()
m.AddMessage(ctx, "key1", providers.Message{Role: "user", Content: "msg1"})
m.AddMessage(ctx, "key1", providers.Message{Role: "assistant", Content: "msg2"})
m.AddMessage(ctx, "key1", providers.Message{Role: "user", Content: "msg3"})
history := m.GetHistory(ctx, "key1")
if len(history) != 3 {
t.Fatalf("expected 3 messages, got %d", len(history))
}
}
// --- Concurrent AddMessage: messages must not be lost ---
func TestAddMessage_ConcurrentSafety(t *testing.T) {
m := NewManager("")
ctx := context.Background()
key := "agent:concurrent:test"
const n = 100
var wg sync.WaitGroup
wg.Add(n)
for i := 0; i < n; i++ {
go func(i int) {
defer wg.Done()
m.AddMessage(ctx, key, providers.Message{Role: "user", Content: "msg"})
}(i)
}
wg.Wait()
history := m.GetHistory(ctx, key)
if len(history) != n {
t.Fatalf("expected %d messages (no loss), got %d", n, len(history))
}
}
// --- GetHistory returns defensive copy ---
func TestGetHistory_DefensiveCopy(t *testing.T) {
m := NewManager("")
ctx := context.Background()
key := "agent:a1:s1"
m.AddMessage(ctx, key, providers.Message{Role: "user", Content: "original"})
history := m.GetHistory(ctx, key)
history[0].Content = "mutated" // mutate the copy
// Original should be unchanged
original := m.GetHistory(ctx, key)
if original[0].Content != "original" {
t.Fatal("GetHistory should return a copy, not a reference to internal slice")
}
}
func TestGetHistory_NonExistentSession(t *testing.T) {
m := NewManager("")
history := m.GetHistory(context.Background(), "nonexistent")
if history != nil {
t.Fatalf("expected nil for non-existent session, got %v", history)
}
}
// --- SetLabel / SetSummary on missing session: silent no-op ---
func TestSetLabel_MissingSession_SilentNoOp(t *testing.T) {
m := NewManager("")
// Should not panic
m.SetLabel(context.Background(), "nonexistent", "label")
}
func TestSetSummary_MissingSession_SilentNoOp(t *testing.T) {
m := NewManager("")
m.SetSummary(context.Background(), "nonexistent", "summary")
}
// --- Metadata accumulation ---
func TestAccumulateTokens(t *testing.T) {
m := NewManager("")
ctx := context.Background()
key := "agent:a1:s1"
m.GetOrCreate(ctx, key)
m.AccumulateTokens(ctx, key, 100, 50)
m.AccumulateTokens(ctx, key, 200, 80)
m.mu.RLock()
s := m.sessions[key]
m.mu.RUnlock()
if s.InputTokens != 300 {
t.Fatalf("input tokens: got %d, want 300", s.InputTokens)
}
if s.OutputTokens != 130 {
t.Fatalf("output tokens: got %d, want 130", s.OutputTokens)
}
}
// --- CompactionCount tracking ---
func TestCompactionCount(t *testing.T) {
m := NewManager("")
ctx := context.Background()
key := "agent:a1:s1"
m.GetOrCreate(ctx, key)
if m.GetCompactionCount(ctx, key) != 0 {
t.Fatal("initial compaction count should be 0")
}
m.IncrementCompaction(ctx, key)
m.IncrementCompaction(ctx, key)
if m.GetCompactionCount(ctx, key) != 2 {
t.Fatalf("compaction count: got %d, want 2", m.GetCompactionCount(ctx, key))
}
// MemoryFlushCompactionCount tracks separately
if m.GetMemoryFlushCompactionCount(ctx, key) != 0 {
t.Fatal("memory flush compaction should start at 0")
}
m.SetMemoryFlushDone(ctx, key)
if m.GetMemoryFlushCompactionCount(ctx, key) != 2 {
t.Fatalf("memory flush should match current compaction count: got %d", m.GetMemoryFlushCompactionCount(ctx, key))
}
}
func TestGetMemoryFlushCompactionCount_NonExistent(t *testing.T) {
m := NewManager("")
got := m.GetMemoryFlushCompactionCount(context.Background(), "nonexistent")
if got != -1 {
t.Fatalf("expected -1 for non-existent session, got %d", got)
}
}
// --- TruncateHistory ---
func TestTruncateHistory(t *testing.T) {
m := NewManager("")
ctx := context.Background()
key := "agent:a1:s1"
for i := 0; i < 10; i++ {
m.AddMessage(ctx, key, providers.Message{Role: "user", Content: "msg"})
}
m.TruncateHistory(ctx, key, 3)
history := m.GetHistory(ctx, key)
if len(history) != 3 {
t.Fatalf("expected 3 messages after truncate, got %d", len(history))
}
}
func TestTruncateHistory_KeepZero(t *testing.T) {
m := NewManager("")
ctx := context.Background()
key := "agent:a1:s1"
m.AddMessage(ctx, key, providers.Message{Role: "user", Content: "msg"})
m.TruncateHistory(ctx, key, 0)
history := m.GetHistory(ctx, key)
if len(history) != 0 {
t.Fatalf("expected 0 messages, got %d", len(history))
}
}
// --- Reset ---
func TestReset_ClearsHistoryAndSummary(t *testing.T) {
m := NewManager("")
ctx := context.Background()
key := "agent:a1:s1"
m.AddMessage(ctx, key, providers.Message{Role: "user", Content: "msg"})
m.SetSummary(ctx, key, "summary text")
m.Reset(ctx, key)
history := m.GetHistory(ctx, key)
if len(history) != 0 {
t.Fatalf("expected empty history after reset, got %d messages", len(history))
}
if m.GetSummary(ctx, key) != "" {
t.Fatal("expected empty summary after reset")
}
}
// --- Delete ---
func TestDelete_RemovesSession(t *testing.T) {
m := NewManager("")
ctx := context.Background()
key := "agent:a1:s1"
m.GetOrCreate(ctx, key)
if err := m.Delete(ctx, key); err != nil {
t.Fatalf("delete error: %v", err)
}
history := m.GetHistory(ctx, key)
if history != nil {
t.Fatal("expected nil history after delete")
}
}
// --- List ---
func TestList_FiltersByAgentID(t *testing.T) {
m := NewManager("")
ctx := context.Background()
m.GetOrCreate(ctx, "agent:a1:scope1")
m.GetOrCreate(ctx, "agent:a1:scope2")
m.GetOrCreate(ctx, "agent:a2:scope1")
all := m.List(ctx, "")
if len(all) != 3 {
t.Fatalf("expected 3 sessions, got %d", len(all))
}
filtered := m.List(ctx, "a1")
if len(filtered) != 2 {
t.Fatalf("expected 2 sessions for agent a1, got %d", len(filtered))
}
}
// --- ContextWindow ---
func TestContextWindow_SetAndGet(t *testing.T) {
m := NewManager("")
ctx := context.Background()
key := "agent:a1:s1"
m.GetOrCreate(ctx, key)
if m.GetContextWindow(ctx, key) != 0 {
t.Fatal("initial context window should be 0")
}
m.SetContextWindow(ctx, key, 200000)
if m.GetContextWindow(ctx, key) != 200000 {
t.Fatalf("context window mismatch: got %d", m.GetContextWindow(ctx, key))
}
}
// --- Persistence: Save and load roundtrip ---
func TestSave_LoadAll_Roundtrip(t *testing.T) {
dir := t.TempDir()
// Create manager and save a session
m1 := NewManager(dir)
ctx := context.Background()
key := "agent:a1:scope1"
m1.AddMessage(ctx, key, providers.Message{Role: "user", Content: "hello"})
m1.AddMessage(ctx, key, providers.Message{Role: "assistant", Content: "hi"})
m1.SetLabel(ctx, key, "test-label")
m1.AccumulateTokens(ctx, key, 100, 50)
if err := m1.Save(ctx, key); err != nil {
t.Fatalf("save error: %v", err)
}
// Create new manager from same directory → should load the session
m2 := NewManager(dir)
history := m2.GetHistory(ctx, key)
if len(history) != 2 {
t.Fatalf("expected 2 messages after load, got %d", len(history))
}
if history[0].Content != "hello" || history[1].Content != "hi" {
t.Fatal("message content mismatch after load")
}
m2.mu.RLock()
s := m2.sessions[key]
m2.mu.RUnlock()
if s.Label != "test-label" {
t.Fatalf("label mismatch: got %q", s.Label)
}
if s.InputTokens != 100 {
t.Fatalf("input tokens mismatch: got %d", s.InputTokens)
}
}
func TestDelete_RemovesFile(t *testing.T) {
dir := t.TempDir()
m := NewManager(dir)
ctx := context.Background()
key := "agent:a1:scope1"
m.AddMessage(ctx, key, providers.Message{Role: "user", Content: "msg"})
_ = m.Save(ctx, key)
// File should exist
files, _ := os.ReadDir(dir)
if len(files) == 0 {
t.Fatal("expected session file to exist")
}
_ = m.Delete(ctx, key)
// File should be gone
files, _ = os.ReadDir(dir)
jsonCount := 0
for _, f := range files {
if !f.IsDir() {
jsonCount++
}
}
if jsonCount != 0 {
t.Fatalf("expected no session files after delete, got %d", jsonCount)
}
}
// --- UpdateMetadata ---
func TestUpdateMetadata_PartialUpdate(t *testing.T) {
m := NewManager("")
ctx := context.Background()
key := "agent:a1:s1"
m.GetOrCreate(ctx, key)
m.UpdateMetadata(ctx, key, "claude-3", "anthropic", "telegram")
m.mu.RLock()
s := m.sessions[key]
m.mu.RUnlock()
if s.Model != "claude-3" || s.Provider != "anthropic" || s.Channel != "telegram" {
t.Fatalf("metadata mismatch: model=%q provider=%q channel=%q", s.Model, s.Provider, s.Channel)
}
// Partial update: only model, rest unchanged
m.UpdateMetadata(ctx, key, "gpt-4", "", "")
m.mu.RLock()
s = m.sessions[key]
m.mu.RUnlock()
if s.Model != "gpt-4" {
t.Fatalf("model should update: got %q", s.Model)
}
if s.Provider != "anthropic" {
t.Fatalf("provider should stay unchanged: got %q", s.Provider)
}
}
// --- SetSpawnInfo ---
func TestSetSpawnInfo(t *testing.T) {
m := NewManager("")
ctx := context.Background()
key := "agent:a1:s1"
m.GetOrCreate(ctx, key)
m.SetSpawnInfo(ctx, key, "parent-run-123", 2)
m.mu.RLock()
s := m.sessions[key]
m.mu.RUnlock()
if s.SpawnedBy != "parent-run-123" {
t.Fatalf("spawnedBy: got %q", s.SpawnedBy)
}
if s.SpawnDepth != 2 {
t.Fatalf("spawnDepth: got %d", s.SpawnDepth)
}
}
// --- LastPromptTokens ---
func TestLastPromptTokens(t *testing.T) {
m := NewManager("")
ctx := context.Background()
key := "agent:a1:s1"
m.GetOrCreate(ctx, key)
tokens, msgCount := m.GetLastPromptTokens(ctx, key)
if tokens != 0 || msgCount != 0 {
t.Fatal("initial values should be 0")
}
m.SetLastPromptTokens(ctx, key, 5000, 42)
tokens, msgCount = m.GetLastPromptTokens(ctx, key)
if tokens != 5000 || msgCount != 42 {
t.Fatalf("got tokens=%d, msgCount=%d", tokens, msgCount)
}
}
// --- sanitizeFilename ---
func TestSanitizeFilename(t *testing.T) {
got := sanitizeFilename("agent:a1:telegram:direct:123")
want := "agent_a1_telegram_direct_123"
if got != want {
t.Fatalf("sanitizeFilename = %q, want %q", got, want)
}
}