diff --git a/server/internal/backup/repository.go b/server/internal/backup/repository.go index 54b3f49..da0659c 100644 --- a/server/internal/backup/repository.go +++ b/server/internal/backup/repository.go @@ -16,6 +16,7 @@ import ( "errors" "fmt" "io" + "io/fs" "os" "path" "path/filepath" @@ -34,6 +35,9 @@ const ( repositoryPackPrefix = repositoryRoot + "/packs" repositorySnapshotRoot = repositoryRoot + "/snapshots" repositoryDefaultPack = int64(32 << 20) + repositoryMaxIndexSize = int64(64 << 20) + repositoryMaxSnapshot = int64(256 << 20) + repositoryMaxEncoded = int64(repositoryChunkMax + (256 << 10)) ) type RepositoryStore struct { @@ -321,6 +325,9 @@ func (s *RepositoryStore) BuildPlan(ctx context.Context, task TaskSpec, writer L return nil, fmt.Errorf("close repository chunk spool: %w", err) } spoolClosed = true + if err := s.validateSnapshot(&plan.snapshot); err != nil { + return nil, err + } writer.WriteLine(fmt.Sprintf("CDC 扫描完成:%d 个条目,逻辑数据 %d bytes,任务内唯一数据 %d bytes", len(plan.snapshot.Entries), plan.LogicalSize, plan.UniqueSize)) return plan, nil } @@ -399,11 +406,14 @@ func (s *RepositoryStore) EstimateUploadSize(ctx context.Context, provider stora return estimate, nil } -func (s *RepositoryStore) Restore(ctx context.Context, provider storage.StorageProvider, snapshotKey string, task TaskSpec, writer LogWriter) error { +func (s *RepositoryStore) Restore(ctx context.Context, provider storage.StorageProvider, snapshotKey, expectedChecksum string, task TaskSpec, writer LogWriter) (err error) { if writer == nil { writer = NopLogWriter{} } - snapshot, _, err := s.loadSnapshot(ctx, provider, snapshotKey) + if strings.TrimSpace(expectedChecksum) == "" { + return fmt.Errorf("repository snapshot checksum is required for restore") + } + snapshot, _, err := s.loadSnapshot(ctx, provider, snapshotKey, expectedChecksum) if err != nil { return err } @@ -415,19 +425,25 @@ func (s *RepositoryStore) Restore(ctx context.Context, provider storage.StorageP if len(task.SourcePaths) > 0 { restoreSource = strings.TrimSpace(task.SourcePaths[0]) } - if restoreSource == "" && len(snapshot.SourcePaths) > 0 { - restoreSource = snapshot.SourcePaths[0] - } - targetRoot := filepath.Dir(filepath.Clean(restoreSource)) - if strings.TrimSpace(task.RestoreTargetPath) != "" { - targetRoot = filepath.Clean(task.RestoreTargetPath) + targetRoot := strings.TrimSpace(task.RestoreTargetPath) + if targetRoot == "" { + if restoreSource == "" { + return fmt.Errorf("repository restore source path is required when no restore target is provided") + } + targetRoot = filepath.Dir(filepath.Clean(restoreSource)) + } else { + targetRoot = filepath.Clean(targetRoot) } if !filepath.IsAbs(targetRoot) { return fmt.Errorf("repository restore target must be absolute: %s", targetRoot) } - if err := os.MkdirAll(targetRoot, 0o755); err != nil { - return fmt.Errorf("create repository restore root: %w", err) + restoreRoot, err := s.openRestoreRoot(targetRoot) + if err != nil { + return fmt.Errorf("open repository restore root: %w", err) } + defer func() { + err = errors.Join(err, restoreRoot.Close()) + }() restored := 0 directories := make([]repositoryEntry, 0) @@ -438,44 +454,62 @@ func (s *RepositoryStore) Restore(ctx context.Context, provider storage.StorageP if len(task.SelectedPaths) > 0 && !pathSelected(entry.Path, task.SelectedPaths) { continue } - targetPath, ok := resolveWithinParent(targetRoot, entry.Path) - if !ok { + entryPath, localizeErr := filepath.Localize(entry.Path) + if localizeErr != nil || !filepath.IsLocal(entryPath) { return fmt.Errorf("unsafe repository restore path %q", entry.Path) } - if err := s.ensureSafeParent(targetRoot, targetPath); err != nil { - return err + if parent := filepath.Dir(entryPath); parent != "." { + if err := restoreRoot.MkdirAll(parent, 0o755); err != nil { + return fmt.Errorf("create restore parent for %s: %w", entry.Path, err) + } } switch entry.Kind { case "directory": - if info, statErr := os.Lstat(targetPath); statErr == nil { + if info, statErr := restoreRoot.Lstat(entryPath); statErr == nil { if info.Mode()&os.ModeSymlink != 0 { - return fmt.Errorf("restore directory crosses existing symlink: %s", targetPath) + return fmt.Errorf("restore directory crosses existing symlink: %s", entry.Path) } } else if !errors.Is(statErr, os.ErrNotExist) { - return fmt.Errorf("inspect restore directory %s: %w", targetPath, statErr) + return fmt.Errorf("inspect restore directory %s: %w", entry.Path, statErr) } - if err := os.MkdirAll(targetPath, os.FileMode(entry.Mode)); err != nil { - return fmt.Errorf("create restore directory %s: %w", targetPath, err) + if err := restoreRoot.MkdirAll(entryPath, os.FileMode(entry.Mode)); err != nil { + return fmt.Errorf("create restore directory %s: %w", entry.Path, err) } directories = append(directories, entry) case "symlink": - if err := os.RemoveAll(targetPath); err != nil { - return fmt.Errorf("replace restore symlink %s: %w", targetPath, err) + resolvedLinkTarget := path.Clean(path.Join(path.Dir(entry.Path), strings.ReplaceAll(entry.LinkTarget, "\\", "/"))) + localizedLinkTarget, localizeTargetErr := filepath.Localize(resolvedLinkTarget) + if localizeTargetErr != nil || !filepath.IsLocal(localizedLinkTarget) { + return fmt.Errorf("unsafe repository symlink target %q", entry.LinkTarget) } - if err := os.Symlink(entry.LinkTarget, targetPath); err != nil { - return fmt.Errorf("create restore symlink %s: %w", targetPath, err) + linkTarget, relativeTargetErr := filepath.Rel(filepath.Dir(entryPath), localizedLinkTarget) + if relativeTargetErr != nil { + return fmt.Errorf("resolve restore symlink target %s: %w", entry.Path, relativeTargetErr) } - case "file": - if info, statErr := os.Lstat(targetPath); statErr == nil { - if info.Mode()&os.ModeSymlink != 0 { - return fmt.Errorf("restore file would overwrite existing symlink: %s", targetPath) + if info, statErr := restoreRoot.Lstat(entryPath); statErr == nil { + if info.IsDir() { + return fmt.Errorf("refuse to replace restore directory with symlink: %s", entry.Path) + } + if err := restoreRoot.Remove(entryPath); err != nil { + return fmt.Errorf("replace restore symlink %s: %w", entry.Path, err) } } else if !errors.Is(statErr, os.ErrNotExist) { - return fmt.Errorf("inspect restore file %s: %w", targetPath, statErr) + return fmt.Errorf("inspect restore symlink %s: %w", entry.Path, statErr) } - file, openErr := os.OpenFile(targetPath, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, os.FileMode(entry.Mode)) + if err := restoreRoot.Symlink(linkTarget, entryPath); err != nil { + return fmt.Errorf("create restore symlink %s: %w", entry.Path, err) + } + case "file": + if info, statErr := restoreRoot.Lstat(entryPath); statErr == nil { + if info.Mode()&os.ModeSymlink != 0 { + return fmt.Errorf("restore file would overwrite existing symlink: %s", entry.Path) + } + } else if !errors.Is(statErr, os.ErrNotExist) { + return fmt.Errorf("inspect restore file %s: %w", entry.Path, statErr) + } + file, openErr := restoreRoot.OpenFile(entryPath, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, os.FileMode(entry.Mode)) if openErr != nil { - return fmt.Errorf("create restore file %s: %w", targetPath, openErr) + return fmt.Errorf("create restore file %s: %w", entry.Path, openErr) } var written int64 for _, chunkID := range entry.Chunks { @@ -486,54 +520,55 @@ func (s *RepositoryStore) Restore(ctx context.Context, provider storage.StorageP count, writeErr := file.Write(raw) written += int64(count) if writeErr != nil { - return errors.Join(fmt.Errorf("write restore file %s: %w", targetPath, writeErr), file.Close()) + return errors.Join(fmt.Errorf("write restore file %s: %w", entry.Path, writeErr), file.Close()) } if count != len(raw) { return errors.Join(io.ErrShortWrite, file.Close()) } } - if closeErr := file.Close(); closeErr != nil { - return fmt.Errorf("close restore file %s: %w", targetPath, closeErr) - } if written != entry.Size { - return fmt.Errorf("restored size mismatch for %s: expected %d, got %d", entry.Path, entry.Size, written) + return errors.Join(fmt.Errorf("restored size mismatch for %s: expected %d, got %d", entry.Path, entry.Size, written), file.Close()) } - if err := os.Chmod(targetPath, os.FileMode(entry.Mode)); err != nil { - return fmt.Errorf("restore mode for %s: %w", targetPath, err) + if err := file.Chmod(os.FileMode(entry.Mode)); err != nil { + return errors.Join(fmt.Errorf("restore mode for %s: %w", entry.Path, err), file.Close()) + } + if closeErr := file.Close(); closeErr != nil { + return fmt.Errorf("close restore file %s: %w", entry.Path, closeErr) + } + modTime := time.Unix(0, entry.ModTime) + if err := restoreRoot.Chtimes(entryPath, modTime, modTime); err != nil { + return fmt.Errorf("restore timestamp for %s: %w", entry.Path, err) } default: return fmt.Errorf("unsupported repository entry kind %q", entry.Kind) } - if entry.Kind == "file" { - modTime := time.Unix(0, entry.ModTime) - if err := os.Chtimes(targetPath, modTime, modTime); err != nil { - return fmt.Errorf("restore timestamp for %s: %w", targetPath, err) - } - } restored++ } // Children modify their parent directory timestamps, so directory metadata // is restored only after every selected file has been written. for index := len(directories) - 1; index >= 0; index-- { entry := directories[index] - targetPath, ok := resolveWithinParent(targetRoot, entry.Path) - if !ok { + entryPath, localizeErr := filepath.Localize(entry.Path) + if localizeErr != nil || !filepath.IsLocal(entryPath) { return fmt.Errorf("unsafe repository restore path %q", entry.Path) } - if err := os.Chmod(targetPath, os.FileMode(entry.Mode)); err != nil { - return fmt.Errorf("restore mode for %s: %w", targetPath, err) + if err := restoreRoot.Chmod(entryPath, os.FileMode(entry.Mode)); err != nil { + return fmt.Errorf("restore mode for %s: %w", entry.Path, err) } modTime := time.Unix(0, entry.ModTime) - if err := os.Chtimes(targetPath, modTime, modTime); err != nil { - return fmt.Errorf("restore timestamp for %s: %w", targetPath, err) + if err := restoreRoot.Chtimes(entryPath, modTime, modTime); err != nil { + return fmt.Errorf("restore timestamp for %s: %w", entry.Path, err) } } writer.WriteLine(fmt.Sprintf("CDC 仓库恢复完成:%d 个条目", restored)) return nil } -func (s *RepositoryStore) ExportTar(ctx context.Context, provider storage.StorageProvider, snapshotKey, destination string) (err error) { - snapshot, _, err := s.loadSnapshot(ctx, provider, snapshotKey) +func (s *RepositoryStore) ExportTar(ctx context.Context, provider storage.StorageProvider, snapshotKey, expectedChecksum, destination string) (err error) { + if strings.TrimSpace(expectedChecksum) == "" { + return fmt.Errorf("repository snapshot checksum is required for export") + } + snapshot, _, err := s.loadSnapshot(ctx, provider, snapshotKey, expectedChecksum) if err != nil { return err } @@ -609,16 +644,13 @@ func (s *RepositoryStore) ExportTar(ctx context.Context, provider storage.Storag } func (s *RepositoryStore) Verify(ctx context.Context, provider storage.StorageProvider, snapshotKey, expectedChecksum string) (*RepositoryVerifyResult, error) { - snapshot, snapshotBytes, err := s.loadSnapshot(ctx, provider, snapshotKey) + if strings.TrimSpace(expectedChecksum) == "" { + return nil, fmt.Errorf("repository snapshot checksum is required for verification") + } + snapshot, _, err := s.loadSnapshot(ctx, provider, snapshotKey, expectedChecksum) if err != nil { return nil, err } - if strings.TrimSpace(expectedChecksum) != "" { - digest := sha256.Sum256(snapshotBytes) - if !strings.EqualFold(hex.EncodeToString(digest[:]), strings.TrimSpace(expectedChecksum)) { - return nil, fmt.Errorf("repository snapshot checksum mismatch") - } - } locations, err := s.loadIndex(ctx, provider) if err != nil { return nil, err @@ -653,7 +685,7 @@ func (s *RepositoryStore) Prune(ctx context.Context, provider storage.StoragePro } sort.Slice(snapshots, func(i, j int) bool { return snapshots[i].Key < snapshots[j].Key }) for _, object := range snapshots { - snapshot, _, loadErr := s.loadSnapshot(ctx, provider, object.Key) + snapshot, _, loadErr := s.loadSnapshot(ctx, provider, object.Key, "") if loadErr != nil { return nil, fmt.Errorf("refuse to prune with unreadable snapshot %s: %w", object.Key, loadErr) } @@ -883,9 +915,6 @@ func (s *RepositoryStore) loadIndex(ctx context.Context, provider storage.Storag return nil, readErr } for chunkID, location := range segment.Chunks { - if location.Pack == "" { - location.Pack = segment.Pack - } if _, exists := locations[chunkID]; !exists { locations[chunkID] = location } @@ -895,11 +924,19 @@ func (s *RepositoryStore) loadIndex(ctx context.Context, provider storage.Storag } func (s *RepositoryStore) readIndexSegment(ctx context.Context, provider storage.StorageProvider, key string) (*repositoryIndexSegment, error) { + indexPrefix := repositoryIndexPrefix + "/" + indexName := strings.TrimPrefix(key, indexPrefix) + indexID := strings.TrimSuffix(indexName, ".json") + decodedIndexID, decodeIndexErr := hex.DecodeString(indexID) + if indexName == key || indexID == indexName || indexID != strings.ToLower(indexID) || strings.Contains(indexName, "/") || strings.Contains(indexName, "\\") || decodeIndexErr != nil || len(decodedIndexID) != sha256.Size { + return nil, fmt.Errorf("invalid repository index key %q", key) + } + expectedPack := fmt.Sprintf("%s/%s/%s.pack", repositoryPackPrefix, indexID[:2], indexID) reader, err := provider.Download(ctx, key) if err != nil { return nil, fmt.Errorf("download repository index %s: %w", key, err) } - data, readErr := io.ReadAll(io.LimitReader(reader, 64<<20)) + data, readErr := io.ReadAll(io.LimitReader(reader, repositoryMaxIndexSize+1)) closeErr := reader.Close() if readErr != nil { return nil, fmt.Errorf("read repository index %s: %w", key, readErr) @@ -907,27 +944,45 @@ func (s *RepositoryStore) readIndexSegment(ctx context.Context, provider storage if closeErr != nil { return nil, fmt.Errorf("close repository index %s: %w", key, closeErr) } + if int64(len(data)) > repositoryMaxIndexSize { + return nil, fmt.Errorf("repository index %s exceeds %d bytes", key, repositoryMaxIndexSize) + } var segment repositoryIndexSegment if err := json.Unmarshal(data, &segment); err != nil { return nil, fmt.Errorf("decode repository index %s: %w", key, err) } - if segment.Version != repositoryFormatVersion || strings.TrimSpace(segment.Pack) == "" { + if segment.Version != repositoryFormatVersion || segment.Pack != expectedPack || len(segment.Chunks) == 0 { return nil, fmt.Errorf("unsupported repository index %s", key) } + packLimit := s.packSize + if packLimit <= 0 { + packLimit = repositoryDefaultPack + } + if packLimit < repositoryMaxEncoded { + packLimit = repositoryMaxEncoded + } for chunkID, location := range segment.Chunks { - if chunkID == "" || location.Offset < 0 || location.Length <= 0 || location.PlainSize < 0 { + expectedPrefix := "p-" + if location.Encrypted { + expectedPrefix = "e-" + } + encodedID := strings.TrimPrefix(chunkID, expectedPrefix) + decodedID, decodeErr := hex.DecodeString(encodedID) + validPrefix := encodedID != chunkID + validCompression := location.Compression == "none" || location.Compression == "gzip" || location.Compression == "zstd" + if len(decodedID) != sha256.Size || decodeErr != nil || encodedID != strings.ToLower(encodedID) || !validPrefix || !validCompression || location.Pack != expectedPack || location.Offset < 0 || location.Length <= 0 || location.Length > repositoryMaxEncoded || location.Length > packLimit || location.PlainSize <= 0 || location.PlainSize > repositoryChunkMax || location.Offset > packLimit-location.Length { return nil, fmt.Errorf("invalid chunk location in repository index %s", key) } } return &segment, nil } -func (s *RepositoryStore) loadSnapshot(ctx context.Context, provider storage.StorageProvider, key string) (*repositorySnapshot, []byte, error) { +func (s *RepositoryStore) loadSnapshot(ctx context.Context, provider storage.StorageProvider, key, expectedChecksum string) (*repositorySnapshot, []byte, error) { reader, err := provider.Download(ctx, key) if err != nil { return nil, nil, fmt.Errorf("download repository snapshot %s: %w", key, err) } - data, readErr := io.ReadAll(io.LimitReader(reader, 256<<20)) + data, readErr := io.ReadAll(io.LimitReader(reader, repositoryMaxSnapshot+1)) closeErr := reader.Close() if readErr != nil { return nil, nil, fmt.Errorf("read repository snapshot %s: %w", key, readErr) @@ -935,6 +990,19 @@ func (s *RepositoryStore) loadSnapshot(ctx context.Context, provider storage.Sto if closeErr != nil { return nil, nil, fmt.Errorf("close repository snapshot %s: %w", key, closeErr) } + if int64(len(data)) > repositoryMaxSnapshot { + return nil, nil, fmt.Errorf("repository snapshot %s exceeds %d bytes", key, repositoryMaxSnapshot) + } + if expected := strings.TrimSpace(expectedChecksum); expected != "" { + expectedBytes, decodeErr := hex.DecodeString(expected) + if decodeErr != nil || len(expectedBytes) != sha256.Size { + return nil, nil, fmt.Errorf("repository snapshot checksum is invalid") + } + digest := sha256.Sum256(data) + if !strings.EqualFold(hex.EncodeToString(digest[:]), expected) { + return nil, nil, fmt.Errorf("repository snapshot checksum mismatch") + } + } var envelope repositorySnapshotEnvelope if err := json.Unmarshal(data, &envelope); err != nil { return nil, nil, fmt.Errorf("decode repository snapshot envelope %s: %w", key, err) @@ -960,9 +1028,79 @@ func (s *RepositoryStore) loadSnapshot(ctx context.Context, provider storage.Sto if snapshot.Version != repositoryFormatVersion { return nil, nil, fmt.Errorf("unsupported repository snapshot payload version %d", snapshot.Version) } + if snapshot.Encrypted != envelope.Encrypted { + return nil, nil, fmt.Errorf("repository snapshot encryption metadata mismatch") + } + if err := s.validateSnapshot(&snapshot); err != nil { + return nil, nil, err + } return &snapshot, data, nil } +func (*RepositoryStore) validateSnapshot(snapshot *repositorySnapshot) error { + if snapshot == nil { + return fmt.Errorf("repository snapshot is required") + } + if snapshot.Version != repositoryFormatVersion { + return fmt.Errorf("unsupported repository snapshot payload version %d", snapshot.Version) + } + if snapshot.Compression != "none" && snapshot.Compression != "gzip" && snapshot.Compression != "zstd" { + return fmt.Errorf("unsupported repository snapshot compression %q", snapshot.Compression) + } + seenPaths := make(map[string]struct{}, len(snapshot.Entries)) + symlinkPaths := make(map[string]struct{}) + for _, entry := range snapshot.Entries { + if entry.Path == "." || !fs.ValidPath(entry.Path) || path.Clean(entry.Path) != entry.Path || strings.Contains(entry.Path, "\\") || entry.Mode & ^uint32(0o777) != 0 || entry.Size < 0 { + return fmt.Errorf("invalid repository snapshot entry %q", entry.Path) + } + if _, exists := seenPaths[entry.Path]; exists { + return fmt.Errorf("duplicate repository snapshot entry %q", entry.Path) + } + seenPaths[entry.Path] = struct{}{} + switch entry.Kind { + case "directory": + if len(entry.Chunks) != 0 || entry.LinkTarget != "" { + return fmt.Errorf("invalid repository directory entry %q", entry.Path) + } + case "symlink": + if entry.LinkTarget == "" || len(entry.Chunks) != 0 { + return fmt.Errorf("invalid repository symlink entry %q", entry.Path) + } + normalizedTarget := strings.ReplaceAll(entry.LinkTarget, "\\", "/") + resolvedTarget := path.Clean(path.Join(path.Dir(entry.Path), normalizedTarget)) + if strings.ContainsRune(entry.LinkTarget, 0) || path.IsAbs(normalizedTarget) || filepath.IsAbs(entry.LinkTarget) || filepath.VolumeName(entry.LinkTarget) != "" || resolvedTarget == ".." || strings.HasPrefix(resolvedTarget, "../") { + return fmt.Errorf("repository symlink %q escapes the restore root", entry.Path) + } + symlinkPaths[entry.Path] = struct{}{} + case "file": + if entry.LinkTarget != "" || (entry.Size == 0 && len(entry.Chunks) != 0) || (entry.Size > 0 && len(entry.Chunks) == 0) { + return fmt.Errorf("invalid repository file entry %q", entry.Path) + } + expectedPrefix := "p-" + if snapshot.Encrypted { + expectedPrefix = "e-" + } + for _, chunkID := range entry.Chunks { + encodedID := strings.TrimPrefix(chunkID, expectedPrefix) + decodedID, decodeErr := hex.DecodeString(encodedID) + if decodeErr != nil || len(decodedID) != sha256.Size || encodedID == chunkID || encodedID != strings.ToLower(encodedID) { + return fmt.Errorf("invalid chunk id in repository entry %q", entry.Path) + } + } + default: + return fmt.Errorf("unsupported repository entry kind %q", entry.Kind) + } + } + for entryPath := range seenPaths { + for parent := path.Dir(entryPath); parent != "."; parent = path.Dir(parent) { + if _, crossesSymlink := symlinkPaths[parent]; crossesSymlink { + return fmt.Errorf("repository entry %q crosses symlink %q", entryPath, parent) + } + } + } + return nil +} + func (s *RepositoryStore) encodeSnapshot(snapshot repositorySnapshot) ([]byte, error) { payload, err := json.Marshal(snapshot) if err != nil { @@ -1005,7 +1143,7 @@ func (s *RepositoryStore) readChunk(ctx context.Context, provider storage.Storag if err != nil { return nil, fmt.Errorf("read repository pack %s: %w", location.Pack, err) } - encoded := make([]byte, location.Length) + encoded := make([]byte, int(location.Length)) _, readErr := io.ReadFull(reader, encoded) closeErr := reader.Close() if readErr != nil { @@ -1079,7 +1217,7 @@ func (s *RepositoryStore) decodeChunk(encoded []byte, location repositoryChunkLo if err != nil { return nil, fmt.Errorf("open repository gzip chunk: %w", err) } - raw, readErr := io.ReadAll(reader) + raw, readErr := io.ReadAll(io.LimitReader(reader, location.PlainSize+1)) closeErr := reader.Close() if readErr != nil { return nil, fmt.Errorf("decompress repository gzip chunk: %w", readErr) @@ -1089,14 +1227,14 @@ func (s *RepositoryStore) decodeChunk(encoded []byte, location repositoryChunkLo } return raw, nil case "zstd": - reader, err := zstd.NewReader(nil) + reader, err := zstd.NewReader(bytes.NewReader(payload), zstd.WithDecoderMaxMemory(uint64(repositoryChunkMax+(1<<20)))) if err != nil { return nil, fmt.Errorf("create repository zstd decoder: %w", err) } - raw, err := reader.DecodeAll(payload, nil) + raw, readErr := io.ReadAll(io.LimitReader(reader, location.PlainSize+1)) reader.Close() - if err != nil { - return nil, fmt.Errorf("decompress repository zstd chunk: %w", err) + if readErr != nil { + return nil, fmt.Errorf("decompress repository zstd chunk: %w", readErr) } return raw, nil default: @@ -1171,33 +1309,40 @@ func (s *RepositoryStore) normalizeCompression(value string) (string, error) { } } -func (s *RepositoryStore) ensureSafeParent(root, target string) error { - parent := filepath.Dir(target) - relative, err := filepath.Rel(root, parent) - if err != nil || relative == ".." || strings.HasPrefix(relative, ".."+string(filepath.Separator)) { - return fmt.Errorf("restore target escapes repository root: %s", target) +func (s *RepositoryStore) openRestoreRoot(target string) (*os.Root, error) { + cleanTarget := filepath.Clean(target) + if !filepath.IsAbs(cleanTarget) { + return nil, fmt.Errorf("restore target must be absolute: %s", target) } - current := root - if relative == "." { - return nil + volume := filepath.VolumeName(cleanTarget) + if strings.Contains(volume, "..") || strings.ContainsRune(volume, 0) { + return nil, fmt.Errorf("invalid restore target volume: %s", target) } - for _, segment := range strings.Split(relative, string(filepath.Separator)) { - current = filepath.Join(current, segment) - info, statErr := os.Lstat(current) - if errors.Is(statErr, os.ErrNotExist) { - if err := os.Mkdir(current, 0o755); err != nil && !errors.Is(err, os.ErrExist) { - return fmt.Errorf("create restore parent %s: %w", current, err) - } - continue - } - if statErr != nil { - return fmt.Errorf("inspect restore parent %s: %w", current, statErr) - } - if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() { - return fmt.Errorf("unsafe restore parent %s", current) - } + volumeRootPath := string(filepath.Separator) + relativeTarget := strings.TrimLeft(cleanTarget, string(filepath.Separator)) + if volume != "" { + volumeRootPath = volume + string(filepath.Separator) + relativeTarget = strings.TrimLeft(strings.TrimPrefix(cleanTarget, volume), string(filepath.Separator)) } - return nil + if relativeTarget != "" && !filepath.IsLocal(relativeTarget) { + return nil, fmt.Errorf("restore target is not local to its volume: %s", target) + } + volumeRoot, err := os.OpenRoot(volumeRootPath) + if err != nil { + return nil, fmt.Errorf("open restore volume: %w", err) + } + if relativeTarget == "" { + return volumeRoot, nil + } + if err := volumeRoot.MkdirAll(relativeTarget, 0o755); err != nil { + return nil, errors.Join(fmt.Errorf("create repository restore root: %w", err), volumeRoot.Close()) + } + restoreRoot, openErr := volumeRoot.OpenRoot(relativeTarget) + closeErr := volumeRoot.Close() + if openErr != nil || closeErr != nil { + return nil, errors.Join(openErr, closeErr) + } + return restoreRoot, nil } func compactPaths(items []string) []string { diff --git a/server/internal/backup/repository_test.go b/server/internal/backup/repository_test.go index c17d670..282e39d 100644 --- a/server/internal/backup/repository_test.go +++ b/server/internal/backup/repository_test.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "crypto/sha256" + "encoding/json" "fmt" "io" "math/rand" @@ -149,7 +150,13 @@ func TestRepositoryRoundTripDedupAndPrune(t *testing.T) { restoreRoot := filepath.Join(tempDir, "restore") restoreTask := task restoreTask.RestoreTargetPath = restoreRoot - if err := store.Restore(ctx, provider, secondKey, restoreTask, NopLogWriter{}); err != nil { + if err := store.Restore(ctx, provider, secondKey, strings.Repeat("0", sha256.Size*2), restoreTask, NopLogWriter{}); err == nil { + t.Fatal("restore accepted a mismatched snapshot checksum") + } + if err := store.Restore(ctx, provider, secondKey, "", restoreTask, NopLogWriter{}); err == nil { + t.Fatal("restore accepted a missing snapshot checksum") + } + if err := store.Restore(ctx, provider, secondKey, secondResult.Checksum, restoreTask, NopLogWriter{}); err != nil { t.Fatalf("restore snapshot: %v", err) } restored, err := os.ReadFile(filepath.Join(restoreRoot, filepath.Base(sourceDir), "primary.bin")) @@ -184,6 +191,153 @@ func TestRepositoryRoundTripDedupAndPrune(t *testing.T) { } } +func TestRepositoryRestoreRejectsUnsafeSnapshotMetadata(t *testing.T) { + cases := []struct { + name string + entries []repositoryEntry + }{ + { + name: "path traversal", + entries: []repositoryEntry{{Path: "../escape", Kind: "directory", Mode: 0o755}}, + }, + { + name: "entry below symlink", + entries: []repositoryEntry{ + {Path: "link", Kind: "symlink", Mode: 0o777, LinkTarget: "inside"}, + {Path: "link/payload", Kind: "file", Mode: 0o600, Size: 1, Chunks: []string{"p-" + strings.Repeat("0", sha256.Size*2)}}, + }, + }, + { + name: "escaping symlink target", + entries: []repositoryEntry{{Path: "escape", Kind: "symlink", Mode: 0o777, LinkTarget: "../outside"}}, + }, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + ctx := context.Background() + store := NewRepositoryStore(nil) + provider := newMemoryRepositoryProvider() + snapshot := repositorySnapshot{ + Version: repositoryFormatVersion, + TaskID: 1, + CreatedAt: time.Now().UTC(), + Compression: "none", + Entries: tc.entries, + } + data, err := store.encodeSnapshot(snapshot) + if err != nil { + t.Fatalf("encodeSnapshot returned error: %v", err) + } + key := store.SnapshotKey(1, 1, snapshot.CreatedAt) + if err := provider.Upload(ctx, key, bytes.NewReader(data), int64(len(data)), nil); err != nil { + t.Fatalf("Upload snapshot returned error: %v", err) + } + digest := sha256.Sum256(data) + task := TaskSpec{SourcePath: filepath.Join(t.TempDir(), "source"), RestoreTargetPath: filepath.Join(t.TempDir(), "restore")} + if err := store.Restore(ctx, provider, key, fmt.Sprintf("%x", digest[:]), task, NopLogWriter{}); err == nil { + t.Fatal("restore accepted unsafe snapshot metadata") + } + }) + } +} + +func TestRepositoryRestorePreservesDirectoryWhenSnapshotContainsSymlink(t *testing.T) { + ctx := context.Background() + store := NewRepositoryStore(nil) + provider := newMemoryRepositoryProvider() + snapshot := repositorySnapshot{ + Version: repositoryFormatVersion, + TaskID: 1, + CreatedAt: time.Now().UTC(), + Compression: "none", + Entries: []repositoryEntry{{Path: "link", Kind: "symlink", Mode: 0o777, LinkTarget: "inside"}}, + } + data, err := store.encodeSnapshot(snapshot) + if err != nil { + t.Fatalf("encodeSnapshot returned error: %v", err) + } + key := store.SnapshotKey(1, 1, snapshot.CreatedAt) + if err := provider.Upload(ctx, key, bytes.NewReader(data), int64(len(data)), nil); err != nil { + t.Fatalf("Upload snapshot returned error: %v", err) + } + digest := sha256.Sum256(data) + restoreRoot := filepath.Join(t.TempDir(), "restore") + markerPath := filepath.Join(restoreRoot, "link", "keep.txt") + if err := os.MkdirAll(filepath.Dir(markerPath), 0o755); err != nil { + t.Fatalf("MkdirAll marker parent: %v", err) + } + if err := os.WriteFile(markerPath, []byte("keep"), 0o600); err != nil { + t.Fatalf("WriteFile marker: %v", err) + } + task := TaskSpec{SourcePath: filepath.Join(t.TempDir(), "source"), RestoreTargetPath: restoreRoot} + if err := store.Restore(ctx, provider, key, fmt.Sprintf("%x", digest[:]), task, NopLogWriter{}); err == nil { + t.Fatal("restore replaced an existing directory with a symlink") + } + if data, err := os.ReadFile(markerPath); err != nil || string(data) != "keep" { + t.Fatalf("existing directory content changed: data=%q err=%v", data, err) + } +} + +func TestRepositoryRejectsOversizedChunkLocation(t *testing.T) { + ctx := context.Background() + store := NewRepositoryStore(nil) + provider := newMemoryRepositoryProvider() + chunkID := "p-" + strings.Repeat("0", sha256.Size*2) + packID := strings.Repeat("a", sha256.Size*2) + segment := repositoryIndexSegment{ + Version: repositoryFormatVersion, + Pack: fmt.Sprintf("%s/%s/%s.pack", repositoryPackPrefix, packID[:2], packID), + Chunks: map[string]repositoryChunkLocation{ + chunkID: {Pack: fmt.Sprintf("%s/%s/%s.pack", repositoryPackPrefix, packID[:2], packID), Offset: 0, Length: repositoryMaxEncoded + 1, PlainSize: 1, Compression: "none"}, + }, + } + data, err := json.Marshal(segment) + if err != nil { + t.Fatalf("Marshal index returned error: %v", err) + } + indexKey := fmt.Sprintf("%s/%s.json", repositoryIndexPrefix, packID) + if err := provider.Upload(ctx, indexKey, bytes.NewReader(data), int64(len(data)), nil); err != nil { + t.Fatalf("Upload index returned error: %v", err) + } + if _, err := store.loadIndex(ctx, provider); err == nil { + t.Fatal("loadIndex accepted an oversized encoded chunk") + } +} + +func TestRepositoryRejectsIndexPackMismatch(t *testing.T) { + ctx := context.Background() + store := NewRepositoryStore(nil) + provider := newMemoryRepositoryProvider() + chunkID := "p-" + strings.Repeat("0", sha256.Size*2) + indexID := strings.Repeat("a", sha256.Size*2) + otherPackID := strings.Repeat("b", sha256.Size*2) + expectedPack := fmt.Sprintf("%s/%s/%s.pack", repositoryPackPrefix, indexID[:2], indexID) + segment := repositoryIndexSegment{ + Version: repositoryFormatVersion, + Pack: expectedPack, + Chunks: map[string]repositoryChunkLocation{ + chunkID: { + Pack: fmt.Sprintf("%s/%s/%s.pack", repositoryPackPrefix, otherPackID[:2], otherPackID), + Offset: 0, + Length: 1, + PlainSize: 1, + Compression: "none", + }, + }, + } + data, err := json.Marshal(segment) + if err != nil { + t.Fatalf("Marshal index returned error: %v", err) + } + indexKey := fmt.Sprintf("%s/%s.json", repositoryIndexPrefix, indexID) + if err := provider.Upload(ctx, indexKey, bytes.NewReader(data), int64(len(data)), nil); err != nil { + t.Fatalf("Upload index returned error: %v", err) + } + if _, err := store.loadIndex(ctx, provider); err == nil { + t.Fatal("loadIndex accepted a chunk location pointing to a different pack") + } +} + type memoryRepositoryProvider struct { mu sync.RWMutex objects map[string][]byte diff --git a/server/internal/service/backup_execution_service.go b/server/internal/service/backup_execution_service.go index 3c90055..f880682 100644 --- a/server/internal/service/backup_execution_service.go +++ b/server/internal/service/backup_execution_service.go @@ -237,7 +237,7 @@ func (s *BackupExecutionService) DownloadRecord(ctx context.Context, recordID ui exportName := fmt.Sprintf("backupx-record-%d.tar", record.ID) exportPath := filepath.Join(tempDir, exportName) store := backup.NewRepositoryStore(s.cipher.Key()) - if err := store.ExportTar(ctx, provider, record.StoragePath, exportPath); err != nil { + if err := store.ExportTar(ctx, provider, record.StoragePath, record.Checksum, exportPath); err != nil { cleanupErr := os.RemoveAll(tempDir) return nil, apperror.Internal("BACKUP_RECORD_DOWNLOAD_FAILED", "无法从 CDC 仓库导出归档", errors.Join(err, cleanupErr)) } @@ -276,7 +276,7 @@ func (s *BackupExecutionService) RestoreRecord(ctx context.Context, recordID uin if specErr != nil { return specErr } - if err := backup.NewRepositoryStore(s.cipher.Key()).Restore(ctx, provider, record.StoragePath, spec, backup.NopLogWriter{}); err != nil { + if err := backup.NewRepositoryStore(s.cipher.Key()).Restore(ctx, provider, record.StoragePath, record.Checksum, spec, backup.NopLogWriter{}); err != nil { return apperror.Internal("BACKUP_RECORD_RESTORE_FAILED", "从 CDC 仓库恢复备份失败", err) } return nil diff --git a/server/internal/service/restore_service.go b/server/internal/service/restore_service.go index e5802fe..5d67531 100644 --- a/server/internal/service/restore_service.go +++ b/server/internal/service/restore_service.go @@ -324,7 +324,7 @@ func (s *RestoreService) restoreArtifact(ctx context.Context, record *model.Back } if record.BackupKind == model.BackupKindRepository { logger.Infof("读取 CDC 仓库快照:%s", record.StoragePath) - if err := backup.NewRepositoryStore(s.cipher.Key()).Restore(ctx, provider, record.StoragePath, spec, logger); err != nil { + if err := backup.NewRepositoryStore(s.cipher.Key()).Restore(ctx, provider, record.StoragePath, record.Checksum, spec, logger); err != nil { return fmt.Errorf("恢复 CDC 仓库快照失败:%w", err) } return nil