diff --git a/internal/adapters/local/backend.go b/internal/adapters/local/backend.go index cc268ed..d4fe466 100644 --- a/internal/adapters/local/backend.go +++ b/internal/adapters/local/backend.go @@ -161,31 +161,10 @@ func (b *Backend) Walk(ctx context.Context, prefix string, opts storage.WalkOpti } return b.translateError(storage.OpWalk, prefix, err) } - visited := 0 - emit := func(entry storage.Entry) error { - if err := ctx.Err(); err != nil { - return err - } - if opts.Limit > 0 && visited >= opts.Limit { - return storage.ErrStopWalk - } - visited++ - if err := fn(entry); err != nil { - if errors.Is(err, storage.ErrStopWalk) { - return storage.ErrStopWalk - } - return storage.NewError(storage.OpWalk, backendName, entry.Path, storage.ErrUnknown, err) - } - return nil - } + emitter := storage.NewWalkEmitter(ctx, backendName, opts, fn) if !info.IsDir() { - if err := emit(entryFromInfo(prefix, info)); errors.Is(err, storage.ErrStopWalk) { - return nil - } else if err != nil { - return err - } - return nil + return storage.FinishWalk(emitter.Emit(entryFromInfo(prefix, info))) } walkErr := filepath.WalkDir(nativePrefix, func(nativePath string, dirEntry fs.DirEntry, err error) error { @@ -210,24 +189,13 @@ func (b *Backend) Walk(ctx context.Context, prefix string, opts storage.WalkOpti if err != nil { return b.translateError(storage.OpWalk, logicalPath, err) } - return emit(entryFromInfo(logicalPath, info)) + return emitter.Emit(entryFromInfo(logicalPath, info)) }) - if errors.Is(walkErr, storage.ErrStopWalk) { - return nil - } - return walkErr + return storage.FinishWalk(walkErr) } func (b *Backend) HasAny(ctx context.Context, prefix string) (bool, error) { - found := false - err := b.Walk(ctx, prefix, storage.WalkOptions{Recursive: false, Limit: 1}, func(storage.Entry) error { - found = true - return storage.ErrStopWalk - }) - if err != nil { - return false, err - } - return found, nil + return storage.HasAny(ctx, b, prefix) } func (b *Backend) DeleteManagedBundle(ctx context.Context, bundlePath string, managedOutputPaths []string, opts storage.DeleteOptions) error { diff --git a/internal/adapters/s3/backend.go b/internal/adapters/s3/backend.go index b86a61d..576c15d 100644 --- a/internal/adapters/s3/backend.go +++ b/internal/adapters/s3/backend.go @@ -183,55 +183,25 @@ func (b *Backend) Walk(ctx context.Context, logicalPrefix string, opts storage.W if err := storage.ValidatePrefix(logicalPrefix); err != nil { return err } - visited := 0 - emit := func(entry storage.Entry) error { - if err := ctx.Err(); err != nil { - return err - } - if opts.Limit > 0 && visited >= opts.Limit { - return storage.ErrStopWalk - } - visited++ - if err := fn(entry); err != nil { - if errors.Is(err, storage.ErrStopWalk) { - return storage.ErrStopWalk - } - return storage.NewError(storage.OpWalk, BackendName, entry.Path, storage.ErrUnknown, err) - } - return nil - } + emitter := storage.NewWalkEmitter(ctx, BackendName, opts, fn) if logicalPrefix != "" { entry, err := b.Stat(ctx, logicalPrefix) if err == nil { - if err := emit(entry); errors.Is(err, storage.ErrStopWalk) { - return nil - } else if err != nil { - return err + if err := emitter.Emit(entry); err != nil { + return storage.FinishWalk(err) } - if opts.Limit > 0 && visited >= opts.Limit { + if emitter.LimitReached() { return nil } } else if !storage.IsNotFound(err) { return err } } - err := b.walkObjects(ctx, logicalPrefix, opts, emit) - if errors.Is(err, storage.ErrStopWalk) { - return nil - } - return err + return storage.FinishWalk(b.walkObjects(ctx, logicalPrefix, opts, emitter.Emit)) } func (b *Backend) HasAny(ctx context.Context, logicalPrefix string) (bool, error) { - found := false - err := b.Walk(ctx, logicalPrefix, storage.WalkOptions{Recursive: false, Limit: 1}, func(storage.Entry) error { - found = true - return storage.ErrStopWalk - }) - if err != nil { - return false, err - } - return found, nil + return storage.HasAny(ctx, b, logicalPrefix) } func (b *Backend) DeleteManagedBundle(ctx context.Context, bundlePath string, managedOutputPaths []string, opts storage.DeleteOptions) error { diff --git a/internal/adapters/ssh/backend.go b/internal/adapters/ssh/backend.go index cbd911f..dae5236 100644 --- a/internal/adapters/ssh/backend.go +++ b/internal/adapters/ssh/backend.go @@ -209,51 +209,17 @@ func (b *Backend) Walk(ctx context.Context, prefix string, opts storage.WalkOpti return b.translateError(storage.OpWalk, prefix, err) } - visited := 0 - emit := func(entry storage.Entry) error { - if err := ctx.Err(); err != nil { - return err - } - if opts.Limit > 0 && visited >= opts.Limit { - return storage.ErrStopWalk - } - visited++ - if err := fn(entry); err != nil { - if errors.Is(err, storage.ErrStopWalk) { - return storage.ErrStopWalk - } - return storage.NewError(storage.OpWalk, BackendName, entry.Path, storage.ErrUnknown, err) - } - return nil - } + emitter := storage.NewWalkEmitter(ctx, BackendName, opts, fn) if !info.IsDir() { - if err := emit(entryFromInfo(prefix, info)); errors.Is(err, storage.ErrStopWalk) { - return nil - } else if err != nil { - return err - } - return nil + return storage.FinishWalk(emitter.Emit(entryFromInfo(prefix, info))) } - if err := b.walkDirectory(ctx, prefix, nativePrefix, opts, emit); errors.Is(err, storage.ErrStopWalk) { - return nil - } else if err != nil { - return err - } - return nil + return storage.FinishWalk(b.walkDirectory(ctx, prefix, nativePrefix, opts, emitter.Emit)) } func (b *Backend) HasAny(ctx context.Context, prefix string) (bool, error) { - found := false - err := b.Walk(ctx, prefix, storage.WalkOptions{Recursive: false, Limit: 1}, func(storage.Entry) error { - found = true - return storage.ErrStopWalk - }) - if err != nil { - return false, err - } - return found, nil + return storage.HasAny(ctx, b, prefix) } func (b *Backend) DeleteManagedBundle(ctx context.Context, bundlePath string, managedOutputPaths []string, opts storage.DeleteOptions) error { diff --git a/internal/storage/walk.go b/internal/storage/walk.go new file mode 100644 index 0000000..62cb15a --- /dev/null +++ b/internal/storage/walk.go @@ -0,0 +1,63 @@ +package storage + +import ( + "context" + "errors" +) + +type WalkEmitter struct { + ctx context.Context + backend string + opts WalkOptions + fn WalkFunc + count int +} + +func NewWalkEmitter(ctx context.Context, backend string, opts WalkOptions, fn WalkFunc) *WalkEmitter { + return &WalkEmitter{ + ctx: ctx, + backend: backend, + opts: opts, + fn: fn, + } +} + +func (e *WalkEmitter) Emit(entry Entry) error { + if err := e.ctx.Err(); err != nil { + return err + } + if e.opts.Limit > 0 && e.count >= e.opts.Limit { + return ErrStopWalk + } + e.count++ + if err := e.fn(entry); err != nil { + if errors.Is(err, ErrStopWalk) { + return ErrStopWalk + } + return NewError(OpWalk, e.backend, entry.Path, ErrUnknown, err) + } + return nil +} + +func (e *WalkEmitter) LimitReached() bool { + return e.opts.Limit > 0 && e.count >= e.opts.Limit +} + +func FinishWalk(err error) error { + if errors.Is(err, ErrStopWalk) { + return nil + } + return err +} + +func HasAny(ctx context.Context, backend Backend, prefix string) (bool, error) { + found := false + err := backend.Walk(ctx, prefix, WalkOptions{Recursive: false, Limit: 1}, func(Entry) error { + found = true + return ErrStopWalk + }) + if err != nil { + return false, err + } + return found, nil +} diff --git a/internal/storage/walk_test.go b/internal/storage/walk_test.go new file mode 100644 index 0000000..1b0c4a4 --- /dev/null +++ b/internal/storage/walk_test.go @@ -0,0 +1,107 @@ +package storage + +import ( + "context" + "errors" + "testing" +) + +func TestWalkEmitterHonorsLimit(t *testing.T) { + emitter := NewWalkEmitter(context.Background(), "test", WalkOptions{Limit: 2}, func(Entry) error { + return nil + }) + + if err := emitter.Emit(Entry{Path: "one"}); err != nil { + t.Fatalf("first Emit() error = %v", err) + } + if err := emitter.Emit(Entry{Path: "two"}); err != nil { + t.Fatalf("second Emit() error = %v", err) + } + if err := emitter.Emit(Entry{Path: "three"}); !errors.Is(err, ErrStopWalk) { + t.Fatalf("third Emit() error = %v, want ErrStopWalk", err) + } + if !emitter.LimitReached() { + t.Fatal("LimitReached() = false, want true") + } +} + +func TestWalkEmitterStopsWithoutError(t *testing.T) { + emitter := NewWalkEmitter(context.Background(), "test", WalkOptions{}, func(Entry) error { + return ErrStopWalk + }) + + err := emitter.Emit(Entry{Path: "one"}) + if !errors.Is(err, ErrStopWalk) { + t.Fatalf("Emit() error = %v, want ErrStopWalk", err) + } + if err := FinishWalk(err); err != nil { + t.Fatalf("FinishWalk() error = %v, want nil", err) + } +} + +func TestWalkEmitterWrapsCallbackErrors(t *testing.T) { + callbackErr := errors.New("callback failed") + emitter := NewWalkEmitter(context.Background(), "test", WalkOptions{}, func(Entry) error { + return callbackErr + }) + + err := emitter.Emit(Entry{Path: "one"}) + if !errors.Is(err, callbackErr) { + t.Fatalf("Emit() error = %v, want callback error", err) + } + var storageErr *Error + if !errors.As(err, &storageErr) { + t.Fatalf("Emit() error type = %T, want *Error", err) + } + if storageErr.Op != OpWalk || storageErr.Backend != "test" || storageErr.Path != "one" || storageErr.Kind != ErrUnknown { + t.Fatalf("wrapped error = %#v", storageErr) + } +} + +func TestWalkEmitterHonorsContextCancellation(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + called := false + emitter := NewWalkEmitter(ctx, "test", WalkOptions{}, func(Entry) error { + called = true + return nil + }) + + err := emitter.Emit(Entry{Path: "one"}) + if !errors.Is(err, context.Canceled) { + t.Fatalf("Emit() error = %v, want context.Canceled", err) + } + if called { + t.Fatal("callback was called after context cancellation") + } +} + +func TestHasAnyUsesNonRecursiveLimitOneWalk(t *testing.T) { + backend := &recordingBackend{} + + found, err := HasAny(context.Background(), backend, "bundle") + if err != nil { + t.Fatalf("HasAny() error = %v", err) + } + if !found { + t.Fatal("HasAny() = false, want true") + } + if backend.prefix != "bundle" { + t.Fatalf("walk prefix = %q, want bundle", backend.prefix) + } + if backend.opts != (WalkOptions{Recursive: false, Limit: 1}) { + t.Fatalf("walk options = %#v, want non-recursive limit one", backend.opts) + } +} + +type recordingBackend struct { + Backend + prefix string + opts WalkOptions +} + +func (b *recordingBackend) Walk(_ context.Context, prefix string, opts WalkOptions, fn WalkFunc) error { + b.prefix = prefix + b.opts = opts + return FinishWalk(fn(Entry{Path: "bundle/file.txt", Type: EntryTypeFile})) +}