mirror of
https://github.com/Awuqing/BackupX.git
synced 2026-09-05 15:37:03 +08:00
fix(backup): 加固 CDC 仓库恢复安全
This commit is contained in:
@@ -16,6 +16,7 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
|
"io/fs"
|
||||||
"os"
|
"os"
|
||||||
"path"
|
"path"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
@@ -34,6 +35,9 @@ const (
|
|||||||
repositoryPackPrefix = repositoryRoot + "/packs"
|
repositoryPackPrefix = repositoryRoot + "/packs"
|
||||||
repositorySnapshotRoot = repositoryRoot + "/snapshots"
|
repositorySnapshotRoot = repositoryRoot + "/snapshots"
|
||||||
repositoryDefaultPack = int64(32 << 20)
|
repositoryDefaultPack = int64(32 << 20)
|
||||||
|
repositoryMaxIndexSize = int64(64 << 20)
|
||||||
|
repositoryMaxSnapshot = int64(256 << 20)
|
||||||
|
repositoryMaxEncoded = int64(repositoryChunkMax + (256 << 10))
|
||||||
)
|
)
|
||||||
|
|
||||||
type RepositoryStore struct {
|
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)
|
return nil, fmt.Errorf("close repository chunk spool: %w", err)
|
||||||
}
|
}
|
||||||
spoolClosed = true
|
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))
|
writer.WriteLine(fmt.Sprintf("CDC 扫描完成:%d 个条目,逻辑数据 %d bytes,任务内唯一数据 %d bytes", len(plan.snapshot.Entries), plan.LogicalSize, plan.UniqueSize))
|
||||||
return plan, nil
|
return plan, nil
|
||||||
}
|
}
|
||||||
@@ -399,11 +406,14 @@ func (s *RepositoryStore) EstimateUploadSize(ctx context.Context, provider stora
|
|||||||
return estimate, nil
|
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 {
|
if writer == nil {
|
||||||
writer = NopLogWriter{}
|
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 {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -415,19 +425,25 @@ func (s *RepositoryStore) Restore(ctx context.Context, provider storage.StorageP
|
|||||||
if len(task.SourcePaths) > 0 {
|
if len(task.SourcePaths) > 0 {
|
||||||
restoreSource = strings.TrimSpace(task.SourcePaths[0])
|
restoreSource = strings.TrimSpace(task.SourcePaths[0])
|
||||||
}
|
}
|
||||||
if restoreSource == "" && len(snapshot.SourcePaths) > 0 {
|
targetRoot := strings.TrimSpace(task.RestoreTargetPath)
|
||||||
restoreSource = snapshot.SourcePaths[0]
|
if targetRoot == "" {
|
||||||
}
|
if restoreSource == "" {
|
||||||
targetRoot := filepath.Dir(filepath.Clean(restoreSource))
|
return fmt.Errorf("repository restore source path is required when no restore target is provided")
|
||||||
if strings.TrimSpace(task.RestoreTargetPath) != "" {
|
}
|
||||||
targetRoot = filepath.Clean(task.RestoreTargetPath)
|
targetRoot = filepath.Dir(filepath.Clean(restoreSource))
|
||||||
|
} else {
|
||||||
|
targetRoot = filepath.Clean(targetRoot)
|
||||||
}
|
}
|
||||||
if !filepath.IsAbs(targetRoot) {
|
if !filepath.IsAbs(targetRoot) {
|
||||||
return fmt.Errorf("repository restore target must be absolute: %s", targetRoot)
|
return fmt.Errorf("repository restore target must be absolute: %s", targetRoot)
|
||||||
}
|
}
|
||||||
if err := os.MkdirAll(targetRoot, 0o755); err != nil {
|
restoreRoot, err := s.openRestoreRoot(targetRoot)
|
||||||
return fmt.Errorf("create repository restore root: %w", err)
|
if err != nil {
|
||||||
|
return fmt.Errorf("open repository restore root: %w", err)
|
||||||
}
|
}
|
||||||
|
defer func() {
|
||||||
|
err = errors.Join(err, restoreRoot.Close())
|
||||||
|
}()
|
||||||
|
|
||||||
restored := 0
|
restored := 0
|
||||||
directories := make([]repositoryEntry, 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) {
|
if len(task.SelectedPaths) > 0 && !pathSelected(entry.Path, task.SelectedPaths) {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
targetPath, ok := resolveWithinParent(targetRoot, entry.Path)
|
entryPath, localizeErr := filepath.Localize(entry.Path)
|
||||||
if !ok {
|
if localizeErr != nil || !filepath.IsLocal(entryPath) {
|
||||||
return fmt.Errorf("unsafe repository restore path %q", entry.Path)
|
return fmt.Errorf("unsafe repository restore path %q", entry.Path)
|
||||||
}
|
}
|
||||||
if err := s.ensureSafeParent(targetRoot, targetPath); err != nil {
|
if parent := filepath.Dir(entryPath); parent != "." {
|
||||||
return err
|
if err := restoreRoot.MkdirAll(parent, 0o755); err != nil {
|
||||||
|
return fmt.Errorf("create restore parent for %s: %w", entry.Path, err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
switch entry.Kind {
|
switch entry.Kind {
|
||||||
case "directory":
|
case "directory":
|
||||||
if info, statErr := os.Lstat(targetPath); statErr == nil {
|
if info, statErr := restoreRoot.Lstat(entryPath); statErr == nil {
|
||||||
if info.Mode()&os.ModeSymlink != 0 {
|
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) {
|
} 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 {
|
if err := restoreRoot.MkdirAll(entryPath, os.FileMode(entry.Mode)); err != nil {
|
||||||
return fmt.Errorf("create restore directory %s: %w", targetPath, err)
|
return fmt.Errorf("create restore directory %s: %w", entry.Path, err)
|
||||||
}
|
}
|
||||||
directories = append(directories, entry)
|
directories = append(directories, entry)
|
||||||
case "symlink":
|
case "symlink":
|
||||||
if err := os.RemoveAll(targetPath); err != nil {
|
resolvedLinkTarget := path.Clean(path.Join(path.Dir(entry.Path), strings.ReplaceAll(entry.LinkTarget, "\\", "/")))
|
||||||
return fmt.Errorf("replace restore symlink %s: %w", targetPath, err)
|
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 {
|
linkTarget, relativeTargetErr := filepath.Rel(filepath.Dir(entryPath), localizedLinkTarget)
|
||||||
return fmt.Errorf("create restore symlink %s: %w", targetPath, err)
|
if relativeTargetErr != nil {
|
||||||
|
return fmt.Errorf("resolve restore symlink target %s: %w", entry.Path, relativeTargetErr)
|
||||||
}
|
}
|
||||||
case "file":
|
if info, statErr := restoreRoot.Lstat(entryPath); statErr == nil {
|
||||||
if info, statErr := os.Lstat(targetPath); statErr == nil {
|
if info.IsDir() {
|
||||||
if info.Mode()&os.ModeSymlink != 0 {
|
return fmt.Errorf("refuse to replace restore directory with symlink: %s", entry.Path)
|
||||||
return fmt.Errorf("restore file would overwrite existing symlink: %s", targetPath)
|
}
|
||||||
|
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) {
|
} 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 {
|
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
|
var written int64
|
||||||
for _, chunkID := range entry.Chunks {
|
for _, chunkID := range entry.Chunks {
|
||||||
@@ -486,54 +520,55 @@ func (s *RepositoryStore) Restore(ctx context.Context, provider storage.StorageP
|
|||||||
count, writeErr := file.Write(raw)
|
count, writeErr := file.Write(raw)
|
||||||
written += int64(count)
|
written += int64(count)
|
||||||
if writeErr != nil {
|
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) {
|
if count != len(raw) {
|
||||||
return errors.Join(io.ErrShortWrite, file.Close())
|
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 {
|
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 {
|
if err := file.Chmod(os.FileMode(entry.Mode)); err != nil {
|
||||||
return fmt.Errorf("restore mode for %s: %w", targetPath, err)
|
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:
|
default:
|
||||||
return fmt.Errorf("unsupported repository entry kind %q", entry.Kind)
|
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++
|
restored++
|
||||||
}
|
}
|
||||||
// Children modify their parent directory timestamps, so directory metadata
|
// Children modify their parent directory timestamps, so directory metadata
|
||||||
// is restored only after every selected file has been written.
|
// is restored only after every selected file has been written.
|
||||||
for index := len(directories) - 1; index >= 0; index-- {
|
for index := len(directories) - 1; index >= 0; index-- {
|
||||||
entry := directories[index]
|
entry := directories[index]
|
||||||
targetPath, ok := resolveWithinParent(targetRoot, entry.Path)
|
entryPath, localizeErr := filepath.Localize(entry.Path)
|
||||||
if !ok {
|
if localizeErr != nil || !filepath.IsLocal(entryPath) {
|
||||||
return fmt.Errorf("unsafe repository restore path %q", entry.Path)
|
return fmt.Errorf("unsafe repository restore path %q", entry.Path)
|
||||||
}
|
}
|
||||||
if err := os.Chmod(targetPath, os.FileMode(entry.Mode)); err != nil {
|
if err := restoreRoot.Chmod(entryPath, os.FileMode(entry.Mode)); err != nil {
|
||||||
return fmt.Errorf("restore mode for %s: %w", targetPath, err)
|
return fmt.Errorf("restore mode for %s: %w", entry.Path, err)
|
||||||
}
|
}
|
||||||
modTime := time.Unix(0, entry.ModTime)
|
modTime := time.Unix(0, entry.ModTime)
|
||||||
if err := os.Chtimes(targetPath, modTime, modTime); err != nil {
|
if err := restoreRoot.Chtimes(entryPath, modTime, modTime); err != nil {
|
||||||
return fmt.Errorf("restore timestamp for %s: %w", targetPath, err)
|
return fmt.Errorf("restore timestamp for %s: %w", entry.Path, err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
writer.WriteLine(fmt.Sprintf("CDC 仓库恢复完成:%d 个条目", restored))
|
writer.WriteLine(fmt.Sprintf("CDC 仓库恢复完成:%d 个条目", restored))
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *RepositoryStore) ExportTar(ctx context.Context, provider storage.StorageProvider, snapshotKey, destination string) (err error) {
|
func (s *RepositoryStore) ExportTar(ctx context.Context, provider storage.StorageProvider, snapshotKey, expectedChecksum, destination string) (err error) {
|
||||||
snapshot, _, err := s.loadSnapshot(ctx, provider, snapshotKey)
|
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 {
|
if err != nil {
|
||||||
return err
|
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) {
|
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 {
|
if err != nil {
|
||||||
return nil, err
|
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)
|
locations, err := s.loadIndex(ctx, provider)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
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 })
|
sort.Slice(snapshots, func(i, j int) bool { return snapshots[i].Key < snapshots[j].Key })
|
||||||
for _, object := range snapshots {
|
for _, object := range snapshots {
|
||||||
snapshot, _, loadErr := s.loadSnapshot(ctx, provider, object.Key)
|
snapshot, _, loadErr := s.loadSnapshot(ctx, provider, object.Key, "")
|
||||||
if loadErr != nil {
|
if loadErr != nil {
|
||||||
return nil, fmt.Errorf("refuse to prune with unreadable snapshot %s: %w", object.Key, loadErr)
|
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
|
return nil, readErr
|
||||||
}
|
}
|
||||||
for chunkID, location := range segment.Chunks {
|
for chunkID, location := range segment.Chunks {
|
||||||
if location.Pack == "" {
|
|
||||||
location.Pack = segment.Pack
|
|
||||||
}
|
|
||||||
if _, exists := locations[chunkID]; !exists {
|
if _, exists := locations[chunkID]; !exists {
|
||||||
locations[chunkID] = location
|
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) {
|
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)
|
reader, err := provider.Download(ctx, key)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("download repository index %s: %w", key, err)
|
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()
|
closeErr := reader.Close()
|
||||||
if readErr != nil {
|
if readErr != nil {
|
||||||
return nil, fmt.Errorf("read repository index %s: %w", key, readErr)
|
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 {
|
if closeErr != nil {
|
||||||
return nil, fmt.Errorf("close repository index %s: %w", key, closeErr)
|
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
|
var segment repositoryIndexSegment
|
||||||
if err := json.Unmarshal(data, &segment); err != nil {
|
if err := json.Unmarshal(data, &segment); err != nil {
|
||||||
return nil, fmt.Errorf("decode repository index %s: %w", key, err)
|
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)
|
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 {
|
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 nil, fmt.Errorf("invalid chunk location in repository index %s", key)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return &segment, nil
|
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)
|
reader, err := provider.Download(ctx, key)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, nil, fmt.Errorf("download repository snapshot %s: %w", key, err)
|
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()
|
closeErr := reader.Close()
|
||||||
if readErr != nil {
|
if readErr != nil {
|
||||||
return nil, nil, fmt.Errorf("read repository snapshot %s: %w", key, readErr)
|
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 {
|
if closeErr != nil {
|
||||||
return nil, nil, fmt.Errorf("close repository snapshot %s: %w", key, closeErr)
|
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
|
var envelope repositorySnapshotEnvelope
|
||||||
if err := json.Unmarshal(data, &envelope); err != nil {
|
if err := json.Unmarshal(data, &envelope); err != nil {
|
||||||
return nil, nil, fmt.Errorf("decode repository snapshot envelope %s: %w", key, err)
|
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 {
|
if snapshot.Version != repositoryFormatVersion {
|
||||||
return nil, nil, fmt.Errorf("unsupported repository snapshot payload version %d", snapshot.Version)
|
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
|
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) {
|
func (s *RepositoryStore) encodeSnapshot(snapshot repositorySnapshot) ([]byte, error) {
|
||||||
payload, err := json.Marshal(snapshot)
|
payload, err := json.Marshal(snapshot)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -1005,7 +1143,7 @@ func (s *RepositoryStore) readChunk(ctx context.Context, provider storage.Storag
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("read repository pack %s: %w", location.Pack, err)
|
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)
|
_, readErr := io.ReadFull(reader, encoded)
|
||||||
closeErr := reader.Close()
|
closeErr := reader.Close()
|
||||||
if readErr != nil {
|
if readErr != nil {
|
||||||
@@ -1079,7 +1217,7 @@ func (s *RepositoryStore) decodeChunk(encoded []byte, location repositoryChunkLo
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("open repository gzip chunk: %w", err)
|
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()
|
closeErr := reader.Close()
|
||||||
if readErr != nil {
|
if readErr != nil {
|
||||||
return nil, fmt.Errorf("decompress repository gzip chunk: %w", readErr)
|
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
|
return raw, nil
|
||||||
case "zstd":
|
case "zstd":
|
||||||
reader, err := zstd.NewReader(nil)
|
reader, err := zstd.NewReader(bytes.NewReader(payload), zstd.WithDecoderMaxMemory(uint64(repositoryChunkMax+(1<<20))))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("create repository zstd decoder: %w", err)
|
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()
|
reader.Close()
|
||||||
if err != nil {
|
if readErr != nil {
|
||||||
return nil, fmt.Errorf("decompress repository zstd chunk: %w", err)
|
return nil, fmt.Errorf("decompress repository zstd chunk: %w", readErr)
|
||||||
}
|
}
|
||||||
return raw, nil
|
return raw, nil
|
||||||
default:
|
default:
|
||||||
@@ -1171,33 +1309,40 @@ func (s *RepositoryStore) normalizeCompression(value string) (string, error) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *RepositoryStore) ensureSafeParent(root, target string) error {
|
func (s *RepositoryStore) openRestoreRoot(target string) (*os.Root, error) {
|
||||||
parent := filepath.Dir(target)
|
cleanTarget := filepath.Clean(target)
|
||||||
relative, err := filepath.Rel(root, parent)
|
if !filepath.IsAbs(cleanTarget) {
|
||||||
if err != nil || relative == ".." || strings.HasPrefix(relative, ".."+string(filepath.Separator)) {
|
return nil, fmt.Errorf("restore target must be absolute: %s", target)
|
||||||
return fmt.Errorf("restore target escapes repository root: %s", target)
|
|
||||||
}
|
}
|
||||||
current := root
|
volume := filepath.VolumeName(cleanTarget)
|
||||||
if relative == "." {
|
if strings.Contains(volume, "..") || strings.ContainsRune(volume, 0) {
|
||||||
return nil
|
return nil, fmt.Errorf("invalid restore target volume: %s", target)
|
||||||
}
|
}
|
||||||
for _, segment := range strings.Split(relative, string(filepath.Separator)) {
|
volumeRootPath := string(filepath.Separator)
|
||||||
current = filepath.Join(current, segment)
|
relativeTarget := strings.TrimLeft(cleanTarget, string(filepath.Separator))
|
||||||
info, statErr := os.Lstat(current)
|
if volume != "" {
|
||||||
if errors.Is(statErr, os.ErrNotExist) {
|
volumeRootPath = volume + string(filepath.Separator)
|
||||||
if err := os.Mkdir(current, 0o755); err != nil && !errors.Is(err, os.ErrExist) {
|
relativeTarget = strings.TrimLeft(strings.TrimPrefix(cleanTarget, volume), string(filepath.Separator))
|
||||||
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)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
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 {
|
func compactPaths(items []string) []string {
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"crypto/sha256"
|
"crypto/sha256"
|
||||||
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"math/rand"
|
"math/rand"
|
||||||
@@ -149,7 +150,13 @@ func TestRepositoryRoundTripDedupAndPrune(t *testing.T) {
|
|||||||
restoreRoot := filepath.Join(tempDir, "restore")
|
restoreRoot := filepath.Join(tempDir, "restore")
|
||||||
restoreTask := task
|
restoreTask := task
|
||||||
restoreTask.RestoreTargetPath = restoreRoot
|
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)
|
t.Fatalf("restore snapshot: %v", err)
|
||||||
}
|
}
|
||||||
restored, err := os.ReadFile(filepath.Join(restoreRoot, filepath.Base(sourceDir), "primary.bin"))
|
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 {
|
type memoryRepositoryProvider struct {
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
objects map[string][]byte
|
objects map[string][]byte
|
||||||
|
|||||||
@@ -237,7 +237,7 @@ func (s *BackupExecutionService) DownloadRecord(ctx context.Context, recordID ui
|
|||||||
exportName := fmt.Sprintf("backupx-record-%d.tar", record.ID)
|
exportName := fmt.Sprintf("backupx-record-%d.tar", record.ID)
|
||||||
exportPath := filepath.Join(tempDir, exportName)
|
exportPath := filepath.Join(tempDir, exportName)
|
||||||
store := backup.NewRepositoryStore(s.cipher.Key())
|
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)
|
cleanupErr := os.RemoveAll(tempDir)
|
||||||
return nil, apperror.Internal("BACKUP_RECORD_DOWNLOAD_FAILED", "无法从 CDC 仓库导出归档", errors.Join(err, cleanupErr))
|
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 {
|
if specErr != nil {
|
||||||
return specErr
|
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 apperror.Internal("BACKUP_RECORD_RESTORE_FAILED", "从 CDC 仓库恢复备份失败", err)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -324,7 +324,7 @@ func (s *RestoreService) restoreArtifact(ctx context.Context, record *model.Back
|
|||||||
}
|
}
|
||||||
if record.BackupKind == model.BackupKindRepository {
|
if record.BackupKind == model.BackupKindRepository {
|
||||||
logger.Infof("读取 CDC 仓库快照:%s", record.StoragePath)
|
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 fmt.Errorf("恢复 CDC 仓库快照失败:%w", err)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
Reference in New Issue
Block a user