mirror of
https://github.com/tiennm99/goclaw.git
synced 2026-10-11 12:18:59 +00:00
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:
1 parent
a7800604b6
commit
90b7396d74
11 files changed
+2915
-4
No files matched your search
@@ -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
@@ -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)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
|
||||
@@ -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 }
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user