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
})