mirror of
https://github.com/tiennm99/goclaw.git
synced 2026-10-11 03:13:24 +00:00
test(cache,sessions,knowledgegraph): slim wave 3 coverage push
Add Wave 3 package coverage: - cache/permission_cache_test.go: PermissionCache 9 methods + invalidation - sessions/key_extra_test.go: session key builders - sessions/manager_extra_test.go: SetHistory, Save, loadAll - knowledgegraph/extractor_helpers_test.go: Extract with mock provider, splitChunks, mergeResults
This commit is contained in:
1 parent
bdbed9a35e
commit
83a6b3c68f
4 files changed
+1180
No files matched your search
+275
@@ -0,0 +1,275 @@
|
||||
package cache
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/bus"
|
||||
)
|
||||
|
||||
func TestPermissionCache_TenantRole(t *testing.T) {
|
||||
pc := NewPermissionCache()
|
||||
defer pc.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
tenantID := uuid.New()
|
||||
userID := "user-1"
|
||||
|
||||
// miss before set
|
||||
_, ok := pc.GetTenantRole(ctx, tenantID, userID)
|
||||
if ok {
|
||||
t.Fatal("expected cache miss before SetTenantRole")
|
||||
}
|
||||
|
||||
// set then hit
|
||||
pc.SetTenantRole(ctx, tenantID, userID, "admin")
|
||||
role, ok := pc.GetTenantRole(ctx, tenantID, userID)
|
||||
if !ok {
|
||||
t.Fatal("expected cache hit after SetTenantRole")
|
||||
}
|
||||
if role != "admin" {
|
||||
t.Fatalf("expected role 'admin', got %q", role)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPermissionCache_AgentAccess(t *testing.T) {
|
||||
pc := NewPermissionCache()
|
||||
defer pc.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
agentID := uuid.New()
|
||||
userID := "user-2"
|
||||
|
||||
// miss
|
||||
_, _, ok := pc.GetAgentAccess(ctx, agentID, userID)
|
||||
if ok {
|
||||
t.Fatal("expected cache miss before SetAgentAccess")
|
||||
}
|
||||
|
||||
// set with allowed=true, role=editor
|
||||
pc.SetAgentAccess(ctx, agentID, userID, true, "editor")
|
||||
|
||||
allowed, role, ok := pc.GetAgentAccess(ctx, agentID, userID)
|
||||
if !ok {
|
||||
t.Fatal("expected cache hit after SetAgentAccess")
|
||||
}
|
||||
if !allowed {
|
||||
t.Fatal("expected allowed=true")
|
||||
}
|
||||
if role != "editor" {
|
||||
t.Fatalf("expected role 'editor', got %q", role)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPermissionCache_AgentAccess_Denied(t *testing.T) {
|
||||
pc := NewPermissionCache()
|
||||
defer pc.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
agentID := uuid.New()
|
||||
userID := "user-denied"
|
||||
|
||||
pc.SetAgentAccess(ctx, agentID, userID, false, "")
|
||||
|
||||
allowed, role, ok := pc.GetAgentAccess(ctx, agentID, userID)
|
||||
if !ok {
|
||||
t.Fatal("expected cache hit")
|
||||
}
|
||||
if allowed {
|
||||
t.Fatal("expected allowed=false")
|
||||
}
|
||||
if role != "" {
|
||||
t.Fatalf("expected empty role, got %q", role)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPermissionCache_TeamAccess(t *testing.T) {
|
||||
pc := NewPermissionCache()
|
||||
defer pc.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
teamID := uuid.New()
|
||||
userID := "user-3"
|
||||
|
||||
// miss
|
||||
_, ok := pc.GetTeamAccess(ctx, teamID, userID)
|
||||
if ok {
|
||||
t.Fatal("expected cache miss before SetTeamAccess")
|
||||
}
|
||||
|
||||
// set true
|
||||
pc.SetTeamAccess(ctx, teamID, userID, true)
|
||||
|
||||
allowed, ok := pc.GetTeamAccess(ctx, teamID, userID)
|
||||
if !ok {
|
||||
t.Fatal("expected cache hit after SetTeamAccess")
|
||||
}
|
||||
if !allowed {
|
||||
t.Fatal("expected allowed=true")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPermissionCache_TeamAccess_False(t *testing.T) {
|
||||
pc := NewPermissionCache()
|
||||
defer pc.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
teamID := uuid.New()
|
||||
userID := "user-4"
|
||||
|
||||
pc.SetTeamAccess(ctx, teamID, userID, false)
|
||||
|
||||
allowed, ok := pc.GetTeamAccess(ctx, teamID, userID)
|
||||
if !ok {
|
||||
t.Fatal("expected cache hit")
|
||||
}
|
||||
if allowed {
|
||||
t.Fatal("expected allowed=false")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPermissionCache_Close_Idempotent(t *testing.T) {
|
||||
pc := NewPermissionCache()
|
||||
pc.Close()
|
||||
pc.Close() // must not panic
|
||||
}
|
||||
|
||||
func TestPermissionCache_HandleInvalidation_TenantUsers(t *testing.T) {
|
||||
pc := NewPermissionCache()
|
||||
defer pc.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
tenantID := uuid.New()
|
||||
|
||||
// populate tenant roles for two users
|
||||
pc.SetTenantRole(ctx, tenantID, "u1", "admin")
|
||||
pc.SetTenantRole(ctx, tenantID, "u2", "viewer")
|
||||
|
||||
// invalidate tenant_users → clears all tenant roles
|
||||
pc.HandleInvalidation(bus.CacheInvalidatePayload{Kind: bus.CacheKindTenantUsers, Key: "u1"})
|
||||
|
||||
// both should be gone
|
||||
if _, ok := pc.GetTenantRole(ctx, tenantID, "u1"); ok {
|
||||
t.Error("u1 tenant role should be cleared after tenant_users invalidation")
|
||||
}
|
||||
if _, ok := pc.GetTenantRole(ctx, tenantID, "u2"); ok {
|
||||
t.Error("u2 tenant role should be cleared after tenant_users invalidation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPermissionCache_HandleInvalidation_AgentAccess_WithKey(t *testing.T) {
|
||||
pc := NewPermissionCache()
|
||||
defer pc.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
agentID1 := uuid.New()
|
||||
agentID2 := uuid.New()
|
||||
|
||||
pc.SetAgentAccess(ctx, agentID1, "u1", true, "admin")
|
||||
pc.SetAgentAccess(ctx, agentID1, "u2", true, "viewer")
|
||||
pc.SetAgentAccess(ctx, agentID2, "u1", true, "admin")
|
||||
|
||||
// invalidate agent_access for agentID1 only
|
||||
pc.HandleInvalidation(bus.CacheInvalidatePayload{Kind: bus.CacheKindAgentAccess, Key: agentID1.String()})
|
||||
|
||||
// agentID1 entries should be gone
|
||||
if _, _, ok := pc.GetAgentAccess(ctx, agentID1, "u1"); ok {
|
||||
t.Error("agentID1:u1 access should be cleared")
|
||||
}
|
||||
if _, _, ok := pc.GetAgentAccess(ctx, agentID1, "u2"); ok {
|
||||
t.Error("agentID1:u2 access should be cleared")
|
||||
}
|
||||
// agentID2 should still be cached
|
||||
if _, _, ok := pc.GetAgentAccess(ctx, agentID2, "u1"); !ok {
|
||||
t.Error("agentID2:u1 access should still be cached")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPermissionCache_HandleInvalidation_AgentAccess_ClearAll(t *testing.T) {
|
||||
pc := NewPermissionCache()
|
||||
defer pc.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
agentID := uuid.New()
|
||||
|
||||
pc.SetAgentAccess(ctx, agentID, "u1", true, "admin")
|
||||
|
||||
// empty Key → clear all
|
||||
pc.HandleInvalidation(bus.CacheInvalidatePayload{Kind: bus.CacheKindAgentAccess, Key: ""})
|
||||
|
||||
if _, _, ok := pc.GetAgentAccess(ctx, agentID, "u1"); ok {
|
||||
t.Error("agent access should be cleared when Key is empty")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPermissionCache_HandleInvalidation_TeamAccess_WithKey(t *testing.T) {
|
||||
pc := NewPermissionCache()
|
||||
defer pc.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
teamID1 := uuid.New()
|
||||
teamID2 := uuid.New()
|
||||
|
||||
pc.SetTeamAccess(ctx, teamID1, "u1", true)
|
||||
pc.SetTeamAccess(ctx, teamID1, "u2", false)
|
||||
pc.SetTeamAccess(ctx, teamID2, "u1", true)
|
||||
|
||||
// invalidate team_access for teamID1 only
|
||||
pc.HandleInvalidation(bus.CacheInvalidatePayload{Kind: bus.CacheKindTeamAccess, Key: teamID1.String()})
|
||||
|
||||
if _, ok := pc.GetTeamAccess(ctx, teamID1, "u1"); ok {
|
||||
t.Error("teamID1:u1 access should be cleared")
|
||||
}
|
||||
if _, ok := pc.GetTeamAccess(ctx, teamID1, "u2"); ok {
|
||||
t.Error("teamID1:u2 access should be cleared")
|
||||
}
|
||||
if _, ok := pc.GetTeamAccess(ctx, teamID2, "u1"); !ok {
|
||||
t.Error("teamID2:u1 access should still be cached")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPermissionCache_HandleInvalidation_TeamAccess_ClearAll(t *testing.T) {
|
||||
pc := NewPermissionCache()
|
||||
defer pc.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
teamID := uuid.New()
|
||||
|
||||
pc.SetTeamAccess(ctx, teamID, "u1", true)
|
||||
|
||||
pc.HandleInvalidation(bus.CacheInvalidatePayload{Kind: bus.CacheKindTeamAccess, Key: ""})
|
||||
|
||||
if _, ok := pc.GetTeamAccess(ctx, teamID, "u1"); ok {
|
||||
t.Error("team access should be cleared when Key is empty")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPermissionCache_HandleInvalidation_UnknownKind(t *testing.T) {
|
||||
pc := NewPermissionCache()
|
||||
defer pc.Close()
|
||||
|
||||
// unknown kind should not panic
|
||||
pc.HandleInvalidation(bus.CacheInvalidatePayload{Kind: "unknown_kind", Key: "anything"})
|
||||
}
|
||||
|
||||
func TestPermissionCache_MultipleUsers_Isolation(t *testing.T) {
|
||||
pc := NewPermissionCache()
|
||||
defer pc.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
tenantID := uuid.New()
|
||||
|
||||
pc.SetTenantRole(ctx, tenantID, "user-a", "admin")
|
||||
pc.SetTenantRole(ctx, tenantID, "user-b", "viewer")
|
||||
|
||||
roleA, okA := pc.GetTenantRole(ctx, tenantID, "user-a")
|
||||
roleB, okB := pc.GetTenantRole(ctx, tenantID, "user-b")
|
||||
|
||||
if !okA || roleA != "admin" {
|
||||
t.Errorf("user-a: expected admin, got %q ok=%v", roleA, okA)
|
||||
}
|
||||
if !okB || roleB != "viewer" {
|
||||
t.Errorf("user-b: expected viewer, got %q ok=%v", roleB, okB)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,394 @@
|
||||
package knowledgegraph
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/store"
|
||||
)
|
||||
|
||||
// --- mockProvider implements providers.Provider for unit tests (no network) ---
|
||||
|
||||
type mockProvider struct {
|
||||
response providers.ChatResponse
|
||||
err error
|
||||
}
|
||||
|
||||
func (m *mockProvider) Chat(_ context.Context, _ providers.ChatRequest) (*providers.ChatResponse, error) {
|
||||
if m.err != nil {
|
||||
return nil, m.err
|
||||
}
|
||||
r := m.response
|
||||
return &r, nil
|
||||
}
|
||||
|
||||
func (m *mockProvider) ChatStream(_ context.Context, _ providers.ChatRequest, _ func(providers.StreamChunk)) (*providers.ChatResponse, error) {
|
||||
if m.err != nil {
|
||||
return nil, m.err
|
||||
}
|
||||
r := m.response
|
||||
return &r, nil
|
||||
}
|
||||
|
||||
func (m *mockProvider) DefaultModel() string { return "mock-model" }
|
||||
func (m *mockProvider) Name() string { return "mock" }
|
||||
|
||||
// --- NewExtractor ---
|
||||
|
||||
func TestNewExtractor_DefaultConfidence(t *testing.T) {
|
||||
e := NewExtractor(&mockProvider{}, "model", 0)
|
||||
if e.minConfidence != 0.75 {
|
||||
t.Errorf("expected default minConfidence=0.75, got %f", e.minConfidence)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewExtractor_CustomConfidence(t *testing.T) {
|
||||
e := NewExtractor(&mockProvider{}, "model", 0.9)
|
||||
if e.minConfidence != 0.9 {
|
||||
t.Errorf("expected minConfidence=0.9, got %f", e.minConfidence)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewExtractor_NegativeConfidence(t *testing.T) {
|
||||
e := NewExtractor(&mockProvider{}, "model", -1)
|
||||
if e.minConfidence != 0.75 {
|
||||
t.Errorf("expected default minConfidence=0.75 for negative input, got %f", e.minConfidence)
|
||||
}
|
||||
}
|
||||
|
||||
// --- splitChunks ---
|
||||
|
||||
func TestSplitChunks_ShortText(t *testing.T) {
|
||||
text := "short text"
|
||||
chunks := splitChunks(text, 100)
|
||||
if len(chunks) != 1 {
|
||||
t.Fatalf("expected 1 chunk, got %d", len(chunks))
|
||||
}
|
||||
if chunks[0] != text {
|
||||
t.Errorf("chunk mismatch: got %q", chunks[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitChunks_ExactLimit(t *testing.T) {
|
||||
text := strings.Repeat("a", 100)
|
||||
chunks := splitChunks(text, 100)
|
||||
if len(chunks) != 1 {
|
||||
t.Fatalf("expected 1 chunk for exact-limit text, got %d", len(chunks))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitChunks_LongTextSplitsAtParagraph(t *testing.T) {
|
||||
// Two paragraphs each >50 chars; max=100 → should split.
|
||||
para1 := strings.Repeat("a", 60)
|
||||
para2 := strings.Repeat("b", 60)
|
||||
text := para1 + "\n\n" + para2
|
||||
|
||||
chunks := splitChunks(text, 100)
|
||||
if len(chunks) < 2 {
|
||||
t.Fatalf("expected at least 2 chunks for text > maxChars, got %d", len(chunks))
|
||||
}
|
||||
for i, ch := range chunks {
|
||||
if strings.TrimSpace(ch) == "" {
|
||||
t.Errorf("chunk %d is empty", i)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitChunks_NoParagraphBreak(t *testing.T) {
|
||||
// No \n\n → hard cut at maxChars boundary.
|
||||
text := strings.Repeat("x", 200)
|
||||
chunks := splitChunks(text, 100)
|
||||
if len(chunks) < 2 {
|
||||
t.Fatalf("expected at least 2 chunks, got %d", len(chunks))
|
||||
}
|
||||
// Verify no data loss.
|
||||
var total int
|
||||
for _, ch := range chunks {
|
||||
total += len(ch)
|
||||
}
|
||||
if total != len(text) {
|
||||
t.Errorf("chunks total length %d != original %d", total, len(text))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitChunks_Empty(t *testing.T) {
|
||||
chunks := splitChunks("", 100)
|
||||
// len("") <= maxChars → single (empty) chunk.
|
||||
if len(chunks) != 1 {
|
||||
t.Fatalf("expected 1 chunk for empty string, got %d", len(chunks))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitChunks_MultipleChunks(t *testing.T) {
|
||||
// 5 paragraphs of 40 chars each; max=50 → multiple chunks.
|
||||
var paras []string
|
||||
for i := 0; i < 5; i++ {
|
||||
paras = append(paras, strings.Repeat(fmt.Sprintf("%d", i), 40))
|
||||
}
|
||||
text := strings.Join(paras, "\n\n")
|
||||
chunks := splitChunks(text, 50)
|
||||
if len(chunks) < 3 {
|
||||
t.Fatalf("expected multiple chunks, got %d", len(chunks))
|
||||
}
|
||||
}
|
||||
|
||||
// --- mergeResults ---
|
||||
|
||||
func TestMergeResults_Empty(t *testing.T) {
|
||||
result := mergeResults(&ExtractionResult{}, &ExtractionResult{})
|
||||
if len(result.Entities) != 0 || len(result.Relations) != 0 {
|
||||
t.Error("merging two empty results should produce empty result")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMergeResults_DeduplicatesEntitiesByExternalID(t *testing.T) {
|
||||
a := &ExtractionResult{
|
||||
Entities: []store.Entity{
|
||||
{ExternalID: "alice", Name: "Alice", Confidence: 0.8},
|
||||
{ExternalID: "bob", Name: "Bob", Confidence: 0.7},
|
||||
},
|
||||
}
|
||||
b := &ExtractionResult{
|
||||
Entities: []store.Entity{
|
||||
// Same external_id as alice but higher confidence — should win.
|
||||
{ExternalID: "alice", Name: "Alice Updated", Confidence: 0.95},
|
||||
{ExternalID: "carol", Name: "Carol", Confidence: 0.9},
|
||||
},
|
||||
}
|
||||
|
||||
result := mergeResults(a, b)
|
||||
entityMap := make(map[string]store.Entity)
|
||||
for _, e := range result.Entities {
|
||||
entityMap[e.ExternalID] = e
|
||||
}
|
||||
|
||||
if len(entityMap) != 3 {
|
||||
t.Fatalf("expected 3 unique entities, got %d", len(result.Entities))
|
||||
}
|
||||
if entityMap["alice"].Confidence != 0.95 {
|
||||
t.Errorf("alice: expected confidence 0.95 (higher wins), got %f", entityMap["alice"].Confidence)
|
||||
}
|
||||
if _, ok := entityMap["bob"]; !ok {
|
||||
t.Error("bob should be in merged result")
|
||||
}
|
||||
if _, ok := entityMap["carol"]; !ok {
|
||||
t.Error("carol should be in merged result")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMergeResults_KeepsHigherConfidence(t *testing.T) {
|
||||
a := &ExtractionResult{
|
||||
Entities: []store.Entity{{ExternalID: "alice", Confidence: 0.95}},
|
||||
}
|
||||
b := &ExtractionResult{
|
||||
Entities: []store.Entity{{ExternalID: "alice", Confidence: 0.5}},
|
||||
}
|
||||
|
||||
result := mergeResults(a, b)
|
||||
if len(result.Entities) != 1 {
|
||||
t.Fatalf("expected 1 entity, got %d", len(result.Entities))
|
||||
}
|
||||
if result.Entities[0].Confidence != 0.95 {
|
||||
t.Errorf("expected higher confidence to be kept, got %f", result.Entities[0].Confidence)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMergeResults_DeduplicatesRelations(t *testing.T) {
|
||||
rel := store.Relation{SourceEntityID: "alice", RelationType: "knows", TargetEntityID: "bob", Confidence: 0.8}
|
||||
relHigher := store.Relation{SourceEntityID: "alice", RelationType: "knows", TargetEntityID: "bob", Confidence: 0.95}
|
||||
unrelated := store.Relation{SourceEntityID: "carol", RelationType: "works_at", TargetEntityID: "acme", Confidence: 0.9}
|
||||
|
||||
result := mergeResults(
|
||||
&ExtractionResult{Relations: []store.Relation{rel}},
|
||||
&ExtractionResult{Relations: []store.Relation{relHigher, unrelated}},
|
||||
)
|
||||
|
||||
if len(result.Relations) != 2 {
|
||||
t.Fatalf("expected 2 unique relations, got %d", len(result.Relations))
|
||||
}
|
||||
for _, r := range result.Relations {
|
||||
if r.SourceEntityID == "alice" && r.RelationType == "knows" && r.Confidence != 0.95 {
|
||||
t.Errorf("alice→knows→bob: expected 0.95, got %f", r.Confidence)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMergeResults_OneSideEmpty(t *testing.T) {
|
||||
a := &ExtractionResult{
|
||||
Entities: []store.Entity{{ExternalID: "alice", Confidence: 0.9}},
|
||||
}
|
||||
result := mergeResults(a, &ExtractionResult{})
|
||||
if len(result.Entities) != 1 {
|
||||
t.Errorf("expected 1 entity from non-empty side, got %d", len(result.Entities))
|
||||
}
|
||||
}
|
||||
|
||||
// --- Extract with mock provider ---
|
||||
|
||||
func TestExtract_ShortText_Success(t *testing.T) {
|
||||
respJSON := `{"entities":[{"external_id":"alice","name":"Alice","entity_type":"person","confidence":0.9}],"relations":[]}`
|
||||
e := NewExtractor(&mockProvider{response: providers.ChatResponse{Content: respJSON, FinishReason: "stop"}}, "m", 0.8)
|
||||
|
||||
result, err := e.Extract(context.Background(), "Alice works at Acme.")
|
||||
if err != nil {
|
||||
t.Fatalf("Extract returned error: %v", err)
|
||||
}
|
||||
if len(result.Entities) != 1 {
|
||||
t.Fatalf("expected 1 entity, got %d", len(result.Entities))
|
||||
}
|
||||
if result.Entities[0].ExternalID != "alice" {
|
||||
t.Errorf("expected external_id 'alice', got %q", result.Entities[0].ExternalID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtract_FiltersLowConfidence(t *testing.T) {
|
||||
respJSON := `{"entities":[
|
||||
{"external_id":"high","name":"High","entity_type":"person","confidence":0.9},
|
||||
{"external_id":"low","name":"Low","entity_type":"person","confidence":0.5}
|
||||
],"relations":[]}`
|
||||
e := NewExtractor(&mockProvider{response: providers.ChatResponse{Content: respJSON, FinishReason: "stop"}}, "m", 0.8)
|
||||
|
||||
result, err := e.Extract(context.Background(), "text")
|
||||
if err != nil {
|
||||
t.Fatalf("Extract returned error: %v", err)
|
||||
}
|
||||
if len(result.Entities) != 1 || result.Entities[0].ExternalID != "high" {
|
||||
t.Errorf("expected 1 high-confidence entity, got %d: %+v", len(result.Entities), result.Entities)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtract_NormalizesFields(t *testing.T) {
|
||||
respJSON := `{"entities":[{"external_id":" ALICE ","name":" Alice ","entity_type":" PERSON ","confidence":0.9}],"relations":[]}`
|
||||
e := NewExtractor(&mockProvider{response: providers.ChatResponse{Content: respJSON, FinishReason: "stop"}}, "m", 0.8)
|
||||
|
||||
result, err := e.Extract(context.Background(), "text")
|
||||
if err != nil {
|
||||
t.Fatalf("Extract returned error: %v", err)
|
||||
}
|
||||
if len(result.Entities) != 1 {
|
||||
t.Fatalf("expected 1 entity, got %d", len(result.Entities))
|
||||
}
|
||||
ent := result.Entities[0]
|
||||
if ent.ExternalID != "alice" {
|
||||
t.Errorf("external_id not lowercased/trimmed: %q", ent.ExternalID)
|
||||
}
|
||||
if ent.Name != "Alice" {
|
||||
t.Errorf("name not trimmed: %q", ent.Name)
|
||||
}
|
||||
if ent.EntityType != "person" {
|
||||
t.Errorf("entity_type not lowercased: %q", ent.EntityType)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtract_ProviderError(t *testing.T) {
|
||||
e := NewExtractor(&mockProvider{err: fmt.Errorf("connection refused")}, "m", 0.8)
|
||||
_, err := e.Extract(context.Background(), "text")
|
||||
if err == nil {
|
||||
t.Fatal("expected error when provider fails, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtract_InvalidJSON(t *testing.T) {
|
||||
e := NewExtractor(&mockProvider{response: providers.ChatResponse{Content: "not json at all", FinishReason: "stop"}}, "m", 0.8)
|
||||
_, err := e.Extract(context.Background(), "text")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for invalid JSON response, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtract_CodeBlockStripped(t *testing.T) {
|
||||
respJSON := "```json\n{\"entities\":[{\"external_id\":\"alice\",\"name\":\"Alice\",\"entity_type\":\"person\",\"confidence\":0.9}],\"relations\":[]}\n```"
|
||||
e := NewExtractor(&mockProvider{response: providers.ChatResponse{Content: respJSON, FinishReason: "stop"}}, "m", 0.8)
|
||||
|
||||
result, err := e.Extract(context.Background(), "Alice.")
|
||||
if err != nil {
|
||||
t.Fatalf("Extract returned error: %v", err)
|
||||
}
|
||||
if len(result.Entities) != 1 {
|
||||
t.Fatalf("expected 1 entity, got %d", len(result.Entities))
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtract_RelationsNormalized(t *testing.T) {
|
||||
respJSON := `{"entities":[],"relations":[
|
||||
{"source_entity_id":" ALICE ","relation_type":" KNOWS ","target_entity_id":" BOB ","confidence":0.9}
|
||||
]}`
|
||||
e := NewExtractor(&mockProvider{response: providers.ChatResponse{Content: respJSON, FinishReason: "stop"}}, "m", 0.8)
|
||||
|
||||
result, err := e.Extract(context.Background(), "text")
|
||||
if err != nil {
|
||||
t.Fatalf("Extract returned error: %v", err)
|
||||
}
|
||||
if len(result.Relations) != 1 {
|
||||
t.Fatalf("expected 1 relation, got %d", len(result.Relations))
|
||||
}
|
||||
rel := result.Relations[0]
|
||||
if rel.SourceEntityID != "alice" {
|
||||
t.Errorf("source_entity_id not normalized: %q", rel.SourceEntityID)
|
||||
}
|
||||
if rel.RelationType != "knows" {
|
||||
t.Errorf("relation_type not normalized: %q", rel.RelationType)
|
||||
}
|
||||
if rel.TargetEntityID != "bob" {
|
||||
t.Errorf("target_entity_id not normalized: %q", rel.TargetEntityID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtract_LongText_SplitsIntoChunks(t *testing.T) {
|
||||
// Build text longer than maxChunkChars (12000) with paragraph breaks.
|
||||
para := strings.Repeat("word ", 200) // ~1000 chars per para
|
||||
var paras []string
|
||||
for i := 0; i < 15; i++ {
|
||||
paras = append(paras, para)
|
||||
}
|
||||
longText := strings.Join(paras, "\n\n") // ~15000+ chars
|
||||
|
||||
callCount := 0
|
||||
respJSON := `{"entities":[{"external_id":"e1","name":"E1","entity_type":"person","confidence":0.9}],"relations":[]}`
|
||||
mock := &countingMockProvider{
|
||||
response: providers.ChatResponse{Content: respJSON, FinishReason: "stop"},
|
||||
onCall: func() { callCount++ },
|
||||
}
|
||||
e := NewExtractor(mock, "m", 0.8)
|
||||
|
||||
result, err := e.Extract(context.Background(), longText)
|
||||
if err != nil {
|
||||
t.Fatalf("Extract returned error: %v", err)
|
||||
}
|
||||
if callCount < 2 {
|
||||
t.Errorf("expected multiple LLM calls for long text, got %d", callCount)
|
||||
}
|
||||
// All chunks return same external_id → merged to 1 entity.
|
||||
if len(result.Entities) != 1 {
|
||||
t.Errorf("expected 1 deduplicated entity, got %d", len(result.Entities))
|
||||
}
|
||||
}
|
||||
|
||||
// countingMockProvider counts Chat calls for the long-text split test.
|
||||
type countingMockProvider struct {
|
||||
response providers.ChatResponse
|
||||
err error
|
||||
onCall func()
|
||||
}
|
||||
|
||||
func (m *countingMockProvider) Chat(_ context.Context, _ providers.ChatRequest) (*providers.ChatResponse, error) {
|
||||
if m.onCall != nil {
|
||||
m.onCall()
|
||||
}
|
||||
if m.err != nil {
|
||||
return nil, m.err
|
||||
}
|
||||
r := m.response
|
||||
return &r, nil
|
||||
}
|
||||
|
||||
func (m *countingMockProvider) ChatStream(_ context.Context, _ providers.ChatRequest, _ func(providers.StreamChunk)) (*providers.ChatResponse, error) {
|
||||
r := m.response
|
||||
return &r, nil
|
||||
}
|
||||
|
||||
func (m *countingMockProvider) DefaultModel() string { return "counting-mock" }
|
||||
func (m *countingMockProvider) Name() string { return "counting-mock" }
|
||||
@@ -0,0 +1,277 @@
|
||||
package sessions
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// TestBuildSessionKey covers the canonical DM and group formats.
|
||||
func TestBuildSessionKey(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
agentID string
|
||||
channel string
|
||||
kind PeerKind
|
||||
chatID string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "DM session",
|
||||
agentID: "default",
|
||||
channel: "telegram",
|
||||
kind: PeerDirect,
|
||||
chatID: "386246614",
|
||||
want: "agent:default:telegram:direct:386246614",
|
||||
},
|
||||
{
|
||||
name: "group session",
|
||||
agentID: "default",
|
||||
channel: "telegram",
|
||||
kind: PeerGroup,
|
||||
chatID: "-100123456",
|
||||
want: "agent:default:telegram:group:-100123456",
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := BuildSessionKey(tt.agentID, tt.channel, tt.kind, tt.chatID)
|
||||
if got != tt.want {
|
||||
t.Errorf("BuildSessionKey = %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildGroupTopicSessionKey covers forum topic key format.
|
||||
func TestBuildGroupTopicSessionKey(t *testing.T) {
|
||||
got := BuildGroupTopicSessionKey("default", "telegram", "-100123456", 99)
|
||||
want := "agent:default:telegram:group:-100123456:topic:99"
|
||||
if got != want {
|
||||
t.Errorf("BuildGroupTopicSessionKey = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildDMThreadSessionKey covers DM thread key format.
|
||||
func TestBuildDMThreadSessionKey(t *testing.T) {
|
||||
got := BuildDMThreadSessionKey("my-agent", "telegram", "386246614", 7)
|
||||
want := "agent:my-agent:telegram:direct:386246614:thread:7"
|
||||
if got != want {
|
||||
t.Errorf("BuildDMThreadSessionKey = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildScopedThreadSessionKey covers string-based thread IDs (Slack timestamps).
|
||||
func TestBuildScopedThreadSessionKey(t *testing.T) {
|
||||
got := BuildScopedThreadSessionKey("bot", "slack", PeerDirect, "U12345", "1712345678.000100")
|
||||
want := "agent:bot:slack:direct:U12345:thread:1712345678.000100"
|
||||
if got != want {
|
||||
t.Errorf("BuildScopedThreadSessionKey = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildSubagentSessionKey covers the subagent key format.
|
||||
func TestBuildSubagentSessionKey(t *testing.T) {
|
||||
got := BuildSubagentSessionKey("default", "my-task")
|
||||
want := "agent:default:subagent:my-task"
|
||||
if got != want {
|
||||
t.Errorf("BuildSubagentSessionKey = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildTeamSessionKey covers team session key format.
|
||||
func TestBuildTeamSessionKey(t *testing.T) {
|
||||
got := BuildTeamSessionKey("my-agent", "team-42", "chat-99")
|
||||
want := "agent:my-agent:team:team-42:chat-99"
|
||||
if got != want {
|
||||
t.Errorf("BuildTeamSessionKey = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestIsTeamSession distinguishes team vs non-team keys.
|
||||
func TestIsTeamSession(t *testing.T) {
|
||||
tests := []struct {
|
||||
key string
|
||||
want bool
|
||||
}{
|
||||
{"agent:my-agent:team:t1:c1", true},
|
||||
{"agent:my-agent:cron:job-1", false},
|
||||
{"agent:my-agent:subagent:label", false},
|
||||
{"agent:my-agent:heartbeat", false},
|
||||
{"not-a-session-key", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.key, func(t *testing.T) {
|
||||
if got := IsTeamSession(tt.key); got != tt.want {
|
||||
t.Errorf("IsTeamSession(%q) = %v, want %v", tt.key, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildCronSessionKey_DoublePrefix guards against double-prefixing.
|
||||
func TestBuildCronSessionKey_DoublePrefix(t *testing.T) {
|
||||
// If jobID is already a canonical session key, only the rest part is used.
|
||||
canonical := "agent:my-agent:cron:existing-job"
|
||||
got := BuildCronSessionKey("other-agent", canonical)
|
||||
// Should use "cron:existing-job" as rest, not re-wrap the whole canonical key.
|
||||
if strings.Contains(got, "agent:my-agent") {
|
||||
t.Errorf("double-prefix not guarded: got %q", got)
|
||||
}
|
||||
if !strings.HasPrefix(got, "agent:other-agent:") {
|
||||
t.Errorf("expected other-agent prefix, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildAgentMainSessionKey covers default and custom main keys.
|
||||
func TestBuildAgentMainSessionKey(t *testing.T) {
|
||||
tests := []struct {
|
||||
agentID string
|
||||
mainKey string
|
||||
want string
|
||||
}{
|
||||
{"my-agent", "", "agent:my-agent:main"},
|
||||
{"my-agent", "custom-main", "agent:my-agent:custom-main"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.mainKey, func(t *testing.T) {
|
||||
got := BuildAgentMainSessionKey(tt.agentID, tt.mainKey)
|
||||
if got != tt.want {
|
||||
t.Errorf("BuildAgentMainSessionKey = %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildScopedSessionKey delegates to BuildSessionKey; verify output matches.
|
||||
func TestBuildScopedSessionKey(t *testing.T) {
|
||||
got := BuildScopedSessionKey("default", "telegram", PeerGroup, "-100123")
|
||||
want := BuildSessionKey("default", "telegram", PeerGroup, "-100123")
|
||||
if got != want {
|
||||
t.Errorf("BuildScopedSessionKey = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestIsSubagentSession distinguishes subagent vs non-subagent keys.
|
||||
func TestIsSubagentSession(t *testing.T) {
|
||||
tests := []struct {
|
||||
key string
|
||||
want bool
|
||||
}{
|
||||
{"agent:default:subagent:my-label", true},
|
||||
{"agent:default:SUBAGENT:label", true}, // case-insensitive
|
||||
{"agent:default:cron:job-1", false},
|
||||
{"agent:default:team:t1:c1", false},
|
||||
{"invalid", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.key, func(t *testing.T) {
|
||||
if got := IsSubagentSession(tt.key); got != tt.want {
|
||||
t.Errorf("IsSubagentSession(%q) = %v, want %v", tt.key, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestIsCronSession distinguishes cron vs non-cron keys.
|
||||
func TestIsCronSession(t *testing.T) {
|
||||
tests := []struct {
|
||||
key string
|
||||
want bool
|
||||
}{
|
||||
{"agent:default:cron:reminder-123", true},
|
||||
{"agent:default:CRON:job", true}, // case-insensitive
|
||||
{"agent:default:subagent:x", false},
|
||||
{"agent:default:heartbeat", false},
|
||||
{"invalid", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.key, func(t *testing.T) {
|
||||
if got := IsCronSession(tt.key); got != tt.want {
|
||||
t.Errorf("IsCronSession(%q) = %v, want %v", tt.key, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestIsHeartbeatSession distinguishes heartbeat vs other keys.
|
||||
func TestIsHeartbeatSession(t *testing.T) {
|
||||
tests := []struct {
|
||||
key string
|
||||
want bool
|
||||
}{
|
||||
{"agent:default:heartbeat", true},
|
||||
{"agent:default:heartbeat:1712345678000", true},
|
||||
{"agent:default:cron:job", false},
|
||||
{"invalid", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.key, func(t *testing.T) {
|
||||
if got := IsHeartbeatSession(tt.key); got != tt.want {
|
||||
t.Errorf("IsHeartbeatSession(%q) = %v, want %v", tt.key, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildWSSessionKey and TestIsWSSession cover WS key helpers.
|
||||
func TestBuildWSSessionKey(t *testing.T) {
|
||||
got := BuildWSSessionKey("default", "conv-abc")
|
||||
want := "agent:default:ws:direct:conv-abc"
|
||||
if got != want {
|
||||
t.Errorf("BuildWSSessionKey = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsWSSession(t *testing.T) {
|
||||
tests := []struct {
|
||||
key string
|
||||
want bool
|
||||
}{
|
||||
{"agent:default:ws:direct:conv-1", true},
|
||||
{"agent:default:ws-legacy:room-2", true}, // legacy ws- prefix
|
||||
{"agent:default:telegram:direct:123", false},
|
||||
{"invalid", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.key, func(t *testing.T) {
|
||||
if got := IsWSSession(tt.key); got != tt.want {
|
||||
t.Errorf("IsWSSession(%q) = %v, want %v", tt.key, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestPeerKindFromGroup covers the helper.
|
||||
func TestPeerKindFromGroup(t *testing.T) {
|
||||
if PeerKindFromGroup(true) != PeerGroup {
|
||||
t.Error("expected PeerGroup for isGroup=true")
|
||||
}
|
||||
if PeerKindFromGroup(false) != PeerDirect {
|
||||
t.Error("expected PeerDirect for isGroup=false")
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseSessionKey_InvalidFormats covers non-canonical keys.
|
||||
func TestParseSessionKey_InvalidFormats(t *testing.T) {
|
||||
tests := []string{
|
||||
"",
|
||||
"noprefix",
|
||||
"agent:",
|
||||
"agent:only-one-part",
|
||||
"other:prefix:rest",
|
||||
}
|
||||
for _, key := range tests {
|
||||
agentID, rest := ParseSessionKey(key)
|
||||
if key == "agent:only-one-part" {
|
||||
// Only 2 parts when split by ":" with N=3 — depends on SplitN behavior
|
||||
// SplitN("agent:only-one-part", ":", 3) → ["agent", "only-one-part"] len=2 < 3 → ("","")
|
||||
if agentID != "" || rest != "" {
|
||||
t.Errorf("ParseSessionKey(%q) = (%q, %q), want ('', '')", key, agentID, rest)
|
||||
}
|
||||
continue
|
||||
}
|
||||
if agentID != "" || rest != "" {
|
||||
t.Errorf("ParseSessionKey(%q) = (%q, %q), want ('', '')", key, agentID, rest)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,234 @@
|
||||
package sessions
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
)
|
||||
|
||||
// TestSetHistory replaces session history entirely.
|
||||
func TestSetHistory(t *testing.T) {
|
||||
m := NewManager("")
|
||||
ctx := context.Background()
|
||||
key := "agent:a1:s1"
|
||||
|
||||
m.GetOrCreate(ctx, key)
|
||||
m.AddMessage(ctx, key, providers.Message{Role: "user", Content: "old"})
|
||||
|
||||
newMsgs := []providers.Message{
|
||||
{Role: "user", Content: "new1"},
|
||||
{Role: "assistant", Content: "new2"},
|
||||
}
|
||||
m.SetHistory(ctx, key, newMsgs)
|
||||
|
||||
history := m.GetHistory(ctx, key)
|
||||
if len(history) != 2 {
|
||||
t.Fatalf("expected 2 messages after SetHistory, got %d", len(history))
|
||||
}
|
||||
if history[0].Content != "new1" || history[1].Content != "new2" {
|
||||
t.Errorf("unexpected history content: %+v", history)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSetHistory_NonExistentSession is a silent no-op.
|
||||
func TestSetHistory_NonExistentSession(t *testing.T) {
|
||||
m := NewManager("")
|
||||
// Must not panic
|
||||
m.SetHistory(context.Background(), "nonexistent", []providers.Message{{Role: "user", Content: "x"}})
|
||||
}
|
||||
|
||||
// TestGetSummary_NonExistentSession returns empty string.
|
||||
func TestGetSummary_NonExistentSession(t *testing.T) {
|
||||
m := NewManager("")
|
||||
got := m.GetSummary(context.Background(), "nonexistent")
|
||||
if got != "" {
|
||||
t.Errorf("expected empty string, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestGetCompactionCount_NonExistentSession returns 0.
|
||||
func TestGetCompactionCount_NonExistentSession(t *testing.T) {
|
||||
m := NewManager("")
|
||||
got := m.GetCompactionCount(context.Background(), "nonexistent")
|
||||
if got != 0 {
|
||||
t.Errorf("expected 0, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestGetContextWindow_NonExistentSession returns 0.
|
||||
func TestGetContextWindow_NonExistentSession(t *testing.T) {
|
||||
m := NewManager("")
|
||||
got := m.GetContextWindow(context.Background(), "nonexistent")
|
||||
if got != 0 {
|
||||
t.Errorf("expected 0, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestGetLastPromptTokens_NonExistentSession returns (0, 0).
|
||||
func TestGetLastPromptTokens_NonExistentSession(t *testing.T) {
|
||||
m := NewManager("")
|
||||
tokens, msgCount := m.GetLastPromptTokens(context.Background(), "nonexistent")
|
||||
if tokens != 0 || msgCount != 0 {
|
||||
t.Errorf("expected (0, 0), got (%d, %d)", tokens, msgCount)
|
||||
}
|
||||
}
|
||||
|
||||
// TestTruncateHistory_NonExistentSession is a silent no-op.
|
||||
func TestTruncateHistory_NonExistentSession(t *testing.T) {
|
||||
m := NewManager("")
|
||||
// Must not panic
|
||||
m.TruncateHistory(context.Background(), "nonexistent", 5)
|
||||
}
|
||||
|
||||
// TestTruncateHistory_LessThanKeep keeps all messages when count < keepLast.
|
||||
func TestTruncateHistory_LessThanKeep(t *testing.T) {
|
||||
m := NewManager("")
|
||||
ctx := context.Background()
|
||||
key := "agent:a1:s1"
|
||||
|
||||
m.AddMessage(ctx, key, providers.Message{Role: "user", Content: "only"})
|
||||
m.TruncateHistory(ctx, key, 10) // keepLast > actual count
|
||||
|
||||
history := m.GetHistory(ctx, key)
|
||||
if len(history) != 1 {
|
||||
t.Fatalf("expected 1 message (no truncation needed), got %d", len(history))
|
||||
}
|
||||
}
|
||||
|
||||
// TestLastUsedChannel_Empty returns ("","") for an empty manager.
|
||||
func TestLastUsedChannel_Empty(t *testing.T) {
|
||||
m := NewManager("")
|
||||
ch, chatID := m.LastUsedChannel(context.Background(), "my-agent")
|
||||
if ch != "" || chatID != "" {
|
||||
t.Errorf("expected ('', ''), got (%q, %q)", ch, chatID)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLastUsedChannel_FindsMostRecent returns the most recently updated channel session.
|
||||
func TestLastUsedChannel_FindsMostRecent(t *testing.T) {
|
||||
m := NewManager("")
|
||||
ctx := context.Background()
|
||||
|
||||
key1 := "agent:my-agent:telegram:direct:111"
|
||||
key2 := "agent:my-agent:telegram:direct:222"
|
||||
|
||||
s1 := m.GetOrCreate(ctx, key1)
|
||||
s1.Updated = time.Now().Add(-10 * time.Second)
|
||||
|
||||
s2 := m.GetOrCreate(ctx, key2)
|
||||
s2.Updated = time.Now()
|
||||
|
||||
ch, chatID := m.LastUsedChannel(ctx, "my-agent")
|
||||
if ch != "telegram" {
|
||||
t.Errorf("expected channel 'telegram', got %q", ch)
|
||||
}
|
||||
if chatID != "222" {
|
||||
t.Errorf("expected chatID '222', got %q", chatID)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLastUsedChannel_SkipsCronAndSubagent skips non-channel sessions.
|
||||
func TestLastUsedChannel_SkipsCronAndSubagent(t *testing.T) {
|
||||
m := NewManager("")
|
||||
ctx := context.Background()
|
||||
|
||||
m.GetOrCreate(ctx, "agent:my-agent:cron:job-1")
|
||||
m.GetOrCreate(ctx, "agent:my-agent:subagent:worker")
|
||||
|
||||
ch, chatID := m.LastUsedChannel(ctx, "my-agent")
|
||||
if ch != "" || chatID != "" {
|
||||
t.Errorf("expected ('', '') for only cron/subagent sessions, got (%q, %q)", ch, chatID)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLastUsedChannel_AgentIsolation ensures different agents don't cross-contaminate.
|
||||
func TestLastUsedChannel_AgentIsolation(t *testing.T) {
|
||||
m := NewManager("")
|
||||
ctx := context.Background()
|
||||
|
||||
m.GetOrCreate(ctx, "agent:agent-A:telegram:direct:100")
|
||||
m.GetOrCreate(ctx, "agent:agent-B:telegram:direct:200")
|
||||
|
||||
ch, chatID := m.LastUsedChannel(ctx, "agent-A")
|
||||
if ch != "telegram" || chatID != "100" {
|
||||
t.Errorf("agent-A: expected (telegram, 100), got (%q, %q)", ch, chatID)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSave_MissingSession_NoError returns nil for a key that doesn't exist.
|
||||
func TestSave_MissingSession_NoError(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
m := NewManager(dir)
|
||||
ctx := context.Background()
|
||||
|
||||
if err := m.Save(ctx, "agent:ghost:s1"); err != nil {
|
||||
t.Fatalf("expected nil error for missing session, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSave_NoStorage_NoError returns nil when storage path is "".
|
||||
func TestSave_NoStorage_NoError(t *testing.T) {
|
||||
m := NewManager("")
|
||||
ctx := context.Background()
|
||||
key := "agent:a1:s1"
|
||||
|
||||
m.GetOrCreate(ctx, key)
|
||||
if err := m.Save(ctx, key); err != nil {
|
||||
t.Fatalf("expected nil error with no storage, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestDelete_NonExistentKey returns nil (idempotent).
|
||||
func TestDelete_NonExistentKey(t *testing.T) {
|
||||
m := NewManager("")
|
||||
if err := m.Delete(context.Background(), "nonexistent"); err != nil {
|
||||
t.Fatalf("expected nil error deleting non-existent session, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestDelete_WithStorage_NonExistentFile returns nil when the file was never saved.
|
||||
func TestDelete_WithStorage_NonExistentFile(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
m := NewManager(dir)
|
||||
ctx := context.Background()
|
||||
|
||||
m.GetOrCreate(ctx, "agent:a1:s1")
|
||||
// Delete without prior Save — file does not exist on disk.
|
||||
if err := m.Delete(ctx, "agent:a1:s1"); err != nil {
|
||||
t.Fatalf("expected nil for non-existent file, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewManager_LoadAll_SkipsBadFiles verifies corrupt JSON files are silently skipped.
|
||||
func TestNewManager_LoadAll_SkipsBadFiles(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
|
||||
// Write a corrupt JSON file directly.
|
||||
if err := os.WriteFile(dir+"/corrupt.json", []byte("not valid json"), 0644); err != nil {
|
||||
t.Fatalf("setup: write corrupt file: %v", err)
|
||||
}
|
||||
|
||||
// NewManager must not crash; corrupt file is skipped.
|
||||
m := NewManager(dir)
|
||||
all := m.List(context.Background(), "")
|
||||
if len(all) != 0 {
|
||||
t.Errorf("expected 0 sessions (corrupt file skipped), got %d", len(all))
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewManager_LoadAll_SkipsNonJSON verifies non-JSON files are ignored.
|
||||
func TestNewManager_LoadAll_SkipsNonJSON(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
|
||||
if err := os.WriteFile(dir+"/readme.txt", []byte("ignore me"), 0644); err != nil {
|
||||
t.Fatalf("setup: %v", err)
|
||||
}
|
||||
|
||||
m := NewManager(dir)
|
||||
if len(m.List(context.Background(), "")) != 0 {
|
||||
t.Error("expected 0 sessions (non-JSON skipped)")
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user