From c84d8868d15a14ca3c82cf45b335c19c9f5038fb Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Sat, 13 Jun 2026 13:57:50 -0500 Subject: [PATCH] Fix a bug in the SFTP backend that would cause an error when overwriting existing files --- internal/adapters/ssh/backend.go | 31 ++++++- internal/adapters/ssh/backend_test.go | 125 ++++++++++++++++++++++++++ 2 files changed, 155 insertions(+), 1 deletion(-) create mode 100644 internal/adapters/ssh/backend_test.go diff --git a/internal/adapters/ssh/backend.go b/internal/adapters/ssh/backend.go index 713fd7e..0eed0b0 100644 --- a/internal/adapters/ssh/backend.go +++ b/internal/adapters/ssh/backend.go @@ -170,7 +170,7 @@ func (b *Backend) WriteFrom(ctx context.Context, logicalPath string, r io.Reader return storage.Entry{}, storage.NewError(storage.OpWriteFrom, BackendName, logicalPath, storage.ErrConflict, fmt.Errorf("stream size %d does not match expected size %d", written, opts.Size)) } if opts.PreferAtomic { - if err := b.client.Rename(writePath, nativePath); err != nil { + if err := renamePromotedFile(b.client, writePath, nativePath, opts.Overwrite); err != nil { return storage.Entry{}, b.translateError(storage.OpWriteFrom, logicalPath, err) } cleanup = false @@ -178,6 +178,27 @@ func (b *Backend) WriteFrom(ctx context.Context, logicalPath string, r io.Reader return b.Stat(ctx, logicalPath) } +type sftpRenamer interface { + PosixRename(oldname, newname string) error + Rename(oldname, newname string) error + Remove(path string) error +} + +func renamePromotedFile(client sftpRenamer, oldname, newname string, overwrite bool) error { + if !overwrite { + return client.Rename(oldname, newname) + } + if err := client.PosixRename(oldname, newname); err == nil { + return nil + } else if !isReplaceRenameFallbackError(err) { + return err + } + if err := client.Remove(newname); err != nil && !isNotExist(err) { + return err + } + return client.Rename(oldname, newname) +} + func (b *Backend) Stat(ctx context.Context, logicalPath string) (storage.Entry, error) { if err := ctx.Err(); err != nil { return storage.Entry{}, err @@ -464,6 +485,14 @@ func isNotExist(err error) bool { return errors.Is(err, fs.ErrNotExist) || errors.Is(err, os.ErrNotExist) || errors.Is(err, sftp.ErrSSHFxNoSuchFile) } +func isReplaceRenameFallbackError(err error) bool { + if errors.Is(err, sftp.ErrSSHFxFailure) || errors.Is(err, sftp.ErrSSHFxOpUnsupported) { + return true + } + var statusErr *sftp.StatusError + return errors.As(err, &statusErr) && (statusErr.FxCode() == sftp.ErrSSHFxFailure || statusErr.FxCode() == sftp.ErrSSHFxOpUnsupported) +} + func (b *Backend) translateError(op, logicalPath string, err error) error { kind := storage.ErrUnknown switch { diff --git a/internal/adapters/ssh/backend_test.go b/internal/adapters/ssh/backend_test.go new file mode 100644 index 0000000..78f166c --- /dev/null +++ b/internal/adapters/ssh/backend_test.go @@ -0,0 +1,125 @@ +package ssh + +import ( + "errors" + "os" + "testing" + + "github.com/pkg/sftp" +) + +func TestRenamePromotedFileUsesPlainRenameWithoutOverwrite(t *testing.T) { + client := &recordingRenamer{} + if err := renamePromotedFile(client, "temp", "index.html", false); err != nil { + t.Fatalf("renamePromotedFile() error = %v", err) + } + if got, want := client.calls, []string{"rename temp index.html"}; !equalStrings(got, want) { + t.Fatalf("calls = %q, want %q", got, want) + } +} + +func TestRenamePromotedFileUsesPosixRenameForOverwrite(t *testing.T) { + client := &recordingRenamer{} + if err := renamePromotedFile(client, "temp", "index.html", true); err != nil { + t.Fatalf("renamePromotedFile() error = %v", err) + } + if got, want := client.calls, []string{"posix temp index.html"}; !equalStrings(got, want) { + t.Fatalf("calls = %q, want %q", got, want) + } +} + +func TestRenamePromotedFileFallsBackWhenReplaceRenameUnsupported(t *testing.T) { + for _, err := range []error{ + sftp.ErrSSHFxOpUnsupported, + sftp.ErrSSHFxFailure, + &sftp.StatusError{Code: uint32(sftp.ErrSSHFxOpUnsupported)}, + &sftp.StatusError{Code: uint32(sftp.ErrSSHFxFailure)}, + } { + t.Run(err.Error(), func(t *testing.T) { + client := &recordingRenamer{posixErr: err} + if err := renamePromotedFile(client, "temp", "index.html", true); err != nil { + t.Fatalf("renamePromotedFile() error = %v", err) + } + want := []string{"posix temp index.html", "remove index.html", "rename temp index.html"} + if got := client.calls; !equalStrings(got, want) { + t.Fatalf("calls = %q, want %q", got, want) + } + }) + } +} + +func TestRenamePromotedFileIgnoresMissingTargetDuringFallback(t *testing.T) { + client := &recordingRenamer{ + posixErr: sftp.ErrSSHFxOpUnsupported, + removeErr: &os.PathError{ + Op: "remove", + Path: "index.html", + Err: os.ErrNotExist, + }, + } + if err := renamePromotedFile(client, "temp", "index.html", true); err != nil { + t.Fatalf("renamePromotedFile() error = %v", err) + } + want := []string{"posix temp index.html", "remove index.html", "rename temp index.html"} + if got := client.calls; !equalStrings(got, want) { + t.Fatalf("calls = %q, want %q", got, want) + } +} + +func TestRenamePromotedFileDoesNotFallbackForPermissionError(t *testing.T) { + client := &recordingRenamer{posixErr: sftp.ErrSSHFxPermissionDenied} + if err := renamePromotedFile(client, "temp", "index.html", true); !errors.Is(err, sftp.ErrSSHFxPermissionDenied) { + t.Fatalf("renamePromotedFile() error = %v, want permission denied", err) + } + if got, want := client.calls, []string{"posix temp index.html"}; !equalStrings(got, want) { + t.Fatalf("calls = %q, want %q", got, want) + } +} + +func TestRenamePromotedFileReturnsRemoveFallbackError(t *testing.T) { + client := &recordingRenamer{ + posixErr: sftp.ErrSSHFxOpUnsupported, + removeErr: sftp.ErrSSHFxPermissionDenied, + } + if err := renamePromotedFile(client, "temp", "index.html", true); !errors.Is(err, sftp.ErrSSHFxPermissionDenied) { + t.Fatalf("renamePromotedFile() error = %v, want permission denied", err) + } + want := []string{"posix temp index.html", "remove index.html"} + if got := client.calls; !equalStrings(got, want) { + t.Fatalf("calls = %q, want %q", got, want) + } +} + +type recordingRenamer struct { + calls []string + posixErr error + renameErr error + removeErr error +} + +func (r *recordingRenamer) PosixRename(oldname, newname string) error { + r.calls = append(r.calls, "posix "+oldname+" "+newname) + return r.posixErr +} + +func (r *recordingRenamer) Rename(oldname, newname string) error { + r.calls = append(r.calls, "rename "+oldname+" "+newname) + return r.renameErr +} + +func (r *recordingRenamer) Remove(path string) error { + r.calls = append(r.calls, "remove "+path) + return r.removeErr +} + +func equalStrings(a, b []string) bool { + if len(a) != len(b) { + return false + } + for index := range a { + if a[index] != b[index] { + return false + } + } + return true +}