diff --git a/internal/ufs/fs_unix.go b/internal/ufs/fs_unix.go index f311e1d5..e5df5e04 100644 --- a/internal/ufs/fs_unix.go +++ b/internal/ufs/fs_unix.go @@ -856,3 +856,29 @@ func (fs *UnixFS) unsafeIsPathInsideOfBase(path string) bool { fs.basePath+"/", ) } +// Readlink returns the destination of the named symbolic link. +// If there is an error, it will be of type *PathError. +func (fs *UnixFS) Readlink(name string) (string, error) { + dirfd, name, closeFd, err := fs.safePath(name) + defer closeFd() + if err != nil { + return "", err + } + return fs.Readlinkat(dirfd, name) +} + +// Readlinkat is like Readlink but allows passing an existing directory file +// descriptor rather than needing to resolve one. +func (fs *UnixFS) Readlinkat(dirfd int, name string) (string, error) { + for size := 128; ; size *= 2 { + buf := make([]byte, size) + n, err := unix.Readlinkat(dirfd, name, buf) + if err != nil { + return "", ensurePathError(err, "readlinkat", name) + } + if n < size { + return string(buf[:n]), nil + } + } +} + diff --git a/internal/ufs/fs_unix_test.go b/internal/ufs/fs_unix_test.go index 3abf47c3..650ccb46 100644 --- a/internal/ufs/fs_unix_test.go +++ b/internal/ufs/fs_unix_test.go @@ -610,6 +610,118 @@ func TestUnixFS_Lstat(t *testing.T) { // TODO: implement } +func TestUnixFS_Readlink(t *testing.T) { + t.Parallel() + fs, err := newTestUnixFS() + if err != nil { + t.Fatal(err) + return + } + defer fs.Cleanup() + + t.Run("reads a relative symlink target", func(t *testing.T) { + f, err := fs.Create("target_file") + if err != nil { + t.Error(err) + return + } + _ = f.Close() + + if err := fs.Symlink("target_file", filepath.Join(fs.Root, "relative_link")); err != nil { + t.Error(err) + return + } + + target, err := fs.Readlink("relative_link") + if err != nil { + t.Errorf("expected no error, but got: %v", err) + return + } + if target != "target_file" { + t.Errorf("expected link target %q, got %q", "target_file", target) + } + }) + + t.Run("reads a symlink target pointing outside the base directory", func(t *testing.T) { + outsideTarget := filepath.Join(fs.TmpDir, "outside_file") + if err := os.WriteFile(outsideTarget, []byte("outside"), 0o644); err != nil { + t.Error(err) + return + } + if err := fs.Symlink(outsideTarget, filepath.Join(fs.Root, "outside_link")); err != nil { + t.Error(err) + return + } + + target, err := fs.Readlink("outside_link") + if err != nil { + t.Errorf("expected no error, but got: %v", err) + return + } + if target != outsideTarget { + t.Errorf("expected link target %q, got %q", outsideTarget, target) + } + }) + + t.Run("errors when the file is not a symlink", func(t *testing.T) { + f, err := fs.Create("not_a_link") + if err != nil { + t.Error(err) + return + } + _ = f.Close() + + if _, err := fs.Readlink("not_a_link"); err == nil { + t.Error("expected an error when reading a non-symlink as a link") + } + }) + + t.Run("errors when the file does not exist", func(t *testing.T) { + if _, err := fs.Readlink("does_not_exist"); err == nil { + t.Error("expected an error when reading a non-existent symlink") + } + }) +} + +func TestUnixFS_Readlinkat(t *testing.T) { + t.Parallel() + fs, err := newTestUnixFS() + if err != nil { + t.Fatal(err) + return + } + defer fs.Cleanup() + + if err := fs.Mkdir("nested", 0o755); err != nil { + t.Error(err) + return + } + f, err := fs.Create("nested/target_file") + if err != nil { + t.Error(err) + return + } + _ = f.Close() + if err := fs.Symlink("target_file", filepath.Join(fs.Root, "nested/link")); err != nil { + t.Error(err) + return + } + + dirfd, name, closeFd, err := fs.SafePath("nested/link") + defer closeFd() + if err != nil { + t.Fatal(err) + } + + target, err := fs.Readlinkat(dirfd, name) + if err != nil { + t.Fatalf("expected no error, but got: %v", err) + } + if target != "target_file" { + t.Errorf("expected link target %q, got %q", "target_file", target) + } +} + func TestUnixFS_Symlink(t *testing.T) { t.Parallel() fs, err := newTestUnixFS() diff --git a/server/backup.go b/server/backup.go index 0479cde1..c694585a 100644 --- a/server/backup.go +++ b/server/backup.go @@ -4,6 +4,7 @@ import ( "io" "io/fs" "os" + "path/filepath" "time" "emperror.dev/errors" @@ -151,12 +152,19 @@ func (s *Server) RestoreBackup(b backup.BackupInterface, reader io.ReadCloser) ( // Attempt to restore the backup to the server by running through each entry // in the file one at a time and writing them to the disk. s.Log().Debug("starting file writing process for backup restoration") - err = b.Restore(s.Context(), reader, func(file string, info fs.FileInfo, r io.ReadCloser) error { - defer r.Close() + err = b.Restore(s.Context(), reader, func(file string, info fs.FileInfo, linkTarget string, r io.ReadCloser) error { + if r != nil { + defer r.Close() + } s.Events().Publish(DaemonMessageEvent, "(restoring): "+file) - // TODO: since this will be called a lot, it may be worth adding an optimized - // Write with Chtimes method to the UnixFS that is able to re-use the - // same dirfd and file name. + + if info.Mode()&fs.ModeSymlink != 0 { + if err := s.Filesystem().CreateDirectory(filepath.Dir(file), ""); err != nil { + return err + } + return s.Filesystem().Symlink(linkTarget, file) + } + if err := s.Filesystem().Write(file, r, info.Size(), info.Mode()); err != nil { return err } diff --git a/server/backup/backup.go b/server/backup/backup.go index 823750e7..628829ae 100644 --- a/server/backup/backup.go +++ b/server/backup/backup.go @@ -36,7 +36,7 @@ const ( // RestoreCallback is a generic restoration callback that exists for both local // and remote backups allowing the files to be restored. -type RestoreCallback func(file string, info fs.FileInfo, r io.ReadCloser) error +type RestoreCallback func(file string, info fs.FileInfo, linkTarget string, r io.ReadCloser) error // noinspection GoNameStartsWithPackageName type BackupInterface interface { diff --git a/server/backup/backup_local.go b/server/backup/backup_local.go index dfdbf98b..e4a8c6ce 100644 --- a/server/backup/backup_local.go +++ b/server/backup/backup_local.go @@ -128,13 +128,16 @@ func (b *LocalBackup) Restore(ctx context.Context, _ io.Reader, callback Restore reader = ratelimit.Reader(f, ratelimit.NewBucketWithRate(float64(writeLimit), writeLimit)) } if err := format.Extract(ctx, reader, func(ctx context.Context, f archives.FileInfo) error { - r, err := f.Open() - if err != nil { - return err - } - defer r.Close() + if f.LinkTarget != "" { + return callback(f.NameInArchive, f.FileInfo, f.LinkTarget, nil) + } + r, err := f.Open() + if err != nil { + return err + } + defer r.Close() - return callback(f.NameInArchive, f.FileInfo, r) + return callback(f.NameInArchive, f.FileInfo, "", r) }); err != nil { return err } diff --git a/server/backup/backup_s3.go b/server/backup/backup_s3.go index 3f4df5fa..fe0fb26f 100644 --- a/server/backup/backup_s3.go +++ b/server/backup/backup_s3.go @@ -108,13 +108,16 @@ func (s *S3Backup) Restore(ctx context.Context, r io.Reader, callback RestoreCal reader = ratelimit.Reader(r, ratelimit.NewBucketWithRate(float64(writeLimit), writeLimit)) } if err := format.Extract(ctx, reader, func(ctx context.Context, f archives.FileInfo) error { - r, err := f.Open() - if err != nil { - return err - } - defer r.Close() + if f.LinkTarget != "" { + return callback(f.NameInArchive, f.FileInfo, f.LinkTarget, nil) + } + r, err := f.Open() + if err != nil { + return err + } + defer r.Close() - return callback(f.NameInArchive, f.FileInfo, r) + return callback(f.NameInArchive, f.FileInfo, "", r) }); err != nil { return err } diff --git a/server/backup/backup_test.go b/server/backup/backup_test.go index 6c6138bb..0b8a6f48 100644 --- a/server/backup/backup_test.go +++ b/server/backup/backup_test.go @@ -1,8 +1,12 @@ package backup import ( + "archive/tar" "bytes" + "compress/gzip" "context" + "io" + "io/fs" "os" "path/filepath" "strings" @@ -54,6 +58,122 @@ func TestBackupPathUsesBackupDirectory(t *testing.T) { } } +func TestBackupRestoreDoesNotSkipSymlinks(t *testing.T) { + backupDir := t.TempDir() + config.Set(&config.Configuration{ + AuthenticationToken: "test-token", + System: config.SystemConfiguration{ + BackupDirectory: backupDir, + }, + }) + + archiveData := buildTestArchive(t, + map[string]string{"real_file.txt": "hello, world!\n"}, + map[string]string{"link_to_file.txt": "real_file.txt"}, + ) + + t.Run("local", func(t *testing.T) { + b := NewLocal(nil, "11111111-1111-1111-1111-111111111111", "ce6ee345-6729-4aed-8fed-c866c535a69d", "") + if err := os.MkdirAll(filepath.Dir(b.Path()), 0o700); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(b.Path(), archiveData, 0o600); err != nil { + t.Fatal(err) + } + assertRestoreHandlesSymlinks(t, b.Restore, nil) + }) + + t.Run("s3", func(t *testing.T) { + b := NewS3(nil, "22222222-2222-2222-2222-222222222222", "ce6ee345-6729-4aed-8fed-c866c535a69d", "") + assertRestoreHandlesSymlinks(t, b.Restore, bytes.NewReader(archiveData)) + }) +} + +func assertRestoreHandlesSymlinks(t *testing.T, restore func(context.Context, io.Reader, RestoreCallback) error, reader io.Reader) { + t.Helper() + + type restoredEntry struct { + linkTarget string + hasReader bool + } + got := map[string]restoredEntry{} + + err := restore(context.Background(), reader, func(file string, info fs.FileInfo, linkTarget string, r io.ReadCloser) error { + got[file] = restoredEntry{linkTarget: linkTarget, hasReader: r != nil} + if r != nil { + _ = r.Close() + } + return nil + }) + if err != nil { + t.Fatal(err) + } + + f, ok := got["real_file.txt"] + if !ok { + t.Fatal("expected callback to be invoked for the regular file") + } + if !f.hasReader { + t.Error("expected a reader for the regular file entry") + } + if f.linkTarget != "" { + t.Errorf("expected no link target for the regular file entry, got %q", f.linkTarget) + } + + link, ok := got["link_to_file.txt"] + if !ok { + t.Fatal("expected callback to be invoked for the symlink instead of silently skipping it") + } + if link.hasReader { + t.Error("expected no reader for the symlink entry") + } + if link.linkTarget != "real_file.txt" { + t.Errorf("expected link target %q, got %q", "real_file.txt", link.linkTarget) + } +} + +func buildTestArchive(t *testing.T, files map[string]string, symlinks map[string]string) []byte { + t.Helper() + + var buf bytes.Buffer + gw := gzip.NewWriter(&buf) + tw := tar.NewWriter(gw) + + for name, contents := range files { + hdr := &tar.Header{ + Name: name, + Mode: 0o644, + Size: int64(len(contents)), + } + if err := tw.WriteHeader(hdr); err != nil { + t.Fatal(err) + } + if _, err := tw.Write([]byte(contents)); err != nil { + t.Fatal(err) + } + } + + for name, target := range symlinks { + hdr := &tar.Header{ + Name: name, + Typeflag: tar.TypeSymlink, + Linkname: target, + Mode: 0o777, + } + if err := tw.WriteHeader(hdr); err != nil { + t.Fatal(err) + } + } + + if err := tw.Close(); err != nil { + t.Fatal(err) + } + if err := gw.Close(); err != nil { + t.Fatal(err) + } + return buf.Bytes() +} + func testBackupGenerateRequiresUuidIdentifier(t *testing.T, createBackup func(string) BackupInterface) { t.Helper() diff --git a/server/filesystem/archive.go b/server/filesystem/archive.go index 34256521..41a6053c 100644 --- a/server/filesystem/archive.go +++ b/server/filesystem/archive.go @@ -273,17 +273,11 @@ func (a *Archive) addToArchive(dirfd int, name, relative string, entry ufs.DirEn return nil } - // Resolve the symlink target if the file is a symlink. var target string if s.Mode()&fs.ModeSymlink != 0 { - // Read the target of the symlink. If there are any errors we will dump them out to - // the logs, but we're not going to stop the backup. There are far too many cases of - // symlinks causing all sorts of unnecessary pain in this process. Sucks to suck if - // it doesn't work. - target, err = os.Readlink(s.Name()) + target, err = a.Filesystem.unixFS.Readlinkat(dirfd, name) if err != nil { - // Ignore the not exist errors specifically, since there is nothing important about that. - if !os.IsNotExist(err) { + if !errors.Is(err, ufs.ErrNotExist) { log.WithField("name", name).WithField("readlink_err", err.Error()).Warn("failed reading symlink for target path; skipping...") } return nil diff --git a/server/filesystem/archive_test.go b/server/filesystem/archive_test.go index 26e3fe96..7048019e 100644 --- a/server/filesystem/archive_test.go +++ b/server/filesystem/archive_test.go @@ -1,7 +1,10 @@ package filesystem import ( + "archive/tar" + "compress/gzip" "context" + "io" iofs "io/fs" "os" "path/filepath" @@ -84,6 +87,42 @@ func TestArchive_Stream(t *testing.T) { g.Assert(files).Equal(expected) }) + + g.It("includes symlinks in the archive instead of silently skipping them", func() { + r := strings.NewReader("hello, world!\n") + err := fs.Write("real_file.txt", r, r.Size(), 0o644) + g.Assert(err).IsNil() + + g.Assert(fs.Symlink("real_file.txt", "link_to_file.txt")).IsNil() + g.Assert(fs.Symlink("/etc/passwd", "link_outside_root.txt")).IsNil() + + a := &Archive{ + Filesystem: fs, + } + + archivePath := filepath.Join(rfs.root, "archive_symlinks.tar.gz") + g.Assert(a.Create(context.Background(), archivePath)).IsNil() + + _, err = os.Stat(archivePath) + g.Assert(err).IsNil() + + entries, err := readTarHeaders(archivePath) + g.Assert(err).IsNil() + + link, ok := entries["link_to_file.txt"] + g.Assert(ok).IsTrue() + g.Assert(link.Typeflag).Equal(byte(tar.TypeSymlink)) + g.Assert(link.Linkname).Equal("real_file.txt") + + outside, ok := entries["link_outside_root.txt"] + g.Assert(ok).IsTrue() + g.Assert(outside.Typeflag).Equal(byte(tar.TypeSymlink)) + g.Assert(outside.Linkname).Equal("/etc/passwd") + + file, ok := entries["real_file.txt"] + g.Assert(ok).IsTrue() + g.Assert(file.Typeflag).Equal(byte(tar.TypeReg)) + }) }) } @@ -120,3 +159,31 @@ func getFiles(f iofs.ReadDirFS, name string) ([]string, error) { return v, nil } + +func readTarHeaders(path string) (map[string]*tar.Header, error) { + f, err := os.Open(path) + if err != nil { + return nil, err + } + defer f.Close() + + gz, err := gzip.NewReader(f) + if err != nil { + return nil, err + } + defer gz.Close() + + tr := tar.NewReader(gz) + entries := make(map[string]*tar.Header) + for { + hdr, err := tr.Next() + if err == io.EOF { + break + } + if err != nil { + return nil, err + } + entries[hdr.Name] = hdr + } + return entries, nil +} \ No newline at end of file