From 1d9862e362decce58e235c09124aa35088bebd37 Mon Sep 17 00:00:00 2001 From: Hammad Majid Date: Sat, 5 Sep 2026 13:02:42 +0500 Subject: [PATCH] fix(gitengine): thread context through history walking The entire history-walk phase was uncancellable, making it the primary Ctrl-C dead zone. `collectOrderedCommits` BFS-walked the whole commit DAG -- a packfile seek plus zlib inflate per commit -- before a single occurrence reached the caller, and `cmd/grg/main.go` checked `ctx.Err()` only inside the per-occurrence callback, which was never invoked until that phase completed. With `--all` every ref is seeded, so on a large repository SIGINT had literally no effect. `traverseExclude` was a second unbounded uncancellable walk, reached once per exclude of an `A..B` rev-range. Even during tree traversal the callback was the sole cancellation point while long stretches did real I/O and emitted nothing: subtree-OID pruning, the merge-commit `FindTreeEntry` probe per extra parent, and full-tree recursion in `traverseTreeHelper`. `HistoryWalker.Walk` now takes a leading `context.Context` and threads it through `collectOrderedCommits`, `traverseExclude` (which gains an error return), `walkCommitBlobs`, `diffTreesAndEmit`, and `emitBlobOccurrence`, plus `TraverseTree`/`traverseTreeHelper` and `FindTreeEntry` in tree.go. Guards are plain `ctx.Err()` loads rather than `select`, placed at loop heads where they gate the `ReadObject` that follows and never inside the innermost per-entry comparison, so they cost nothing next to the I/O they guard. The redundant `ctx.Err()` check inside the callback in `runContext` is gone; the surrounding `errors.Is(err, context.Canceled)` handling is unchanged. Signature call sites in existing tests were updated mechanically. New tests assert on work actually performed via a counting `ObjectReader`: cancellation during commit collection stops within a few reads of the trip point and emits nothing (the phase that was completely deaf), mid-walk cancellation stops after the emitting commit and reads far fewer objects than a full drain, and the exclude walk and tree traversal both abort instead of draining. Removing only the `collectOrderedCommits` guard makes the collection test fail on read count. Closes #13 --- cmd/grg/main.go | 5 +- internal/gitengine/hardening_test.go | 7 +- internal/gitengine/tree.go | 21 ++- internal/gitengine/tree_test.go | 5 +- internal/gitengine/walker.go | 53 +++++-- internal/gitengine/walker_test.go | 224 ++++++++++++++++++++++++++- test/leak_test.go | 2 +- 7 files changed, 285 insertions(+), 32 deletions(-) diff --git a/cmd/grg/main.go b/cmd/grg/main.go index f789af0..7a509dc 100644 --- a/cmd/grg/main.go +++ b/cmd/grg/main.go @@ -123,10 +123,7 @@ func runContext(ctx context.Context, args []string, stdout, stderr io.Writer) er walker := gitengine.NewHistoryWalker(repo, reader, cfg, pathFilter) var occurrences []model.BlobOccurrence - err = walker.Walk(func(occ model.BlobOccurrence) error { - if ctxErr := ctx.Err(); ctxErr != nil { - return ctxErr - } + err = walker.Walk(ctx, func(occ model.BlobOccurrence) error { occurrences = append(occurrences, occ) return nil }) diff --git a/internal/gitengine/hardening_test.go b/internal/gitengine/hardening_test.go index 98a0308..baeff52 100644 --- a/internal/gitengine/hardening_test.go +++ b/internal/gitengine/hardening_test.go @@ -2,6 +2,7 @@ package gitengine import ( "bytes" + "context" "encoding/binary" "encoding/hex" "errors" @@ -146,7 +147,7 @@ func TestSafety_Tree_CycleAndMaxDepth(t *testing.T) { reader.objects[treeAOID] = &Object{OID: treeAOID, Type: TypeTree, Data: treeAData} reader.objects[treeBOID] = &Object{OID: treeBOID, Type: TypeTree, Data: treeBData} - err := TraverseTree(reader, treeAOID, func(path string, entry TreeEntry) error { + err := TraverseTree(context.Background(), reader, treeAOID, func(path string, entry TreeEntry) error { return nil }) if err == nil { @@ -187,7 +188,9 @@ func TestSafety_Walker_DeepLinearHistory(t *testing.T) { // Calling traverseExclude on the tip commit traverses all 5000 parents // In the recursive implementation, this risked stack overflow. // With the iterative stack, it completes instantly and safely. - walker.traverseExclude(prevSHA, excluded) + if err := walker.traverseExclude(context.Background(), prevSHA, excluded); err != nil { + t.Fatalf("traverseExclude failed: %v", err) + } if len(excluded) != numCommits { t.Fatalf("expected %d excluded commits, got %d", numCommits, len(excluded)) diff --git a/internal/gitengine/tree.go b/internal/gitengine/tree.go index fe7adca..52e1817 100644 --- a/internal/gitengine/tree.go +++ b/internal/gitengine/tree.go @@ -2,6 +2,7 @@ package gitengine import ( "bytes" + "context" "encoding/hex" "errors" "fmt" @@ -87,12 +88,13 @@ const maxTreeDepth = 256 // TraverseTree recursively traverses a Git tree and calls callback for each entry. // If callback returns ErrSkipDir for a tree entry, the subtree will not be descended into. -func TraverseTree(reader ObjectReader, rootOID string, callback func(path string, entry TreeEntry) error) error { +// Traversal aborts with ctx.Err() once ctx is cancelled. +func TraverseTree(ctx context.Context, reader ObjectReader, rootOID string, callback func(path string, entry TreeEntry) error) error { visited := make(map[string]bool) - return traverseTreeHelper(reader, rootOID, "", 0, visited, callback) + return traverseTreeHelper(ctx, reader, rootOID, "", 0, visited, callback) } -func traverseTreeHelper(reader ObjectReader, treeOID, prefix string, depth int, visited map[string]bool, callback func(path string, entry TreeEntry) error) error { +func traverseTreeHelper(ctx context.Context, reader ObjectReader, treeOID, prefix string, depth int, visited map[string]bool, callback func(path string, entry TreeEntry) error) error { if depth > maxTreeDepth { return fmt.Errorf("%w: tree depth exceeds maximum of %d", ErrCorruptObject, maxTreeDepth) } @@ -116,6 +118,10 @@ func traverseTreeHelper(reader ObjectReader, treeOID, prefix string, depth int, } for _, entry := range entries { + if err := ctx.Err(); err != nil { + return err + } + var entryPath string if prefix == "" { entryPath = entry.Name @@ -132,7 +138,7 @@ func traverseTreeHelper(reader ObjectReader, treeOID, prefix string, depth int, } if entry.IsTree() { - if err := traverseTreeHelper(reader, entry.OID, entryPath, depth+1, visited, callback); err != nil { + if err := traverseTreeHelper(ctx, reader, entry.OID, entryPath, depth+1, visited, callback); err != nil { return err } } @@ -143,7 +149,8 @@ func traverseTreeHelper(reader ObjectReader, treeOID, prefix string, depth int, // FindTreeEntry searches rootOID for the tree entry at path (slash-separated). // Returns the entry and true if found, or a zero TreeEntry and false if not found. -func FindTreeEntry(reader ObjectReader, rootOID, path string) (TreeEntry, bool) { +// A cancelled ctx aborts the descent and reports not-found. +func FindTreeEntry(ctx context.Context, reader ObjectReader, rootOID, path string) (TreeEntry, bool) { if rootOID == "" || path == "" { return TreeEntry{}, false } @@ -152,6 +159,10 @@ func FindTreeEntry(reader ObjectReader, rootOID, path string) (TreeEntry, bool) currOID := rootOID for i, part := range parts { + if ctx.Err() != nil { + return TreeEntry{}, false + } + obj, err := reader.ReadObject(currOID) if err != nil || obj.Type != TypeTree { return TreeEntry{}, false diff --git a/internal/gitengine/tree_test.go b/internal/gitengine/tree_test.go index 42d5b77..7adbccf 100644 --- a/internal/gitengine/tree_test.go +++ b/internal/gitengine/tree_test.go @@ -2,6 +2,7 @@ package gitengine import ( "bytes" + "context" "crypto/sha1" "encoding/hex" "errors" @@ -96,7 +97,7 @@ func TestParseTreeAndTraverse(t *testing.T) { // Test TraverseTree visited := make(map[string]string) - err = TraverseTree(reader, rootTreeOID, func(path string, entry TreeEntry) error { + err = TraverseTree(context.Background(), reader, rootTreeOID, func(path string, entry TreeEntry) error { visited[path] = entry.OID return nil }) @@ -113,7 +114,7 @@ func TestParseTreeAndTraverse(t *testing.T) { // Test ErrSkipDir visitedSkipped := make(map[string]string) - err = TraverseTree(reader, rootTreeOID, func(path string, entry TreeEntry) error { + err = TraverseTree(context.Background(), reader, rootTreeOID, func(path string, entry TreeEntry) error { visitedSkipped[path] = entry.OID if entry.IsTree() && entry.Name == "src" { return ErrSkipDir diff --git a/internal/gitengine/walker.go b/internal/gitengine/walker.go index 3a95a5e..c4ddc20 100644 --- a/internal/gitengine/walker.go +++ b/internal/gitengine/walker.go @@ -2,6 +2,7 @@ package gitengine import ( "container/heap" + "context" "fmt" "regexp" "slices" @@ -28,8 +29,10 @@ func NewHistoryWalker(repo *RepoInfo, reader ObjectReader, cfg *model.Config, pa } // Walk traverses repository history, pruning subtrees by OID and invoking fn on each blob occurrence. -func (w *HistoryWalker) Walk(fn func(occ model.BlobOccurrence) error) error { - commits, err := w.collectOrderedCommits() +// Walk aborts with ctx.Err() once ctx is cancelled, including during the initial commit-DAG collection +// that runs before the first occurrence is emitted. +func (w *HistoryWalker) Walk(ctx context.Context, fn func(occ model.BlobOccurrence) error) error { + commits, err := w.collectOrderedCommits(ctx) if err != nil { return err } @@ -37,7 +40,10 @@ func (w *HistoryWalker) Walk(fn func(occ model.BlobOccurrence) error) error { seenBlobOcc := make(map[string]bool) for _, commit := range commits { - if err := w.walkCommitBlobs(commit, seenBlobOcc, fn); err != nil { + if err := ctx.Err(); err != nil { + return err + } + if err := w.walkCommitBlobs(ctx, commit, seenBlobOcc, fn); err != nil { return err } } @@ -45,7 +51,7 @@ func (w *HistoryWalker) Walk(fn func(occ model.BlobOccurrence) error) error { return nil } -func (w *HistoryWalker) walkCommitBlobs(commit *model.CommitMetadata, seenBlobOcc map[string]bool, fn func(occ model.BlobOccurrence) error) error { +func (w *HistoryWalker) walkCommitBlobs(ctx context.Context, commit *model.CommitMetadata, seenBlobOcc map[string]bool, fn func(occ model.BlobOccurrence) error) error { // If ExpandCommits is false and commit has a parent, perform tree diff against first parent if !w.cfg.ExpandCommits && len(commit.Parents) > 0 { parentObj, err := w.reader.ReadObject(commit.Parents[0]) @@ -56,13 +62,13 @@ func (w *HistoryWalker) walkCommitBlobs(commit *model.CommitMetadata, seenBlobOc // Identical root tree: zero files introduced/modified return nil } - return w.diffTreesAndEmit(parentCommit.TreeOID, commit.TreeOID, "", commit, seenBlobOcc, fn) + return w.diffTreesAndEmit(ctx, parentCommit.TreeOID, commit.TreeOID, "", commit, seenBlobOcc, fn) } } } // Full tree traversal for root commits or when ExpandCommits is enabled - return TraverseTree(w.reader, commit.TreeOID, func(path string, entry TreeEntry) error { + return TraverseTree(ctx, w.reader, commit.TreeOID, func(path string, entry TreeEntry) error { if entry.IsTree() { return nil } @@ -140,12 +146,16 @@ func compareTreeEntries(a, b TreeEntry) int { // diffTreesAndEmit performs subtree-pruned tree comparison between oldTreeOID and newTreeOID. // Leverages canonical sorted tree entry order with a two-pointer merge to eliminate map allocations. -func (w *HistoryWalker) diffTreesAndEmit(oldTreeOID, newTreeOID, prefix string, commit *model.CommitMetadata, seenBlobOcc map[string]bool, fn func(occ model.BlobOccurrence) error) error { +func (w *HistoryWalker) diffTreesAndEmit(ctx context.Context, oldTreeOID, newTreeOID, prefix string, commit *model.CommitMetadata, seenBlobOcc map[string]bool, fn func(occ model.BlobOccurrence) error) error { if oldTreeOID == newTreeOID { // Subtree OID match: prune subtree descending entirely! return nil } + if err := ctx.Err(); err != nil { + return err + } + var oldEntries []TreeEntry if oldTreeOID != "" { if oldObj, err := w.reader.ReadObject(oldTreeOID); err == nil && oldObj.Type == TypeTree { @@ -207,11 +217,11 @@ func (w *HistoryWalker) diffTreesAndEmit(oldTreeOID, newTreeOID, prefix string, if hasOld && oldEntry.IsTree() { oldSubOID = oldEntry.OID } - if err := w.diffTreesAndEmit(oldSubOID, newEntry.OID, entryPath, commit, seenBlobOcc, fn); err != nil { + if err := w.diffTreesAndEmit(ctx, oldSubOID, newEntry.OID, entryPath, commit, seenBlobOcc, fn); err != nil { return err } } else if newEntry.IsBlob() { - if err := w.emitBlobOccurrence(newEntry, entryPath, commit, seenBlobOcc, fn); err != nil { + if err := w.emitBlobOccurrence(ctx, newEntry, entryPath, commit, seenBlobOcc, fn); err != nil { return err } } @@ -220,7 +230,7 @@ func (w *HistoryWalker) diffTreesAndEmit(oldTreeOID, newTreeOID, prefix string, return nil } -func (w *HistoryWalker) emitBlobOccurrence(newEntry TreeEntry, entryPath string, commit *model.CommitMetadata, seenBlobOcc map[string]bool, fn func(occ model.BlobOccurrence) error) error { +func (w *HistoryWalker) emitBlobOccurrence(ctx context.Context, newEntry TreeEntry, entryPath string, commit *model.CommitMetadata, seenBlobOcc map[string]bool, fn func(occ model.BlobOccurrence) error) error { if w.pathFilter != nil && !w.pathFilter(entryPath) { return nil } @@ -230,9 +240,12 @@ func (w *HistoryWalker) emitBlobOccurrence(newEntry TreeEntry, entryPath string, if len(commit.Parents) > 1 && !w.cfg.FirstParent { inOtherParent := false for _, pSHA := range commit.Parents[1:] { + if err := ctx.Err(); err != nil { + return err + } pMeta := w.readCommit(pSHA) if pMeta != nil { - if pe, ok := FindTreeEntry(w.reader, pMeta.TreeOID, entryPath); ok && pe.OID == newEntry.OID { + if pe, ok := FindTreeEntry(ctx, w.reader, pMeta.TreeOID, entryPath); ok && pe.OID == newEntry.OID { inOtherParent = true break } @@ -286,7 +299,7 @@ func (h *commitHeap) Pop() any { return x } -func (w *HistoryWalker) collectOrderedCommits() ([]*model.CommitMetadata, error) { +func (w *HistoryWalker) collectOrderedCommits(ctx context.Context) ([]*model.CommitMetadata, error) { spec, err := ParseRevSpec(w.repo, w.reader, w.cfg.RevRange) if err != nil { return nil, err @@ -312,7 +325,9 @@ func (w *HistoryWalker) collectOrderedCommits() ([]*model.CommitMetadata, error) // Build exclude set from spec.Exclude excluded := make(map[string]bool) for _, exclOID := range spec.Exclude { - w.traverseExclude(exclOID, excluded) + if err := w.traverseExclude(ctx, exclOID, excluded); err != nil { + return nil, err + } } var authorRe *regexp.Regexp @@ -350,6 +365,10 @@ func (w *HistoryWalker) collectOrderedCommits() ([]*model.CommitMetadata, error) var matchedCommits []*model.CommitMetadata for pq.Len() > 0 { + if err := ctx.Err(); err != nil { + return nil, err + } + popped := heap.Pop(pq) commit, ok := popped.(*model.CommitMetadata) if !ok { @@ -389,9 +408,13 @@ func (w *HistoryWalker) readCommit(sha string) *model.CommitMetadata { return meta } -func (w *HistoryWalker) traverseExclude(startSHA string, excluded map[string]bool) { +func (w *HistoryWalker) traverseExclude(ctx context.Context, startSHA string, excluded map[string]bool) error { stack := []string{startSHA} for len(stack) > 0 { + if err := ctx.Err(); err != nil { + return err + } + n := len(stack) - 1 sha := stack[n] stack = stack[:n] @@ -410,4 +433,6 @@ func (w *HistoryWalker) traverseExclude(startSHA string, excluded map[string]boo } } } + + return nil } diff --git a/internal/gitengine/walker_test.go b/internal/gitengine/walker_test.go index e209cc8..60ceba4 100644 --- a/internal/gitengine/walker_test.go +++ b/internal/gitengine/walker_test.go @@ -1,6 +1,8 @@ package gitengine import ( + "context" + "errors" "fmt" "os" "path/filepath" @@ -74,7 +76,7 @@ Update hello.txt walker := NewHistoryWalker(repo, reader, cfg, nil) var occurrences []model.BlobOccurrence - err := walker.Walk(func(occ model.BlobOccurrence) error { + err := walker.Walk(context.Background(), func(occ model.BlobOccurrence) error { occurrences = append(occurrences, occ) return nil }) @@ -114,7 +116,7 @@ Update hello.txt } walkerAlice := NewHistoryWalker(repo, reader, cfgAlice, nil) var aliceOcc []model.BlobOccurrence - _ = walkerAlice.Walk(func(occ model.BlobOccurrence) error { + _ = walkerAlice.Walk(context.Background(), func(occ model.BlobOccurrence) error { aliceOcc = append(aliceOcc, occ) return nil }) @@ -130,7 +132,7 @@ Update hello.txt return filepath.Ext(path) == ".go" }) var goOcc []model.BlobOccurrence - _ = walkerGo.Walk(func(occ model.BlobOccurrence) error { + _ = walkerGo.Walk(context.Background(), func(occ model.BlobOccurrence) error { goOcc = append(goOcc, occ) return nil }) @@ -144,7 +146,7 @@ Update hello.txt } walkerRange := NewHistoryWalker(repo, reader, cfgRange, nil) var rangeOcc []model.BlobOccurrence - _ = walkerRange.Walk(func(occ model.BlobOccurrence) error { + _ = walkerRange.Walk(context.Background(), func(occ model.BlobOccurrence) error { rangeOcc = append(rangeOcc, occ) return nil }) @@ -175,3 +177,217 @@ func TestHistoryWalkerDateFilters(t *testing.T) { t.Errorf("expected commit before sinceTime to be rejected") } } + +// countingReader wraps an ObjectReader, counting reads and optionally cancelling a +// context once a given number of reads has been served. It lets a test assert on the +// I/O the walker actually performed rather than only on the error it returned. +type countingReader struct { + inner ObjectReader + reads int + tripAt int + cancel context.CancelFunc +} + +func (c *countingReader) ReadObject(oid string) (*Object, error) { + c.reads++ + if c.tripAt > 0 && c.reads == c.tripAt && c.cancel != nil { + c.cancel() + } + return c.inner.ReadObject(oid) +} + +func (c *countingReader) HasObject(oid string) bool { return c.inner.HasObject(oid) } + +func (c *countingReader) Close() error { return c.inner.Close() } + +// buildLinearHistory creates n chained commits, each replacing f.txt with a fresh blob, +// so every commit contributes exactly one blob occurrence. Returns the repo and the tip SHA. +func buildLinearHistory(t *testing.T, reader *mockObjectReader, n int) (*RepoInfo, string) { + t.Helper() + + prevSHA := "" + for i := 0; i < n; i++ { + blobOID := reader.put(TypeBlob, []byte(fmt.Sprintf("version %d", i))) + treeOID := reader.put(TypeTree, buildTreePayload([]TreeEntry{ + {Mode: 0100644, Name: "f.txt", OID: blobOID}, + })) + + var raw string + if prevSHA == "" { + raw = fmt.Sprintf("tree %s\nauthor A %d +0000\ncommitter A %d +0000\n\ncommit %d\n", + treeOID, 1600000000+i, 1600000000+i, i) + } else { + raw = fmt.Sprintf("tree %s\nparent %s\nauthor A %d +0000\ncommitter A %d +0000\n\ncommit %d\n", + treeOID, prevSHA, 1600000000+i, 1600000000+i, i) + } + prevSHA = reader.put(TypeCommit, []byte(raw)) + } + + tmpDir := t.TempDir() + gitDir := filepath.Join(tmpDir, ".git") + if err := os.MkdirAll(filepath.Join(gitDir, "refs", "heads"), 0755); err != nil { + t.Fatalf("mkdir: %v", err) + } + if err := os.WriteFile(filepath.Join(gitDir, "HEAD"), []byte("ref: refs/heads/main\n"), 0644); err != nil { + t.Fatalf("write HEAD: %v", err) + } + if err := os.WriteFile(filepath.Join(gitDir, "refs", "heads", "main"), []byte(prevSHA+"\n"), 0644); err != nil { + t.Fatalf("write ref: %v", err) + } + + return &RepoInfo{WorkTree: tmpDir, GitDir: gitDir, CommonGitDir: gitDir}, prevSHA +} + +// Cancellation during commit-DAG collection must abort the walk. This phase reads a +// commit object per DAG node and emits nothing, so before ctx was threaded through +// collectOrderedCommits it was entirely deaf to cancellation. +func TestHistoryWalkerCancelDuringCommitCollection(t *testing.T) { + const numCommits = 400 + + mock := newMockReader() + repo, _ := buildLinearHistory(t, mock, numCommits) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + reader := &countingReader{inner: mock, tripAt: 20, cancel: cancel} + walker := NewHistoryWalker(repo, reader, &model.Config{}, nil) + + occurrences := 0 + err := walker.Walk(ctx, func(occ model.BlobOccurrence) error { + occurrences++ + return nil + }) + + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected context.Canceled, got %v", err) + } + if occurrences != 0 { + t.Fatalf("cancellation happened during commit collection, before any emission, but %d occurrences were emitted", occurrences) + } + // The collection loop pops one commit per iteration and reads one object per pop, + // so an abort at the loop head must stop within a handful of reads of the trip point. + if reader.reads > reader.tripAt+5 { + t.Fatalf("walk kept reading after cancellation: %d reads (cancelled at %d) out of %d commits", + reader.reads, reader.tripAt, numCommits) + } +} + +// Cancellation once the walk has begun emitting must stop promptly. The callback used +// to be the only cancellation point, and it returns nil here, so the walker itself has +// to observe the cancelled context. +func TestHistoryWalkerCancelMidWalk(t *testing.T) { + const numCommits = 400 + + mock := newMockReader() + repo, _ := buildLinearHistory(t, mock, numCommits) + + // Baseline: a full, uncancelled walk, to know what "draining the whole DAG" costs. + baseline := &countingReader{inner: mock} + baselineOcc := 0 + if err := NewHistoryWalker(repo, baseline, &model.Config{}, nil).Walk( + context.Background(), + func(occ model.BlobOccurrence) error { + baselineOcc++ + return nil + }, + ); err != nil { + t.Fatalf("baseline walk failed: %v", err) + } + if baselineOcc != numCommits { + t.Fatalf("baseline walk emitted %d occurrences, want %d", baselineOcc, numCommits) + } + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + reader := &countingReader{inner: mock} + walker := NewHistoryWalker(repo, reader, &model.Config{}, nil) + + occurrences := 0 + readsAtCancel := 0 + err := walker.Walk(ctx, func(occ model.BlobOccurrence) error { + occurrences++ + if occurrences == 1 { + cancel() + readsAtCancel = reader.reads + } + return nil + }) + + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected context.Canceled, got %v", err) + } + if occurrences != 1 { + t.Fatalf("expected the walk to stop after the occurrence that cancelled it, got %d occurrences", occurrences) + } + if after := reader.reads - readsAtCancel; after > 5 { + t.Fatalf("walk performed %d object reads after cancellation (a full drain costs %d)", after, baseline.reads) + } + if reader.reads >= baseline.reads { + t.Fatalf("cancelled walk read %d objects, no better than the full walk's %d", reader.reads, baseline.reads) + } +} + +// The exclude side of an A..B rev-range is a second unbounded commit walk. It must +// abort on cancellation instead of draining the whole ancestry of A. +func TestTraverseExcludeCancellation(t *testing.T) { + const numCommits = 400 + + mock := newMockReader() + repo, tipSHA := buildLinearHistory(t, mock, numCommits) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + reader := &countingReader{inner: mock, tripAt: 10, cancel: cancel} + walker := NewHistoryWalker(repo, reader, &model.Config{}, nil) + + excluded := make(map[string]bool) + err := walker.traverseExclude(ctx, tipSHA, excluded) + + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected context.Canceled, got %v", err) + } + if len(excluded) >= numCommits { + t.Fatalf("traverseExclude drained the whole ancestry (%d commits) despite cancellation", len(excluded)) + } + if reader.reads > reader.tripAt+5 { + t.Fatalf("traverseExclude kept reading after cancellation: %d reads (cancelled at %d)", reader.reads, reader.tripAt) + } +} + +// Tree traversal must also honour cancellation: a root commit with no parent walks the +// full tree with the callback as its only former cancellation point. +func TestTraverseTreeCancellation(t *testing.T) { + reader := newMockReader() + + var entries []TreeEntry + for i := 0; i < 64; i++ { + entries = append(entries, TreeEntry{ + Mode: 0100644, + Name: fmt.Sprintf("f%02d.txt", i), + OID: reader.put(TypeBlob, []byte(fmt.Sprintf("blob %d", i))), + }) + } + rootOID := reader.put(TypeTree, buildTreePayload(entries)) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + visited := 0 + err := TraverseTree(ctx, reader, rootOID, func(path string, entry TreeEntry) error { + visited++ + if visited == 3 { + cancel() + } + return nil + }) + + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected context.Canceled, got %v", err) + } + if visited != 3 { + t.Fatalf("expected traversal to stop at the entry that cancelled it, visited %d of %d", visited, len(entries)) + } +} diff --git a/test/leak_test.go b/test/leak_test.go index 48ad2d9..82a67a8 100644 --- a/test/leak_test.go +++ b/test/leak_test.go @@ -379,7 +379,7 @@ func TestGoroutineLeak_EndToEndLifecycle(t *testing.T) { walker := gitengine.NewHistoryWalker(repoInfo, repoReader, cfg, nil) var occurrences []model.BlobOccurrence - err = walker.Walk(func(occ model.BlobOccurrence) error { + err = walker.Walk(context.Background(), func(occ model.BlobOccurrence) error { occurrences = append(occurrences, occ) return nil })