完善大文件备份分卷上传与恢复

This commit is contained in:
lyen1688
2026-08-09 20:47:07 +08:00
committed by lyen1688
parent 48eb3766d2
commit bbc8b6e906
12 changed files with 1420 additions and 156 deletions
@@ -153,12 +153,12 @@ func (h *BackupHandler) GetDownloadURL(c *gin.Context) {
response.BadRequest(c, "backup ID is required")
return
}
url, err := h.backupService.GetBackupDownloadURL(c.Request.Context(), backupID)
download, err := h.backupService.GetBackupDownloadURL(c.Request.Context(), backupID)
if err != nil {
response.ErrorFrom(c, err)
return
}
response.Success(c, gin.H{"url": url})
response.Success(c, download)
}
// ─── 恢复操作(需要重新输入管理员密码) ───
@@ -5,6 +5,7 @@ import (
"context"
"fmt"
"io"
"os"
"path"
"time"
@@ -59,6 +60,34 @@ func (s *S3BackupStore) Upload(ctx context.Context, key string, body io.Reader,
return int64(len(data)), nil
}
func (s *S3BackupStore) UploadFile(ctx context.Context, key string, filePath string, contentType string) (int64, error) {
file, err := os.Open(filePath)
if err != nil {
return 0, fmt.Errorf("open upload file: %w", err)
}
defer func() { _ = file.Close() }()
info, err := file.Stat()
if err != nil {
return 0, fmt.Errorf("stat upload file: %w", err)
}
sizeBytes := info.Size()
finish := servertiming.ObserveDependency(ctx, "s3")
_, err = s.client.PutObject(ctx, &s3.PutObjectInput{
Bucket: &s.bucket,
Key: &key,
Body: file,
ContentLength: &sizeBytes,
ContentType: &contentType,
})
finish()
if err != nil {
return 0, fmt.Errorf("S3 PutObject file: %w", err)
}
return sizeBytes, nil
}
func (s *S3BackupStore) Download(ctx context.Context, key string) (io.ReadCloser, error) {
finish := servertiming.ObserveDependency(ctx, "s3")
result, err := s.client.GetObject(ctx, &s3.GetObjectInput{
@@ -0,0 +1,48 @@
//go:build unit
package repository
import (
"context"
"io"
"net/http"
"net/http/httptest"
"os"
"testing"
"github.com/stretchr/testify/require"
)
func TestS3BackupStore_UploadFile(t *testing.T) {
var received []byte
var receivedLength int64
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
require.Equal(t, http.MethodPut, r.Method)
receivedLength = r.ContentLength
var err error
received, err = io.ReadAll(r.Body)
require.NoError(t, err)
w.WriteHeader(http.StatusOK)
}))
defer server.Close()
client, err := newS3Client(context.Background(), s3ClientParams{
Endpoint: server.URL,
Region: "auto",
AccessKeyID: "test-ak",
SecretAccessKey: "test-sk",
ForcePathStyle: true,
})
require.NoError(t, err)
content := []byte("streamed backup payload")
filePath := t.TempDir() + "/part.gz"
require.NoError(t, os.WriteFile(filePath, content, 0o600))
store := &S3BackupStore{client: client, bucket: "backup-bucket"}
size, err := store.UploadFile(context.Background(), "backup/part-1", filePath, "application/octet-stream")
require.NoError(t, err)
require.Equal(t, int64(len(content)), size)
require.Equal(t, int64(len(content)), receivedLength)
require.Equal(t, content, received)
}
+106
View File
@@ -0,0 +1,106 @@
package service
import (
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"io"
"os"
)
const defaultBackupPartSizeBytes int64 = 4 * 1024 * 1024 * 1024
// BackupPart 描述一个 gzip 字节分卷。
type BackupPart struct {
Index int `json:"index"`
S3Key string `json:"s3_key"`
SizeBytes int64 `json:"size_bytes"`
SHA256 string `json:"sha256,omitempty"`
}
type localBackupPart struct {
Index int
Path string
SizeBytes int64
SHA256 string
}
func splitBackupFile(srcPath string, partSize int64) (parts []localBackupPart, err error) {
if partSize <= 0 {
return nil, fmt.Errorf("backup part size must be positive")
}
src, err := os.Open(srcPath)
if err != nil {
return nil, fmt.Errorf("open backup archive: %w", err)
}
defer func() {
if closeErr := src.Close(); err == nil && closeErr != nil {
err = fmt.Errorf("close backup archive: %w", closeErr)
}
if err != nil {
paths := make([]string, 0, len(parts))
for _, part := range parts {
paths = append(paths, part.Path)
}
_ = cleanupBackupFiles(paths...)
}
}()
info, err := src.Stat()
if err != nil {
return nil, fmt.Errorf("stat backup archive: %w", err)
}
if info.Size() <= 0 {
return nil, errors.New("backup archive is empty")
}
remaining := info.Size()
for index := 1; remaining > 0; index++ {
partFile, createErr := os.CreateTemp("", "sub2api-backup-part-*")
if createErr != nil {
return nil, fmt.Errorf("create backup part: %w", createErr)
}
partPath := partFile.Name()
partBytes := partSize
if remaining < partBytes {
partBytes = remaining
}
hash := sha256.New()
written, copyErr := io.CopyN(io.MultiWriter(partFile, hash), src, partBytes)
closeErr := partFile.Close()
if copyErr != nil {
_ = os.Remove(partPath)
return nil, fmt.Errorf("write backup part %d: %w", index, copyErr)
}
if closeErr != nil {
_ = os.Remove(partPath)
return nil, fmt.Errorf("close backup part %d: %w", index, closeErr)
}
parts = append(parts, localBackupPart{
Index: index,
Path: partPath,
SizeBytes: written,
SHA256: hex.EncodeToString(hash.Sum(nil)),
})
remaining -= written
}
return parts, nil
}
func cleanupBackupFiles(paths ...string) error {
var errs []error
for _, path := range paths {
if path == "" {
continue
}
if err := os.Remove(path); err != nil && !errors.Is(err, os.ErrNotExist) {
errs = append(errs, fmt.Errorf("remove %s: %w", path, err))
}
}
return errors.Join(errs...)
}
@@ -0,0 +1,62 @@
//go:build unit
package service
import (
"bytes"
"crypto/sha256"
"fmt"
"os"
"testing"
"github.com/stretchr/testify/require"
)
func TestSplitBackupFile_ReassemblesExactBytes(t *testing.T) {
src := writeBackupArchiveFixture(t, []byte("0123456789abcdefg"))
parts, err := splitBackupFile(src, 5)
require.NoError(t, err)
require.Len(t, parts, 4)
var got bytes.Buffer
for i, part := range parts {
require.Equal(t, i+1, part.Index)
require.LessOrEqual(t, part.SizeBytes, int64(5))
data, readErr := os.ReadFile(part.Path)
require.NoError(t, readErr)
require.Equal(t, fmt.Sprintf("%x", sha256.Sum256(data)), part.SHA256)
got.Write(data)
}
require.Equal(t, []byte("0123456789abcdefg"), got.Bytes())
}
func TestSplitBackupFile_RejectsInvalidInput(t *testing.T) {
src := writeBackupArchiveFixture(t, []byte("data"))
_, err := splitBackupFile(src, 0)
require.Error(t, err)
empty := writeBackupArchiveFixture(t, nil)
_, err = splitBackupFile(empty, 5)
require.Error(t, err)
_, err = splitBackupFile(filepathForMissingBackupArchive(t), 5)
require.Error(t, err)
}
func writeBackupArchiveFixture(t *testing.T, content []byte) string {
t.Helper()
path := filepathForBackupArchive(t)
require.NoError(t, os.WriteFile(path, content, 0o600))
return path
}
func filepathForBackupArchive(t *testing.T) string {
t.Helper()
return t.TempDir() + "/archive.gz"
}
func filepathForMissingBackupArchive(t *testing.T) string {
t.Helper()
return t.TempDir() + "/missing.gz"
}
+410 -146
View File
@@ -3,10 +3,13 @@ package service
import (
"compress/gzip"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"sort"
"strings"
"sync"
@@ -26,7 +29,8 @@ const (
settingKeyBackupSchedule = "backup_schedule"
settingKeyBackupRecords = "backup_records"
maxBackupRecords = 100
maxBackupRecords = 100
backupObjectCleanupTimeout = 2 * time.Minute
)
var (
@@ -61,6 +65,7 @@ type DBDumper interface {
// BackupObjectStore abstracts object storage for backup files
type BackupObjectStore interface {
Upload(ctx context.Context, key string, body io.Reader, contentType string) (sizeBytes int64, err error)
UploadFile(ctx context.Context, key string, filePath string, contentType string) (sizeBytes int64, err error)
Download(ctx context.Context, key string) (io.ReadCloser, error)
Delete(ctx context.Context, key string) error
PresignURL(ctx context.Context, key string, expiry time.Duration) (string, error)
@@ -98,21 +103,35 @@ type BackupScheduleConfig struct {
// BackupRecord 备份记录
type BackupRecord struct {
ID string `json:"id"`
Status string `json:"status"` // pending, running, completed, failed
BackupType string `json:"backup_type"` // postgres
FileName string `json:"file_name"`
S3Key string `json:"s3_key"`
SizeBytes int64 `json:"size_bytes"`
TriggeredBy string `json:"triggered_by"` // manual, scheduled
ErrorMsg string `json:"error_message,omitempty"`
StartedAt string `json:"started_at"`
FinishedAt string `json:"finished_at,omitempty"`
ExpiresAt string `json:"expires_at,omitempty"` // 过期时间
Progress string `json:"progress,omitempty"` // "dumping", "uploading", ""
RestoreStatus string `json:"restore_status,omitempty"` // "", "running", "completed", "failed"
RestoreError string `json:"restore_error,omitempty"`
RestoredAt string `json:"restored_at,omitempty"`
ID string `json:"id"`
Status string `json:"status"` // pending, running, completed, failed
BackupType string `json:"backup_type"` // postgres
FileName string `json:"file_name"`
S3Key string `json:"s3_key"`
Parts []BackupPart `json:"parts,omitempty"`
SizeBytes int64 `json:"size_bytes"`
TriggeredBy string `json:"triggered_by"` // manual, scheduled
ErrorMsg string `json:"error_message,omitempty"`
StartedAt string `json:"started_at"`
FinishedAt string `json:"finished_at,omitempty"`
ExpiresAt string `json:"expires_at,omitempty"` // 过期时间
Progress string `json:"progress,omitempty"` // "dumping", "uploading", ""
RestoreStatus string `json:"restore_status,omitempty"` // "", "running", "completed", "failed"
RestoreError string `json:"restore_error,omitempty"`
RestoredAt string `json:"restored_at,omitempty"`
}
// BackupDownloadPart 描述一个可下载的备份分卷。
type BackupDownloadPart struct {
Index int `json:"index"`
SizeBytes int64 `json:"size_bytes"`
URL string `json:"url"`
}
// BackupDownloadResponse 是单文件和分卷下载响应的兼容表示。
type BackupDownloadResponse struct {
URL string `json:"url,omitempty"`
Parts []BackupDownloadPart `json:"parts,omitempty"`
}
// BackupService 数据库备份恢复服务
@@ -142,10 +161,11 @@ type BackupService struct {
cronSched *cron.Cron
cronEntryID cron.EntryID
wg sync.WaitGroup // 追踪活跃的备份/恢复 goroutine
shuttingDown atomic.Bool // 阻止新备份启动
bgCtx context.Context // 所有后台操作的 parent context
bgCancel context.CancelFunc // 取消所有活跃后台操作
wg sync.WaitGroup // 追踪活跃的备份/恢复 goroutine
shuttingDown atomic.Bool // 阻止新备份启动
bgCtx context.Context // 所有后台操作的 parent context
bgCancel context.CancelFunc // 取消所有活跃后台操作
partSizeBytes int64 // 分卷阈值;生产使用 4 GiB,测试可注入更小值
}
func NewBackupService(
@@ -165,6 +185,7 @@ func NewBackupService(
dumper: dumper,
bgCtx: bgCtx,
bgCancel: bgCancel,
partSizeBytes: defaultBackupPartSizeBytes,
}
}
@@ -191,33 +212,57 @@ func (s *BackupService) Start() {
}
}
// recoverStaleRecords 启动时将孤立的 running 记录标记为 failed
// recoverStaleRecords 启动时将孤立的 running 记录标记为 failed,并清理已上传对象。
func (s *BackupService) recoverStaleRecords() {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
loadCtx, loadCancel := context.WithTimeout(context.Background(), 10*time.Second)
defer loadCancel()
records, err := s.loadRecords(ctx)
records, err := s.loadRecords(loadCtx)
if err != nil {
return
}
for i := range records {
if records[i].Status == "running" {
staleRecord := records[i]
records[i].Status = "failed"
records[i].ErrorMsg = "interrupted by server restart"
records[i].Progress = ""
records[i].FinishedAt = time.Now().Format(time.RFC3339)
_ = s.saveRecord(ctx, &records[i])
s.saveRecoveredRecord(&records[i])
if cleanupErr := s.cleanupStaleBackupObjects(&staleRecord); cleanupErr != nil {
records[i].ErrorMsg = fmt.Sprintf("interrupted by server restart; cleanup failed, manual deletion may be required: %v", cleanupErr)
s.saveRecoveredRecord(&records[i])
logger.LegacyPrintf("service.backup", "[Backup] failed to clean stale backup objects for %s: %v", records[i].ID, cleanupErr)
}
logger.LegacyPrintf("service.backup", "[Backup] recovered stale running record: %s", records[i].ID)
}
if records[i].RestoreStatus == "running" {
records[i].RestoreStatus = "failed"
records[i].RestoreError = "interrupted by server restart"
_ = s.saveRecord(ctx, &records[i])
s.saveRecoveredRecord(&records[i])
logger.LegacyPrintf("service.backup", "[Backup] recovered stale restoring record: %s", records[i].ID)
}
}
}
func (s *BackupService) saveRecoveredRecord(record *BackupRecord) {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
if err := s.saveRecord(ctx, record); err != nil {
logger.LegacyPrintf("service.backup", "[Backup] 保存恢复后的备份记录失败 %s: %v", record.ID, err)
}
}
func (s *BackupService) cleanupStaleBackupObjects(record *BackupRecord) error {
if len(backupObjectKeys(record)) == 0 {
return nil
}
ctx, cancel := context.WithTimeout(context.Background(), backupObjectCleanupTimeout)
defer cancel()
return s.deleteBackupObjects(ctx, record)
}
// Stop 停止定时备份并等待活跃操作完成
func (s *BackupService) Stop() {
s.shuttingDown.Store(true)
@@ -452,7 +497,7 @@ func (s *BackupService) runScheduledBackup() {
// ─── 备份/恢复核心 ───
// CreateBackup 创建全量数据库备份并上传到 S3(流式处理)
// CreateBackup 创建全量数据库备份并上传到 S3
// expireDays: 备份过期天数,0=永不过期,默认14天
func (s *BackupService) CreateBackup(ctx context.Context, triggeredBy string, expireDays int) (*BackupRecord, error) {
if s.shuttingDown.Load() {
@@ -506,61 +551,27 @@ func (s *BackupService) CreateBackup(ctx context.Context, triggeredBy string, ex
ExpiresAt: expiresAt,
}
// 流式执行: pg_dump -> gzip -> S3 upload
dumpReader, err := s.dumper.Dump(ctx)
archivePath, sizeBytes, err := s.createCompressedBackupFile(ctx)
if err != nil {
record.Status = "failed"
record.ErrorMsg = fmt.Sprintf("pg_dump failed: %v", err)
record.ErrorMsg = err.Error()
record.FinishedAt = time.Now().Format(time.RFC3339)
_ = s.saveRecord(ctx, record)
return record, fmt.Errorf("pg_dump: %w", err)
return record, err
}
// 使用 io.Pipe 将 gzip 压缩数据流式传递给 S3 上传
pr, pw := io.Pipe()
gzipDone := make(chan error, 1)
go func() {
defer func() {
if r := recover(); r != nil {
pw.CloseWithError(fmt.Errorf("gzip goroutine panic: %v", r)) //nolint:errcheck
gzipDone <- fmt.Errorf("gzip goroutine panic: %v", r)
}
}()
gzWriter := gzip.NewWriter(pw)
var gzErr error
_, gzErr = io.Copy(gzWriter, dumpReader)
if closeErr := gzWriter.Close(); closeErr != nil && gzErr == nil {
gzErr = closeErr
}
if closeErr := dumpReader.Close(); closeErr != nil && gzErr == nil {
gzErr = closeErr
}
if gzErr != nil {
_ = pw.CloseWithError(gzErr)
} else {
_ = pw.Close()
}
gzipDone <- gzErr
}()
contentType := "application/gzip"
sizeBytes, err := objectStore.Upload(ctx, s3Key, pr, contentType)
if err != nil {
_ = pr.CloseWithError(err) // 确保 gzip goroutine 不会悬挂
gzErr := <-gzipDone // 安全等待 gzip goroutine 完成
record.Status = "failed"
errMsg := fmt.Sprintf("S3 upload failed: %v", err)
if gzErr != nil {
errMsg = fmt.Sprintf("gzip/dump failed: %v", gzErr)
}
record.ErrorMsg = errMsg
record.FinishedAt = time.Now().Format(time.RFC3339)
_ = s.saveRecord(ctx, record)
return record, fmt.Errorf("backup upload: %w", err)
}
<-gzipDone // 确保 gzip goroutine 已退出
defer func() { _ = cleanupBackupFiles(archivePath) }()
record.SizeBytes = sizeBytes
if err := s.saveRecord(ctx, record); err != nil {
return nil, fmt.Errorf("save initial record: %w", err)
}
if err := s.uploadBackupArchive(ctx, record, objectStore, s3Cfg, archivePath); err != nil {
record.Status = "failed"
record.ErrorMsg = err.Error()
record.FinishedAt = time.Now().Format(time.RFC3339)
_ = s.saveRecord(ctx, record)
return record, err
}
record.Status = "completed"
record.FinishedAt = time.Now().Format(time.RFC3339)
if err := s.saveRecord(ctx, record); err != nil {
@@ -656,78 +667,43 @@ func (s *BackupService) StartBackup(ctx context.Context, triggeredBy string, exp
_ = s.saveRecord(context.Background(), record)
}
}()
s.executeBackup(record, objectStore)
s.executeBackup(record, objectStore, s3Cfg)
}()
return &result, nil
}
// executeBackup 后台执行备份(独立于 HTTP context
func (s *BackupService) executeBackup(record *BackupRecord, objectStore BackupObjectStore) {
func (s *BackupService) executeBackup(record *BackupRecord, objectStore BackupObjectStore, s3Cfg *BackupS3Config) {
ctx, cancel := context.WithTimeout(s.bgCtx, 30*time.Minute)
defer cancel()
// 阶段1: pg_dump
// 阶段1: pg_dump -> gzip 临时文件
record.Progress = "dumping"
_ = s.saveRecord(ctx, record)
dumpReader, err := s.dumper.Dump(ctx)
archivePath, sizeBytes, err := s.createCompressedBackupFile(ctx)
if err != nil {
record.Status = "failed"
record.ErrorMsg = fmt.Sprintf("pg_dump failed: %v", err)
record.ErrorMsg = err.Error()
record.Progress = ""
record.FinishedAt = time.Now().Format(time.RFC3339)
_ = s.saveRecord(context.Background(), record)
return
}
defer func() { _ = cleanupBackupFiles(archivePath) }()
record.SizeBytes = sizeBytes
// 阶段2: gzip + upload
// 阶段2: 单对象或分卷上传
record.Progress = "uploading"
_ = s.saveRecord(ctx, record)
pr, pw := io.Pipe()
gzipDone := make(chan error, 1)
go func() {
defer func() {
if r := recover(); r != nil {
pw.CloseWithError(fmt.Errorf("gzip goroutine panic: %v", r)) //nolint:errcheck
gzipDone <- fmt.Errorf("gzip goroutine panic: %v", r)
}
}()
gzWriter := gzip.NewWriter(pw)
var gzErr error
_, gzErr = io.Copy(gzWriter, dumpReader)
if closeErr := gzWriter.Close(); closeErr != nil && gzErr == nil {
gzErr = closeErr
}
if closeErr := dumpReader.Close(); closeErr != nil && gzErr == nil {
gzErr = closeErr
}
if gzErr != nil {
_ = pw.CloseWithError(gzErr)
} else {
_ = pw.Close()
}
gzipDone <- gzErr
}()
contentType := "application/gzip"
sizeBytes, err := objectStore.Upload(ctx, record.S3Key, pr, contentType)
if err != nil {
_ = pr.CloseWithError(err) // 确保 gzip goroutine 不会悬挂
gzErr := <-gzipDone // 安全等待 gzip goroutine 完成
if err := s.uploadBackupArchive(ctx, record, objectStore, s3Cfg, archivePath); err != nil {
record.Status = "failed"
errMsg := fmt.Sprintf("S3 upload failed: %v", err)
if gzErr != nil {
errMsg = fmt.Sprintf("gzip/dump failed: %v", gzErr)
}
record.ErrorMsg = errMsg
record.ErrorMsg = err.Error()
record.Progress = ""
record.FinishedAt = time.Now().Format(time.RFC3339)
_ = s.saveRecord(context.Background(), record)
return
}
<-gzipDone // 确保 gzip goroutine 已退出
record.SizeBytes = sizeBytes
record.Status = "completed"
@@ -738,6 +714,108 @@ func (s *BackupService) executeBackup(record *BackupRecord, objectStore BackupOb
}
}
func (s *BackupService) createCompressedBackupFile(ctx context.Context) (string, int64, error) {
dumpReader, err := s.dumper.Dump(ctx)
if err != nil {
return "", 0, fmt.Errorf("pg_dump: %w", err)
}
archive, err := os.CreateTemp("", "sub2api-backup-*.sql.gz")
if err != nil {
_ = dumpReader.Close()
return "", 0, fmt.Errorf("create backup archive: %w", err)
}
archivePath := archive.Name()
gzWriter := gzip.NewWriter(archive)
_, copyErr := io.Copy(gzWriter, dumpReader)
if closeErr := gzWriter.Close(); copyErr == nil && closeErr != nil {
copyErr = closeErr
}
if closeErr := dumpReader.Close(); copyErr == nil && closeErr != nil {
copyErr = closeErr
}
if closeErr := archive.Close(); copyErr == nil && closeErr != nil {
copyErr = closeErr
}
if copyErr != nil {
_ = cleanupBackupFiles(archivePath)
return "", 0, fmt.Errorf("gzip/dump failed: %w", copyErr)
}
info, err := os.Stat(archivePath)
if err != nil {
_ = cleanupBackupFiles(archivePath)
return "", 0, fmt.Errorf("stat backup archive: %w", err)
}
return archivePath, info.Size(), nil
}
func (s *BackupService) uploadBackupArchive(ctx context.Context, record *BackupRecord, objectStore BackupObjectStore, cfg *BackupS3Config, archivePath string) error {
info, err := os.Stat(archivePath)
if err != nil {
return fmt.Errorf("stat backup archive: %w", err)
}
partSize := s.partSizeBytes
if partSize <= 0 {
partSize = defaultBackupPartSizeBytes
}
if info.Size() <= partSize {
if _, err := objectStore.UploadFile(ctx, record.S3Key, archivePath, "application/gzip"); err != nil {
cleanupCtx, cleanupCancel := context.WithTimeout(context.Background(), backupObjectCleanupTimeout)
cleanupErr := deleteBackupObjectKeys(cleanupCtx, objectStore, record)
cleanupCancel()
return errors.Join(fmt.Errorf("backup upload: %w", err), cleanupErr)
}
record.Parts = nil
return nil
}
localParts, err := splitBackupFile(archivePath, partSize)
if err != nil {
return fmt.Errorf("split backup archive: %w", err)
}
defer func() {
paths := make([]string, 0, len(localParts))
for _, part := range localParts {
paths = append(paths, part.Path)
}
_ = cleanupBackupFiles(paths...)
}()
if cfg == nil {
return errors.New("backup S3 config is unavailable for split upload")
}
record.S3Key = ""
record.Parts = make([]BackupPart, 0, len(localParts))
partRoot := strings.TrimRight(s.buildS3Key(cfg, record.ID), "/")
for _, part := range localParts {
record.Parts = append(record.Parts, BackupPart{
Index: part.Index,
S3Key: s.buildBackupPartKey(partRoot, part.Index),
SizeBytes: part.SizeBytes,
SHA256: part.SHA256,
})
}
if err := s.saveRecord(ctx, record); err != nil {
return fmt.Errorf("save split backup plan: %w", err)
}
for i, part := range localParts {
if _, err := objectStore.UploadFile(ctx, record.Parts[i].S3Key, part.Path, "application/octet-stream"); err != nil {
// PUT 可能已经在对象存储端成功、但客户端因超时收到错误;
// 因此失败时清理整份分卷计划,而不只清理此前返回成功的卷。
cleanupCtx, cleanupCancel := context.WithTimeout(context.Background(), backupObjectCleanupTimeout)
cleanupErr := deleteBackupObjectKeys(cleanupCtx, objectStore, record)
cleanupCancel()
return errors.Join(fmt.Errorf("upload backup part %d: %w", part.Index, err), cleanupErr)
}
}
return nil
}
func (s *BackupService) buildBackupPartKey(root string, index int) string {
return fmt.Sprintf("%s/payload.part-%06d", strings.TrimRight(root, "/"), index)
}
// RestoreBackup 从 S3 下载备份并流式恢复到数据库
func (s *BackupService) RestoreBackup(ctx context.Context, backupID string) error {
s.opMu.Lock()
@@ -770,7 +848,16 @@ func (s *BackupService) RestoreBackup(ctx context.Context, backupID string) erro
return fmt.Errorf("init object store: %w", err)
}
// 从 S3 流式下载
if len(record.Parts) > 0 {
archivePath, err := s.downloadBackupParts(ctx, objectStore, record.Parts)
if err != nil {
return err
}
defer func() { _ = cleanupBackupFiles(archivePath) }()
return s.restoreArchive(ctx, archivePath)
}
// 旧记录从 S3 流式下载
body, err := objectStore.Download(ctx, record.S3Key)
if err != nil {
return fmt.Errorf("S3 download failed: %w", err)
@@ -866,6 +953,29 @@ func (s *BackupService) executeRestore(record *BackupRecord, objectStore BackupO
ctx, cancel := context.WithTimeout(s.bgCtx, 30*time.Minute)
defer cancel()
if len(record.Parts) > 0 {
archivePath, err := s.downloadBackupParts(ctx, objectStore, record.Parts)
if err != nil {
record.RestoreStatus = "failed"
record.RestoreError = err.Error()
_ = s.saveRecord(context.Background(), record)
return
}
defer func() { _ = cleanupBackupFiles(archivePath) }()
if err := s.restoreArchive(ctx, archivePath); err != nil {
record.RestoreStatus = "failed"
record.RestoreError = fmt.Sprintf("pg restore: %v", err)
_ = s.saveRecord(context.Background(), record)
return
}
record.RestoreStatus = "completed"
record.RestoredAt = time.Now().Format(time.RFC3339)
if err := s.saveRecord(context.Background(), record); err != nil {
logger.LegacyPrintf("service.backup", "[Backup] 保存恢复记录失败: %v", err)
}
return
}
body, err := objectStore.Download(ctx, record.S3Key)
if err != nil {
record.RestoreStatus = "failed"
@@ -898,6 +1008,79 @@ func (s *BackupService) executeRestore(record *BackupRecord, objectStore BackupO
}
}
func (s *BackupService) downloadBackupParts(ctx context.Context, objectStore BackupObjectStore, parts []BackupPart) (path string, err error) {
if len(parts) == 0 {
return "", errors.New("backup parts are empty")
}
ordered := append([]BackupPart(nil), parts...)
sort.Slice(ordered, func(i, j int) bool { return ordered[i].Index < ordered[j].Index })
for i, part := range ordered {
if part.Index != i+1 || part.S3Key == "" || part.SizeBytes <= 0 {
return "", fmt.Errorf("invalid backup part metadata at index %d", i+1)
}
}
archive, err := os.CreateTemp("", "sub2api-restore-*.sql.gz")
if err != nil {
return "", fmt.Errorf("create restore archive: %w", err)
}
path = archive.Name()
cleanup := func() {
_ = archive.Close()
_ = cleanupBackupFiles(path)
}
for _, part := range ordered {
body, downloadErr := objectStore.Download(ctx, part.S3Key)
if downloadErr != nil {
cleanup()
return "", fmt.Errorf("download backup part %d: %w", part.Index, downloadErr)
}
hash := sha256.New()
written, copyErr := io.Copy(io.MultiWriter(archive, hash), body)
closeErr := body.Close()
if copyErr != nil {
cleanup()
return "", fmt.Errorf("read backup part %d: %w", part.Index, copyErr)
}
if closeErr != nil {
cleanup()
return "", fmt.Errorf("close backup part %d: %w", part.Index, closeErr)
}
if written != part.SizeBytes {
cleanup()
return "", fmt.Errorf("backup part %d size mismatch: got %d, want %d", part.Index, written, part.SizeBytes)
}
if part.SHA256 != "" && !strings.EqualFold(part.SHA256, hex.EncodeToString(hash.Sum(nil))) {
cleanup()
return "", fmt.Errorf("backup part %d checksum mismatch", part.Index)
}
}
if err := archive.Close(); err != nil {
_ = cleanupBackupFiles(path)
return "", fmt.Errorf("close restore archive: %w", err)
}
return path, nil
}
func (s *BackupService) restoreArchive(ctx context.Context, archivePath string) error {
archive, err := os.Open(archivePath)
if err != nil {
return fmt.Errorf("open restore archive: %w", err)
}
defer func() { _ = archive.Close() }()
gzReader, err := gzip.NewReader(archive)
if err != nil {
return fmt.Errorf("gzip reader: %w", err)
}
defer func() { _ = gzReader.Close() }()
if err := s.dumper.Restore(ctx, gzReader); err != nil {
return fmt.Errorf("pg restore: %w", err)
}
return nil
}
// ─── 备份记录管理 ───
func (s *BackupService) ListBackups(ctx context.Context) ([]BackupRecord, error) {
@@ -946,45 +1129,67 @@ func (s *BackupService) DeleteBackup(ctx context.Context, backupID string) error
if found == nil {
return ErrBackupNotFound
}
if found.Status == "running" {
// 后台上传仍可能依赖 Parts 计划;删除对象会让随后完成的记录引用失效卷。
return ErrBackupInProgress
}
// 从 S3 删除
if found.S3Key != "" && found.Status == "completed" {
s3Cfg, err := s.loadS3Config(ctx)
if err == nil && s3Cfg != nil && s3Cfg.IsConfigured() {
objectStore, err := s.getOrCreateStore(ctx, s3Cfg)
if err == nil {
_ = objectStore.Delete(ctx, found.S3Key)
}
}
// 从对象存储删除所有单文件或分卷对象。删除不完整时保留记录,便于重试。
if err := s.deleteBackupObjects(ctx, found); err != nil {
return err
}
return s.saveRecordsLocked(ctx, remaining)
}
// GetBackupDownloadURL 获取备份文件预签名下载 URL
func (s *BackupService) GetBackupDownloadURL(ctx context.Context, backupID string) (string, error) {
func (s *BackupService) GetBackupDownloadURL(ctx context.Context, backupID string) (BackupDownloadResponse, error) {
var download BackupDownloadResponse
record, err := s.GetBackupRecord(ctx, backupID)
if err != nil {
return "", err
return download, err
}
if record.Status != "completed" {
return "", infraerrors.BadRequest("BACKUP_NOT_COMPLETED", "backup is not completed")
return download, infraerrors.BadRequest("BACKUP_NOT_COMPLETED", "backup is not completed")
}
s3Cfg, err := s.loadS3Config(ctx)
if err != nil {
return "", err
return download, err
}
objectStore, err := s.getOrCreateStore(ctx, s3Cfg)
if err != nil {
return "", err
return download, err
}
if len(record.Parts) > 0 {
parts := append([]BackupPart(nil), record.Parts...)
sort.Slice(parts, func(i, j int) bool { return parts[i].Index < parts[j].Index })
for i, part := range parts {
if part.Index != i+1 || part.S3Key == "" || part.SizeBytes <= 0 {
return download, fmt.Errorf("invalid backup part metadata at index %d", i+1)
}
url, presignErr := objectStore.PresignURL(ctx, part.S3Key, 1*time.Hour)
if presignErr != nil {
return download, fmt.Errorf("presign backup part %d: %w", part.Index, presignErr)
}
download.Parts = append(download.Parts, BackupDownloadPart{
Index: part.Index,
SizeBytes: part.SizeBytes,
URL: url,
})
}
return download, nil
}
if record.S3Key == "" {
return download, errors.New("backup object key is empty")
}
url, err := objectStore.PresignURL(ctx, record.S3Key, 1*time.Hour)
if err != nil {
return "", fmt.Errorf("presign url: %w", err)
return download, fmt.Errorf("presign url: %w", err)
}
return url, nil
download.URL = url
return download, nil
}
// ─── 内部方法 ───
@@ -1141,28 +1346,87 @@ func (s *BackupService) cleanupOldBackups(ctx context.Context, schedule *BackupS
}
}
// 删除 S3 上的文件
var cleanupErrs []error
deletedCount := 0
for _, r := range toDelete {
if r.S3Key != "" {
_ = s.deleteS3Object(ctx, r.S3Key)
if err := s.deleteBackupObjects(ctx, &r); err != nil {
// 对象删除失败时保留记录,避免丢失后续重试所需的 key。
toKeep = append(toKeep, r)
cleanupErrs = append(cleanupErrs, fmt.Errorf("cleanup backup %s: %w", r.ID, err))
continue
}
deletedCount++
}
if len(toDelete) > 0 {
logger.LegacyPrintf("service.backup", "[Backup] 自动清理了 %d 个过期备份", len(toDelete))
return s.saveRecordsLocked(ctx, toKeep)
if err := s.saveRecordsLocked(ctx, toKeep); err != nil {
cleanupErrs = append(cleanupErrs, fmt.Errorf("save backup records after cleanup: %w", err))
}
if deletedCount > 0 {
logger.LegacyPrintf("service.backup", "[Backup] 自动清理了 %d 个过期备份", deletedCount)
}
return errors.Join(cleanupErrs...)
}
return nil
}
func (s *BackupService) deleteS3Object(ctx context.Context, key string) error {
// backupObjectKeys 返回一条备份记录关联的全部对象 key。
// 新记录使用 Parts,旧记录使用 S3Key;两者同时存在时也全部返回,便于清理异常残留对象。
func backupObjectKeys(record *BackupRecord) []string {
if record == nil {
return nil
}
keys := make([]string, 0, len(record.Parts)+1)
seen := make(map[string]struct{}, len(record.Parts)+1)
appendKey := func(key string) {
if key == "" {
return
}
if _, ok := seen[key]; ok {
return
}
seen[key] = struct{}{}
keys = append(keys, key)
}
appendKey(record.S3Key)
parts := append([]BackupPart(nil), record.Parts...)
sort.Slice(parts, func(i, j int) bool { return parts[i].Index < parts[j].Index })
for _, part := range parts {
appendKey(part.S3Key)
}
return keys
}
// deleteBackupObjects 尝试删除记录关联的所有对象,并聚合删除错误。
func (s *BackupService) deleteBackupObjects(ctx context.Context, record *BackupRecord) error {
if len(backupObjectKeys(record)) == 0 {
return nil
}
s3Cfg, err := s.loadS3Config(ctx)
if err != nil || s3Cfg == nil {
if err != nil {
return err
}
if s3Cfg == nil || !s3Cfg.IsConfigured() {
// 兼容没有配置对象存储的旧记录:记录仍可被删除。
return nil
}
objectStore, err := s.getOrCreateStore(ctx, s3Cfg)
if err != nil {
return err
}
return objectStore.Delete(ctx, key)
return deleteBackupObjectKeys(ctx, objectStore, record)
}
func deleteBackupObjectKeys(ctx context.Context, objectStore BackupObjectStore, record *BackupRecord) error {
keys := backupObjectKeys(record)
if len(keys) == 0 {
return nil
}
var errs []error
for _, key := range keys {
if deleteErr := objectStore.Delete(ctx, key); deleteErr != nil {
errs = append(errs, fmt.Errorf("delete backup object %q: %w", key, deleteErr))
}
}
return errors.Join(errs...)
}
+516 -5
View File
@@ -4,10 +4,13 @@ package service
import (
"bytes"
"compress/gzip"
"context"
"crypto/sha256"
"encoding/json"
"fmt"
"io"
"os"
"strings"
"sync"
"testing"
@@ -159,12 +162,49 @@ func (d *blockingDumper) Restore(_ context.Context, data io.Reader) error {
}
type mockObjectStore struct {
objects map[string][]byte
mu sync.Mutex
objects map[string][]byte
mu sync.Mutex
failUploadFileAt int
uploadFileCalls int
deletedKeys []string
failDeleteKeys map[string]error
}
type cancelingUploadFailureStore struct {
*mockObjectStore
cancel context.CancelFunc
}
func (m *cancelingUploadFailureStore) UploadFile(_ context.Context, key string, filePath string, _ string) (int64, error) {
file, err := os.Open(filePath)
if err != nil {
return 0, err
}
data, readErr := io.ReadAll(file)
closeErr := file.Close()
if readErr != nil {
return 0, readErr
}
if closeErr != nil {
return 0, closeErr
}
m.mu.Lock()
m.objects[key] = data
m.mu.Unlock()
m.cancel()
return 0, fmt.Errorf("injected upload failure after object landed")
}
func (m *cancelingUploadFailureStore) Delete(ctx context.Context, key string) error {
if err := ctx.Err(); err != nil {
return err
}
return m.mockObjectStore.Delete(ctx, key)
}
func newMockObjectStore() *mockObjectStore {
return &mockObjectStore{objects: make(map[string][]byte)}
return &mockObjectStore{objects: make(map[string][]byte), failDeleteKeys: make(map[string]error)}
}
func (m *mockObjectStore) Upload(_ context.Context, key string, body io.Reader, _ string) (int64, error) {
@@ -178,6 +218,23 @@ func (m *mockObjectStore) Upload(_ context.Context, key string, body io.Reader,
return int64(len(data)), nil
}
func (m *mockObjectStore) UploadFile(ctx context.Context, key string, filePath string, contentType string) (int64, error) {
m.mu.Lock()
m.uploadFileCalls++
call := m.uploadFileCalls
failAt := m.failUploadFileAt
m.mu.Unlock()
if failAt > 0 && call == failAt {
return 0, fmt.Errorf("injected upload failure at call %d", call)
}
file, err := os.Open(filePath)
if err != nil {
return 0, err
}
defer func() { _ = file.Close() }()
return m.Upload(ctx, key, file, contentType)
}
func (m *mockObjectStore) Download(_ context.Context, key string) (io.ReadCloser, error) {
m.mu.Lock()
data, ok := m.objects[key]
@@ -190,6 +247,11 @@ func (m *mockObjectStore) Download(_ context.Context, key string) (io.ReadCloser
func (m *mockObjectStore) Delete(_ context.Context, key string) error {
m.mu.Lock()
m.deletedKeys = append(m.deletedKeys, key)
if err, ok := m.failDeleteKeys[key]; ok {
m.mu.Unlock()
return err
}
delete(m.objects, key)
m.mu.Unlock()
return nil
@@ -405,6 +467,122 @@ func TestBackupService_CreateBackup_Streaming(t *testing.T) {
store.mu.Unlock()
}
func TestBackupService_CreateBackup_SplitsCompressedArchive(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
dumpContent := entropyBackupFixture(512)
dumper := &mockDumper{dumpData: dumpContent}
store := newMockObjectStore()
svc := newTestBackupService(repo, dumper, store)
svc.partSizeBytes = 32
record, err := svc.CreateBackup(context.Background(), "manual", 14)
require.NoError(t, err)
require.Equal(t, "completed", record.Status)
require.Greater(t, len(record.Parts), 1)
require.Empty(t, record.S3Key)
var compressed bytes.Buffer
store.mu.Lock()
for _, part := range record.Parts {
data, ok := store.objects[part.S3Key]
require.True(t, ok)
require.LessOrEqual(t, len(data), 32)
compressed.Write(data)
}
store.mu.Unlock()
gzReader, err := gzip.NewReader(bytes.NewReader(compressed.Bytes()))
require.NoError(t, err)
decompressed, err := io.ReadAll(gzReader)
require.NoError(t, err)
require.NoError(t, gzReader.Close())
require.Equal(t, dumpContent, decompressed)
}
func TestBackupService_StartBackup_SplitsCompressedArchive(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
store := newMockObjectStore()
svc := newTestBackupService(repo, &mockDumper{dumpData: entropyBackupFixture(512)}, store)
svc.partSizeBytes = 32
record, err := svc.StartBackup(context.Background(), "manual", 14)
require.NoError(t, err)
svc.wg.Wait()
final, err := svc.GetBackupRecord(context.Background(), record.ID)
require.NoError(t, err)
require.Equal(t, "completed", final.Status)
require.Greater(t, len(final.Parts), 1)
require.Empty(t, final.S3Key)
}
func TestBackupService_StartBackup_UploadFailureCleansUploadedParts(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
store := newMockObjectStore()
store.failUploadFileAt = 2
svc := newTestBackupService(repo, &mockDumper{dumpData: entropyBackupFixture(512)}, store)
svc.partSizeBytes = 32
record, err := svc.StartBackup(context.Background(), "manual", 14)
require.NoError(t, err)
svc.wg.Wait()
final, err := svc.GetBackupRecord(context.Background(), record.ID)
require.NoError(t, err)
require.Equal(t, "failed", final.Status)
require.NotEmpty(t, final.Parts)
store.mu.Lock()
deletedKeys := append([]string(nil), store.deletedKeys...)
store.mu.Unlock()
for _, part := range final.Parts {
require.Contains(t, deletedKeys, part.S3Key)
}
}
func TestBackupService_UploadFailureCleanupUsesDetachedContext(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
svc := newTestBackupService(repo, &mockDumper{}, newMockObjectStore())
svc.partSizeBytes = 4
archive, err := os.CreateTemp("", "backup-upload-context-*.gz")
require.NoError(t, err)
archivePath := archive.Name()
defer func() { _ = os.Remove(archivePath) }()
_, err = archive.Write([]byte("0123456789"))
require.NoError(t, err)
require.NoError(t, archive.Close())
ctx, cancel := context.WithCancel(context.Background())
store := &cancelingUploadFailureStore{
mockObjectStore: newMockObjectStore(),
cancel: cancel,
}
record := &BackupRecord{ID: "cancel-cleanup", S3Key: "backups/cancel-cleanup.sql.gz"}
err = svc.uploadBackupArchive(ctx, record, store, &BackupS3Config{Prefix: "backups"}, archivePath)
require.Error(t, err)
require.NotContains(t, err.Error(), "context canceled")
store.mu.Lock()
defer store.mu.Unlock()
for _, part := range record.Parts {
require.Contains(t, store.deletedKeys, part.S3Key)
require.NotContains(t, store.objects, part.S3Key)
}
}
func entropyBackupFixture(size int) []byte {
data := make([]byte, size)
for i := range data {
data[i] = byte((i*31 + 17) % 251)
}
return data
}
func TestBackupService_CreateBackup_DumpFailure(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
@@ -466,6 +644,123 @@ func TestBackupService_RestoreBackup_Streaming(t *testing.T) {
require.Equal(t, dumpContent, string(dumper.restored))
}
func TestBackupService_RestoreBackup_SplitParts(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
dumpContent := entropyBackupFixture(512)
dumper := &mockDumper{}
store := newMockObjectStore()
svc := newTestBackupService(repo, dumper, store)
compressed := gzipBackupBytes(t, dumpContent)
parts := splitBackupBytes(compressed, 11)
recordParts := make([]BackupPart, 0, len(parts))
for i, data := range parts {
key := fmt.Sprintf("backups/split-1/payload.part-%06d", i+1)
store.objects[key] = data
recordParts = append(recordParts, BackupPart{
Index: i + 1,
S3Key: key,
SizeBytes: int64(len(data)),
SHA256: fmt.Sprintf("%x", sha256.Sum256(data)),
})
}
record := &BackupRecord{
ID: "split-1",
Status: "completed",
Parts: recordParts,
SizeBytes: int64(len(compressed)),
}
require.NoError(t, svc.saveRecord(context.Background(), record))
require.NoError(t, svc.RestoreBackup(context.Background(), record.ID))
require.Equal(t, dumpContent, dumper.restored)
}
func TestBackupService_RestoreBackup_SplitPartsMissingPartDoesNotRestore(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
dumpContent := entropyBackupFixture(256)
dumper := &mockDumper{}
store := newMockObjectStore()
svc := newTestBackupService(repo, dumper, store)
compressed := gzipBackupBytes(t, dumpContent)
parts := splitBackupBytes(compressed, 11)
recordParts := make([]BackupPart, 0, len(parts))
for i, data := range parts {
key := fmt.Sprintf("backups/split-missing/payload.part-%06d", i+1)
store.objects[key] = data
recordParts = append(recordParts, BackupPart{
Index: i + 1,
S3Key: key,
SizeBytes: int64(len(data)),
SHA256: fmt.Sprintf("%x", sha256.Sum256(data)),
})
}
delete(store.objects, recordParts[1].S3Key)
record := &BackupRecord{ID: "split-missing", Status: "completed", Parts: recordParts}
require.NoError(t, svc.saveRecord(context.Background(), record))
require.Error(t, svc.RestoreBackup(context.Background(), record.ID))
require.Empty(t, dumper.restored)
}
func TestBackupService_DownloadBackupPartsRejectsMismatchedMetadata(t *testing.T) {
tests := []struct {
name string
part BackupPart
want string
}{
{
name: "size",
part: BackupPart{Index: 1, S3Key: "backups/mismatch/size", SizeBytes: 4},
want: "size mismatch",
},
{
name: "checksum",
part: BackupPart{Index: 1, S3Key: "backups/mismatch/checksum", SizeBytes: 3, SHA256: "bad-checksum"},
want: "checksum mismatch",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
store := newMockObjectStore()
store.objects[tt.part.S3Key] = []byte("abc")
svc := newTestBackupService(repo, &mockDumper{}, store)
_, err := svc.downloadBackupParts(context.Background(), store, []BackupPart{tt.part})
require.ErrorContains(t, err, tt.want)
})
}
}
func gzipBackupBytes(t *testing.T, content []byte) []byte {
t.Helper()
var out bytes.Buffer
writer := gzip.NewWriter(&out)
_, err := writer.Write(content)
require.NoError(t, err)
require.NoError(t, writer.Close())
return out.Bytes()
}
func splitBackupBytes(data []byte, partSize int) [][]byte {
parts := make([][]byte, 0, (len(data)+partSize-1)/partSize)
for len(data) > 0 {
size := partSize
if len(data) < size {
size = len(data)
}
parts = append(parts, append([]byte(nil), data[:size]...))
data = data[size:]
}
return parts
}
func TestBackupService_RestoreBackup_NotCompleted(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
@@ -512,6 +807,34 @@ func TestBackupService_DeleteBackup(t *testing.T) {
require.ErrorIs(t, err, ErrBackupNotFound)
}
func TestBackupService_DeleteBackup_RunningKeepsUploadObjects(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
store := newMockObjectStore()
svc := newTestBackupService(repo, &mockDumper{}, store)
parts := []BackupPart{
{Index: 1, S3Key: "backups/running/payload.part-000001", SizeBytes: 3},
{Index: 2, S3Key: "backups/running/payload.part-000002", SizeBytes: 3},
}
for _, part := range parts {
store.objects[part.S3Key] = []byte("abc")
}
record := &BackupRecord{ID: "running-delete", Status: "running", Parts: parts}
require.NoError(t, svc.saveRecord(context.Background(), record))
err := svc.DeleteBackup(context.Background(), record.ID)
require.ErrorIs(t, err, ErrBackupInProgress)
store.mu.Lock()
require.Empty(t, store.deletedKeys)
for _, part := range parts {
require.Contains(t, store.objects, part.S3Key)
}
store.mu.Unlock()
got, getErr := svc.GetBackupRecord(context.Background(), record.ID)
require.NoError(t, getErr)
require.Equal(t, "running", got.Status)
}
func TestBackupService_GetDownloadURL(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
@@ -523,9 +846,103 @@ func TestBackupService_GetDownloadURL(t *testing.T) {
record, err := svc.CreateBackup(context.Background(), "manual", 14)
require.NoError(t, err)
url, err := svc.GetBackupDownloadURL(context.Background(), record.ID)
download, err := svc.GetBackupDownloadURL(context.Background(), record.ID)
require.NoError(t, err)
require.Contains(t, url, "https://presigned.example.com/")
require.Contains(t, download.URL, "https://presigned.example.com/")
}
func TestBackupService_DeleteBackup_SplitPartsFailureKeepsRecord(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
store := newMockObjectStore()
svc := newTestBackupService(repo, &mockDumper{}, store)
parts := []BackupPart{
{Index: 1, S3Key: "backups/split/payload.part-000001", SizeBytes: 3},
{Index: 2, S3Key: "backups/split/payload.part-000002", SizeBytes: 3},
{Index: 3, S3Key: "backups/split/payload.part-000003", SizeBytes: 3},
}
for _, part := range parts {
store.objects[part.S3Key] = []byte("abc")
}
store.failDeleteKeys[parts[1].S3Key] = fmt.Errorf("delete failed")
record := &BackupRecord{ID: "split-delete", Status: "completed", Parts: parts}
require.NoError(t, svc.saveRecord(context.Background(), record))
err := svc.DeleteBackup(context.Background(), record.ID)
require.Error(t, err)
store.mu.Lock()
deleted := append([]string(nil), store.deletedKeys...)
store.mu.Unlock()
for _, part := range parts {
require.Contains(t, deleted, part.S3Key)
}
got, getErr := svc.GetBackupRecord(context.Background(), record.ID)
require.NoError(t, getErr)
require.Equal(t, record.ID, got.ID)
store.mu.Lock()
require.Contains(t, store.objects, parts[1].S3Key)
store.mu.Unlock()
}
func TestBackupService_GetDownloadURL_SplitParts(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
store := newMockObjectStore()
svc := newTestBackupService(repo, &mockDumper{}, store)
parts := []BackupPart{
{Index: 2, S3Key: "backups/split/payload.part-000002", SizeBytes: 7},
{Index: 1, S3Key: "backups/split/payload.part-000001", SizeBytes: 5},
}
record := &BackupRecord{ID: "split-download", Status: "completed", Parts: parts}
require.NoError(t, svc.saveRecord(context.Background(), record))
download, err := svc.GetBackupDownloadURL(context.Background(), record.ID)
require.NoError(t, err)
require.Empty(t, download.URL)
require.Len(t, download.Parts, 2)
require.Equal(t, 1, download.Parts[0].Index)
require.Equal(t, int64(5), download.Parts[0].SizeBytes)
require.Equal(t, "https://presigned.example.com/backups/split/payload.part-000001", download.Parts[0].URL)
require.Equal(t, 2, download.Parts[1].Index)
}
func TestBackupService_CleanupOldBackups_SplitParts(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
store := newMockObjectStore()
svc := newTestBackupService(repo, &mockDumper{}, store)
now := time.Now()
parts := []BackupPart{
{Index: 1, S3Key: "backups/old/payload.part-000001", SizeBytes: 3},
{Index: 2, S3Key: "backups/old/payload.part-000002", SizeBytes: 3},
}
for _, part := range parts {
store.objects[part.S3Key] = []byte("abc")
}
require.NoError(t, svc.saveRecord(context.Background(), &BackupRecord{
ID: "new",
Status: "completed",
StartedAt: now.Format(time.RFC3339),
}))
require.NoError(t, svc.saveRecord(context.Background(), &BackupRecord{
ID: "old",
Status: "completed",
StartedAt: now.Add(-2 * time.Hour).Format(time.RFC3339),
Parts: parts,
}))
err := svc.cleanupOldBackups(context.Background(), &BackupScheduleConfig{RetainCount: 1})
require.NoError(t, err)
_, err = svc.GetBackupRecord(context.Background(), "old")
require.ErrorIs(t, err, ErrBackupNotFound)
store.mu.Lock()
for _, part := range parts {
require.NotContains(t, store.objects, part.S3Key)
}
store.mu.Unlock()
}
func TestBackupService_ListBackups_Sorted(t *testing.T) {
@@ -690,6 +1107,65 @@ func TestRecoverStaleRecords(t *testing.T) {
require.Contains(t, r2.RestoreError, "server restart")
}
func TestBackupService_RecoverStaleRecords_CleansBackupObjects(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
store := newMockObjectStore()
svc := newTestBackupService(repo, &mockDumper{}, store)
parts := []BackupPart{
{Index: 1, S3Key: "backups/stale/payload.part-000001", SizeBytes: 3},
{Index: 2, S3Key: "backups/stale/payload.part-000002", SizeBytes: 3},
}
for _, part := range parts {
store.objects[part.S3Key] = []byte("abc")
}
require.NoError(t, svc.saveRecord(context.Background(), &BackupRecord{
ID: "stale-parts",
Status: "running",
Parts: parts,
StartedAt: time.Now().Add(-time.Hour).Format(time.RFC3339),
}))
svc.recoverStaleRecords()
record, err := svc.GetBackupRecord(context.Background(), "stale-parts")
require.NoError(t, err)
require.Equal(t, "failed", record.Status)
store.mu.Lock()
defer store.mu.Unlock()
for _, part := range parts {
require.Contains(t, store.deletedKeys, part.S3Key)
require.NotContains(t, store.objects, part.S3Key)
}
}
func TestBackupService_RecoverStaleRecords_PreservesKeysWhenCleanupFails(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
store := newMockObjectStore()
svc := newTestBackupService(repo, &mockDumper{}, store)
part := BackupPart{Index: 1, S3Key: "backups/stale-failed/payload.part-000001", SizeBytes: 3}
store.objects[part.S3Key] = []byte("abc")
store.failDeleteKeys[part.S3Key] = fmt.Errorf("delete failed")
require.NoError(t, svc.saveRecord(context.Background(), &BackupRecord{
ID: "stale-cleanup-failed",
Status: "running",
Parts: []BackupPart{part},
StartedAt: time.Now().Add(-time.Hour).Format(time.RFC3339),
}))
svc.recoverStaleRecords()
record, err := svc.GetBackupRecord(context.Background(), "stale-cleanup-failed")
require.NoError(t, err)
require.Equal(t, "failed", record.Status)
require.Contains(t, record.ErrorMsg, "cleanup failed")
require.Equal(t, part.S3Key, record.Parts[0].S3Key)
store.mu.Lock()
defer store.mu.Unlock()
require.Contains(t, store.objects, part.S3Key)
}
func TestGracefulShutdown(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
@@ -753,3 +1229,38 @@ func TestStartRestore_Async(t *testing.T) {
require.NoError(t, err)
require.Equal(t, "completed", final.RestoreStatus)
}
func TestBackupService_StartRestore_SplitParts(t *testing.T) {
repo := newMockSettingRepo()
seedS3Config(t, repo)
dumpContent := entropyBackupFixture(384)
dumper := &mockDumper{}
store := newMockObjectStore()
svc := newTestBackupService(repo, dumper, store)
compressed := gzipBackupBytes(t, dumpContent)
parts := splitBackupBytes(compressed, 13)
recordParts := make([]BackupPart, 0, len(parts))
for i, data := range parts {
key := fmt.Sprintf("backups/split-async/payload.part-%06d", i+1)
store.objects[key] = data
recordParts = append(recordParts, BackupPart{
Index: i + 1,
S3Key: key,
SizeBytes: int64(len(data)),
SHA256: fmt.Sprintf("%x", sha256.Sum256(data)),
})
}
record := &BackupRecord{ID: "split-async", Status: "completed", Parts: recordParts}
require.NoError(t, svc.saveRecord(context.Background(), record))
started, err := svc.StartRestore(context.Background(), record.ID)
require.NoError(t, err)
require.Equal(t, "running", started.RestoreStatus)
svc.wg.Wait()
final, err := svc.GetBackupRecord(context.Background(), record.ID)
require.NoError(t, err)
require.Equal(t, "completed", final.RestoreStatus)
require.Equal(t, dumpContent, dumper.restored)
}
+21 -2
View File
@@ -23,6 +23,7 @@ export interface BackupRecord {
backup_type: string
file_name: string
s3_key: string
parts?: BackupPart[]
size_bytes: number
triggered_by: string
error_message?: string
@@ -35,6 +36,24 @@ export interface BackupRecord {
restored_at?: string
}
export interface BackupPart {
index: number
s3_key: string
size_bytes: number
sha256?: string
}
export interface BackupDownloadPart {
index: number
size_bytes: number
url: string
}
export interface BackupDownloadResponse {
url?: string
parts?: BackupDownloadPart[]
}
export interface CreateBackupRequest {
expire_days?: number
}
@@ -137,8 +156,8 @@ export async function deleteBackup(id: string): Promise<void> {
await apiClient.delete(`/admin/backups/${id}`)
}
export async function getDownloadURL(id: string): Promise<{ url: string }> {
const { data } = await apiClient.get<{ url: string }>(`/admin/backups/${id}/download-url`)
export async function getDownloadURL(id: string): Promise<BackupDownloadResponse> {
const { data } = await apiClient.get<BackupDownloadResponse>(`/admin/backups/${id}/download-url`)
return data
}
@@ -144,6 +144,7 @@ export default {
status: 'Status',
fileName: 'File Name',
size: 'Size',
parts: 'Parts',
expiresAt: 'Expires At',
triggeredBy: 'Triggered By',
startedAt: 'Started At',
@@ -168,6 +169,10 @@ export default {
empty: 'No backup records',
actions: {
download: 'Download',
downloadParts: 'Download Parts',
downloadPartsHint: 'Download every part in order and concatenate the gzip bytes: on Linux/macOS run cat payload.part-* > backup.sql.gz; on Windows run copy /b payload.part-000001+payload.part-000002 backup.sql.gz.',
partLabel: 'Part {index}',
downloadFailed: 'Download URL is empty',
restore: 'Restore',
restoreConfirm: 'Are you sure you want to restore from this backup? This will overwrite the current database!',
restorePasswordPrompt: 'Please enter your admin password to confirm the restore operation',
@@ -144,6 +144,7 @@ export default {
status: '状态',
fileName: '文件名',
size: '大小',
parts: '分卷数',
expiresAt: '过期时间',
triggeredBy: '触发方式',
startedAt: '开始时间',
@@ -168,6 +169,10 @@ export default {
empty: '暂无备份记录',
actions: {
download: '下载',
downloadParts: '下载分卷',
downloadPartsHint: '请按顺序下载全部分卷后拼接 gzip 字节流:Linux/macOS 使用 cat payload.part-* > backup.sql.gzWindows 使用 copy /b payload.part-000001+payload.part-000002 backup.sql.gz。',
partLabel: '第 {index} 卷',
downloadFailed: '下载地址为空',
restore: '恢复',
restoreConfirm: '确定要从此备份恢复吗?这将覆盖当前数据库!',
restorePasswordPrompt: '请输入管理员密码以确认恢复操作',
+62 -1
View File
@@ -200,6 +200,7 @@
<th class="py-2 pr-4">{{ t('admin.backup.columns.status') }}</th>
<th class="py-2 pr-4">{{ t('admin.backup.columns.fileName') }}</th>
<th class="py-2 pr-4">{{ t('admin.backup.columns.size') }}</th>
<th class="py-2 pr-4">{{ t('admin.backup.columns.parts') }}</th>
<th class="py-2 pr-4">{{ t('admin.backup.columns.expiresAt') }}</th>
<th class="py-2 pr-4">{{ t('admin.backup.columns.triggeredBy') }}</th>
<th class="py-2 pr-4">{{ t('admin.backup.columns.startedAt') }}</th>
@@ -221,6 +222,7 @@
</td>
<td class="py-3 pr-4 text-xs">{{ record.file_name }}</td>
<td class="py-3 pr-4 text-xs">{{ formatSize(record.size_bytes) }}</td>
<td class="py-3 pr-4 text-xs">{{ record.parts?.length || (record.status === 'running' ? '-' : 1) }}</td>
<td class="py-3 pr-4 text-xs">
{{ record.expires_at ? formatDate(record.expires_at) : t('admin.backup.neverExpire') }}
</td>
@@ -248,6 +250,7 @@
{{ restoringId === record.id ? t('common.loading') : t('admin.backup.actions.restore') }}
</button>
<button
v-if="record.status !== 'running'"
type="button"
class="btn btn-danger btn-xs"
@click="removeBackup(record.id)"
@@ -258,7 +261,7 @@
</td>
</tr>
<tr v-if="backups.length === 0">
<td colspan="8" class="py-6 text-center text-sm text-gray-500 dark:text-gray-400">
<td colspan="9" class="py-6 text-center text-sm text-gray-500 dark:text-gray-400">
{{ t('admin.backup.empty') }}
</td>
</tr>
@@ -351,6 +354,48 @@
</div>
</transition>
</teleport>
<!-- 分卷下载链接 -->
<teleport to="body">
<transition name="modal">
<div
v-if="downloadPartsModalOpen"
class="fixed inset-0 z-50 flex items-center justify-center p-4"
@mousedown.self="closeDownloadParts"
>
<div class="fixed inset-0 bg-black/50" @click="closeDownloadParts"></div>
<div class="relative max-h-[85vh] w-full max-w-lg overflow-y-auto rounded-xl bg-white p-6 shadow-2xl dark:bg-dark-800">
<button
type="button"
class="absolute right-4 top-4 text-gray-400 hover:text-gray-600 dark:hover:text-gray-200"
:aria-label="t('common.close')"
@click="closeDownloadParts"
>
<svg class="h-5 w-5" fill="none" viewBox="0 0 24 24" stroke="currentColor" stroke-width="2"><path stroke-linecap="round" stroke-linejoin="round" d="M6 18L18 6M6 6l12 12" /></svg>
</button>
<h2 class="mb-1 text-lg font-bold text-gray-900 dark:text-white">{{ t('admin.backup.actions.downloadParts') }}</h2>
<p class="mb-4 text-sm text-gray-500 dark:text-gray-400">{{ t('admin.backup.actions.downloadPartsHint') }}</p>
<div class="space-y-2">
<div
v-for="part in downloadParts"
:key="part.index"
class="flex items-center justify-between gap-3 rounded-lg border border-gray-200 px-3 py-2 dark:border-dark-600"
>
<span class="text-sm text-gray-700 dark:text-gray-300">
{{ t('admin.backup.actions.partLabel', { index: part.index }) }}
<span class="ml-2 text-xs text-gray-500 dark:text-gray-400">{{ formatSize(part.size_bytes) }}</span>
</span>
<a :href="part.url" class="btn btn-secondary btn-xs" rel="noopener">
{{ t('admin.backup.actions.download') }}
</a>
</div>
</div>
<div class="mt-4 text-right">
<button type="button" class="btn btn-primary btn-sm" @click="closeDownloadParts">{{ t('common.close') }}</button>
</div>
</div>
</div>
</transition>
</teleport>
<TotpStepUpDialog :controller="backupStepUp" />
</template>
@@ -363,6 +408,7 @@ import type {
BackupS3Config,
BackupScheduleConfig,
BackupRecord,
BackupDownloadPart,
ImageStorageConfig,
} from '@/api/admin/backup'
import { useStepUp, isStepUpBlocked, isStepUpCancelled, stepUpBlockReason } from '@/composables/useStepUp'
@@ -432,6 +478,8 @@ const loadingBackups = ref(false)
const creatingBackup = ref(false)
const restoringId = ref('')
const manualExpireDays = ref(14)
const downloadParts = ref<BackupDownloadPart[]>([])
const downloadPartsModalOpen = ref(false)
// Polling
const pollingTimer = ref<ReturnType<typeof setInterval> | null>(null)
@@ -715,6 +763,14 @@ async function createBackup() {
async function downloadBackup(id: string) {
try {
const result = await backupStepUp.run(() => adminAPI.backup.getDownloadURL(id))
if (result.parts && result.parts.length > 0) {
downloadParts.value = result.parts
downloadPartsModalOpen.value = true
return
}
if (!result.url) {
throw new Error(t('admin.backup.actions.downloadFailed'))
}
// 预签名 URL 带 attachment disposition,同页 anchor 导航直接触发下载;
// 不用 window.openstep-up 弹窗 await 会耗尽瞬态用户激活,新标签页会被浏览器拦截。
const link = document.createElement('a')
@@ -728,6 +784,11 @@ async function downloadBackup(id: string) {
}
}
function closeDownloadParts() {
downloadPartsModalOpen.value = false
downloadParts.value = []
}
async function restoreBackup(id: string) {
if (!window.confirm(t('admin.backup.actions.restoreConfirm'))) return
const password = window.prompt(t('admin.backup.actions.restorePasswordPrompt'))
@@ -0,0 +1,154 @@
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
import { flushPromises, mount } from '@vue/test-utils'
import BackupView from '../BackupView.vue'
const {
getS3Config,
getImageStorageConfig,
getSchedule,
listBackups,
getDownloadURL,
} = vi.hoisted(() => ({
getS3Config: vi.fn(),
getImageStorageConfig: vi.fn(),
getSchedule: vi.fn(),
listBackups: vi.fn(),
getDownloadURL: vi.fn(),
}))
vi.mock('@/api', () => ({
adminAPI: {
backup: {
getS3Config,
updateS3Config: vi.fn(),
testS3Connection: vi.fn(),
getImageStorageConfig,
updateImageStorageConfig: vi.fn(),
testImageStorageConnection: vi.fn(),
getSchedule,
updateSchedule: vi.fn(),
createBackup: vi.fn(),
listBackups,
getBackup: vi.fn(),
deleteBackup: vi.fn(),
getDownloadURL,
restoreBackup: vi.fn(),
},
},
}))
vi.mock('@/stores', () => ({
useAppStore: () => ({
showError: vi.fn(),
showSuccess: vi.fn(),
showWarning: vi.fn(),
}),
}))
vi.mock('@/composables/useStepUp', () => ({
useStepUp: () => ({ run: (fn: () => unknown) => fn() }),
isStepUpBlocked: () => false,
isStepUpCancelled: () => false,
stepUpBlockReason: () => '',
}))
vi.mock('vue-i18n', () => ({
useI18n: () => ({
t: (key: string, params?: Record<string, unknown>) =>
params?.index === undefined ? key : `${key}:${params.index}`,
}),
}))
const baseRecord = (id: string, parts?: unknown[]) => ({
id,
status: 'completed',
backup_type: 'postgres',
file_name: `${id}.sql.gz`,
s3_key: `backups/${id}.sql.gz`,
parts,
size_bytes: 10,
triggered_by: 'manual',
started_at: '2026-08-09T00:00:00Z',
})
function mountBackupView() {
return mount(BackupView, {
global: {
stubs: {
TotpStepUpDialog: true,
},
},
})
}
describe('admin BackupView 分卷备份', () => {
beforeEach(() => {
getS3Config.mockResolvedValue({})
getImageStorageConfig.mockResolvedValue({ config: {}, secret_configured: false })
getSchedule.mockResolvedValue({ enabled: false, cron_expr: '', retain_days: 14, retain_count: 10 })
getDownloadURL.mockReset()
vi.spyOn(HTMLAnchorElement.prototype, 'click').mockImplementation(() => {})
})
afterEach(() => {
vi.restoreAllMocks()
document.body.innerHTML = ''
})
it('显示分卷数并在下载时列出每个分卷链接', async () => {
listBackups.mockResolvedValue({
items: [baseRecord('split', [{ index: 1 }, { index: 2 }, { index: 3 }])],
})
getDownloadURL.mockResolvedValue({
parts: [
{ index: 1, size_bytes: 5, url: 'https://example.test/part-1' },
{ index: 2, size_bytes: 6, url: 'https://example.test/part-2' },
{ index: 3, size_bytes: 7, url: 'https://example.test/part-3' },
],
})
const wrapper = mountBackupView()
await flushPromises()
expect(wrapper.text()).toContain('3')
const downloadButton = wrapper.findAll('button').find(button =>
button.text().includes('admin.backup.actions.download'),
)
expect(downloadButton).toBeDefined()
await downloadButton!.trigger('click')
await flushPromises()
expect(document.body.textContent).toContain('admin.backup.actions.partLabel:1')
expect(document.body.textContent).toContain('admin.backup.actions.partLabel:3')
expect(document.body.querySelector('a[href="https://example.test/part-2"]')).not.toBeNull()
})
it('旧单文件记录仍使用单个下载地址', async () => {
listBackups.mockResolvedValue({ items: [baseRecord('legacy')] })
getDownloadURL.mockResolvedValue({ url: 'https://example.test/legacy.sql.gz' })
const wrapper = mountBackupView()
await flushPromises()
const downloadButton = wrapper.findAll('button').find(button =>
button.text().includes('admin.backup.actions.download'),
)
await downloadButton!.trigger('click')
await flushPromises()
expect(getDownloadURL).toHaveBeenCalledWith('legacy')
expect(document.body.textContent).not.toContain('admin.backup.actions.downloadParts')
})
it('运行中的备份不显示删除入口', async () => {
listBackups.mockResolvedValue({
items: [{ ...baseRecord('running'), status: 'running', progress: 'uploading' }],
})
const wrapper = mountBackupView()
await flushPromises()
expect(wrapper.find('tbody tr td:nth-child(5)').text()).toBe('-')
expect(wrapper.findAll('button').some(button => button.text() === 'common.delete')).toBe(false)
})
})