Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 1 addition & 4 deletions cmd/grg/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
})
Expand Down
7 changes: 5 additions & 2 deletions internal/gitengine/hardening_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package gitengine

import (
"bytes"
"context"
"encoding/binary"
"encoding/hex"
"errors"
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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))
Expand Down
21 changes: 16 additions & 5 deletions internal/gitengine/tree.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package gitengine

import (
"bytes"
"context"
"encoding/hex"
"errors"
"fmt"
Expand Down Expand Up @@ -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)
}
Expand All @@ -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
Expand All @@ -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
}
}
Expand All @@ -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
}
Expand All @@ -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
Expand Down
5 changes: 3 additions & 2 deletions internal/gitengine/tree_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package gitengine

import (
"bytes"
"context"
"crypto/sha1"
"encoding/hex"
"errors"
Expand Down Expand Up @@ -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
})
Expand All @@ -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
Expand Down
53 changes: 39 additions & 14 deletions internal/gitengine/walker.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package gitengine

import (
"container/heap"
"context"
"fmt"
"regexp"
"slices"
Expand All @@ -28,24 +29,29 @@ 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
}

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

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])
Expand All @@ -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
}
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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
}
}
Expand All @@ -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
}
Expand All @@ -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
}
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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]
Expand All @@ -410,4 +433,6 @@ func (w *HistoryWalker) traverseExclude(startSHA string, excluded map[string]boo
}
}
}

return nil
}
Loading
Loading