diff --git a/tests/invariants/helpers_test.go b/tests/invariants/helpers_test.go new file mode 100644 index 00000000..bda317c7 --- /dev/null +++ b/tests/invariants/helpers_test.go @@ -0,0 +1,226 @@ +//go:build integration + +package invariants + +import ( + "context" + "database/sql" + "os" + "reflect" + "sync" + "testing" + + "github.com/golang-migrate/migrate/v4" + _ "github.com/golang-migrate/migrate/v4/database/postgres" + _ "github.com/golang-migrate/migrate/v4/source/file" + "github.com/google/uuid" + _ "github.com/jackc/pgx/v5/stdlib" + + "github.com/nextlevelbuilder/goclaw/internal/store" + "github.com/nextlevelbuilder/goclaw/internal/store/pg" +) + +const defaultTestDSN = "postgres://postgres:test@localhost:5433/goclaw_test?sslmode=disable" + +var ( + sharedDB *sql.DB + sharedDBOnce sync.Once + sharedDBErr error +) + +// testDB connects to the test PG instance, runs migrations once, and returns +// a shared *sql.DB. Skips test if PG is unreachable. +func testDB(t *testing.T) *sql.DB { + t.Helper() + + sharedDBOnce.Do(func() { + dsn := os.Getenv("TEST_DATABASE_URL") + if dsn == "" { + dsn = defaultTestDSN + } + + db, err := sql.Open("pgx", dsn) + if err != nil { + sharedDBErr = err + return + } + if err := db.Ping(); err != nil { + sharedDBErr = err + return + } + + m, err := migrate.New("file://../../migrations", dsn) + if err != nil { + sharedDBErr = err + return + } + if err := m.Up(); err != nil && err != migrate.ErrNoChange { + sharedDBErr = err + return + } + m.Close() + + pg.InitSqlx(db) + sharedDB = db + }) + + if sharedDBErr != nil { + t.Skipf("test PG not available: %v", sharedDBErr) + } + return sharedDB +} + +// seedTenantAgent creates a minimal tenant + agent for FK satisfaction. +// Agent insert includes all columns expected by agentSelectCols (37 columns). +func seedTenantAgent(t *testing.T, db *sql.DB) (tenantID, agentID uuid.UUID) { + t.Helper() + + tenantID = uuid.New() + agentID = uuid.New() + agentKey := "inv-" + agentID.String()[:8] + + _, err := db.Exec( + `INSERT INTO tenants (id, name, slug, status) VALUES ($1, $2, $3, 'active') + ON CONFLICT DO NOTHING`, + tenantID, "inv-tenant-"+tenantID.String()[:8], "i"+tenantID.String()[:8]) + if err != nil { + t.Fatalf("seed tenant: %v", err) + } + + // Insert with all 37 columns expected by agentSelectCols + _, err = db.Exec( + `INSERT INTO agents ( + id, agent_key, display_name, frontmatter, owner_id, provider, model, + context_window, max_tool_iterations, workspace, restrict_to_workspace, + tools_config, sandbox_config, subagents_config, memory_config, + compaction_config, context_pruning, other_config, + emoji, agent_description, thinking_level, max_tokens, + self_evolve, skill_evolve, skill_nudge_interval, + reasoning_config, workspace_sharing, chatgpt_oauth_routing, + shell_deny_groups, kg_dedup_config, + agent_type, is_default, status, budget_monthly_cents, created_at, updated_at, tenant_id + ) VALUES ( + $1, $2, $3, NULL, $4, $5, $6, + 8192, 10, '', false, + '{}', NULL, NULL, NULL, + NULL, NULL, '{}', + '', '', '', 4096, + false, false, 0, + '{}', '{}', '{}', + '{}', '{}', + 'predefined', false, 'active', 0, NOW(), NOW(), $7 + ) ON CONFLICT DO NOTHING`, + agentID, agentKey, "Test Agent "+agentKey, "test-owner", "test", "test-model", tenantID) + if err != nil { + t.Fatalf("seed agent: %v", err) + } + + t.Cleanup(func() { + cleanupTenant(db, tenantID, agentID) + }) + + return tenantID, agentID +} + +// seedTwoTenants creates 2 independent tenants with agents for isolation testing. +func seedTwoTenants(t *testing.T, db *sql.DB) (tenantA, agentA, tenantB, agentB uuid.UUID) { + t.Helper() + tenantA, agentA = seedTenantAgent(t, db) + tenantB, agentB = seedTenantAgent(t, db) + return +} + +// tenantCtx returns a context with tenant ID set for store scoping. +func tenantCtx(tenantID uuid.UUID) context.Context { + return store.WithTenantID(context.Background(), tenantID) +} + +// userCtx returns a context with both tenant ID and user ID set. +func userCtx(tenantID uuid.UUID, userID string) context.Context { + ctx := store.WithTenantID(context.Background(), tenantID) + return store.WithUserID(ctx, userID) +} + +// agentCtx returns a context with tenant, agent type and agent ID set. +func agentCtx(tenantID, agentID uuid.UUID, agentType string) context.Context { + ctx := store.WithTenantID(context.Background(), tenantID) + ctx = store.WithAgentID(ctx, agentID) + ctx = store.WithAgentType(ctx, agentType) + return ctx +} + +// assertAccessDenied verifies that a cross-tenant access returns nil or error. +// INVARIANT: Cross-tenant access MUST NOT return data. +func assertAccessDenied(t *testing.T, result any, err error, msg string) { + t.Helper() + // Access denied can manifest as: + // 1. Non-nil error (explicit denial) + // 2. Nil result with nil error (no data found - implicit denial) + // Use reflection to check for nil because interface{} with typed nil pointer is not nil. + isNil := result == nil || (reflect.ValueOf(result).Kind() == reflect.Ptr && reflect.ValueOf(result).IsNil()) + if err == nil && !isNil { + t.Errorf("INVARIANT VIOLATION: %s - expected nil or error, got data", msg) + } +} + +// assertNotEmpty verifies that the result is not nil/empty (valid access). +func assertNotEmpty(t *testing.T, result any, msg string) { + t.Helper() + if result == nil { + t.Errorf("%s: expected non-nil result", msg) + } +} + +// cleanupTenant removes all tenant data in FK order. +func cleanupTenant(db *sql.DB, tenantID, agentID uuid.UUID) { + // Team-related + db.Exec("DELETE FROM team_task_comments WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM team_task_events WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM team_task_attachments WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM team_tasks WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM agent_team_members WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM agent_teams WHERE tenant_id = $1", tenantID) + + // Knowledge stores + db.Exec("DELETE FROM episodic_summaries WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM vault_links WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM vault_documents WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM kg_dedup_candidates WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM kg_relations WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM kg_entities WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM memory_chunks WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM memory_documents WHERE tenant_id = $1", tenantID) + + // Sessions + db.Exec("DELETE FROM sessions WHERE tenant_id = $1", tenantID) + + // Skills + db.Exec("DELETE FROM skill_agent_grants WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM skill_user_grants WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM skills WHERE tenant_id = $1", tenantID) + + // Cron + db.Exec("DELETE FROM cron_run_logs WHERE job_id IN (SELECT id FROM cron_jobs WHERE tenant_id = $1)", tenantID) + db.Exec("DELETE FROM cron_jobs WHERE tenant_id = $1", tenantID) + + // Security stores + db.Exec("DELETE FROM mcp_user_credentials WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM mcp_access_requests WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM mcp_user_grants WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM mcp_agent_grants WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM mcp_servers WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM api_keys WHERE tenant_id = $1", tenantID) + db.Exec("DELETE FROM agent_config_permissions WHERE agent_id = $1", agentID) + db.Exec("DELETE FROM channel_contacts WHERE tenant_id = $1", tenantID) + + // Agent-related + db.Exec("DELETE FROM agent_shares WHERE agent_id = $1", agentID) + db.Exec("DELETE FROM agent_context_files WHERE agent_id = $1", agentID) + db.Exec("DELETE FROM user_context_files WHERE agent_id = $1", agentID) + db.Exec("DELETE FROM user_agent_overrides WHERE agent_id = $1", agentID) + db.Exec("DELETE FROM agent_user_profiles WHERE agent_id = $1", agentID) + db.Exec("DELETE FROM agent_evolution_suggestions WHERE agent_id = $1", agentID) + db.Exec("DELETE FROM agent_evolution_metrics WHERE agent_id = $1", agentID) + db.Exec("DELETE FROM agents WHERE id = $1", agentID) + db.Exec("DELETE FROM tenants WHERE id = $1", tenantID) +} diff --git a/tests/invariants/permission_enforcement_test.go b/tests/invariants/permission_enforcement_test.go new file mode 100644 index 00000000..8526ed24 --- /dev/null +++ b/tests/invariants/permission_enforcement_test.go @@ -0,0 +1,269 @@ +//go:build integration + +package invariants + +import ( + "testing" + + "github.com/nextlevelbuilder/goclaw/internal/permissions" + "github.com/nextlevelbuilder/goclaw/internal/store" + "github.com/nextlevelbuilder/goclaw/internal/store/pg" + "github.com/nextlevelbuilder/goclaw/pkg/protocol" +) + +// INVARIANT: Admin-only operations MUST reject Operator role. +func TestPermission_AdminOnlyRejectsOperator(t *testing.T) { + pe := permissions.NewPolicyEngine(nil) + + adminMethods := []string{ + protocol.MethodConfigApply, + protocol.MethodConfigPatch, + protocol.MethodAgentsCreate, + protocol.MethodAgentsUpdate, + protocol.MethodAgentsDelete, + protocol.MethodTeamsCreate, + protocol.MethodTeamsDelete, + protocol.MethodAPIKeysCreate, + protocol.MethodAPIKeysRevoke, + } + + for _, method := range adminMethods { + t.Run(method, func(t *testing.T) { + // INVARIANT: Operator MUST NOT access admin methods + if pe.CanAccess(permissions.RoleOperator, method) { + t.Errorf("INVARIANT VIOLATION: Operator can access admin method %s", method) + } + + // INVARIANT: Viewer MUST NOT access admin methods + if pe.CanAccess(permissions.RoleViewer, method) { + t.Errorf("INVARIANT VIOLATION: Viewer can access admin method %s", method) + } + + // Admin and Owner should have access + if !pe.CanAccess(permissions.RoleAdmin, method) { + t.Errorf("Admin should access %s", method) + } + if !pe.CanAccess(permissions.RoleOwner, method) { + t.Errorf("Owner should access %s", method) + } + }) + } +} + +// INVARIANT: Write operations MUST reject Viewer role. +func TestPermission_WriteRejectsViewer(t *testing.T) { + pe := permissions.NewPolicyEngine(nil) + + writeMethods := []string{ + protocol.MethodChatSend, + protocol.MethodChatAbort, + protocol.MethodSessionsDelete, + protocol.MethodSessionsReset, + protocol.MethodCronCreate, + protocol.MethodCronUpdate, + protocol.MethodCronDelete, + } + + for _, method := range writeMethods { + t.Run(method, func(t *testing.T) { + // INVARIANT: Viewer MUST NOT access write methods + if pe.CanAccess(permissions.RoleViewer, method) { + t.Errorf("INVARIANT VIOLATION: Viewer can access write method %s", method) + } + + // Operator, Admin, and Owner should have access + if !pe.CanAccess(permissions.RoleOperator, method) { + t.Errorf("Operator should access %s", method) + } + }) + } +} + +// INVARIANT: Owner MUST have access to all methods (superset of Admin). +func TestPermission_OwnerSupersetOfAdmin(t *testing.T) { + pe := permissions.NewPolicyEngine(nil) + + allMethods := []string{ + // Admin methods + protocol.MethodConfigApply, + protocol.MethodAgentsCreate, + protocol.MethodTeamsCreate, + // Write methods + protocol.MethodChatSend, + protocol.MethodSessionsDelete, + // Read methods + protocol.MethodSessionsList, + protocol.MethodAgentsList, + } + + for _, method := range allMethods { + t.Run(method, func(t *testing.T) { + adminCan := pe.CanAccess(permissions.RoleAdmin, method) + ownerCan := pe.CanAccess(permissions.RoleOwner, method) + + // INVARIANT: If Admin can access, Owner MUST be able to access + if adminCan && !ownerCan { + t.Errorf("INVARIANT VIOLATION: Admin can access %s but Owner cannot", method) + } + + // INVARIANT: Owner MUST always have access + if !ownerCan { + t.Errorf("INVARIANT VIOLATION: Owner cannot access %s", method) + } + }) + } +} + +// INVARIANT: Role hierarchy MUST be strictly ordered: Owner > Admin > Operator > Viewer. +func TestPermission_RoleHierarchy(t *testing.T) { + tests := []struct { + name string + higher permissions.Role + lower permissions.Role + required permissions.Role + }{ + {"owner_beats_admin", permissions.RoleOwner, permissions.RoleAdmin, permissions.RoleAdmin}, + {"admin_beats_operator", permissions.RoleAdmin, permissions.RoleOperator, permissions.RoleAdmin}, + {"operator_beats_viewer", permissions.RoleOperator, permissions.RoleViewer, permissions.RoleOperator}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + higherMeets := permissions.HasMinRole(tt.higher, tt.required) + lowerMeets := permissions.HasMinRole(tt.lower, tt.required) + + // INVARIANT: Higher role MUST meet requirement that lower cannot + if lowerMeets && !higherMeets { + t.Errorf("INVARIANT VIOLATION: %s cannot meet %s but %s can", + tt.higher, tt.required, tt.lower) + } + }) + } +} + +// INVARIANT: API key scopes MUST map correctly to roles. +func TestPermission_ScopesToRoleMapping(t *testing.T) { + tests := []struct { + name string + scopes []permissions.Scope + expected permissions.Role + }{ + {"admin_scope_is_admin", []permissions.Scope{permissions.ScopeAdmin}, permissions.RoleAdmin}, + {"write_scope_is_operator", []permissions.Scope{permissions.ScopeWrite}, permissions.RoleOperator}, + {"read_scope_is_viewer", []permissions.Scope{permissions.ScopeRead}, permissions.RoleViewer}, + {"empty_scope_is_viewer", []permissions.Scope{}, permissions.RoleViewer}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := permissions.RoleFromScopes(tt.scopes) + if got != tt.expected { + t.Errorf("INVARIANT VIOLATION: scopes %v should map to %s, got %s", + tt.scopes, tt.expected, got) + } + }) + } +} + +// INVARIANT: Owner IDs MUST be consistently recognized. +func TestPermission_OwnerRecognition(t *testing.T) { + owners := []string{"owner-alice", "owner-bob"} + pe := permissions.NewPolicyEngine(owners) + + // INVARIANT: Listed owners MUST be recognized + for _, owner := range owners { + if !pe.IsOwner(owner) { + t.Errorf("INVARIANT VIOLATION: %s should be recognized as owner", owner) + } + } + + // INVARIANT: Non-owners MUST NOT be recognized as owners + nonOwners := []string{"charlie", "admin", "", "owner"} + for _, nonOwner := range nonOwners { + if pe.IsOwner(nonOwner) { + t.Errorf("INVARIANT VIOLATION: %s should NOT be recognized as owner", nonOwner) + } + } +} + +// INVARIANT: Config permission store MUST enforce tenant scoping. +func TestPermission_ConfigPermissionTenantIsolation(t *testing.T) { + db := testDB(t) + tenantA, agentA, tenantB, _ := seedTwoTenants(t, db) + ctxA := tenantCtx(tenantA) + ctxB := tenantCtx(tenantB) + + cps := pg.NewPGConfigPermissionStore(db) + + // Grant permission in tenant A using colon-delimited scope (matchWildcard requires :* format) + err := cps.Grant(ctxA, &store.ConfigPermission{ + AgentID: agentA, + Scope: "group:*", + ConfigType: "file_writer", + Permission: "allow", + UserID: "user-a", + }) + if err != nil { + t.Fatalf("Grant: %v", err) + } + t.Cleanup(func() { + db.Exec("DELETE FROM agent_config_permissions WHERE agent_id = $1", agentA) + }) + + // Verify tenant A has permission (scope must match pattern: group:* matches group:telegram) + allowed, err := cps.CheckPermission(ctxA, agentA, "group:telegram", "file_writer", "user-a") + if err != nil { + t.Fatalf("CheckPermission A: %v", err) + } + if !allowed { + t.Error("tenant A should have permission on own agent") + } + + // INVARIANT: Tenant B MUST NOT see tenant A's permissions + allowedB, err := cps.CheckPermission(ctxB, agentA, "group:telegram", "file_writer", "user-a") + if err != nil { + t.Fatalf("CheckPermission B: %v", err) + } + if allowedB { + t.Errorf("INVARIANT VIOLATION: tenant B sees tenant A's permission") + } +} + +// INVARIANT: Scope-based access MUST require correct scope for method. +func TestPermission_ScopeBasedAccess(t *testing.T) { + pe := permissions.NewPolicyEngine(nil) + + tests := []struct { + name string + scopes []permissions.Scope + method string + shouldAllow bool + }{ + // Admin methods require admin scope + {"read_cannot_create_agent", []permissions.Scope{permissions.ScopeRead}, protocol.MethodAgentsCreate, false}, + {"write_cannot_create_agent", []permissions.Scope{permissions.ScopeWrite}, protocol.MethodAgentsCreate, false}, + {"admin_can_create_agent", []permissions.Scope{permissions.ScopeAdmin}, protocol.MethodAgentsCreate, true}, + + // Write methods require write or admin scope + {"read_cannot_send_chat", []permissions.Scope{permissions.ScopeRead}, protocol.MethodChatSend, false}, + {"write_can_send_chat", []permissions.Scope{permissions.ScopeWrite}, protocol.MethodChatSend, true}, + {"admin_can_send_chat", []permissions.Scope{permissions.ScopeAdmin}, protocol.MethodChatSend, true}, + + // Approval methods require approvals or admin scope + {"read_cannot_list_approvals", []permissions.Scope{permissions.ScopeRead}, "approvals.list", false}, + {"approvals_can_list_approvals", []permissions.Scope{permissions.ScopeApprovals}, "approvals.list", true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := pe.CanAccessWithScopes(tt.scopes, tt.method) + if got != tt.shouldAllow { + if tt.shouldAllow { + t.Errorf("INVARIANT VIOLATION: scopes %v should allow %s", tt.scopes, tt.method) + } else { + t.Errorf("INVARIANT VIOLATION: scopes %v should NOT allow %s", tt.scopes, tt.method) + } + } + }) + } +} diff --git a/tests/invariants/session_boundary_test.go b/tests/invariants/session_boundary_test.go new file mode 100644 index 00000000..917b8a7b --- /dev/null +++ b/tests/invariants/session_boundary_test.go @@ -0,0 +1,264 @@ +//go:build integration + +package invariants + +import ( + "testing" + + "github.com/google/uuid" + "github.com/nextlevelbuilder/goclaw/internal/providers" + "github.com/nextlevelbuilder/goclaw/internal/store/pg" +) + +// INVARIANT: Session A's message history MUST NOT leak to Session B. +func TestSessionBoundary_MessageHistory(t *testing.T) { + db := testDB(t) + tenantID, _ := seedTenantAgent(t, db) + ctx := tenantCtx(tenantID) + + ss := pg.NewPGSessionStore(db) + sessionA := "inv-sess-a-" + uuid.New().String()[:8] + sessionB := "inv-sess-b-" + uuid.New().String()[:8] + + // Create messages in session A + ss.GetOrCreate(ctx, sessionA) + ss.AddMessage(ctx, sessionA, providers.Message{Role: "user", Content: "secret message A"}) + ss.AddMessage(ctx, sessionA, providers.Message{Role: "assistant", Content: "response A"}) + if err := ss.Save(ctx, sessionA); err != nil { + t.Fatalf("Save A: %v", err) + } + + // Create session B with different messages + ss.GetOrCreate(ctx, sessionB) + ss.AddMessage(ctx, sessionB, providers.Message{Role: "user", Content: "message B"}) + if err := ss.Save(ctx, sessionB); err != nil { + t.Fatalf("Save B: %v", err) + } + + // INVARIANT: Session A's history must be separate from session B + histA := ss.GetHistory(ctx, sessionA) + histB := ss.GetHistory(ctx, sessionB) + + if len(histA) != 2 { + t.Errorf("session A should have 2 messages, got %d", len(histA)) + } + if len(histB) != 1 { + t.Errorf("session B should have 1 message, got %d", len(histB)) + } + + // INVARIANT: Content must not leak between sessions + for _, m := range histB { + if m.Content == "secret message A" || m.Content == "response A" { + t.Errorf("INVARIANT VIOLATION: session A's message leaked to session B: %q", m.Content) + } + } +} + +// INVARIANT: Session A's summary MUST NOT leak to Session B. +func TestSessionBoundary_Summary(t *testing.T) { + db := testDB(t) + tenantID, _ := seedTenantAgent(t, db) + ctx := tenantCtx(tenantID) + + ss := pg.NewPGSessionStore(db) + sessionA := "inv-sess-sum-a-" + uuid.New().String()[:8] + sessionB := "inv-sess-sum-b-" + uuid.New().String()[:8] + + // Set summary in session A + ss.GetOrCreate(ctx, sessionA) + ss.SetSummary(ctx, sessionA, "secret summary for session A") + if err := ss.Save(ctx, sessionA); err != nil { + t.Fatalf("Save A: %v", err) + } + + // Create session B without summary + ss.GetOrCreate(ctx, sessionB) + if err := ss.Save(ctx, sessionB); err != nil { + t.Fatalf("Save B: %v", err) + } + + // INVARIANT: Session B MUST NOT see session A's summary + summaryB := ss.GetSummary(ctx, sessionB) + if summaryB == "secret summary for session A" { + t.Errorf("INVARIANT VIOLATION: session A's summary leaked to session B") + } + if summaryB != "" { + t.Errorf("session B should have empty summary, got %q", summaryB) + } + + // Verify session A still has its summary + summaryA := ss.GetSummary(ctx, sessionA) + if summaryA != "secret summary for session A" { + t.Errorf("session A's summary should persist, got %q", summaryA) + } +} + +// INVARIANT: Session A's metadata MUST NOT leak to Session B. +func TestSessionBoundary_Metadata(t *testing.T) { + db := testDB(t) + tenantID, _ := seedTenantAgent(t, db) + ctx := tenantCtx(tenantID) + + ss := pg.NewPGSessionStore(db) + sessionA := "inv-sess-meta-a-" + uuid.New().String()[:8] + sessionB := "inv-sess-meta-b-" + uuid.New().String()[:8] + + // Set metadata in session A + ss.GetOrCreate(ctx, sessionA) + ss.SetSessionMetadata(ctx, sessionA, map[string]string{ + "secret_key": "secret_value", + "user_data": "sensitive", + }) + if err := ss.Save(ctx, sessionA); err != nil { + t.Fatalf("Save A: %v", err) + } + + // Create session B with different metadata + ss.GetOrCreate(ctx, sessionB) + ss.SetSessionMetadata(ctx, sessionB, map[string]string{ + "channel": "telegram", + }) + if err := ss.Save(ctx, sessionB); err != nil { + t.Fatalf("Save B: %v", err) + } + + // INVARIANT: Session B MUST NOT see session A's metadata + metaB := ss.GetSessionMetadata(ctx, sessionB) + if metaB["secret_key"] != "" { + t.Errorf("INVARIANT VIOLATION: session A's metadata leaked to session B: secret_key=%q", metaB["secret_key"]) + } + if metaB["user_data"] != "" { + t.Errorf("INVARIANT VIOLATION: session A's metadata leaked to session B: user_data=%q", metaB["user_data"]) + } + + // Verify session B has its own metadata + if metaB["channel"] != "telegram" { + t.Errorf("session B should have its own metadata, got channel=%q", metaB["channel"]) + } + + // Verify session A still has its metadata + metaA := ss.GetSessionMetadata(ctx, sessionA) + if metaA["secret_key"] != "secret_value" { + t.Errorf("session A's metadata should persist, got secret_key=%q", metaA["secret_key"]) + } +} + +// INVARIANT: Session A's label MUST NOT leak to Session B. +func TestSessionBoundary_Label(t *testing.T) { + db := testDB(t) + tenantID, _ := seedTenantAgent(t, db) + ctx := tenantCtx(tenantID) + + ss := pg.NewPGSessionStore(db) + sessionA := "inv-sess-label-a-" + uuid.New().String()[:8] + sessionB := "inv-sess-label-b-" + uuid.New().String()[:8] + + // Set label in session A + ss.GetOrCreate(ctx, sessionA) + ss.SetLabel(ctx, sessionA, "Private Conversation") + if err := ss.Save(ctx, sessionA); err != nil { + t.Fatalf("Save A: %v", err) + } + + // Create session B without label + ss.GetOrCreate(ctx, sessionB) + if err := ss.Save(ctx, sessionB); err != nil { + t.Fatalf("Save B: %v", err) + } + + // INVARIANT: Session B MUST NOT see session A's label + labelB := ss.GetLabel(ctx, sessionB) + if labelB == "Private Conversation" { + t.Errorf("INVARIANT VIOLATION: session A's label leaked to session B") + } +} + +// INVARIANT: Session A's token counts MUST NOT affect Session B. +func TestSessionBoundary_TokenCounts(t *testing.T) { + db := testDB(t) + tenantID, _ := seedTenantAgent(t, db) + ctx := tenantCtx(tenantID) + + ss := pg.NewPGSessionStore(db) + sessionA := "inv-sess-tokens-a-" + uuid.New().String()[:8] + sessionB := "inv-sess-tokens-b-" + uuid.New().String()[:8] + + // Accumulate tokens in session A + ss.GetOrCreate(ctx, sessionA) + ss.AccumulateTokens(ctx, sessionA, 1000, 500) + ss.AccumulateTokens(ctx, sessionA, 2000, 1000) + if err := ss.Save(ctx, sessionA); err != nil { + t.Fatalf("Save A: %v", err) + } + + // Create session B + ss.GetOrCreate(ctx, sessionB) + if err := ss.Save(ctx, sessionB); err != nil { + t.Fatalf("Save B: %v", err) + } + + // Verify session A has accumulated tokens + dataA := ss.Get(ctx, sessionA) + if dataA.InputTokens != 3000 || dataA.OutputTokens != 1500 { + t.Errorf("session A tokens: expected 3000/1500, got %d/%d", dataA.InputTokens, dataA.OutputTokens) + } + + // INVARIANT: Session B MUST NOT have session A's token counts + dataB := ss.Get(ctx, sessionB) + if dataB.InputTokens != 0 || dataB.OutputTokens != 0 { + t.Errorf("INVARIANT VIOLATION: session B has tokens %d/%d, should be 0/0", + dataB.InputTokens, dataB.OutputTokens) + } +} + +// INVARIANT: Resetting Session A MUST NOT affect Session B. +func TestSessionBoundary_ResetIsolation(t *testing.T) { + db := testDB(t) + tenantID, _ := seedTenantAgent(t, db) + ctx := tenantCtx(tenantID) + + ss := pg.NewPGSessionStore(db) + sessionA := "inv-sess-reset-a-" + uuid.New().String()[:8] + sessionB := "inv-sess-reset-b-" + uuid.New().String()[:8] + + // Set up both sessions with data + ss.GetOrCreate(ctx, sessionA) + ss.AddMessage(ctx, sessionA, providers.Message{Role: "user", Content: "message A"}) + ss.SetSummary(ctx, sessionA, "summary A") + if err := ss.Save(ctx, sessionA); err != nil { + t.Fatalf("Save A: %v", err) + } + + ss.GetOrCreate(ctx, sessionB) + ss.AddMessage(ctx, sessionB, providers.Message{Role: "user", Content: "message B"}) + ss.SetSummary(ctx, sessionB, "summary B") + if err := ss.Save(ctx, sessionB); err != nil { + t.Fatalf("Save B: %v", err) + } + + // Reset session A + ss.Reset(ctx, sessionA) + if err := ss.Save(ctx, sessionA); err != nil { + t.Fatalf("Save after reset: %v", err) + } + + // INVARIANT: Session B MUST NOT be affected by session A's reset + histB := ss.GetHistory(ctx, sessionB) + if len(histB) != 1 { + t.Errorf("INVARIANT VIOLATION: session B lost messages after session A reset, got %d", len(histB)) + } + if histB[0].Content != "message B" { + t.Errorf("session B should keep its message, got %q", histB[0].Content) + } + + summaryB := ss.GetSummary(ctx, sessionB) + if summaryB != "summary B" { + t.Errorf("INVARIANT VIOLATION: session B lost summary after session A reset, got %q", summaryB) + } + + // Verify session A was actually reset + histA := ss.GetHistory(ctx, sessionA) + if len(histA) != 0 { + t.Errorf("session A should have 0 messages after reset, got %d", len(histA)) + } +} diff --git a/tests/invariants/tenant_isolation_test.go b/tests/invariants/tenant_isolation_test.go new file mode 100644 index 00000000..5e5fcecd --- /dev/null +++ b/tests/invariants/tenant_isolation_test.go @@ -0,0 +1,372 @@ +//go:build integration + +// Package invariants tests critical system invariants that must never be violated. +// These tests verify security boundaries, not functional behavior. +package invariants + +import ( + "testing" + + "github.com/google/uuid" + "github.com/nextlevelbuilder/goclaw/internal/providers" + "github.com/nextlevelbuilder/goclaw/internal/store" + "github.com/nextlevelbuilder/goclaw/internal/store/pg" +) + +// INVARIANT: Tenant A cannot access Tenant B's sessions. +func TestTenantIsolation_SessionStore(t *testing.T) { + db := testDB(t) + tenantA, _, tenantB, _ := seedTwoTenants(t, db) + ctxA := tenantCtx(tenantA) + ctxB := tenantCtx(tenantB) + + ss := pg.NewPGSessionStore(db) + sessionKey := "inv-sess-" + uuid.New().String()[:8] + + // Create session in tenant A + ss.GetOrCreate(ctxA, sessionKey) + ss.AddMessage(ctxA, sessionKey, providers.Message{Role: "user", Content: "secret data"}) + if err := ss.Save(ctxA, sessionKey); err != nil { + t.Fatalf("Save: %v", err) + } + + // Verify tenant A can access + if got := ss.Get(ctxA, sessionKey); got == nil { + t.Fatal("tenant A should access own session") + } + + // INVARIANT: Tenant B MUST NOT access tenant A's session + ss2 := pg.NewPGSessionStore(db) // new store instance to bypass cache + got := ss2.Get(ctxB, sessionKey) + assertAccessDenied(t, got, nil, "tenant B accessing tenant A's session") + + // INVARIANT: Tenant B MUST NOT see tenant A's messages + hist := ss2.GetHistory(ctxB, sessionKey) + if len(hist) > 0 { + t.Errorf("INVARIANT VIOLATION: tenant B sees %d messages from tenant A", len(hist)) + } +} + +// INVARIANT: Tenant A cannot access Tenant B's agents. +func TestTenantIsolation_AgentStore(t *testing.T) { + db := testDB(t) + tenantA, agentA, tenantB, _ := seedTwoTenants(t, db) + ctxA := tenantCtx(tenantA) + ctxB := tenantCtx(tenantB) + + as := pg.NewPGAgentStore(db) + + // Verify tenant A can access own agent + agent, err := as.GetByID(ctxA, agentA) + if err != nil || agent == nil { + t.Fatal("tenant A should access own agent") + } + + // INVARIANT: Tenant B MUST NOT access tenant A's agent + got, err := as.GetByID(ctxB, agentA) + assertAccessDenied(t, got, err, "tenant B accessing tenant A's agent by ID") + + // INVARIANT: Tenant B's agent list MUST NOT include tenant A's agents + agents, err := as.List(ctxB, "") // empty ownerID = list all + if err != nil { + t.Fatalf("List: %v", err) + } + for _, a := range agents { + if a.ID == agentA { + t.Errorf("INVARIANT VIOLATION: tenant B's list includes tenant A's agent") + } + } +} + +// INVARIANT: Tenant A cannot access Tenant B's memory documents. +func TestTenantIsolation_MemoryStore(t *testing.T) { + db := testDB(t) + tenantA, agentA, tenantB, agentB := seedTwoTenants(t, db) + + // Create memory document in tenant A via direct insert + docID := uuid.New() + _, err := db.Exec( + `INSERT INTO memory_documents (id, tenant_id, agent_id, user_id, path, content, hash, created_at, updated_at) + VALUES ($1, $2, $3, 'user-a', '/test/doc.md', 'secret memory content', 'abc123', NOW(), NOW())`, + docID, tenantA, agentA) + if err != nil { + t.Fatalf("create memory doc: %v", err) + } + t.Cleanup(func() { + db.Exec("DELETE FROM memory_chunks WHERE document_id = $1", docID) + db.Exec("DELETE FROM memory_documents WHERE id = $1", docID) + }) + + // Verify tenant A's document exists + var count int + db.QueryRow("SELECT COUNT(*) FROM memory_documents WHERE id = $1 AND tenant_id = $2", docID, tenantA).Scan(&count) + if count != 1 { + t.Fatal("tenant A should have memory document") + } + + // INVARIANT: Tenant B MUST NOT see tenant A's memory via tenant-scoped query + db.QueryRow("SELECT COUNT(*) FROM memory_documents WHERE id = $1 AND tenant_id = $2", docID, tenantB).Scan(&count) + if count != 0 { + t.Errorf("INVARIANT VIOLATION: tenant B can see tenant A's memory document") + } + + // Test store-level isolation + ms := pg.NewPGMemoryStore(db, pg.PGMemoryConfig{}) + ctxA := agentCtx(tenantA, agentA, "predefined") + ctxB := agentCtx(tenantB, agentB, "predefined") + + // GetDocument uses agentID+userID+path, so we test with a unique path + contentA, err := ms.GetDocument(ctxA, agentA.String(), "user-a", "test-path") + // It's OK if this returns empty - the point is tenant B shouldn't see A's data + _ = contentA + _ = err + + contentB, err := ms.GetDocument(ctxB, agentA.String(), "user-a", "test-path") + if contentB != "" { + t.Errorf("INVARIANT VIOLATION: tenant B got content from tenant A's agent: %s", contentB) + } +} + +// INVARIANT: Tenant A cannot access Tenant B's teams. +func TestTenantIsolation_TeamStore(t *testing.T) { + db := testDB(t) + tenantA, agentA, tenantB, _ := seedTwoTenants(t, db) + ctxA := tenantCtx(tenantA) + ctxB := tenantCtx(tenantB) + + ts := pg.NewPGTeamStore(db) + + // Create team in tenant A + teamID := uuid.New() + _, err := db.Exec( + `INSERT INTO agent_teams (id, tenant_id, name, lead_agent_id, status, settings, created_by) + VALUES ($1, $2, 'test-team', $3, 'active', '{"version": 2}', 'test')`, + teamID, tenantA, agentA) + if err != nil { + t.Fatalf("create team: %v", err) + } + t.Cleanup(func() { + db.Exec("DELETE FROM agent_team_members WHERE team_id = $1", teamID) + db.Exec("DELETE FROM agent_teams WHERE id = $1", teamID) + }) + + // Verify tenant A can access + team, err := ts.GetTeam(ctxA, teamID) + if err != nil || team == nil { + t.Fatal("tenant A should access own team") + } + + // INVARIANT: Tenant B MUST NOT access tenant A's team + got, err := ts.GetTeam(ctxB, teamID) + assertAccessDenied(t, got, err, "tenant B accessing tenant A's team") + + // INVARIANT: Tenant B's team list MUST NOT include tenant A's teams + teams, err := ts.ListTeams(ctxB) + if err != nil { + t.Fatalf("List: %v", err) + } + for _, tm := range teams { + if tm.ID == teamID { + t.Errorf("INVARIANT VIOLATION: tenant B's list includes tenant A's team") + } + } +} + +// INVARIANT: Tenant A cannot access Tenant B's skills. +func TestTenantIsolation_SkillStore(t *testing.T) { + db := testDB(t) + tenantA, _, tenantB, _ := seedTwoTenants(t, db) + ctxA := tenantCtx(tenantA) + ctxB := tenantCtx(tenantB) + + ss := pg.NewPGSkillStore(db, t.TempDir()) + + // Create skill in tenant A + desc := "test skill" + slug := "inv-skill-" + tenantA.String()[:8] + skillID, err := ss.CreateSkillManaged(ctxA, store.SkillCreateParams{ + Name: slug, Slug: slug, Description: &desc, OwnerID: "test-owner", + Visibility: "private", Status: "active", Version: 1, FilePath: "/tmp/" + slug, + }) + if err != nil { + t.Fatalf("CreateSkill: %v", err) + } + t.Cleanup(func() { + db.Exec("DELETE FROM skill_agent_grants WHERE skill_id = $1", skillID) + db.Exec("DELETE FROM skill_user_grants WHERE skill_id = $1", skillID) + db.Exec("DELETE FROM skills WHERE id = $1", skillID) + }) + + // Verify tenant A can access + skill, ok := ss.GetSkillByID(ctxA, skillID) + if !ok { + t.Fatal("tenant A should access own skill") + } + _ = skill + + // INVARIANT: Tenant B MUST NOT access tenant A's skill + gotSkill, gotOK := ss.GetSkillByID(ctxB, skillID) + if gotOK { + t.Errorf("INVARIANT VIOLATION: tenant B can access tenant A's skill: %s", gotSkill.Slug) + } +} + +// INVARIANT: Tenant A cannot access Tenant B's cron jobs. +func TestTenantIsolation_CronStore(t *testing.T) { + db := testDB(t) + tenantA, agentA, tenantB, _ := seedTwoTenants(t, db) + ctxA := tenantCtx(tenantA) + ctxB := tenantCtx(tenantB) + + cs := pg.NewPGCronStore(db) + + // Create cron job in tenant A using AddJob + job, err := cs.AddJob(ctxA, "test-cron", + store.CronSchedule{Kind: "cron", Expr: "0 * * * *"}, + "test prompt", false, "", "", agentA.String(), "test-user") + if err != nil { + t.Fatalf("AddJob: %v", err) + } + t.Cleanup(func() { + db.Exec("DELETE FROM cron_run_logs WHERE job_id = $1", job.ID) + db.Exec("DELETE FROM cron_jobs WHERE id = $1", job.ID) + }) + + // Verify tenant A can access + gotJob, ok := cs.GetJob(ctxA, job.ID) // ID is already a string + if !ok || gotJob == nil { + t.Fatal("tenant A should access own cron job") + } + + // INVARIANT: Tenant B MUST NOT access tenant A's cron job + gotJobB, okB := cs.GetJob(ctxB, job.ID) // ID is already a string + if okB && gotJobB != nil { + t.Errorf("INVARIANT VIOLATION: tenant B can access tenant A's cron job") + } +} + +// INVARIANT: Tenant A cannot access Tenant B's API keys. +func TestTenantIsolation_APIKeyStore(t *testing.T) { + db := testDB(t) + tenantA, _, tenantB, _ := seedTwoTenants(t, db) + ctxA := tenantCtx(tenantA) + ctxB := tenantCtx(tenantB) + + ks := pg.NewPGAPIKeyStore(db) + + // Create API key in tenant A + keyID := uuid.New() + _, err := db.Exec( + `INSERT INTO api_keys (id, tenant_id, name, prefix, key_hash, scopes, created_by) + VALUES ($1, $2, 'test-key', 'gclw_inv', 'hash123', '{}', 'test-user')`, + keyID, tenantA) + if err != nil { + t.Fatalf("create api key: %v", err) + } + t.Cleanup(func() { + db.Exec("DELETE FROM api_keys WHERE id = $1", keyID) + }) + + // Verify tenant A can access + keys, err := ks.List(ctxA, "") // empty ownerID = all keys + if err != nil { + t.Fatalf("List A: %v", err) + } + found := false + for _, k := range keys { + if k.ID == keyID { + found = true + break + } + } + if !found { + t.Fatal("tenant A should see own API key") + } + + // INVARIANT: Tenant B's key list MUST NOT include tenant A's keys + keysB, err := ks.List(ctxB, "") // empty ownerID = all keys + if err != nil { + t.Fatalf("List B: %v", err) + } + for _, k := range keysB { + if k.ID == keyID { + t.Errorf("INVARIANT VIOLATION: tenant B's list includes tenant A's API key") + } + } +} + +// INVARIANT: Tenant A cannot access Tenant B's MCP servers. +func TestTenantIsolation_MCPServerStore(t *testing.T) { + db := testDB(t) + tenantA, _, tenantB, _ := seedTwoTenants(t, db) + ctxA := tenantCtx(tenantA) + ctxB := tenantCtx(tenantB) + + ms := pg.NewPGMCPServerStore(db, "0123456789abcdef0123456789abcdef") // test encryption key + + // Create MCP server in tenant A + serverID := uuid.New() + _, err := db.Exec( + `INSERT INTO mcp_servers (id, tenant_id, name, display_name, transport, enabled, created_by) + VALUES ($1, $2, 'test-mcp', 'Test MCP', 'stdio', true, 'test-user')`, + serverID, tenantA) + if err != nil { + t.Fatalf("create mcp server: %v", err) + } + t.Cleanup(func() { + db.Exec("DELETE FROM mcp_user_credentials WHERE server_id = $1", serverID) + db.Exec("DELETE FROM mcp_access_requests WHERE server_id = $1", serverID) + db.Exec("DELETE FROM mcp_user_grants WHERE server_id = $1", serverID) + db.Exec("DELETE FROM mcp_agent_grants WHERE server_id = $1", serverID) + db.Exec("DELETE FROM mcp_servers WHERE id = $1", serverID) + }) + + // Verify tenant A can access + server, err := ms.GetServer(ctxA, serverID) + if err != nil || server == nil { + t.Fatal("tenant A should access own MCP server") + } + + // INVARIANT: Tenant B MUST NOT access tenant A's MCP server + got, err := ms.GetServer(ctxB, serverID) + assertAccessDenied(t, got, err, "tenant B accessing tenant A's MCP server") +} + +// INVARIANT: Tenant A cannot access Tenant B's vault documents. +func TestTenantIsolation_VaultStore(t *testing.T) { + db := testDB(t) + tenantA, _, tenantB, _ := seedTwoTenants(t, db) + + // Create vault document in tenant A via direct SQL + docID := uuid.New() + _, err := db.Exec( + `INSERT INTO vault_documents (id, tenant_id, title, path, doc_type, content_hash, scope, created_at, updated_at) + VALUES ($1, $2, 'secret-doc', '/tmp/secret.md', 'document', 'abc123', 'personal', NOW(), NOW())`, + docID, tenantA) + if err != nil { + t.Fatalf("create vault doc: %v", err) + } + t.Cleanup(func() { + db.Exec("DELETE FROM vault_links WHERE tenant_id = $1", tenantA) + db.Exec("DELETE FROM vault_documents WHERE id = $1", docID) + }) + + // Verify tenant A can access via direct query + var count int + db.QueryRow("SELECT COUNT(*) FROM vault_documents WHERE id = $1 AND tenant_id = $2", docID, tenantA).Scan(&count) + if count != 1 { + t.Fatal("tenant A should have vault document") + } + + // INVARIANT: Tenant B MUST NOT see tenant A's vault document via tenant-scoped query + db.QueryRow("SELECT COUNT(*) FROM vault_documents WHERE id = $1 AND tenant_id = $2", docID, tenantB).Scan(&count) + if count != 0 { + t.Errorf("INVARIANT VIOLATION: tenant B can see tenant A's vault document") + } + + // INVARIANT: Any query scoped by tenant_id MUST exclude other tenants + db.QueryRow("SELECT COUNT(*) FROM vault_documents WHERE path = $1 AND tenant_id = $2", "/tmp/secret.md", tenantB).Scan(&count) + if count != 0 { + t.Errorf("INVARIANT VIOLATION: tenant B can see tenant A's vault document by path") + } +}