diff --git a/internal/kitchen/analyze.go b/internal/kitchen/analyze.go index bdc24fc..6410d5c 100644 --- a/internal/kitchen/analyze.go +++ b/internal/kitchen/analyze.go @@ -264,31 +264,28 @@ func (h *Handler) handleAnalyze(w http.ResponseWriter, r *http.Request) { h.broadcastAnalysisProgress(req, startedAt, analysisPhaseImport, formatImportPhaseMessage(result), "", result.ReposAnalyzed, result.ReposAnalyzed, len(result.SecretFindings)) if len(result.Findings) > 0 || len(result.Workflows) > 0 || len(result.SecretFindings) > 0 { importStarted := time.Now() - imported := h.importAnalysisToPantry(result) - slog.Info("imported analysis to pantry", - "findings", len(result.Findings), - "workflows", len(result.Workflows), - "secrets", len(result.SecretFindings), - "assets", imported, - "duration", time.Since(importStarted)) + imported := 0 + h.broadcastAnalysisProgress(req, startedAt, analysisPhaseImport, "Persisting attack graph", "", result.ReposAnalyzed, result.ReposAnalyzed, len(result.SecretFindings)) + if err := h.committedPantry().Replace(ctx, func(candidate *pantry.Pantry) error { + imported = h.importAnalysisToPantry(candidate, result) + return nil + }); err != nil { + slog.Warn("failed to commit analysis pantry", "target", req.Target, "error", err) + } else { + slog.Info("committed analysis to pantry", + "findings", len(result.Findings), + "workflows", len(result.Workflows), + "secrets", len(result.SecretFindings), + "assets", imported, + "revision", h.Pantry().Revision(), + "duration", time.Since(importStarted)) + } } if err := h.persistAnalysisLoot(req, result); err != nil { slog.Warn("failed to persist analysis loot", "target", req.Target, "error", err) } - h.broadcastAnalysisProgress(req, startedAt, analysisPhaseImport, "Persisting attack graph", "", result.ReposAnalyzed, result.ReposAnalyzed, len(result.SecretFindings)) - saveStarted := time.Now() - if err := h.SavePantry(); err != nil { - slog.Warn("failed to persist pantry", "error", err) - } else { - slog.Info("analysis pantry persisted", - "target", req.Target, - "assets", h.Pantry().Size(), - "edges", h.Pantry().EdgeCount(), - "duration", time.Since(saveStarted)) - } - h.recordAnalysisCompleted(req, result) // Return result @@ -390,24 +387,21 @@ func (h *Handler) runAnalysisMetadataSync(req AnalyzeRequest, result *poutine.An "duration", time.Since(visibilityStarted)) inventoryStarted := time.Now() - h.importPrivateReposToPantry(req.SessionID) + if err := h.importPrivateReposToPantry(ctx, req.SessionID); err != nil { + slog.Warn("failed to commit analysis private repository inventory", "target", req.Target, "error", err) + h.broadcastAnalysisMetadataSync(req, analysisMetadataFailed, "Repository access update incomplete", repoCount, err) + return + } slog.Info("analysis private repo inventory updated", "target", req.Target, "session", req.SessionID, "duration", time.Since(inventoryStarted)) - saveStarted := time.Now() - if err := h.SavePantry(); err != nil { - slog.Warn("failed to persist pantry after analysis metadata sync", "target", req.Target, "error", err) - h.broadcastAnalysisMetadataSync(req, analysisMetadataFailed, "Repository access update incomplete", repoCount, err) - return - } - slog.Info("analysis metadata sync completed", "target", req.Target, "type", req.TargetType, "repos", repoCount, - "persist_duration", time.Since(saveStarted), + "revision", h.Pantry().Revision(), "duration", time.Since(started)) h.broadcastAnalysisMetadataSync(req, analysisMetadataDone, "Repository access updated", repoCount, nil) } @@ -566,8 +560,7 @@ func (o *analysisProgressObserver) broadcast(message, repo string, completedDelt } // importAnalysisToPantry imports poutine findings into the Kitchen's pantry. -func (h *Handler) importAnalysisToPantry(result *poutine.AnalysisResult) int { - p := h.Pantry() +func (h *Handler) importAnalysisToPantry(p *pantry.Pantry, result *poutine.AnalysisResult) int { imported := 0 orgAssets := make(map[string]string) repoAssets := make(map[string]string) @@ -1106,20 +1099,21 @@ func (h *Handler) recordAnalyzedRepoVisibility(ctx context.Context, req AnalyzeR } } -func (h *Handler) importPrivateReposToPantry(sessionID string) { +func (h *Handler) importPrivateReposToPantry(ctx context.Context, sessionID string) error { repo := db.NewKnownEntityRepository(h.database) entities, err := repo.ListRepos(sessionID) if err != nil { - slog.Warn("failed to list known entities for private repo import", "session", sessionID, "error", err) - return + return err } - p := h.Pantry() - for _, entity := range entities { - if !entity.IsPrivate && entity.SSHPermission == "" && len(entity.Permissions) == 0 { - continue + return h.committedPantry().Update(ctx, func(candidate *pantry.Pantry) error { + for _, entity := range entities { + if !entity.IsPrivate && entity.SSHPermission == "" && len(entity.Permissions) == 0 { + continue + } + upsertKnownRepoAsset(candidate, entity) } - upsertKnownRepoAsset(p, entity) - } + return nil + }) } func upsertKnownRepoAsset(p *pantry.Pantry, entity *db.KnownEntityRow) { diff --git a/internal/kitchen/analyze_perf_test.go b/internal/kitchen/analyze_perf_test.go index 15b1f7e..3339d1a 100644 --- a/internal/kitchen/analyze_perf_test.go +++ b/internal/kitchen/analyze_perf_test.go @@ -19,6 +19,7 @@ import ( "github.com/stretchr/testify/require" + "github.com/boostsecurityio/smokedmeat/internal/pantry" "github.com/boostsecurityio/smokedmeat/internal/poutine" ) @@ -87,7 +88,11 @@ func TestAnalyzePerformanceProfile(t *testing.T) { fmt.Printf("[perf] importing analysis results - elapsed=%s\n", roundPerfDuration(time.Since(totalStarted))) importStarted := time.Now() - importedAssets := h.importAnalysisToPantry(result) + importedAssets := 0 + require.NoError(t, h.committedPantry().Replace(ctx, func(candidate *pantry.Pantry) error { + importedAssets = h.importAnalysisToPantry(candidate, result) + return nil + })) importDuration := time.Since(importStarted) fmt.Printf("[perf] updating repository access - elapsed=%s\n", roundPerfDuration(time.Since(totalStarted))) @@ -97,15 +102,10 @@ func TestAnalyzePerformanceProfile(t *testing.T) { fmt.Printf("[perf] updating private repo inventory - elapsed=%s\n", roundPerfDuration(time.Since(totalStarted))) inventoryStarted := time.Now() - h.importPrivateReposToPantry(config.SessionID) + require.NoError(t, h.importPrivateReposToPantry(ctx, config.SessionID)) inventoryDuration := time.Since(inventoryStarted) - fmt.Printf("[perf] persisting attack graph - elapsed=%s\n", roundPerfDuration(time.Since(totalStarted))) - persistStarted := time.Now() - require.NoError(t, h.SavePantry()) - persistDuration := time.Since(persistStarted) - - tailDuration := importDuration + secretScanDuration + repoAccessDuration + inventoryDuration + persistDuration + tailDuration := importDuration + secretScanDuration + repoAccessDuration + inventoryDuration totalDuration := time.Since(totalStarted) repoCount := len(collectAnalyzedRepos(result)) @@ -121,13 +121,12 @@ func TestAnalyzePerformanceProfile(t *testing.T) { h.Pantry().Size(), h.Pantry().EdgeCount(), ) - t.Logf("analysis timings scan=%s secret_scan=%s import=%s repo_access=%s private_repo_inventory=%s persist=%s tail=%s total=%s", + t.Logf("analysis timings scan=%s secret_scan=%s import_commit=%s repo_access=%s private_repo_inventory_commit=%s tail=%s total=%s", scanDuration, secretScanDuration, importDuration, repoAccessDuration, inventoryDuration, - persistDuration, tailDuration, totalDuration, ) diff --git a/internal/kitchen/analyze_test.go b/internal/kitchen/analyze_test.go index 1abd5a6..94873d3 100644 --- a/internal/kitchen/analyze_test.go +++ b/internal/kitchen/analyze_test.go @@ -853,9 +853,14 @@ func TestImportPrivateReposToPantry_AddsPrivateRepos(t *testing.T) { h := NewHandlerWithPublisher(mock, nil) h.database = database - h.importPrivateReposToPantry("sess1") + require.NoError(t, h.importPrivateReposToPantry(context.Background(), "sess1")) p := h.Pantry() + assert.Equal(t, uint64(1), p.Revision()) + persisted, err := database.LoadPantry() + require.NoError(t, err) + require.NotNil(t, persisted) + assert.Equal(t, uint64(1), persisted.Revision()) repos := p.GetAssetsByType(pantry.AssetRepository) assert.Len(t, repos, 2) @@ -877,7 +882,7 @@ func TestImportPrivateReposToPantry_SkipsPublicRepos(t *testing.T) { h := NewHandlerWithPublisher(mock, nil) h.database = database - h.importPrivateReposToPantry("sess1") + require.NoError(t, h.importPrivateReposToPantry(context.Background(), "sess1")) p := h.Pantry() repos := p.GetAssetsByType(pantry.AssetRepository) @@ -897,7 +902,7 @@ func TestImportPrivateReposToPantry_CreatesOrgAssets(t *testing.T) { h := NewHandlerWithPublisher(mock, nil) h.database = database - h.importPrivateReposToPantry("sess1") + require.NoError(t, h.importPrivateReposToPantry(context.Background(), "sess1")) p := h.Pantry() @@ -936,7 +941,7 @@ func TestImportPrivateReposToPantry_ImportsSSHAccessRepos(t *testing.T) { h := NewHandlerWithPublisher(mock, nil) h.database = database - h.importPrivateReposToPantry("sess1") + require.NoError(t, h.importPrivateReposToPantry(context.Background(), "sess1")) repos := h.Pantry().GetAssetsByType(pantry.AssetRepository) require.Len(t, repos, 1) @@ -1021,7 +1026,7 @@ func TestHandleAnalyze_EmptySessionID_SkipsRepoVisibility(t *testing.T) { // When SessionID is empty, this block is skipped entirely. if req.SessionID != "" && h.database != nil { h.recordAnalyzedRepoVisibility(t.Context(), req, result) - h.importPrivateReposToPantry(req.SessionID) + require.NoError(t, h.importPrivateReposToPantry(context.Background(), req.SessionID)) } // Prove: no known entities recorded @@ -1032,7 +1037,7 @@ func TestHandleAnalyze_EmptySessionID_SkipsRepoVisibility(t *testing.T) { // Prove: pantry has no private property p := h.Pantry() - h.importAnalysisToPantry(result) + h.importAnalysisToPantry(h.Pantry(), result) repos := p.GetAssetsByType(pantry.AssetRepository) for _, repo := range repos { _, hasPrivate := repo.Properties["private"] @@ -1079,7 +1084,7 @@ func TestHandleAnalyze_WithSessionID_RecordsRepoVisibility(t *testing.T) { // Same guard as handleAnalyze if req.SessionID != "" && h.database != nil { h.recordAnalyzedRepoVisibility(t.Context(), req, result) - h.importPrivateReposToPantry(req.SessionID) + require.NoError(t, h.importPrivateReposToPantry(context.Background(), req.SessionID)) } // Prove: entity recorded with IsPrivate=true @@ -1177,7 +1182,7 @@ func TestImportAnalysisToPantry_SetsExploitSupportMetadata(t *testing.T) { }, } - h.importAnalysisToPantry(result) + h.importAnalysisToPantry(h.Pantry(), result) vulns := h.Pantry().FindVulnerabilities() require.Len(t, vulns, 1) @@ -1200,7 +1205,7 @@ func TestImportAnalysisToPantry_SkipsSelfHostedRunnerAnalyzeOnlyVuln(t *testing. }, } - h.importAnalysisToPantry(result) + h.importAnalysisToPantry(h.Pantry(), result) assert.Empty(t, h.Pantry().FindVulnerabilities()) targets := h.Pantry().GetAssetsByType(pantry.AssetSelfHostedRunner) @@ -1238,7 +1243,7 @@ func TestImportAnalysisToPantry_CreatesSelfHostedRunnerTargets(t *testing.T) { }, } - h.importAnalysisToPantry(result) + h.importAnalysisToPantry(h.Pantry(), result) targets := h.Pantry().GetAssetsByType(pantry.AssetSelfHostedRunner) require.Len(t, targets, 1) @@ -1265,7 +1270,7 @@ func TestImportAnalysisToPantry_AttachesGitleaksSecretsToFindingRepo(t *testing. }, } - h.importAnalysisToPantry(result) + h.importAnalysisToPantry(h.Pantry(), result) secrets := h.Pantry().GetAssetsByType(pantry.AssetSecret) require.Len(t, secrets, 1) @@ -1295,7 +1300,7 @@ func TestImportAnalysisToPantry_PersistsBashContextBeforeExploitSupport(t *testi }, } - h.importAnalysisToPantry(result) + h.importAnalysisToPantry(h.Pantry(), result) vulns := h.Pantry().FindVulnerabilities() require.Len(t, vulns, 1) @@ -1338,7 +1343,7 @@ func TestImportAnalysisToPantry_PreservesMultiSourceInjectionFindings(t *testing }, } - h.importAnalysisToPantry(result) + h.importAnalysisToPantry(h.Pantry(), result) vulns := h.Pantry().FindVulnerabilities() require.Len(t, vulns, 4) diff --git a/internal/kitchen/committed_pantry_test.go b/internal/kitchen/committed_pantry_test.go new file mode 100644 index 0000000..2b4f7aa --- /dev/null +++ b/internal/kitchen/committed_pantry_test.go @@ -0,0 +1,61 @@ +// Copyright (C) 2026 boostsecurity.io +// SPDX-License-Identifier: AGPL-3.0-or-later + +package kitchen + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/boostsecurityio/smokedmeat/internal/pantry" +) + +func TestHandlerKnownRepositoryCommitsOnePantryRevision(t *testing.T) { + database := newTestDB(t) + h := NewHandlerWithPublisher(&mockPublisher{}, nil) + h.SetDatabase(database) + body, err := json.Marshal(KnownEntityRequest{ + ID: "repo:acme/api", + EntityType: "repo", + Name: "acme/api", + SessionID: "session-1", + DiscoveredVia: "analysis", + IsPrivate: true, + }) + require.NoError(t, err) + req := httptest.NewRequest(http.MethodPost, "/known-entities", bytes.NewReader(body)) + rec := httptest.NewRecorder() + + h.handlePostKnownEntities(rec, req) + + assert.Equal(t, http.StatusCreated, rec.Code) + assert.Equal(t, uint64(1), h.Pantry().Revision()) + persisted, err := database.LoadPantry() + require.NoError(t, err) + require.NotNil(t, persisted) + assert.Equal(t, uint64(1), persisted.Revision()) + repo, err := persisted.GetAsset("github:acme/api") + require.NoError(t, err) + assert.Equal(t, true, repo.Properties["private"]) +} + +func TestHandlerSetDatabaseRestoresCommittedPantryRevision(t *testing.T) { + database := newTestDB(t) + h := NewHandlerWithPublisher(&mockPublisher{}, nil) + h.SetDatabase(database) + require.NoError(t, h.committedPantry().Update(t.Context(), func(candidate *pantry.Pantry) error { + return candidate.AddAsset(pantry.NewOrganization("acme", "github")) + })) + + restarted := NewHandlerWithPublisher(&mockPublisher{}, nil) + restarted.SetDatabase(database) + + assert.Equal(t, uint64(1), restarted.Pantry().Revision()) + assert.True(t, restarted.Pantry().HasAsset("github:org:acme")) +} diff --git a/internal/kitchen/db/db.go b/internal/kitchen/db/db.go index 1d4d956..4d5346a 100644 --- a/internal/kitchen/db/db.go +++ b/internal/kitchen/db/db.go @@ -36,8 +36,8 @@ var ( var schemaKey = []byte("schema") const ( - currentSchemaMajor = 2 - currentSchemaMinor = 5 + currentSchemaMajor = 3 + currentSchemaMinor = 0 legacySchemaMajor = 1 legacySchemaMinor = 0 // Keep this string stable - quickstart readiness checks grep for it in Kitchen logs. diff --git a/internal/kitchen/db/pantry.go b/internal/kitchen/db/pantry.go index ebaaddf..628c4dc 100644 --- a/internal/kitchen/db/pantry.go +++ b/internal/kitchen/db/pantry.go @@ -13,16 +13,18 @@ import ( var pantryKey = []byte("graph") -// SavePantry persists the attack graph to the database. -func (db *DB) SavePantry(p *pantry.Pantry) error { - data, err := json.Marshal(p) - if err != nil { - return err - } +type PantrySnapshotStore struct { + db *DB +} + +func NewPantrySnapshotStore(database *DB) *PantrySnapshotStore { + return &PantrySnapshotStore{db: database} +} - return db.bolt.Update(func(tx *bolt.Tx) error { +func (s *PantrySnapshotStore) Replace(serialized []byte) error { + return s.db.bolt.Update(func(tx *bolt.Tx) error { b := tx.Bucket(BucketPantry) - return b.Put(pantryKey, data) + return b.Put(pantryKey, serialized) }) } @@ -33,7 +35,7 @@ func (db *DB) LoadPantry() (*pantry.Pantry, error) { err := db.bolt.View(func(tx *bolt.Tx) error { b := tx.Bucket(BucketPantry) - data = b.Get(pantryKey) + data = append([]byte(nil), b.Get(pantryKey)...) return nil }) if err != nil { diff --git a/internal/kitchen/db/pantry_test.go b/internal/kitchen/db/pantry_test.go new file mode 100644 index 0000000..91393e5 --- /dev/null +++ b/internal/kitchen/db/pantry_test.go @@ -0,0 +1,68 @@ +// Copyright (C) 2026 boostsecurity.io +// SPDX-License-Identifier: AGPL-3.0-or-later + +package db + +import ( + "context" + "encoding/json" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/boostsecurityio/smokedmeat/internal/pantry" +) + +func TestPantrySnapshotStoreReopensCommittedRevision(t *testing.T) { + path := filepath.Join(t.TempDir(), "pantry.db") + database, err := Open(Config{Path: path, CreateDir: true}) + require.NoError(t, err) + + live := pantry.New() + state := pantry.NewCommittedState(live, NewPantrySnapshotStore(database)) + require.NoError(t, state.Update(context.Background(), func(candidate *pantry.Pantry) error { + return candidate.AddAsset(pantry.NewOrganization("acme", "github")) + })) + require.NoError(t, database.Close()) + + database, err = Open(Config{Path: path}) + require.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, database.Close()) }) + restored, err := database.LoadPantry() + require.NoError(t, err) + require.NotNil(t, restored) + assert.Equal(t, uint64(1), restored.Revision()) + assert.True(t, restored.HasAsset("github:org:acme")) +} + +func TestPantrySnapshotStoreOwnsPersistedBytes(t *testing.T) { + database, err := Open(Config{Path: filepath.Join(t.TempDir(), "pantry.db"), CreateDir: true}) + require.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, database.Close()) }) + + snapshot := pantry.Snapshot{Revision: 7, Assets: []pantry.Asset{pantry.NewOrganization("acme", "github")}} + serialized, err := json.Marshal(snapshot) + require.NoError(t, err) + require.NoError(t, NewPantrySnapshotStore(database).Replace(serialized)) + for index := range serialized { + serialized[index] = 'x' + } + + restored, err := database.LoadPantry() + require.NoError(t, err) + require.NotNil(t, restored) + assert.Equal(t, uint64(7), restored.Revision()) + assert.True(t, restored.HasAsset("github:org:acme")) +} + +func TestLoadPantryReturnsNilWithoutSnapshot(t *testing.T) { + database, err := Open(Config{Path: filepath.Join(t.TempDir(), "pantry.db"), CreateDir: true}) + require.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, database.Close()) }) + + restored, err := database.LoadPantry() + require.NoError(t, err) + assert.Nil(t, restored) +} diff --git a/internal/kitchen/db/schema_test.go b/internal/kitchen/db/schema_test.go index 7faa5cf..b2801d7 100644 --- a/internal/kitchen/db/schema_test.go +++ b/internal/kitchen/db/schema_test.go @@ -78,6 +78,18 @@ func TestOpen_RejectsIncompatibleSchemaMajor(t *testing.T) { assert.ErrorContains(t, err, "incompatible") } +func TestOpen_RejectsSchemaTwoWithPurgeGuidance(t *testing.T) { + path := createSchemaTestDB(t, &schemaMetadata{Major: 2, Minor: 5}, BucketPantry) + + database, err := Open(Config{Path: path}) + require.Error(t, err) + assert.Nil(t, database) + assert.ErrorContains(t, err, "schema 2.5") + assert.ErrorContains(t, err, "binary schema 3.0") + assert.ErrorContains(t, err, "make quickstart-purge") + assert.ErrorContains(t, err, "make dev-quickstart-purge") +} + func TestOpen_RejectsUnknownUnversionedLayout(t *testing.T) { path := createSchemaTestDB(t, nil, []byte("mystery")) diff --git a/internal/kitchen/graph.go b/internal/kitchen/graph.go index 1fb5fb0..7446e57 100644 --- a/internal/kitchen/graph.go +++ b/internal/kitchen/graph.go @@ -34,7 +34,7 @@ func (h *Handler) handleGraph(w http.ResponseWriter, _ *http.Request) { func (h *Handler) handleGraphData(w http.ResponseWriter, r *http.Request) { writeGraphSecurityHeaders(w) p := h.Pantry() - snapshot := buildGraphSnapshot(p, p.Version(), r.URL.Query().Get("mode")) + snapshot := buildGraphSnapshot(p, p.Revision(), r.URL.Query().Get("mode")) data := graphData{ Mode: snapshot.Mode, LargeGraph: snapshot.LargeGraph, diff --git a/internal/kitchen/graph_hub.go b/internal/kitchen/graph_hub.go index 25c7c56..a0c0177 100644 --- a/internal/kitchen/graph_hub.go +++ b/internal/kitchen/graph_hub.go @@ -40,7 +40,7 @@ type GraphClient struct { conn *websocket.Conn send chan GraphMessage hub *GraphHub - version int64 + version uint64 mode string } @@ -103,7 +103,7 @@ func (h *GraphHub) unregister(client *GraphClient) { } func (h *GraphHub) buildSnapshot(mode string) GraphSnapshot { - return buildGraphSnapshot(h.pantry, h.pantry.Version(), mode) + return buildGraphSnapshot(h.pantry, h.pantry.Revision(), mode) } // broadcast sends a message to all connected clients. @@ -123,16 +123,17 @@ func (h *GraphHub) broadcast(msg GraphMessage) { // flushDelta sends accumulated changes to all clients. func (h *GraphHub) flushDelta() { h.deltaMu.Lock() + defer h.deltaMu.Unlock() + delta := h.pendingDelta h.pendingDelta = nil h.batchTimer = nil - h.deltaMu.Unlock() if delta == nil { return } - delta.Version = h.pantry.Version() + // Keep the publication fence held so a replacement snapshot cannot overtake this older delta. h.broadcast(GraphMessage{Type: "delta", Data: delta}) } @@ -143,76 +144,65 @@ func (h *GraphHub) scheduleDeltaFlush() { } } -// PantryObserver implementation - -func (h *GraphHub) OnAssetAdded(asset pantry.Asset) { - h.deltaMu.Lock() - defer h.deltaMu.Unlock() - - if h.pendingDelta == nil { - h.pendingDelta = &GraphDelta{} +func (h *GraphHub) OnPantryChange(change pantry.ChangeSet) { + if change.Kind == pantry.ChangeCommittedState { + h.deltaMu.Lock() + if h.batchTimer != nil { + h.batchTimer.Stop() + } + h.batchTimer = nil + h.pendingDelta = nil + h.deltaMu.Unlock() + h.broadcastSnapshots() + return } - h.pendingDelta.AddedNodes = append(h.pendingDelta.AddedNodes, AssetToGraphNode(asset)) - h.scheduleDeltaFlush() -} -func (h *GraphHub) OnAssetUpdated(asset pantry.Asset, oldState pantry.AssetState) { h.deltaMu.Lock() defer h.deltaMu.Unlock() if h.pendingDelta == nil { h.pendingDelta = &GraphDelta{} } - node := AssetToGraphNode(asset) - h.pendingDelta.UpdatedNodes = append(h.pendingDelta.UpdatedNodes, NodeUpdate{ - ID: asset.ID, - OldState: string(oldState), - NewState: string(asset.State), - Label: node.Label, - Properties: node.Properties, - TooltipProperties: node.TooltipProperties, - }) - h.scheduleDeltaFlush() -} - -func (h *GraphHub) OnRelationshipAdded(from, to string, rel pantry.Relationship) { - h.deltaMu.Lock() - defer h.deltaMu.Unlock() - - if h.pendingDelta == nil { - h.pendingDelta = &GraphDelta{} + h.pendingDelta.Version = change.Revision + for _, asset := range change.Granular.AddedAssets { + h.pendingDelta.AddedNodes = append(h.pendingDelta.AddedNodes, AssetToGraphNode(asset)) } - h.pendingDelta.AddedEdges = append(h.pendingDelta.AddedEdges, GraphEdge{ - Source: from, - Target: to, - Type: string(rel.Type), - }) - h.scheduleDeltaFlush() -} - -func (h *GraphHub) OnAssetRemoved(id string) { - h.deltaMu.Lock() - defer h.deltaMu.Unlock() - - if h.pendingDelta == nil { - h.pendingDelta = &GraphDelta{} + for _, update := range change.Granular.UpdatedAssets { + node := AssetToGraphNode(update.After) + h.pendingDelta.UpdatedNodes = append(h.pendingDelta.UpdatedNodes, NodeUpdate{ + ID: update.After.ID, + OldState: string(update.Before.State), + NewState: string(update.After.State), + Label: node.Label, + Properties: node.Properties, + TooltipProperties: node.TooltipProperties, + }) + } + for _, edge := range change.Granular.AddedRelationships { + h.pendingDelta.AddedEdges = append(h.pendingDelta.AddedEdges, GraphEdge{ + Source: edge.From, + Target: edge.To, + Type: string(edge.Relationship.Type), + }) + } + h.pendingDelta.RemovedNodes = append(h.pendingDelta.RemovedNodes, change.Granular.RemovedAssetIDs...) + for _, edge := range change.Granular.RemovedRelationships { + h.pendingDelta.RemovedEdges = append(h.pendingDelta.RemovedEdges, EdgeRef{Source: edge.From, Target: edge.To}) } - h.pendingDelta.RemovedNodes = append(h.pendingDelta.RemovedNodes, id) h.scheduleDeltaFlush() } -func (h *GraphHub) OnRelationshipRemoved(from, to string) { - h.deltaMu.Lock() - defer h.deltaMu.Unlock() - - if h.pendingDelta == nil { - h.pendingDelta = &GraphDelta{} +func (h *GraphHub) broadcastSnapshots() { + h.mu.RLock() + defer h.mu.RUnlock() + for client := range h.clients { + snapshot := h.buildSnapshot(client.mode) + select { + case client.send <- GraphMessage{Type: "snapshot", Data: snapshot}: + default: + slog.Warn("graph client buffer full, dropping snapshot") + } } - h.pendingDelta.RemovedEdges = append(h.pendingDelta.RemovedEdges, EdgeRef{ - Source: from, - Target: to, - }) - h.scheduleDeltaFlush() } // ClientCount returns the number of connected graph clients. diff --git a/internal/kitchen/graph_hub_test.go b/internal/kitchen/graph_hub_test.go new file mode 100644 index 0000000..71d19ca --- /dev/null +++ b/internal/kitchen/graph_hub_test.go @@ -0,0 +1,71 @@ +// Copyright (C) 2026 boostsecurity.io +// SPDX-License-Identifier: AGPL-3.0-or-later + +package kitchen + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/boostsecurityio/smokedmeat/internal/pantry" +) + +type graphHubSnapshotStore struct{} + +func (graphHubSnapshotStore) Replace([]byte) error { return nil } + +func TestGraphHubTranslatesCommittedGranularChangeSet(t *testing.T) { + live := pantry.New() + hub := NewGraphHub(live) + client := &GraphClient{send: make(chan GraphMessage, 2), hub: hub, mode: graphModeFull} + hub.register(client) + t.Cleanup(func() { hub.unregister(client) }) + state := pantry.NewCommittedState(live, graphHubSnapshotStore{}) + + require.NoError(t, state.Update(context.Background(), func(candidate *pantry.Pantry) error { + return candidate.AddAsset(pantry.NewOrganization("acme", "github")) + })) + hub.flushDelta() + + message := <-client.send + assert.Equal(t, "delta", message.Type) + delta, ok := message.Data.(*GraphDelta) + require.True(t, ok) + assert.Equal(t, uint64(1), delta.Version) + require.Len(t, delta.AddedNodes, 1) + assert.Equal(t, "github:org:acme", delta.AddedNodes[0].ID) +} + +func TestGraphHubCommittedStateMarkerSupersedesPendingDelta(t *testing.T) { + live := pantry.New() + hub := NewGraphHub(live) + client := &GraphClient{send: make(chan GraphMessage, 2), hub: hub, mode: graphModeFull} + hub.register(client) + t.Cleanup(func() { hub.unregister(client) }) + state := pantry.NewCommittedState(live, graphHubSnapshotStore{}) + + require.NoError(t, state.Update(context.Background(), func(candidate *pantry.Pantry) error { + return candidate.AddAsset(pantry.NewOrganization("acme", "github")) + })) + require.NoError(t, state.Replace(context.Background(), func(candidate *pantry.Pantry) error { + return candidate.AddAsset(pantry.NewOrganization("globex", "github")) + })) + + message := <-client.send + assert.Equal(t, "snapshot", message.Type) + snapshot, ok := message.Data.(GraphSnapshot) + require.True(t, ok) + assert.Equal(t, uint64(2), snapshot.Version) + assert.Equal(t, 2, snapshot.TotalNodes) + + time.Sleep(2 * graphBatchWindow) + select { + case unexpected := <-client.send: + t.Fatalf("unexpected state after committed snapshot: %#v", unexpected) + default: + } +} diff --git a/internal/kitchen/graph_ws.go b/internal/kitchen/graph_ws.go index 3b4f3fc..a4b62b0 100644 --- a/internal/kitchen/graph_ws.go +++ b/internal/kitchen/graph_ws.go @@ -31,7 +31,7 @@ type GraphMessage struct { // GraphSnapshot is the initial full graph state sent on connect. type GraphSnapshot struct { - Version int64 `json:"version"` + Version uint64 `json:"version"` Mode string `json:"mode"` LargeGraph bool `json:"large_graph"` TotalNodes int `json:"total_nodes"` @@ -43,7 +43,7 @@ type GraphSnapshot struct { // GraphDelta contains incremental changes to the graph. type GraphDelta struct { - Version int64 `json:"version"` + Version uint64 `json:"version"` AddedNodes []GraphNode `json:"added_nodes,omitempty"` UpdatedNodes []NodeUpdate `json:"updated_nodes,omitempty"` AddedEdges []GraphEdge `json:"added_edges,omitempty"` @@ -189,10 +189,10 @@ func buildGraphSelection(p *pantry.Pantry, requestedMode string) graphSelection } } -func buildGraphSnapshot(p *pantry.Pantry, version int64, requestedMode string) GraphSnapshot { +func buildGraphSnapshot(p *pantry.Pantry, revision uint64, requestedMode string) GraphSnapshot { selection := buildGraphSelection(p, requestedMode) return GraphSnapshot{ - Version: version, + Version: revision, Mode: selection.mode, LargeGraph: selection.largeGraph, TotalNodes: selection.totalNodes, diff --git a/internal/kitchen/handlers.go b/internal/kitchen/handlers.go index 26f5690..b197198 100644 --- a/internal/kitchen/handlers.go +++ b/internal/kitchen/handlers.go @@ -92,7 +92,11 @@ type Handler struct { sessions *SessionRegistry database *db.DB operators *OperatorHub + pantryMu sync.Mutex pantry *pantry.Pantry + pantryState *pantry.CommittedState + pantryStateFor *pantry.Pantry + pantryStateDB *db.DB auth *auth.Auth preflightCache *deployPreflightCache sourceCache *sourceCache @@ -102,6 +106,10 @@ type Handler struct { analysisRuns map[string]*cachedAnalysisResult } +type volatilePantrySnapshotStore struct{} + +func (volatilePantrySnapshotStore) Replace([]byte) error { return nil } + // NewHandler creates a new Handler. func NewHandler(publisher *pass.Publisher, store *OrderStore, sessions *SessionRegistry) *Handler { return &Handler{ @@ -134,7 +142,12 @@ func NewHandlerWithPublisher(publisher Publisher, store *OrderStore) *Handler { // SetDatabase sets the database for persistence and loads existing pantry. func (h *Handler) SetDatabase(database *db.DB) { + h.pantryMu.Lock() + defer h.pantryMu.Unlock() h.database = database + h.pantryState = nil + h.pantryStateFor = nil + h.pantryStateDB = nil if h.stagerStore != nil { h.stagerStore.config.DeleteHook = h.deleteStager if database == nil { @@ -154,18 +167,33 @@ func (h *Handler) SetDatabase(database *db.DB) { // Pantry returns the attack graph, creating it if needed. func (h *Handler) Pantry() *pantry.Pantry { + h.pantryMu.Lock() + defer h.pantryMu.Unlock() + return h.pantryLocked() +} + +func (h *Handler) pantryLocked() *pantry.Pantry { if h.pantry == nil { h.pantry = pantry.New() } return h.pantry } -// SavePantry persists the pantry to the database. -func (h *Handler) SavePantry() error { - if h.database == nil || h.pantry == nil { - return nil +func (h *Handler) committedPantry() *pantry.CommittedState { + h.pantryMu.Lock() + defer h.pantryMu.Unlock() + live := h.pantryLocked() + if h.pantryState != nil && h.pantryStateFor == live && h.pantryStateDB == h.database { + return h.pantryState } - return h.database.SavePantry(h.pantry) + var store pantry.SnapshotStore = volatilePantrySnapshotStore{} + if h.database != nil { + store = db.NewPantrySnapshotStore(h.database) + } + h.pantryState = pantry.NewCommittedState(live, store) + h.pantryStateFor = live + h.pantryStateDB = h.database + return h.pantryState } // SetOperatorHub sets the operator hub for WebSocket broadcasts. @@ -1512,9 +1540,14 @@ func (h *Handler) handlePostKnownEntities(w http.ResponseWriter, r *http.Request if req.EntityType == "repo" { parts := strings.Split(req.Name, "/") if len(parts) >= 2 { - p := h.Pantry() - upsertKnownRepoAsset(p, row) - _ = h.SavePantry() + if err := h.committedPantry().Update(r.Context(), func(candidate *pantry.Pantry) error { + upsertKnownRepoAsset(candidate, row) + return nil + }); err != nil { + slog.Warn("failed to commit known repository to pantry", "id", req.ID, "error", err) + http.Error(w, "failed to persist entity", http.StatusInternalServerError) + return + } } } diff --git a/internal/kitchen/purge.go b/internal/kitchen/purge.go index 532191e..49d8276 100644 --- a/internal/kitchen/purge.go +++ b/internal/kitchen/purge.go @@ -4,6 +4,7 @@ package kitchen import ( + "context" "encoding/json" "fmt" "net/http" @@ -50,7 +51,7 @@ func (h *Handler) handlePurge(w http.ResponseWriter, r *http.Request) { return } - resp, err := h.runPurge(req.SessionID, scopeType, scopeValue, req.DryRun) + resp, err := h.runPurge(r.Context(), req.SessionID, scopeType, scopeValue, req.DryRun) if err != nil { http.Error(w, "failed to purge state", http.StatusInternalServerError) return @@ -85,7 +86,7 @@ func normalizePurgeScope(scopeType, scopeValue string) (normalizedType, normaliz } } -func (h *Handler) runPurge(sessionID, scopeType, scopeValue string, dryRun bool) (PurgeResponse, error) { +func (h *Handler) runPurge(ctx context.Context, sessionID, scopeType, scopeValue string, dryRun bool) (PurgeResponse, error) { sessionID = strings.TrimSpace(sessionID) if sessionID == "" { return PurgeResponse{}, fmt.Errorf("session_id is required") @@ -99,9 +100,8 @@ func (h *Handler) runPurge(sessionID, scopeType, scopeValue string, dryRun bool) DryRun: dryRun, } - if h.pantry != nil { - resp.PantryAssets = countPurgeAssets(h.pantry, purgeRootAssetID(scopeType, scopeValue)) - } + p := h.Pantry() + resp.PantryAssets = countPurgeAssets(p, purgeRootAssetID(scopeType, scopeValue)) if h.database != nil { entityRepo := db.NewKnownEntityRepository(h.database) count, err := entityRepo.CountByScopeAndSession(db.EntityType(scopeType), scopeValue, sessionID) @@ -115,13 +115,15 @@ func (h *Handler) runPurge(sessionID, scopeType, scopeValue string, dryRun bool) return resp, nil } - if h.pantry != nil && resp.PantryAssets > 0 { - for _, id := range collectPurgeAssetIDs(h.pantry, purgeRootAssetID(scopeType, scopeValue)) { - if err := h.pantry.RemoveAsset(id); err != nil && err != pantry.ErrAssetNotFound { - return PurgeResponse{}, err + if resp.PantryAssets > 0 { + if err := h.committedPantry().Update(ctx, func(candidate *pantry.Pantry) error { + for _, id := range collectPurgeAssetIDs(candidate, purgeRootAssetID(scopeType, scopeValue)) { + if err := candidate.RemoveAsset(id); err != nil && err != pantry.ErrAssetNotFound { + return err + } } - } - if err := h.SavePantry(); err != nil { + return nil + }); err != nil { return PurgeResponse{}, err } } diff --git a/internal/kitchen/purge_test.go b/internal/kitchen/purge_test.go index 746519e..205805f 100644 --- a/internal/kitchen/purge_test.go +++ b/internal/kitchen/purge_test.go @@ -5,6 +5,7 @@ package kitchen import ( "bytes" + "context" "net/http" "net/http/httptest" "testing" @@ -37,7 +38,7 @@ func TestHandler_runPurge_PreviewCountsRepoScopeForRequestingSession(t *testing. h.database = database h.pantry = purgeTestPantry(t) - resp, err := h.runPurge("sess-1", "repo", "acme/api", true) + resp, err := h.runPurge(context.Background(), "sess-1", "repo", "acme/api", true) require.NoError(t, err) assert.Equal(t, "preview", resp.Status) @@ -74,7 +75,7 @@ func TestHandler_runPurge_ExecuteRemovesOrgScopeAndPreservesHistory(t *testing.T h.database = database h.pantry = purgeTestPantry(t) - resp, err := h.runPurge("sess-1", "org", "acme", false) + resp, err := h.runPurge(context.Background(), "sess-1", "org", "acme", false) require.NoError(t, err) assert.Equal(t, "purged", resp.Status) @@ -88,6 +89,12 @@ func TestHandler_runPurge_ExecuteRemovesOrgScopeAndPreservesHistory(t *testing.T assert.False(t, h.Pantry().HasAsset("github:acme/api:workflow:.github/workflows/build.yml")) assert.True(t, h.Pantry().HasAsset("github:org:globex")) assert.True(t, h.Pantry().HasAsset("github:globex/portal")) + assert.Equal(t, uint64(1), h.Pantry().Revision()) + persisted, err := database.LoadPantry() + require.NoError(t, err) + require.NotNil(t, persisted) + assert.Equal(t, uint64(1), persisted.Revision()) + assert.False(t, persisted.HasAsset("github:org:acme")) sess1Entities, err := entityRepo.ListBySession("sess-1") require.NoError(t, err) @@ -113,7 +120,7 @@ func TestHandler_runPurge_ExecuteRemovesOrgScopeAndPreservesHistory(t *testing.T func TestHandler_runPurge_RejectsEmptySessionID(t *testing.T) { h := NewHandlerWithPublisher(&mockPublisher{}, nil) - _, err := h.runPurge("", "repo", "acme/api", true) + _, err := h.runPurge(context.Background(), "", "repo", "acme/api", true) require.Error(t, err) assert.Equal(t, "session_id is required", err.Error()) diff --git a/internal/kitchen/server.go b/internal/kitchen/server.go index 4b4d3b3..f757652 100644 --- a/internal/kitchen/server.go +++ b/internal/kitchen/server.go @@ -451,11 +451,6 @@ func (s *Server) Shutdown(ctx context.Context) error { if s.handler != nil { s.handler.StagerStore().StopCleanup() - if err := s.handler.SavePantry(); err != nil { - slog.Warn("failed to save pantry on shutdown", "error", err) - } else { - slog.Info("pantry saved on shutdown") - } } s.closeDB() diff --git a/internal/pantry/committed_state.go b/internal/pantry/committed_state.go new file mode 100644 index 0000000..5da29ae --- /dev/null +++ b/internal/pantry/committed_state.go @@ -0,0 +1,281 @@ +// Copyright (C) 2026 boostsecurity.io +// SPDX-License-Identifier: AGPL-3.0-or-later + +package pantry + +import ( + "context" + "encoding/json" + "fmt" + "log/slog" + "math" + "reflect" + "sort" + "time" +) + +type SnapshotStore interface { + Replace(serialized []byte) error +} + +type Snapshot struct { + Revision uint64 `json:"revision"` + Assets []Asset `json:"assets"` + Edges []Edge `json:"edges"` +} + +type ChangeKind string + +const ( + ChangeGranular ChangeKind = "granular" + ChangeCommittedState ChangeKind = "committed_state" +) + +type AssetChange struct { + Before Asset + After Asset +} + +type EdgeRef struct { + From string + To string +} + +type GranularChanges struct { + AddedAssets []Asset + UpdatedAssets []AssetChange + RemovedAssetIDs []string + AddedRelationships []Edge + RemovedRelationships []EdgeRef +} + +func (c GranularChanges) count() int { + return len(c.AddedAssets) + len(c.UpdatedAssets) + len(c.RemovedAssetIDs) + len(c.AddedRelationships) + len(c.RemovedRelationships) +} + +type ChangeSet struct { + Kind ChangeKind + BaseRevision uint64 + Revision uint64 + Granular GranularChanges +} + +type Observer interface { + OnPantryChange(change ChangeSet) +} + +type CommittedState struct { + live *Pantry + store SnapshotStore + gate chan struct{} +} + +func NewCommittedState(live *Pantry, store SnapshotStore) *CommittedState { + if live == nil { + panic("pantry: committed state requires a live Pantry") + } + if store == nil { + panic("pantry: committed state requires a SnapshotStore") + } + return &CommittedState{live: live, store: store, gate: live.committedWriterGate()} +} + +func (s *CommittedState) Update(ctx context.Context, change func(candidate *Pantry) error) error { + return s.commit(ctx, ChangeGranular, change) +} + +func (s *CommittedState) Replace(ctx context.Context, change func(candidate *Pantry) error) error { + return s.commit(ctx, ChangeCommittedState, change) +} + +func (s *CommittedState) commit(ctx context.Context, kind ChangeKind, change func(candidate *Pantry) error) error { + if change == nil { + return fmt.Errorf("pantry: committed state change callback is required") + } + + select { + case <-ctx.Done(): + return ctx.Err() + case <-s.gate: + } + defer func() { s.gate <- struct{}{} }() + + totalStarted := time.Now() + candidateStarted := time.Now() + base := s.live.Snapshot() + candidate := pantryFromSnapshot(base) + candidateDuration := time.Since(candidateStarted) + + if err := change(candidate); err != nil { + return err + } + if err := ctx.Err(); err != nil { + return err + } + + diffStarted := time.Now() + next := candidate.Snapshot() + granular := diffSnapshots(base, next) + diffDuration := time.Since(diffStarted) + if granular.count() == 0 { + slog.Debug("pantry state unchanged", + "revision", base.Revision, + "candidate_duration", candidateDuration, + "diff_duration", diffDuration, + "total_duration", time.Since(totalStarted)) + return nil + } + if base.Revision == math.MaxUint64 { + return fmt.Errorf("pantry: committed revision exhausted") + } + + next.Revision = base.Revision + 1 + candidate.setRevision(next.Revision) + + validationStarted := time.Now() + if err := validateSnapshot(next); err != nil { + return err + } + validationDuration := time.Since(validationStarted) + + serializationStarted := time.Now() + serialized, err := json.Marshal(next) + serializationDuration := time.Since(serializationStarted) + if err != nil { + return fmt.Errorf("serialize Pantry snapshot: %w", err) + } + if err := ctx.Err(); err != nil { + return err + } + + persistenceStarted := time.Now() + if err := s.store.Replace(serialized); err != nil { + return fmt.Errorf("persist Pantry snapshot: %w", err) + } + persistenceDuration := time.Since(persistenceStarted) + + replacementStarted := time.Now() + s.live.replaceWith(candidate) + replacementDuration := time.Since(replacementStarted) + + publicationStarted := time.Now() + changeSet := ChangeSet{ + Kind: kind, + BaseRevision: base.Revision, + Revision: next.Revision, + } + if kind == ChangeGranular { + changeSet.Granular = granular + } + s.live.notifyChange(changeSet) + publicationDuration := time.Since(publicationStarted) + + slog.Debug("pantry state committed", + "kind", kind, + "base_revision", base.Revision, + "revision", next.Revision, + "assets", len(next.Assets), + "relationships", len(next.Edges), + "granular_changes", granular.count(), + "serialized_bytes", len(serialized), + "candidate_duration", candidateDuration, + "diff_duration", diffDuration, + "validation_duration", validationDuration, + "serialization_duration", serializationDuration, + "persistence_duration", persistenceDuration, + "replacement_duration", replacementDuration, + "publication_duration", publicationDuration, + "total_duration", time.Since(totalStarted)) + return nil +} + +func validateSnapshot(snapshot Snapshot) error { + assets := make(map[string]struct{}, len(snapshot.Assets)) + for _, asset := range snapshot.Assets { + if asset.ID == "" { + return fmt.Errorf("validate Pantry snapshot: asset ID is empty") + } + if _, exists := assets[asset.ID]; exists { + return fmt.Errorf("validate Pantry snapshot: duplicate asset %q", asset.ID) + } + assets[asset.ID] = struct{}{} + } + edges := make(map[string]struct{}, len(snapshot.Edges)) + for _, edge := range snapshot.Edges { + if _, exists := assets[edge.From]; !exists { + return fmt.Errorf("validate Pantry snapshot: relationship source %q is missing", edge.From) + } + if _, exists := assets[edge.To]; !exists { + return fmt.Errorf("validate Pantry snapshot: relationship target %q is missing", edge.To) + } + key := edgeKey(edge.From, edge.To) + if _, exists := edges[key]; exists { + return fmt.Errorf("validate Pantry snapshot: duplicate relationship %q", key) + } + edges[key] = struct{}{} + } + return nil +} + +func diffSnapshots(before, after Snapshot) GranularChanges { + beforeAssets := make(map[string]Asset, len(before.Assets)) + afterAssets := make(map[string]Asset, len(after.Assets)) + for _, asset := range before.Assets { + beforeAssets[asset.ID] = asset + } + for _, asset := range after.Assets { + afterAssets[asset.ID] = asset + } + + changes := GranularChanges{} + for id, asset := range afterAssets { + previous, exists := beforeAssets[id] + switch { + case !exists: + changes.AddedAssets = append(changes.AddedAssets, cloneAsset(asset)) + case !reflect.DeepEqual(previous, asset): + changes.UpdatedAssets = append(changes.UpdatedAssets, AssetChange{Before: cloneAsset(previous), After: cloneAsset(asset)}) + } + } + for id := range beforeAssets { + if _, exists := afterAssets[id]; !exists { + changes.RemovedAssetIDs = append(changes.RemovedAssetIDs, id) + } + } + + beforeEdges := make(map[string]Edge, len(before.Edges)) + afterEdges := make(map[string]Edge, len(after.Edges)) + for _, edge := range before.Edges { + beforeEdges[edgeKey(edge.From, edge.To)] = edge + } + for _, edge := range after.Edges { + afterEdges[edgeKey(edge.From, edge.To)] = edge + } + for key, edge := range afterEdges { + previous, exists := beforeEdges[key] + if !exists { + changes.AddedRelationships = append(changes.AddedRelationships, cloneEdge(edge)) + continue + } + if !reflect.DeepEqual(previous.Relationship, edge.Relationship) { + changes.RemovedRelationships = append(changes.RemovedRelationships, EdgeRef{From: previous.From, To: previous.To}) + changes.AddedRelationships = append(changes.AddedRelationships, cloneEdge(edge)) + } + } + for key, edge := range beforeEdges { + if _, exists := afterEdges[key]; !exists { + changes.RemovedRelationships = append(changes.RemovedRelationships, EdgeRef{From: edge.From, To: edge.To}) + } + } + + sort.Slice(changes.AddedAssets, func(i, j int) bool { return changes.AddedAssets[i].ID < changes.AddedAssets[j].ID }) + sort.Slice(changes.UpdatedAssets, func(i, j int) bool { return changes.UpdatedAssets[i].After.ID < changes.UpdatedAssets[j].After.ID }) + sort.Strings(changes.RemovedAssetIDs) + sort.Slice(changes.AddedRelationships, func(i, j int) bool { + return edgeKey(changes.AddedRelationships[i].From, changes.AddedRelationships[i].To) < edgeKey(changes.AddedRelationships[j].From, changes.AddedRelationships[j].To) + }) + sort.Slice(changes.RemovedRelationships, func(i, j int) bool { + return edgeKey(changes.RemovedRelationships[i].From, changes.RemovedRelationships[i].To) < edgeKey(changes.RemovedRelationships[j].From, changes.RemovedRelationships[j].To) + }) + return changes +} diff --git a/internal/pantry/committed_state_test.go b/internal/pantry/committed_state_test.go new file mode 100644 index 0000000..4d882bf --- /dev/null +++ b/internal/pantry/committed_state_test.go @@ -0,0 +1,509 @@ +// Copyright (C) 2026 boostsecurity.io +// SPDX-License-Identifier: AGPL-3.0-or-later + +package pantry + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type memorySnapshotStore struct { + mu sync.Mutex + data []byte + err error + before func() + writes int + started chan struct{} + release chan struct{} +} + +func (s *memorySnapshotStore) Replace(serialized []byte) error { + if s.started != nil { + select { + case s.started <- struct{}{}: + default: + } + } + if s.release != nil { + <-s.release + } + if s.before != nil { + s.before() + } + s.mu.Lock() + defer s.mu.Unlock() + if s.err != nil { + return s.err + } + s.data = append([]byte(nil), serialized...) + s.writes++ + return nil +} + +func (s *memorySnapshotStore) snapshot(t *testing.T) Snapshot { + t.Helper() + s.mu.Lock() + defer s.mu.Unlock() + var snapshot Snapshot + require.NoError(t, json.Unmarshal(s.data, &snapshot)) + return snapshot +} + +func (s *memorySnapshotStore) writeCount() int { + s.mu.Lock() + defer s.mu.Unlock() + return s.writes +} + +type recordingObserver struct { + mu sync.Mutex + changes []ChangeSet +} + +type observerFunc func(ChangeSet) + +func (f observerFunc) OnPantryChange(change ChangeSet) { + f(change) +} + +type countingJSONValue struct { + calls *atomic.Int32 +} + +func (v countingJSONValue) MarshalJSON() ([]byte, error) { + v.calls.Add(1) + return []byte(`"counted"`), nil +} + +func (o *recordingObserver) OnPantryChange(change ChangeSet) { + o.mu.Lock() + o.changes = append(o.changes, change) + o.mu.Unlock() +} + +func (o *recordingObserver) all() []ChangeSet { + o.mu.Lock() + defer o.mu.Unlock() + return append([]ChangeSet(nil), o.changes...) +} + +func TestCommittedStateUpdatePersistsBeforeOneGranularPublication(t *testing.T) { + live := New() + store := &memorySnapshotStore{} + observer := &recordingObserver{} + live.AddObserver(observer) + state := NewCommittedState(live, store) + pointer := live + + err := state.Update(context.Background(), func(candidate *Pantry) error { + repo := NewRepository("acme", "api", "github") + workflow := NewWorkflow(repo.ID, ".github/workflows/ci.yml") + require.NoError(t, candidate.AddAsset(repo)) + require.NoError(t, candidate.AddAsset(workflow)) + return candidate.AddRelationship(repo.ID, workflow.ID, Contains()) + }) + require.NoError(t, err) + + assert.Same(t, pointer, live) + assert.Equal(t, uint64(1), live.Revision()) + assert.Equal(t, uint64(1), store.snapshot(t).Revision) + changes := observer.all() + require.Len(t, changes, 1) + assert.Equal(t, ChangeGranular, changes[0].Kind) + assert.Equal(t, uint64(0), changes[0].BaseRevision) + assert.Equal(t, uint64(1), changes[0].Revision) + assert.Len(t, changes[0].Granular.AddedAssets, 2) + assert.Len(t, changes[0].Granular.AddedRelationships, 1) +} + +func TestCommittedStateIsolatesObserverChangeSets(t *testing.T) { + live := New() + store := &memorySnapshotStore{} + live.AddObserver(observerFunc(func(change ChangeSet) { + change.Granular.AddedAssets[0].ID = "mutated" + change.Granular.AddedAssets[0].Properties["nested"].(map[string]any)["value"] = "mutated" + change.Granular.AddedRelationships[0].Relationship.Properties["nested"].(map[string]any)["value"] = "mutated" + })) + observer := &recordingObserver{} + live.AddObserver(observer) + state := NewCommittedState(live, store) + + repo := NewRepository("acme", "api", "github") + repo.Properties["nested"] = map[string]any{"value": "original"} + workflow := NewWorkflow(repo.ID, ".github/workflows/ci.yml") + relationship := Contains().WithProperty("nested", map[string]any{"value": "original"}) + require.NoError(t, state.Update(context.Background(), func(candidate *Pantry) error { + require.NoError(t, candidate.AddAsset(repo)) + require.NoError(t, candidate.AddAsset(workflow)) + return candidate.AddRelationship(repo.ID, workflow.ID, relationship) + })) + + changes := observer.all() + require.Len(t, changes, 1) + require.Len(t, changes[0].Granular.AddedAssets, 2) + assert.Equal(t, repo.ID, changes[0].Granular.AddedAssets[0].ID) + assert.Equal(t, "original", changes[0].Granular.AddedAssets[0].Properties["nested"].(map[string]any)["value"]) + require.Len(t, changes[0].Granular.AddedRelationships, 1) + assert.Equal(t, "original", changes[0].Granular.AddedRelationships[0].Relationship.Properties["nested"].(map[string]any)["value"]) +} + +func TestCommittedStateReplacePublishesOnlyCommittedStateMarker(t *testing.T) { + live := New() + store := &memorySnapshotStore{} + observer := &recordingObserver{} + live.AddObserver(observer) + state := NewCommittedState(live, store) + + require.NoError(t, state.Replace(context.Background(), func(candidate *Pantry) error { + return candidate.AddAsset(NewOrganization("acme", "github")) + })) + + changes := observer.all() + require.Len(t, changes, 1) + assert.Equal(t, ChangeCommittedState, changes[0].Kind) + assert.Equal(t, 0, changes[0].Granular.count()) +} + +func TestCommittedStateNoOpDoesNotPersistAdvanceOrPublish(t *testing.T) { + live := New() + store := &memorySnapshotStore{} + observer := &recordingObserver{} + live.AddObserver(observer) + state := NewCommittedState(live, store) + + require.NoError(t, state.Update(context.Background(), func(*Pantry) error { return nil })) + + assert.Zero(t, live.Revision()) + assert.Zero(t, store.writeCount()) + assert.Empty(t, observer.all()) +} + +func TestCommittedStateSerializesChangedCandidateOnce(t *testing.T) { + live := New() + store := &memorySnapshotStore{} + state := NewCommittedState(live, store) + var calls atomic.Int32 + + require.NoError(t, state.Update(context.Background(), func(candidate *Pantry) error { + asset := NewOrganization("acme", "github") + asset.SetProperty("counted", countingJSONValue{calls: &calls}) + return candidate.AddAsset(asset) + })) + + assert.Equal(t, int32(1), calls.Load()) +} + +func TestCommittedStateRejectsRevisionOverflowBeforePersistence(t *testing.T) { + live := New() + live.revision = ^uint64(0) + store := &memorySnapshotStore{} + state := NewCommittedState(live, store) + + err := state.Update(context.Background(), func(candidate *Pantry) error { + return candidate.AddAsset(NewOrganization("acme", "github")) + }) + + require.ErrorContains(t, err, "revision exhausted") + assert.Equal(t, ^uint64(0), live.Revision()) + assert.Zero(t, store.writeCount()) +} + +func TestCommittedStateFailuresLeavePriorStateAuthoritative(t *testing.T) { + tests := []struct { + name string + store *memorySnapshotStore + change func(context.Context, context.CancelFunc, *Pantry) error + want error + }{ + { + name: "callback", + store: &memorySnapshotStore{}, + change: func(_ context.Context, _ context.CancelFunc, candidate *Pantry) error { + require.NoError(t, candidate.AddAsset(NewOrganization("acme", "github"))) + return errors.New("callback failed") + }, + want: errors.New("callback failed"), + }, + { + name: "serialization", + store: &memorySnapshotStore{}, + change: func(_ context.Context, _ context.CancelFunc, candidate *Pantry) error { + asset := NewOrganization("acme", "github") + asset.SetProperty("invalid", make(chan int)) + return candidate.AddAsset(asset) + }, + }, + { + name: "persistence", + store: &memorySnapshotStore{err: errors.New("disk full")}, + change: func(_ context.Context, _ context.CancelFunc, candidate *Pantry) error { + return candidate.AddAsset(NewOrganization("acme", "github")) + }, + want: errors.New("disk full"), + }, + { + name: "cancellation", + store: &memorySnapshotStore{}, + change: func(_ context.Context, cancel context.CancelFunc, candidate *Pantry) error { + require.NoError(t, candidate.AddAsset(NewOrganization("acme", "github"))) + cancel() + return nil + }, + want: context.Canceled, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + live := New() + observer := &recordingObserver{} + live.AddObserver(observer) + state := NewCommittedState(live, test.store) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + err := state.Update(ctx, func(candidate *Pantry) error { + return test.change(ctx, cancel, candidate) + }) + require.Error(t, err) + if test.want != nil { + assert.Contains(t, err.Error(), test.want.Error()) + } + assert.Zero(t, live.Revision()) + assert.Zero(t, live.Size()) + assert.Zero(t, test.store.writeCount()) + assert.Empty(t, observer.all()) + }) + } +} + +func TestCommittedStateIgnoresCancellationAfterPersistence(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + store := &memorySnapshotStore{before: cancel} + live := New() + state := NewCommittedState(live, store) + + err := state.Update(ctx, func(candidate *Pantry) error { + return candidate.AddAsset(NewOrganization("acme", "github")) + }) + + require.NoError(t, err) + assert.Equal(t, uint64(1), live.Revision()) + assert.True(t, live.HasAsset("github:org:acme")) +} + +func TestCommittedStateReadersSeeOldStateUntilPersistenceCompletes(t *testing.T) { + store := &memorySnapshotStore{started: make(chan struct{}, 1), release: make(chan struct{})} + live := New() + observer := &recordingObserver{} + live.AddObserver(observer) + state := NewCommittedState(live, store) + done := make(chan error, 1) + + go func() { + done <- state.Update(context.Background(), func(candidate *Pantry) error { + return candidate.AddAsset(NewOrganization("acme", "github")) + }) + }() + + <-store.started + assert.Zero(t, live.Revision()) + assert.False(t, live.HasAsset("github:org:acme")) + assert.Empty(t, observer.all()) + close(store.release) + require.NoError(t, <-done) + assert.Equal(t, uint64(1), live.Revision()) + assert.True(t, live.HasAsset("github:org:acme")) + assert.Len(t, observer.all(), 1) +} + +func TestCommittedStateSerializesWritersAgainstLatestRevision(t *testing.T) { + store := &memorySnapshotStore{started: make(chan struct{}, 1), release: make(chan struct{})} + live := New() + observer := &recordingObserver{} + live.AddObserver(observer) + state := NewCommittedState(live, store) + firstDone := make(chan error, 1) + secondDone := make(chan error, 1) + + go func() { + firstDone <- state.Update(context.Background(), func(candidate *Pantry) error { + return candidate.AddAsset(NewOrganization("acme", "github")) + }) + }() + <-store.started + go func() { + secondDone <- state.Update(context.Background(), func(candidate *Pantry) error { + return candidate.AddAsset(NewOrganization("globex", "github")) + }) + }() + close(store.release) + + require.NoError(t, <-firstDone) + require.NoError(t, <-secondDone) + assert.Equal(t, uint64(2), live.Revision()) + assert.True(t, live.HasAsset("github:org:acme")) + assert.True(t, live.HasAsset("github:org:globex")) + changes := observer.all() + require.Len(t, changes, 2) + assert.Equal(t, uint64(1), changes[0].Revision) + assert.Equal(t, uint64(1), changes[1].BaseRevision) + assert.Equal(t, uint64(2), changes[1].Revision) +} + +func TestCommittedStateInstancesShareContextAwareWriterGate(t *testing.T) { + store := &memorySnapshotStore{started: make(chan struct{}, 1), release: make(chan struct{})} + live := New() + first := NewCommittedState(live, store) + second := NewCommittedState(live, store) + firstDone := make(chan error, 1) + + go func() { + firstDone <- first.Update(context.Background(), func(candidate *Pantry) error { + return candidate.AddAsset(NewOrganization("acme", "github")) + }) + }() + <-store.started + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + err := second.Update(ctx, func(candidate *Pantry) error { + return candidate.AddAsset(NewOrganization("globex", "github")) + }) + + assert.ErrorIs(t, err, context.Canceled) + assert.False(t, live.HasAsset("github:org:globex")) + close(store.release) + require.NoError(t, <-firstDone) +} + +func TestCommittedStateDetachesRetainedCandidateAfterCommit(t *testing.T) { + live := New() + store := &memorySnapshotStore{} + observer := &recordingObserver{} + live.AddObserver(observer) + state := NewCommittedState(live, store) + var retained *Pantry + + require.NoError(t, state.Update(context.Background(), func(candidate *Pantry) error { + retained = candidate + return candidate.AddAsset(NewOrganization("acme", "github")) + })) + require.NoError(t, retained.AddAsset(NewOrganization("globex", "github"))) + + assert.False(t, live.HasAsset("github:org:globex")) + assert.Len(t, observer.all(), 1) +} + +func TestPantryDefensivelyCopiesProperties(t *testing.T) { + type nestedProperty struct { + Values []string + } + p := New() + labels := []string{"self-hosted", "linux"} + nested := nestedProperty{Values: []string{"one", "two"}} + asset := NewOrganization("acme", "github") + asset.SetProperty("labels", labels) + asset.SetProperty("nested", nested) + require.NoError(t, p.AddAsset(asset)) + + labels[0] = "mutated" + nested.Values[0] = "mutated" + asset.Properties["org"] = "mutated" + stored, err := p.GetAsset(asset.ID) + require.NoError(t, err) + assert.Equal(t, []string{"self-hosted", "linux"}, stored.Properties["labels"]) + assert.Equal(t, []string{"one", "two"}, stored.Properties["nested"].(nestedProperty).Values) + assert.Equal(t, "acme", stored.Properties["org"]) + + stored.Properties["org"] = "changed through read" + again, err := p.GetAsset(asset.ID) + require.NoError(t, err) + assert.Equal(t, "acme", again.Properties["org"]) +} + +func TestPantrySnapshotRoundTripPreservesRevision(t *testing.T) { + live := New() + store := &memorySnapshotStore{} + state := NewCommittedState(live, store) + require.NoError(t, state.Update(context.Background(), func(candidate *Pantry) error { + return candidate.AddAsset(NewOrganization("acme", "github")) + })) + + serialized, err := json.Marshal(live) + require.NoError(t, err) + restored := New() + require.NoError(t, json.Unmarshal(serialized, restored)) + + assert.Equal(t, uint64(1), restored.Revision()) + assert.True(t, restored.HasAsset("github:org:acme")) +} + +func TestCommittedStateLargeGraphReplacement(t *testing.T) { + if testing.Short() { + t.Skip("large graph performance check") + } + live := largeTestPantry(t, 2_000) + store := &memorySnapshotStore{} + state := NewCommittedState(live, store) + started := time.Now() + + require.NoError(t, state.Replace(context.Background(), func(candidate *Pantry) error { + asset, err := candidate.GetAsset("github:acme/repo-1999") + if err != nil { + return err + } + asset.State = StateValidated + return candidate.AddAsset(asset) + })) + + t.Logf("committed 2,000 assets and 1,999 relationships in %s", time.Since(started)) + assert.Equal(t, uint64(1), live.Revision()) + assert.Equal(t, 2_000, live.Size()) +} + +func BenchmarkCommittedStateReplaceLargeGraph(b *testing.B) { + for _, size := range []int{1_000, 10_000} { + b.Run(fmt.Sprintf("assets-%d", size), func(b *testing.B) { + for range b.N { + live := largeTestPantry(b, size) + state := NewCommittedState(live, &memorySnapshotStore{}) + require.NoError(b, state.Replace(context.Background(), func(candidate *Pantry) error { + asset, err := candidate.GetAsset(fmt.Sprintf("github:acme/repo-%d", size-1)) + if err != nil { + return err + } + asset.State = StateValidated + return candidate.AddAsset(asset) + })) + } + }) + } +} + +func largeTestPantry(tb testing.TB, size int) *Pantry { + tb.Helper() + p := New() + for index := 0; index < size; index++ { + repo := NewRepository("acme", fmt.Sprintf("repo-%d", index), "github") + require.NoError(tb, p.AddAsset(repo)) + if index > 0 { + require.NoError(tb, p.AddRelationship( + fmt.Sprintf("github:acme/repo-%d", index-1), + repo.ID, + LeadsTo("synthetic"), + )) + } + } + return p +} diff --git a/internal/pantry/filter.go b/internal/pantry/filter.go index a36b86c..77876eb 100644 --- a/internal/pantry/filter.go +++ b/internal/pantry/filter.go @@ -12,7 +12,7 @@ func (p *Pantry) VulnBearingSubgraph() *Pantry { defer p.mu.RUnlock() subgraph := New() - subgraph.version = p.version + subgraph.revision = p.revision vulnIDs := p.byType[AssetVulnerability] if len(vulnIDs) == 0 { @@ -109,7 +109,7 @@ func (p *Pantry) VulnBearingSubgraph() *Pantry { _ = subgraph.AddRelationship(fromID, toID, rel) } - subgraph.version = p.version + subgraph.revision = p.revision return subgraph } diff --git a/internal/pantry/graph.go b/internal/pantry/graph.go index 91ba493..3e5153c 100644 --- a/internal/pantry/graph.go +++ b/internal/pantry/graph.go @@ -6,6 +6,8 @@ package pantry import ( "encoding/json" "errors" + "reflect" + "sort" "sync" "github.com/hmdsefi/gograph" @@ -37,22 +39,26 @@ type Pantry struct { // Index for faster lookups by type byType map[AssetType]map[string]struct{} + commitGate chan struct{} + // Observer support for real-time notifications obsMu sync.RWMutex observers []Observer - // Version counter for change tracking - version int64 + revision uint64 } // New creates a new Pantry with an empty directed graph. func New() *Pantry { + commitGate := make(chan struct{}, 1) + commitGate <- struct{}{} return &Pantry{ graph: gograph.New[string](gograph.Directed()), assets: make(map[string]Asset), edges: make(map[string]Relationship), reverseEdges: make(map[string][]string), byType: make(map[AssetType]map[string]struct{}), + commitGate: commitGate, } } @@ -63,36 +69,24 @@ func edgeKey(from, to string) string { // AddAsset adds or updates an asset vertex. func (p *Pantry) AddAsset(asset Asset) error { + asset = cloneAsset(asset) p.mu.Lock() existing := p.graph.GetVertexByID(asset.ID) if existing != nil { oldAsset := p.assets[asset.ID] - oldState := oldAsset.State - hasNewKeys := false - for k := range asset.Properties { - if _, has := oldAsset.Properties[k]; !has { - hasNewKeys = true - break - } - } if len(oldAsset.Properties) > 0 { if asset.Properties == nil { asset.Properties = make(map[string]any) } for k, v := range oldAsset.Properties { if _, has := asset.Properties[k]; !has { - asset.Properties[k] = v + asset.Properties[k] = cloneProperty(v) } } } p.assets[asset.ID] = asset - p.version++ p.mu.Unlock() - - if oldState != asset.State || hasNewKeys { - p.notifyAssetUpdated(asset, oldState) - } return nil } @@ -103,15 +97,13 @@ func (p *Pantry) AddAsset(asset Asset) error { p.byType[asset.Type] = make(map[string]struct{}) } p.byType[asset.Type][asset.ID] = struct{}{} - p.version++ p.mu.Unlock() - - p.notifyAssetAdded(asset) return nil } // AddRelationship adds an edge between assets. func (p *Pantry) AddRelationship(fromID, toID string, rel Relationship) error { + rel = cloneRelationship(rel) p.mu.Lock() fromVertex := p.graph.GetVertexByID(fromID) @@ -133,10 +125,7 @@ func (p *Pantry) AddRelationship(fromID, toID string, rel Relationship) error { p.edges[edgeKey(fromID, toID)] = rel p.reverseEdges[toID] = append(p.reverseEdges[toID], fromID) - p.version++ p.mu.Unlock() - - p.notifyRelationshipAdded(fromID, toID, rel) return nil } @@ -176,10 +165,7 @@ func (p *Pantry) RemoveAsset(id string) error { } delete(p.assets, id) - p.version++ p.mu.Unlock() - - p.notifyAssetRemoved(id) return nil } @@ -203,10 +189,7 @@ func (p *Pantry) RemoveRelationship(fromID, toID string) error { p.graph.RemoveEdges(edge) } - p.version++ p.mu.Unlock() - - p.notifyRelationshipRemoved(fromID, toID) return nil } @@ -219,7 +202,7 @@ func (p *Pantry) GetAsset(id string) (Asset, error) { if !ok { return Asset{}, ErrAssetNotFound } - return asset, nil + return cloneAsset(asset), nil } // HasAsset checks if an asset exists. @@ -244,7 +227,7 @@ func (p *Pantry) GetAssetsByType(assetType AssetType) []Asset { assets := make([]Asset, 0, len(ids)) for id := range ids { if asset, ok := p.assets[id]; ok { - assets = append(assets, asset) + assets = append(assets, cloneAsset(asset)) } } return assets @@ -261,7 +244,7 @@ func (p *Pantry) GetNeighbors(id string, hops int) ([]Asset, error) { if hops <= 0 { asset := p.assets[id] - return []Asset{asset}, nil + return []Asset{cloneAsset(asset)}, nil } visited := make(map[string]int) @@ -294,7 +277,7 @@ func (p *Pantry) GetNeighbors(id string, hops int) ([]Asset, error) { assets := make([]Asset, 0, len(visited)) for assetID := range visited { if asset, ok := p.assets[assetID]; ok { - assets = append(assets, asset) + assets = append(assets, cloneAsset(asset)) } } @@ -308,7 +291,7 @@ func (p *Pantry) AllAssets() []Asset { assets := make([]Asset, 0, len(p.assets)) for _, asset := range p.assets { - assets = append(assets, asset) + assets = append(assets, cloneAsset(asset)) } return assets } @@ -324,7 +307,7 @@ func (p *Pantry) AllRelationships() []Edge { for _, e := range graphEdges { fromID := e.Source().Label() toID := e.Destination().Label() - rel := p.edges[edgeKey(fromID, toID)] + rel := cloneRelationship(p.edges[edgeKey(fromID, toID)]) result = append(result, Edge{ From: fromID, To: toID, @@ -351,13 +334,12 @@ func (p *Pantry) EdgeCount() int { return len(p.graph.AllEdges()) } -// Version returns the current version of the graph. -// Version is incremented on each mutation. -func (p *Pantry) Version() int64 { +// Revision returns the current committed Pantry revision. +func (p *Pantry) Revision() uint64 { p.mu.RLock() defer p.mu.RUnlock() - return p.version + return p.revision } // UpdateAssetState updates the state of an asset. @@ -378,10 +360,7 @@ func (p *Pantry) UpdateAssetState(id string, state AssetState) error { asset.State = state p.assets[id] = asset - p.version++ p.mu.Unlock() - - p.notifyAssetUpdated(asset, oldState) return nil } @@ -403,7 +382,7 @@ func (p *Pantry) FindHighValueTargets() []Asset { var targets []Asset for _, asset := range p.assets { if asset.State == StateHighValue { - targets = append(targets, asset) + targets = append(targets, cloneAsset(asset)) } } return targets @@ -448,7 +427,7 @@ func (p *Pantry) GetAttackPaths(sourceID string, targetTypes []AssetType) [][]As func (p *Pantry) findPath(fromID, toID string) []Asset { if fromID == toID { if asset, ok := p.assets[fromID]; ok { - return []Asset{asset} + return []Asset{cloneAsset(asset)} } return nil } @@ -471,7 +450,7 @@ func (p *Pantry) findPath(fromID, toID string) []Asset { var path []Asset for node := toID; node != ""; node = parent[node] { if asset, ok := p.assets[node]; ok { - path = append([]Asset{asset}, path...) + path = append([]Asset{cloneAsset(asset)}, path...) } if node == fromID { break @@ -510,55 +489,54 @@ func (p *Pantry) Clear() { p.byType = make(map[AssetType]map[string]struct{}) } -// pantryData is the serializable form of Pantry. -type pantryData struct { - Assets []Asset `json:"assets"` - Edges []Edge `json:"edges"` -} - // MarshalJSON serializes the Pantry to JSON. func (p *Pantry) MarshalJSON() ([]byte, error) { - p.mu.RLock() - defer p.mu.RUnlock() + return json.Marshal(p.Snapshot()) +} - data := pantryData{ - Assets: make([]Asset, 0, len(p.assets)), - Edges: make([]Edge, 0, len(p.edges)), +// UnmarshalJSON deserializes JSON into a Pantry. +func (p *Pantry) UnmarshalJSON(data []byte) error { + var snapshot Snapshot + if err := json.Unmarshal(data, &snapshot); err != nil { + return err + } + if err := validateSnapshot(snapshot); err != nil { + return err } + decoded := pantryFromSnapshot(snapshot) + p.replaceWith(decoded) + return nil +} +func (p *Pantry) Snapshot() Snapshot { + p.mu.RLock() + + snapshot := Snapshot{ + Revision: p.revision, + Assets: make([]Asset, 0, len(p.assets)), + Edges: make([]Edge, 0, len(p.edges)), + } for _, asset := range p.assets { - data.Assets = append(data.Assets, asset) + snapshot.Assets = append(snapshot.Assets, cloneAsset(asset)) } - - for key, rel := range p.edges { + for key, relationship := range p.edges { from, to := parseEdgeKey(key) - data.Edges = append(data.Edges, Edge{ - From: from, - To: to, - Relationship: rel, - }) + snapshot.Edges = append(snapshot.Edges, Edge{From: from, To: to, Relationship: cloneRelationship(relationship)}) } + p.mu.RUnlock() - return json.Marshal(data) + sort.Slice(snapshot.Assets, func(i, j int) bool { return snapshot.Assets[i].ID < snapshot.Assets[j].ID }) + sort.Slice(snapshot.Edges, func(i, j int) bool { + return edgeKey(snapshot.Edges[i].From, snapshot.Edges[i].To) < edgeKey(snapshot.Edges[j].From, snapshot.Edges[j].To) + }) + return snapshot } -// UnmarshalJSON deserializes JSON into a Pantry. -func (p *Pantry) UnmarshalJSON(data []byte) error { - var pd pantryData - if err := json.Unmarshal(data, &pd); err != nil { - return err - } - - p.mu.Lock() - defer p.mu.Unlock() - - p.graph = gograph.New[string](gograph.Directed()) - p.assets = make(map[string]Asset) - p.edges = make(map[string]Relationship) - p.reverseEdges = make(map[string][]string) - p.byType = make(map[AssetType]map[string]struct{}) - - for _, asset := range pd.Assets { +func pantryFromSnapshot(snapshot Snapshot) *Pantry { + p := New() + p.revision = snapshot.Revision + for _, original := range snapshot.Assets { + asset := cloneAsset(original) p.graph.AddVertexByLabel(asset.ID) p.assets[asset.ID] = asset if p.byType[asset.Type] == nil { @@ -566,18 +544,59 @@ func (p *Pantry) UnmarshalJSON(data []byte) error { } p.byType[asset.Type][asset.ID] = struct{}{} } - - for _, edge := range pd.Edges { + for _, original := range snapshot.Edges { + edge := cloneEdge(original) fromVertex := p.graph.GetVertexByID(edge.From) toVertex := p.graph.GetVertexByID(edge.To) - if fromVertex != nil && toVertex != nil { - _, _ = p.graph.AddEdge(fromVertex, toVertex) - p.edges[edgeKey(edge.From, edge.To)] = edge.Relationship - p.reverseEdges[edge.To] = append(p.reverseEdges[edge.To], edge.From) + if fromVertex == nil || toVertex == nil { + continue } + _, _ = p.graph.AddEdge(fromVertex, toVertex) + p.edges[edgeKey(edge.From, edge.To)] = edge.Relationship + p.reverseEdges[edge.To] = append(p.reverseEdges[edge.To], edge.From) } + return p +} - return nil +func (p *Pantry) setRevision(revision uint64) { + p.mu.Lock() + p.revision = revision + p.mu.Unlock() +} + +func (p *Pantry) committedWriterGate() chan struct{} { + p.mu.Lock() + defer p.mu.Unlock() + if p.commitGate == nil { + p.commitGate = make(chan struct{}, 1) + p.commitGate <- struct{}{} + } + return p.commitGate +} + +func (p *Pantry) replaceWith(candidate *Pantry) { + candidate.mu.Lock() + p.mu.Lock() + previousGraph := p.graph + previousAssets := p.assets + previousEdges := p.edges + previousReverseEdges := p.reverseEdges + previousByType := p.byType + previousRevision := p.revision + p.graph = candidate.graph + p.assets = candidate.assets + p.edges = candidate.edges + p.reverseEdges = candidate.reverseEdges + p.byType = candidate.byType + p.revision = candidate.revision + candidate.graph = previousGraph + candidate.assets = previousAssets + candidate.edges = previousEdges + candidate.reverseEdges = previousReverseEdges + candidate.byType = previousByType + candidate.revision = previousRevision + p.mu.Unlock() + candidate.mu.Unlock() } // GetPredecessors returns assets with edges TO this node (reverse lookup). @@ -589,7 +608,7 @@ func (p *Pantry) GetPredecessors(id string) []Asset { assets := make([]Asset, 0, len(sourceIDs)) for _, srcID := range sourceIDs { if asset, ok := p.assets[srcID]; ok { - assets = append(assets, asset) + assets = append(assets, cloneAsset(asset)) } } return assets @@ -608,13 +627,104 @@ func (p *Pantry) GetOutgoingEdges(id string) []Edge { result = append(result, Edge{ From: id, To: to, - Relationship: rel, + Relationship: cloneRelationship(rel), }) } } return result } +func cloneAsset(asset Asset) Asset { + asset.Properties = cloneProperties(asset.Properties) + return asset +} + +func cloneRelationship(relationship Relationship) Relationship { + relationship.Properties = cloneProperties(relationship.Properties) + return relationship +} + +func cloneEdge(edge Edge) Edge { + edge.Relationship = cloneRelationship(edge.Relationship) + return edge +} + +func cloneProperties(properties map[string]any) map[string]any { + if properties == nil { + return nil + } + cloned := make(map[string]any, len(properties)) + for key, value := range properties { + cloned[key] = cloneProperty(value) + } + return cloned +} + +func cloneProperty(value any) any { + if value == nil { + return nil + } + return clonePropertyValue(reflect.ValueOf(value)).Interface() +} + +func clonePropertyValue(value reflect.Value) reflect.Value { + switch value.Kind() { + case reflect.Interface: + if value.IsNil() { + return reflect.Zero(value.Type()) + } + cloned := clonePropertyValue(value.Elem()) + result := reflect.New(value.Type()).Elem() + result.Set(cloned) + return result + case reflect.Map: + if value.IsNil() { + return reflect.Zero(value.Type()) + } + result := reflect.MakeMapWithSize(value.Type(), value.Len()) + iterator := value.MapRange() + for iterator.Next() { + result.SetMapIndex(clonePropertyValue(iterator.Key()), clonePropertyValue(iterator.Value())) + } + return result + case reflect.Slice: + if value.IsNil() { + return reflect.Zero(value.Type()) + } + result := reflect.MakeSlice(value.Type(), value.Len(), value.Len()) + for index := 0; index < value.Len(); index++ { + result.Index(index).Set(clonePropertyValue(value.Index(index))) + } + return result + case reflect.Array: + result := reflect.New(value.Type()).Elem() + for index := 0; index < value.Len(); index++ { + result.Index(index).Set(clonePropertyValue(value.Index(index))) + } + return result + case reflect.Struct: + for index := 0; index < value.NumField(); index++ { + if value.Type().Field(index).PkgPath != "" { + return value + } + } + result := reflect.New(value.Type()).Elem() + for index := 0; index < value.NumField(); index++ { + result.Field(index).Set(clonePropertyValue(value.Field(index))) + } + return result + case reflect.Pointer: + if value.IsNil() { + return reflect.Zero(value.Type()) + } + result := reflect.New(value.Type().Elem()) + result.Elem().Set(clonePropertyValue(value.Elem())) + return result + default: + return value + } +} + func (p *Pantry) removeReverseEdge(toID, fromID string) { sources := p.reverseEdges[toID] for i, src := range sources { diff --git a/internal/pantry/graph_test.go b/internal/pantry/graph_test.go index c42c9ea..da1e022 100644 --- a/internal/pantry/graph_test.go +++ b/internal/pantry/graph_test.go @@ -535,44 +535,6 @@ func TestPantry_AddAsset_PreservesProperties(t *testing.T) { assert.Equal(t, "analysis", final.Properties["discovered_by"], "New property should override old") } -func TestPantry_AddAsset_PropertyChangeNotifiesObserver(t *testing.T) { - p := New() - - repo := NewRepository("acme", "api", "github") - require.NoError(t, p.AddAsset(repo)) - - var updatedAssets []Asset - obs := &testObserver{ - onUpdated: func(a Asset, _ AssetState) { updatedAssets = append(updatedAssets, a) }, - } - p.AddObserver(obs) - - updated := NewRepository("acme", "api", "github") - updated.SetProperty("private", true) - require.NoError(t, p.AddAsset(updated)) - - require.Len(t, updatedAssets, 1, "property change should fire OnAssetUpdated") - assert.Equal(t, true, updatedAssets[0].Properties["private"]) -} - -func TestPantry_AddAsset_NoPropertyChangeNoNotification(t *testing.T) { - p := New() - - repo := NewRepository("acme", "api", "github") - require.NoError(t, p.AddAsset(repo)) - - var updateCount int - obs := &testObserver{ - onUpdated: func(_ Asset, _ AssetState) { updateCount++ }, - } - p.AddObserver(obs) - - same := NewRepository("acme", "api", "github") - require.NoError(t, p.AddAsset(same)) - - assert.Equal(t, 0, updateCount, "identical re-add should not fire observer") -} - func TestPantry_AddAsset_SlicePropertyNoPanic(t *testing.T) { p := New() @@ -590,25 +552,6 @@ func TestPantry_AddAsset_SlicePropertyNoPanic(t *testing.T) { assert.Equal(t, []interface{}{"read", "write"}, final.Properties["permissions"]) } -type testObserver struct { - onAdded func(Asset) - onUpdated func(Asset, AssetState) -} - -func (o *testObserver) OnAssetAdded(a Asset) { - if o.onAdded != nil { - o.onAdded(a) - } -} -func (o *testObserver) OnAssetUpdated(a Asset, old AssetState) { - if o.onUpdated != nil { - o.onUpdated(a, old) - } -} -func (o *testObserver) OnRelationshipAdded(_, _ string, _ Relationship) {} -func (o *testObserver) OnAssetRemoved(_ string) {} -func (o *testObserver) OnRelationshipRemoved(_, _ string) {} - func TestPantry_JSON_OrganizationRoundTrip(t *testing.T) { p := New() @@ -647,3 +590,21 @@ func TestPantry_JSON_OrganizationRoundTrip(t *testing.T) { assert.Equal(t, 1, p2.EdgeCount()) } + +func TestPantrySnapshotHasDeterministicOrder(t *testing.T) { + p := New() + second := NewOrganization("second", "github") + first := NewOrganization("first", "github") + require.NoError(t, p.AddAsset(second)) + require.NoError(t, p.AddAsset(first)) + require.NoError(t, p.AddRelationship(second.ID, first.ID, Contains())) + + snapshot := p.Snapshot() + + require.Len(t, snapshot.Assets, 2) + assert.Equal(t, first.ID, snapshot.Assets[0].ID) + assert.Equal(t, second.ID, snapshot.Assets[1].ID) + require.Len(t, snapshot.Edges, 1) + assert.Equal(t, second.ID, snapshot.Edges[0].From) + assert.Equal(t, first.ID, snapshot.Edges[0].To) +} diff --git a/internal/pantry/observer.go b/internal/pantry/observer.go index 2c14a9f..07e5ce9 100644 --- a/internal/pantry/observer.go +++ b/internal/pantry/observer.go @@ -3,17 +3,6 @@ package pantry -// Observer receives notifications of graph changes. -// Implementations must be thread-safe as notifications may come from -// concurrent goroutines. -type Observer interface { - OnAssetAdded(asset Asset) - OnAssetUpdated(asset Asset, oldState AssetState) - OnRelationshipAdded(from, to string, rel Relationship) - OnAssetRemoved(id string) - OnRelationshipRemoved(from, to string) -} - // AddObserver registers an observer to receive change notifications. func (p *Pantry) AddObserver(obs Observer) { p.obsMu.Lock() @@ -33,57 +22,35 @@ func (p *Pantry) RemoveObserver(obs Observer) { } } -func (p *Pantry) notifyAssetAdded(asset Asset) { +func (p *Pantry) notifyChange(change ChangeSet) { p.obsMu.RLock() observers := make([]Observer, len(p.observers)) copy(observers, p.observers) p.obsMu.RUnlock() for _, obs := range observers { - obs.OnAssetAdded(asset) + obs.OnPantryChange(cloneChangeSet(change)) } } -func (p *Pantry) notifyAssetUpdated(asset Asset, oldState AssetState) { - p.obsMu.RLock() - observers := make([]Observer, len(p.observers)) - copy(observers, p.observers) - p.obsMu.RUnlock() - - for _, obs := range observers { - obs.OnAssetUpdated(asset, oldState) - } -} - -func (p *Pantry) notifyRelationshipAdded(from, to string, rel Relationship) { - p.obsMu.RLock() - observers := make([]Observer, len(p.observers)) - copy(observers, p.observers) - p.obsMu.RUnlock() - - for _, obs := range observers { - obs.OnRelationshipAdded(from, to, rel) +func cloneChangeSet(change ChangeSet) ChangeSet { + cloned := change + cloned.Granular.AddedAssets = make([]Asset, len(change.Granular.AddedAssets)) + for index, asset := range change.Granular.AddedAssets { + cloned.Granular.AddedAssets[index] = cloneAsset(asset) } -} - -func (p *Pantry) notifyAssetRemoved(id string) { - p.obsMu.RLock() - observers := make([]Observer, len(p.observers)) - copy(observers, p.observers) - p.obsMu.RUnlock() - - for _, obs := range observers { - obs.OnAssetRemoved(id) + cloned.Granular.UpdatedAssets = make([]AssetChange, len(change.Granular.UpdatedAssets)) + for index, update := range change.Granular.UpdatedAssets { + cloned.Granular.UpdatedAssets[index] = AssetChange{ + Before: cloneAsset(update.Before), + After: cloneAsset(update.After), + } } -} - -func (p *Pantry) notifyRelationshipRemoved(from, to string) { - p.obsMu.RLock() - observers := make([]Observer, len(p.observers)) - copy(observers, p.observers) - p.obsMu.RUnlock() - - for _, obs := range observers { - obs.OnRelationshipRemoved(from, to) + cloned.Granular.RemovedAssetIDs = append([]string(nil), change.Granular.RemovedAssetIDs...) + cloned.Granular.AddedRelationships = make([]Edge, len(change.Granular.AddedRelationships)) + for index, edge := range change.Granular.AddedRelationships { + cloned.Granular.AddedRelationships[index] = cloneEdge(edge) } + cloned.Granular.RemovedRelationships = append([]EdgeRef(nil), change.Granular.RemovedRelationships...) + return cloned }