From acd430836f6d361195e1d08fbc7342a00a33dc9e Mon Sep 17 00:00:00 2001 From: ryan Date: Sat, 13 Jun 2026 15:18:22 +0800 Subject: [PATCH] feat(upload): implement parallel storage migration and integrity check - Parallelize storage migration using `errgroup` with a concurrency limit of 10. - Perform post-copy SHA-256 data integrity validation to prevent silent data corruption. - Add test case verifying migration with both incorrect and correct hashes. --- .../apps/upload/storage_migration_task.go | 111 +++++++++++------ .../upload/storage_migration_task_test.go | 112 ++++++++++++++++++ 2 files changed, 187 insertions(+), 36 deletions(-) diff --git a/internal/apps/upload/storage_migration_task.go b/internal/apps/upload/storage_migration_task.go index 28bbd62d..edb787c5 100644 --- a/internal/apps/upload/storage_migration_task.go +++ b/internal/apps/upload/storage_migration_task.go @@ -5,16 +5,21 @@ package upload import ( "context" + "crypto/sha256" + "encoding/hex" "encoding/json" "errors" "fmt" + "io" "os" "strings" + "sync/atomic" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/storage" "github.com/Rain-kl/Wavelet/internal/task" + "golang.org/x/sync/errgroup" "gorm.io/gorm" ) @@ -187,68 +192,102 @@ func migrateObjects( total int64, ) (int64, error) { const batchSize = 50 + const migrationConcurrency = 10 + const sha256HexLength = 64 var migrated int64 for { if err := ctx.Err(); err != nil { - return migrated, fmt.Errorf("storage migration canceled: %w", err) + return atomic.LoadInt64(&migrated), fmt.Errorf("storage migration canceled: %w", err) } var objects []struct { FilePath string `gorm:"column:file_path"` FileSize int64 `gorm:"column:file_size"` MimeType string `gorm:"column:mime_type"` + Hash string `gorm:"column:hash"` } if err := db.DB(ctx).Model(&model.Upload{}). - Select("file_path, MAX(file_size) AS file_size, MAX(mime_type) AS mime_type"). + Select("file_path, MAX(file_size) AS file_size, MAX(mime_type) AS mime_type, MAX(hash) AS hash"). Where("storage_driver = ? AND status != ?", sourceDriver, model.UploadStatusDeleted). Group("file_path"). Order("file_path ASC"). Limit(batchSize). Scan(&objects).Error; err != nil { - return migrated, fmt.Errorf("query source objects: %w", err) + return atomic.LoadInt64(&migrated), fmt.Errorf("query source objects: %w", err) } if len(objects) == 0 { break } + var g errgroup.Group + g.SetLimit(migrationConcurrency) + for _, object := range objects { - source, err := sourceBackend.Get(ctx, object.FilePath) - if err != nil { - if isNotFoundError(err) { - task.AppendLog(ctx, "警告: 源存储中物理文件不存在,标记为已删除并跳过: %s (错误: %v)", object.FilePath, err) - if updateErr := db.DB(ctx).Model(&model.Upload{}). - Where("storage_driver = ? AND file_path = ?", sourceDriver, object.FilePath). - Updates(map[string]any{ - "status": model.UploadStatusDeleted, - "storage_driver": targetDriver, - }).Error; updateErr != nil { - return migrated, fmt.Errorf("update missing object %q: %w", object.FilePath, updateErr) + obj := object // Capture range variable + g.Go(func() error { + source, err := sourceBackend.Get(ctx, obj.FilePath) + if err != nil { + if isNotFoundError(err) { + task.AppendLog(ctx, "警告: 源存储中物理文件不存在,标记为已删除并跳过: %s (错误: %v)", obj.FilePath, err) + if updateErr := db.DB(ctx).Model(&model.Upload{}). + Where("storage_driver = ? AND file_path = ?", sourceDriver, obj.FilePath). + Updates(map[string]any{ + "status": model.UploadStatusDeleted, + "storage_driver": targetDriver, + }).Error; updateErr != nil { + return fmt.Errorf("update missing object %q: %w", obj.FilePath, updateErr) + } + return nil } - continue + return fmt.Errorf("open source object %q: %w", obj.FilePath, err) } - return migrated, fmt.Errorf("open source object %q: %w", object.FilePath, err) - } - targetPath, putErr := targetBackend.Put(ctx, object.FilePath, source.Body, object.FileSize, object.MimeType) - closeErr := source.Body.Close() - if putErr != nil { - return migrated, fmt.Errorf("copy object %q: %w", object.FilePath, putErr) - } - if closeErr != nil { - return migrated, fmt.Errorf("close source object %q: %w", object.FilePath, closeErr) - } - if err := db.DB(ctx).Model(&model.Upload{}). - Where("storage_driver = ? AND file_path = ?", sourceDriver, object.FilePath). - Updates(map[string]any{ - "storage_driver": targetDriver, - "file_path": targetPath, - }).Error; err != nil { - return migrated, fmt.Errorf("update migrated object %q: %w", object.FilePath, err) - } - migrated++ + targetPath, putErr := targetBackend.Put(ctx, obj.FilePath, source.Body, obj.FileSize, obj.MimeType) + closeErr := source.Body.Close() + if putErr != nil { + return fmt.Errorf("copy object %q: %w", obj.FilePath, putErr) + } + if closeErr != nil { + return fmt.Errorf("close source object %q: %w", obj.FilePath, closeErr) + } + + // Data integrity check (SHA-256 hash verification) + if len(obj.Hash) == sha256HexLength { + targetObj, getErr := targetBackend.Get(ctx, targetPath) + if getErr != nil { + return fmt.Errorf("retrieve target object for verification %q: %w", obj.FilePath, getErr) + } + h := sha256.New() + if _, copyErr := io.Copy(h, targetObj.Body); copyErr != nil { + _ = targetObj.Body.Close() + return fmt.Errorf("read target object for verification %q: %w", obj.FilePath, copyErr) + } + _ = targetObj.Body.Close() + computedHash := hex.EncodeToString(h.Sum(nil)) + if computedHash != obj.Hash { + return fmt.Errorf("integrity check failed for %q: got hash %s, want %s", obj.FilePath, computedHash, obj.Hash) + } + } + + if err := db.DB(ctx).Model(&model.Upload{}). + Where("storage_driver = ? AND file_path = ?", sourceDriver, obj.FilePath). + Updates(map[string]any{ + "storage_driver": targetDriver, + "file_path": targetPath, + }).Error; err != nil { + return fmt.Errorf("update migrated object %q: %w", obj.FilePath, err) + } + atomic.AddInt64(&migrated, 1) + return nil + }) } - task.AppendLog(ctx, "迁移进度: %d/%d", migrated, total) + + if err := g.Wait(); err != nil { + return atomic.LoadInt64(&migrated), err + } + + task.AppendLog(ctx, "迁移进度: %d/%d", atomic.LoadInt64(&migrated), total) } - return migrated, nil + return atomic.LoadInt64(&migrated), nil } func isNotFoundError(err error) bool { diff --git a/internal/apps/upload/storage_migration_task_test.go b/internal/apps/upload/storage_migration_task_test.go index 53a9953e..2f6a428b 100644 --- a/internal/apps/upload/storage_migration_task_test.go +++ b/internal/apps/upload/storage_migration_task_test.go @@ -6,10 +6,13 @@ package upload import ( "bytes" "context" + "crypto/sha256" + "encoding/hex" "encoding/json" "io" "os" "path/filepath" + "strings" "testing" "github.com/Rain-kl/Wavelet/internal/model" @@ -108,3 +111,112 @@ func TestMigrationHandlerExecute(t *testing.T) { t.Errorf("active driver = %q, want %q", current.Driver, storage.DriverS3) } } + +func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) { + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + + sourceRoot := t.TempDir() + sourcePath := filepath.Join(sourceRoot, "uploads", "test-hash.txt") + if err := os.MkdirAll(filepath.Dir(sourcePath), 0755); err != nil { + t.Fatalf("MkdirAll(%q) returned error: %v", sourcePath, err) + } + const content = "storage migration integrity check content" + if err := os.WriteFile(sourcePath, []byte(content), 0644); err != nil { + t.Fatalf("WriteFile(%q) returned error: %v", sourcePath, err) + } + + // Calculate correct SHA-256 hash + h := sha256.New() + h.Write([]byte(content)) + correctHash := hex.EncodeToString(h.Sum(nil)) + + ctx := context.Background() + active := storage.DefaultConfig() + active.Local.Root = sourceRoot + if err := storage.SaveActiveConfig(ctx, active); err != nil { + t.Fatalf("SaveActiveConfig() returned error: %v", err) + } + + target := storage.DefaultConfig() + target.Driver = storage.DriverS3 + target.S3 = storage.ObjectConfig{ + Region: "us-east-1", + Bucket: "target", + AccessKeyID: "key", + SecretAccessKey: "secret", + } + payload, err := json.Marshal(storageMigrationPayload{Target: target}) + if err != nil { + t.Fatalf("Marshal(storageMigrationPayload) returned error: %v", err) + } + + // Case 1: Incorrect Hash (should fail validation) + uploadIncorrect := model.Upload{ + ID: 99102, + UserID: 1, + FileName: "test-hash.txt", + FilePath: "uploads/test-hash.txt", + FileSize: int64(len(content)), + MimeType: "text/plain", + Extension: "txt", + Hash: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", // Invalid hash + StorageDriver: string(storage.DriverLocal), + Type: "attachment", + Status: model.UploadStatusUsed, + } + if err := dbConn.Create(&uploadIncorrect).Error; err != nil { + t.Fatalf("Create(uploadIncorrect) returned error: %v", err) + } + + var copied bytes.Buffer + restore := storage.MockStorage( + func(_ context.Context, _ string, body io.Reader, _ int64, _ string) error { + copied.Reset() + _, err := io.Copy(&copied, body) + return err + }, + func(context.Context, string) (*storage.Object, error) { + return &storage.Object{ + Body: io.NopCloser(bytes.NewBuffer(copied.Bytes())), + ContentLength: int64(copied.Len()), + ContentType: "text/plain", + }, nil + }, + func(context.Context, string) error { + return nil + }, + ) + defer restore() + + // Running execution with incorrect hash should fail with integrity error + _, err = (&MigrationHandler{}).Execute(ctx, payload) + if err == nil { + t.Fatal("Execute() succeeded with incorrect hash, want error") + } + if !strings.Contains(err.Error(), "integrity check failed") { + t.Errorf("expected integrity check failed error, got: %v", err) + } + + // Case 2: Correct Hash (should succeed) + if err := dbConn.Model(&model.Upload{}).Where("id = ?", uploadIncorrect.ID).Update("hash", correctHash).Error; err != nil { + t.Fatalf("Update hash to correct value returned error: %v", err) + } + + // Run execution with correct hash should succeed + result, err := (&MigrationHandler{}).Execute(ctx, payload) + if err != nil { + t.Fatalf("Execute() with correct hash failed: %v", err) + } + if result == nil { + t.Fatal("Execute() result = nil, want non-nil") + } + + var migrated model.Upload + if err := dbConn.First(&migrated, uploadIncorrect.ID).Error; err != nil { + t.Fatalf("First(upload) returned error: %v", err) + } + if migrated.StorageDriver != string(storage.DriverS3) { + t.Errorf("StorageDriver = %q, want %q", migrated.StorageDriver, storage.DriverS3) + } +}