Merge pull request #137 from Awuqing/codex/architecture-simplification

refactor: simplify architecture and harden lifecycle
This commit is contained in:
Wu Qing
2026-08-26 11:48:33 +08:00
committed by GitHub
189 changed files with 8983 additions and 4264 deletions
Vendored
BIN
View File
Binary file not shown.
+81 -3
View File
@@ -26,7 +26,7 @@ jobs:
- name: Set up Go - name: Set up Go
uses: actions/setup-go@v7 uses: actions/setup-go@v7
with: with:
go-version: '1.25' go-version: "1.25"
cache-dependency-path: server/go.sum cache-dependency-path: server/go.sum
- name: Verify modules - name: Verify modules
@@ -54,6 +54,31 @@ jobs:
working-directory: server working-directory: server
run: go test ./... -v run: go test ./... -v
backend-windows:
name: Go Test (Windows)
runs-on: windows-latest
timeout-minutes: 20
steps:
- uses: actions/checkout@v7
- name: Set up Go
uses: actions/setup-go@v7
with:
go-version: "1.25"
cache-dependency-path: server/go.sum
- name: Verify modules
working-directory: server
run: go mod verify
- name: Build
working-directory: server
run: go build ./...
- name: Test
working-directory: server
run: go test ./...
frontend: frontend:
name: React Build & Test name: React Build & Test
runs-on: ubuntu-latest runs-on: ubuntu-latest
@@ -64,14 +89,22 @@ jobs:
- name: Set up Node.js - name: Set up Node.js
uses: actions/setup-node@v7 uses: actions/setup-node@v7
with: with:
node-version: '24' node-version: "24"
cache: 'npm' cache: "npm"
cache-dependency-path: web/package-lock.json cache-dependency-path: web/package-lock.json
- name: Install dependencies - name: Install dependencies
working-directory: web working-directory: web
run: npm ci run: npm ci
- name: Audit production dependencies
working-directory: web
run: npm audit --omit=dev --audit-level=high
- name: Lint and formatting
working-directory: web
run: npm run lint && npm run format:check
- name: Test - name: Test
working-directory: web working-directory: web
run: npm run test run: npm run test
@@ -79,3 +112,48 @@ jobs:
- name: Type Check & Build - name: Type Check & Build
working-directory: web working-directory: web
run: npm run build run: npm run build
docs:
name: Documentation Build
runs-on: ubuntu-latest
timeout-minutes: 20
steps:
- uses: actions/checkout@v7
- name: Set up Node.js
uses: actions/setup-node@v7
with:
node-version: "24"
cache: "npm"
cache-dependency-path: docs-site/package-lock.json
- name: Install dependencies
working-directory: docs-site
run: npm ci
- name: Type Check & Build
working-directory: docs-site
run: npm run typecheck && npm run build
container:
name: Container Build
needs: [backend, backend-windows, frontend, docs]
runs-on: ubuntu-latest
timeout-minutes: 30
steps:
- uses: actions/checkout@v7
- name: Validate Compose configuration
run: docker compose config --quiet
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v4
- name: Build image
uses: docker/build-push-action@v7
with:
context: .
push: false
tags: backupx:ci
cache-from: type=gha,scope=backupx-ci
cache-to: type=gha,mode=max,scope=backupx-ci
+17 -14
View File
@@ -18,11 +18,11 @@ name: Release
on: on:
push: push:
tags: tags:
- 'v*' - "v*"
workflow_dispatch: workflow_dispatch:
inputs: inputs:
version: version:
description: '版本号(如 v1.2.3' description: "版本号(如 v1.2.3"
required: true required: true
type: string type: string
@@ -65,7 +65,7 @@ jobs:
- name: Set up Go - name: Set up Go
uses: actions/setup-go@v7 uses: actions/setup-go@v7
with: with:
go-version: '1.25' go-version: "1.25"
cache-dependency-path: server/go.sum cache-dependency-path: server/go.sum
- name: Verify backend - name: Verify backend
@@ -83,14 +83,17 @@ jobs:
- name: Set up Node.js - name: Set up Node.js
uses: actions/setup-node@v7 uses: actions/setup-node@v7
with: with:
node-version: '24' node-version: "24"
cache: 'npm' cache: "npm"
cache-dependency-path: web/package-lock.json cache-dependency-path: web/package-lock.json
- name: Verify frontend - name: Verify frontend
working-directory: web working-directory: web
run: | run: |
npm ci npm ci
npm audit --omit=dev --audit-level=high
npm run lint
npm run format:check
npm run test npm run test
# ─── Job 1: 构建前端 ─── # ─── Job 1: 构建前端 ───
@@ -105,8 +108,8 @@ jobs:
- name: Set up Node.js - name: Set up Node.js
uses: actions/setup-node@v7 uses: actions/setup-node@v7
with: with:
node-version: '24' node-version: "24"
cache: 'npm' cache: "npm"
cache-dependency-path: web/package-lock.json cache-dependency-path: web/package-lock.json
- name: Install & Build - name: Install & Build
@@ -143,7 +146,7 @@ jobs:
- name: Set up Go - name: Set up Go
uses: actions/setup-go@v7 uses: actions/setup-go@v7
with: with:
go-version: '1.25' go-version: "1.25"
cache-dependency-path: server/go.sum cache-dependency-path: server/go.sum
- name: Download frontend artifact - name: Download frontend artifact
@@ -157,7 +160,7 @@ jobs:
env: env:
GOOS: ${{ matrix.goos }} GOOS: ${{ matrix.goos }}
GOARCH: ${{ matrix.goarch }} GOARCH: ${{ matrix.goarch }}
CGO_ENABLED: '0' CGO_ENABLED: "0"
run: | run: |
go build \ go build \
-trimpath \ -trimpath \
@@ -172,13 +175,13 @@ jobs:
cp backupx "${ARCHIVE_NAME}/" cp backupx "${ARCHIVE_NAME}/"
cp -r web/dist "${ARCHIVE_NAME}/web" cp -r web/dist "${ARCHIVE_NAME}/web"
cp server/config.example.yaml "${ARCHIVE_NAME}/" cp server/config.example.yaml "${ARCHIVE_NAME}/"
cp deploy/install.sh "${ARCHIVE_NAME}/" 2>/dev/null || true cp deploy/install.sh "${ARCHIVE_NAME}/"
cp deploy/backupx.service "${ARCHIVE_NAME}/" 2>/dev/null || true cp deploy/backupx.service "${ARCHIVE_NAME}/"
# v2.2+: 随发布包提供 Grafana dashboard 与 nginx.conf 模板 # v2.2+: 随发布包提供 Grafana dashboard 与 nginx.conf 模板
if [ -d deploy/grafana ]; then if [ -d deploy/grafana ]; then
cp -r deploy/grafana "${ARCHIVE_NAME}/grafana" cp -r deploy/grafana "${ARCHIVE_NAME}/grafana"
fi fi
cp deploy/nginx.conf "${ARCHIVE_NAME}/nginx.conf" 2>/dev/null || true cp deploy/nginx.conf "${ARCHIVE_NAME}/nginx.conf"
tar czf "${ARCHIVE_NAME}.tar.gz" "${ARCHIVE_NAME}" tar czf "${ARCHIVE_NAME}.tar.gz" "${ARCHIVE_NAME}"
cp "${ARCHIVE_NAME}.tar.gz" "backupx-${{ matrix.goos }}-${{ matrix.goarch }}.tar.gz" cp "${ARCHIVE_NAME}.tar.gz" "backupx-${{ matrix.goos }}-${{ matrix.goarch }}.tar.gz"
sha256sum "${ARCHIVE_NAME}.tar.gz" > "${ARCHIVE_NAME}.tar.gz.sha256" sha256sum "${ARCHIVE_NAME}.tar.gz" > "${ARCHIVE_NAME}.tar.gz.sha256"
@@ -199,7 +202,7 @@ jobs:
# ─── Job 3: Docker 多架构 → Docker Hub ─── # ─── Job 3: Docker 多架构 → Docker Hub ───
build-docker: build-docker:
name: Build & Push Docker name: Build & Push Docker
needs: build-web needs: verify
runs-on: ubuntu-latest runs-on: ubuntu-latest
timeout-minutes: 45 timeout-minutes: 45
steps: steps:
@@ -226,7 +229,7 @@ jobs:
build-args: | build-args: |
VERSION=${{ env.VERSION }} VERSION=${{ env.VERSION }}
tags: | tags: |
${{ secrets.DOCKERHUB_USERNAME }}/backupx:latest
${{ secrets.DOCKERHUB_USERNAME }}/backupx:${{ env.VERSION }} ${{ secrets.DOCKERHUB_USERNAME }}/backupx:${{ env.VERSION }}
${{ !contains(env.VERSION, '-') && format('{0}/backupx:latest', secrets.DOCKERHUB_USERNAME) || '' }}
cache-from: type=gha cache-from: type=gha
cache-to: type=gha,mode=max cache-to: type=gha,mode=max
+2 -2
View File
@@ -10,7 +10,7 @@ ARG USE_CHINA_MIRROR=false
# ---- Stage 1: Build frontend ---- # ---- Stage 1: Build frontend ----
FROM node:26-alpine AS web-builder FROM node:24-alpine AS web-builder
ARG USE_CHINA_MIRROR ARG USE_CHINA_MIRROR
# 国内镜像:npm 使用淘宝源 # 国内镜像:npm 使用淘宝源
@@ -26,7 +26,7 @@ RUN npm run build
# ---- Stage 2: Build backend ---- # ---- Stage 2: Build backend ----
FROM golang:1.26-alpine AS server-builder FROM golang:1.25-alpine AS server-builder
ARG USE_CHINA_MIRROR ARG USE_CHINA_MIRROR
ARG VERSION=dev ARG VERSION=dev
+2 -1
View File
@@ -51,7 +51,8 @@ verify: verify-server verify-web verify-docs
verify-server: format-check vet-server test-server build-server verify-server: format-check vet-server test-server build-server
verify-web: test-web build-web verify-web:
cd web && npm run lint && npm run format:check && npm run test && npm run build
verify-docs: check-docs verify-docs: check-docs
+7 -7
View File
@@ -25,17 +25,17 @@ For significant features or refactors, open an issue first to align on scope bef
## Pull requests ## Pull requests
1. Fork and create a topic branch (e.g. `fix/windows-path-escape`) 1. Fork and create a topic branch (e.g. `fix/windows-path-escape`)
2. Run `make test` and make sure everything passes 2. Run `make verify` and make sure formatting, tests, builds, and documentation checks pass
3. Keep changes focused — one concern per PR 3. Keep changes focused — one concern per PR
4. Write commit messages in Chinese following `类型: 简要描述` — examples: 4. Write Conventional Commits with a Chinese subject — examples:
- `功能: 新增审计日志模块` - `feat(audit): 新增审计日志模块`
- `修复: 目录浏览器无法进入子目录` - `fix(browser): 修复目录浏览器无法进入子目录`
- `重构: 简化存储目标解密逻辑` - `refactor(storage): 简化存储目标解密逻辑`
- Types: `功能` / `修复` / `重构` / `文档` / `构建` / `测试` - Types: `feat` / `fix` / `docs` / `style` / `refactor` / `perf` / `test` / `chore`
5. PR title and body in Chinese too. Describe the why and how, not just the what. 5. PR title and body in Chinese too. Describe the why and how, not just the what.
## Coding guidelines ## Coding guidelines
- **Go** — handle every error (no `_ = err`); use the existing logger (`zap`); no `fmt.Println` in production paths - **Go** — handle every error (no `_ = err`); use the existing logger (`zap`); no `fmt.Println` in production paths
- **TypeScript** — strict mode, no implicit any, follow existing ESLint/Prettier configs - **TypeScript** — strict mode, no implicit any, and pass the repository ESLint and Prettier checks
- **Commit scope** — one logical change per commit; don't mix drive-by cleanups with feature work - **Commit scope** — one logical change per commit; don't mix drive-by cleanups with feature work
@@ -25,17 +25,17 @@ BackupX 使用 Apache License 2.0 开源,欢迎提交 Issue 与 Pull Request
## 提交 PR ## 提交 PR
1. Fork 仓库,创建主题分支(如 `fix/windows-path-escape` 1. Fork 仓库,创建主题分支(如 `fix/windows-path-escape`
2. 执行 `make test` 确认本地全通过 2. 执行 `make verify`确认格式、测试、构建和文档检查全部通过
3. 保持每个 PR 只做一件事 3. 保持每个 PR 只做一件事
4. Commit message 使用中文,格式 `类型: 简要描述` 4. Commit message 使用 Conventional Commits,主题使用中文
- `功能: 新增审计日志模块` - `feat(audit): 新增审计日志模块`
- `修复: 目录浏览器无法进入子目录` - `fix(browser): 修复目录浏览器无法进入子目录`
- `重构: 简化存储目标解密逻辑` - `refactor(storage): 简化存储目标解密逻辑`
- 类型:`功能` / `修复` / `重构` / `文档` / `构建` / `测试` - 类型:`feat` / `fix` / `docs` / `style` / `refactor` / `perf` / `test` / `chore`
5. PR 标题和正文同样使用中文,描述"为什么"和"怎么做",而非仅仅"做了什么" 5. PR 标题和正文同样使用中文,描述"为什么"和"怎么做",而非仅仅"做了什么"
## 代码规范 ## 代码规范
- **Go** — 所有错误必须处理(禁止 `_ = err`),日志使用现有 `zap`,禁止生产路径中出现 `fmt.Println` - **Go** — 所有错误必须处理(禁止 `_ = err`),日志使用现有 `zap`,禁止生产路径中出现 `fmt.Println`
- **TypeScript** — 严格模式,禁止隐式 any,遵循现有 ESLint/Prettier 配置 - **TypeScript** — 严格模式,禁止隐式 any,并通过仓库中的 ESLintPrettier 检查
- **Commit 粒度** — 每个 commit 一件事,不要把顺手的小修改和功能代码混在一起 - **Commit 粒度** — 每个 commit 一件事,不要把顺手的小修改和功能代码混在一起
BIN
View File
Binary file not shown.
+3 -8
View File
@@ -1,15 +1,10 @@
APP_NAME=backupx
BUILD_DIR=./bin
VERSION=$(shell git describe --tags --always --dirty 2>/dev/null || echo "dev")
.PHONY: build run test .PHONY: build run test
build: build:
mkdir -p $(BUILD_DIR) $(MAKE) -C .. build-server
go build -trimpath -ldflags "-s -w -X main.version=$(VERSION)" -o $(BUILD_DIR)/$(APP_NAME) ./cmd/backupx
run: run:
go run -ldflags "-X main.version=$(VERSION)" ./cmd/backupx $(MAKE) -C .. dev-server
test: test:
go test ./... $(MAKE) -C .. test-server
+14
View File
@@ -9,6 +9,8 @@ import (
"syscall" "syscall"
"backupx/server/internal/agent" "backupx/server/internal/agent"
"backupx/server/internal/config"
applogger "backupx/server/internal/logger"
) )
// runAgent 是 `backupx agent` 子命令入口。 // runAgent 是 `backupx agent` 子命令入口。
@@ -59,11 +61,23 @@ func runAgent(args []string) {
os.Exit(2) os.Exit(2)
} }
agentLogger, err := applogger.New(config.LogConfig{Level: "info"})
if err != nil {
fmt.Fprintf(os.Stderr, "agent: init logger: %v\n", err)
os.Exit(1)
}
defer func() {
if syncErr := agentLogger.Sync(); syncErr != nil {
fmt.Fprintf(os.Stderr, "agent: flush logger: %v\n", syncErr)
}
}()
a, err := agent.New(cfg, version) a, err := agent.New(cfg, version)
if err != nil { if err != nil {
fmt.Fprintf(os.Stderr, "agent: init: %v\n", err) fmt.Fprintf(os.Stderr, "agent: init: %v\n", err)
os.Exit(1) os.Exit(1)
} }
a.SetLogger(agentLogger)
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
defer stop() defer stop()
+71 -27
View File
@@ -4,7 +4,6 @@ import (
"context" "context"
"encoding/json" "encoding/json"
"fmt" "fmt"
"log"
"net" "net"
"os" "os"
"runtime" "runtime"
@@ -13,6 +12,7 @@ import (
"time" "time"
"backupx/server/internal/backup" "backupx/server/internal/backup"
"go.uber.org/zap"
) )
// Agent 是 Agent 进程的主控制器。 // Agent 是 Agent 进程的主控制器。
@@ -21,6 +21,7 @@ type Agent struct {
client *MasterClient client *MasterClient
executor *Executor executor *Executor
version string version string
logger *zap.Logger
mu sync.Mutex mu sync.Mutex
started bool started bool
@@ -44,9 +45,20 @@ func New(cfg *Config, version string) (*Agent, error) {
client: client, client: client,
executor: executor, executor: executor,
version: version, version: version,
logger: zap.NewNop(),
}, nil }, nil
} }
// SetLogger attaches the process logger used by the Agent runtime loop.
func (a *Agent) SetLogger(logger *zap.Logger) {
if logger != nil {
a.logger = logger
if a.executor != nil {
a.executor.SetLogger(logger)
}
}
}
// Run 启动 Agent 主循环,阻塞直到 ctx 被取消。 // Run 启动 Agent 主循环,阻塞直到 ctx 被取消。
func (a *Agent) Run(ctx context.Context) error { func (a *Agent) Run(ctx context.Context) error {
a.mu.Lock() a.mu.Lock()
@@ -64,7 +76,7 @@ func (a *Agent) Run(ctx context.Context) error {
if err := a.heartbeatOnce(ctx); err != nil { if err := a.heartbeatOnce(ctx); err != nil {
return fmt.Errorf("initial heartbeat failed: %w", err) return fmt.Errorf("initial heartbeat failed: %w", err)
} }
log.Printf("[agent] connected to master %s", a.cfg.Master) a.logger.Info("agent connected to master", zap.String("master", a.cfg.Master))
var wg sync.WaitGroup var wg sync.WaitGroup
wg.Add(2) wg.Add(2)
@@ -90,14 +102,17 @@ func (a *Agent) heartbeatLoop(ctx context.Context, interval time.Duration) {
return return
case <-ticker.C: case <-ticker.C:
if err := a.heartbeatOnce(ctx); err != nil { if err := a.heartbeatOnce(ctx); err != nil {
log.Printf("[agent] heartbeat failed: %v", err) a.logger.Warn("agent heartbeat failed", zap.Error(err))
} }
} }
} }
} }
func (a *Agent) heartbeatOnce(ctx context.Context) error { func (a *Agent) heartbeatOnce(ctx context.Context) error {
hostname, _ := os.Hostname() hostname, err := os.Hostname()
if err != nil {
a.logger.Warn("resolve agent hostname failed", zap.Error(err))
}
req := HeartbeatRequest{ req := HeartbeatRequest{
Hostname: hostname, Hostname: hostname,
IPAddress: detectLocalIP(), IPAddress: detectLocalIP(),
@@ -105,7 +120,7 @@ func (a *Agent) heartbeatOnce(ctx context.Context) error {
OS: runtime.GOOS, OS: runtime.GOOS,
Arch: runtime.GOARCH, Arch: runtime.GOARCH,
} }
_, err := a.client.Heartbeat(ctx, req) _, err = a.client.Heartbeat(ctx, req)
return err return err
} }
@@ -126,13 +141,13 @@ func (a *Agent) pollLoop(ctx context.Context, interval time.Duration) {
func (a *Agent) pollAndHandleOnce(ctx context.Context) { func (a *Agent) pollAndHandleOnce(ctx context.Context) {
cmd, err := a.client.PollCommand(ctx) cmd, err := a.client.PollCommand(ctx)
if err != nil { if err != nil {
log.Printf("[agent] poll command failed: %v", err) a.logger.Warn("poll agent command failed", zap.Error(err))
return return
} }
if cmd == nil { if cmd == nil {
return return
} }
log.Printf("[agent] received command #%d type=%s", cmd.ID, cmd.Type) a.logger.Info("agent command received", zap.Uint("command_id", cmd.ID), zap.String("command_type", cmd.Type))
switch cmd.Type { switch cmd.Type {
case "run_task": case "run_task":
a.handleRunTask(ctx, cmd) a.handleRunTask(ctx, cmd)
@@ -146,8 +161,8 @@ func (a *Agent) pollAndHandleOnce(ctx context.Context) {
a.handleDeleteStorageObject(ctx, cmd) a.handleDeleteStorageObject(ctx, cmd)
default: default:
msg := fmt.Sprintf("unknown command type: %s", cmd.Type) msg := fmt.Sprintf("unknown command type: %s", cmd.Type)
log.Printf("[agent] %s", msg) a.logger.Warn("unknown agent command", zap.Uint("command_id", cmd.ID), zap.String("command_type", cmd.Type))
_ = a.client.SubmitCommandResult(ctx, cmd.ID, false, msg, nil) a.submitCommandResult(ctx, cmd.ID, false, msg, nil)
} }
} }
@@ -158,14 +173,14 @@ func (a *Agent) handleRunTask(ctx context.Context, cmd *CommandPayload) {
RecordID uint `json:"recordId"` RecordID uint `json:"recordId"`
} }
if err := json.Unmarshal(cmd.Payload, &payload); err != nil { if err := json.Unmarshal(cmd.Payload, &payload); err != nil {
_ = a.client.SubmitCommandResult(ctx, cmd.ID, false, "invalid payload: "+err.Error(), nil) a.submitCommandResult(ctx, cmd.ID, false, "invalid payload: "+err.Error(), nil)
return return
} }
if err := a.executor.ExecuteRunTask(ctx, payload.TaskID, payload.RecordID); err != nil { if err := a.executor.ExecuteRunTask(ctx, payload.TaskID, payload.RecordID); err != nil {
_ = a.client.SubmitCommandResult(ctx, cmd.ID, false, err.Error(), nil) a.submitCommandResult(ctx, cmd.ID, false, err.Error(), nil)
return return
} }
_ = a.client.SubmitCommandResult(ctx, cmd.ID, true, "", map[string]any{ a.submitCommandResult(ctx, cmd.ID, true, "", map[string]any{
"taskId": payload.TaskID, "taskId": payload.TaskID,
"recordId": payload.RecordID, "recordId": payload.RecordID,
}) })
@@ -177,18 +192,18 @@ func (a *Agent) handleRestoreRecord(ctx context.Context, cmd *CommandPayload) {
RestoreRecordID uint `json:"restoreRecordId"` RestoreRecordID uint `json:"restoreRecordId"`
} }
if err := json.Unmarshal(cmd.Payload, &payload); err != nil { if err := json.Unmarshal(cmd.Payload, &payload); err != nil {
_ = a.client.SubmitCommandResult(ctx, cmd.ID, false, "invalid payload: "+err.Error(), nil) a.submitCommandResult(ctx, cmd.ID, false, "invalid payload: "+err.Error(), nil)
return return
} }
if payload.RestoreRecordID == 0 { if payload.RestoreRecordID == 0 {
_ = a.client.SubmitCommandResult(ctx, cmd.ID, false, "restoreRecordId is required", nil) a.submitCommandResult(ctx, cmd.ID, false, "restoreRecordId is required", nil)
return return
} }
if err := a.executor.ExecuteRestore(ctx, payload.RestoreRecordID); err != nil { if err := a.executor.ExecuteRestore(ctx, payload.RestoreRecordID); err != nil {
_ = a.client.SubmitCommandResult(ctx, cmd.ID, false, err.Error(), nil) a.submitCommandResult(ctx, cmd.ID, false, err.Error(), nil)
return return
} }
_ = a.client.SubmitCommandResult(ctx, cmd.ID, true, "", map[string]any{ a.submitCommandResult(ctx, cmd.ID, true, "", map[string]any{
"restoreRecordId": payload.RestoreRecordID, "restoreRecordId": payload.RestoreRecordID,
}) })
} }
@@ -202,23 +217,23 @@ func (a *Agent) handleDeleteStorageObject(ctx context.Context, cmd *CommandPaylo
StoragePath string `json:"storagePath"` StoragePath string `json:"storagePath"`
} }
if err := json.Unmarshal(cmd.Payload, &payload); err != nil { if err := json.Unmarshal(cmd.Payload, &payload); err != nil {
_ = a.client.SubmitCommandResult(ctx, cmd.ID, false, "invalid payload: "+err.Error(), nil) a.submitCommandResult(ctx, cmd.ID, false, "invalid payload: "+err.Error(), nil)
return return
} }
if strings.TrimSpace(payload.StoragePath) == "" { if strings.TrimSpace(payload.StoragePath) == "" {
_ = a.client.SubmitCommandResult(ctx, cmd.ID, false, "storagePath is required", nil) a.submitCommandResult(ctx, cmd.ID, false, "storagePath is required", nil)
return return
} }
provider, err := a.executor.storageRegistry.Create(ctx, payload.TargetType, payload.TargetConfig) provider, err := a.executor.storageRegistry.Create(ctx, payload.TargetType, payload.TargetConfig)
if err != nil { if err != nil {
_ = a.client.SubmitCommandResult(ctx, cmd.ID, false, "create provider: "+err.Error(), nil) a.submitCommandResult(ctx, cmd.ID, false, "create provider: "+err.Error(), nil)
return return
} }
if err := provider.Delete(ctx, payload.StoragePath); err != nil { if err := provider.Delete(ctx, payload.StoragePath); err != nil {
_ = a.client.SubmitCommandResult(ctx, cmd.ID, false, "delete object: "+err.Error(), nil) a.submitCommandResult(ctx, cmd.ID, false, "delete object: "+err.Error(), nil)
return return
} }
_ = a.client.SubmitCommandResult(ctx, cmd.ID, true, "", map[string]any{"deleted": true}) a.submitCommandResult(ctx, cmd.ID, true, "", map[string]any{"deleted": true})
} }
// handleDiscoverDB 处理 discover_db 命令:在 Agent 本机执行 mysql/psql 列出数据库。 // handleDiscoverDB 处理 discover_db 命令:在 Agent 本机执行 mysql/psql 列出数据库。
@@ -231,7 +246,7 @@ func (a *Agent) handleDiscoverDB(ctx context.Context, cmd *CommandPayload) {
Password string `json:"password"` Password string `json:"password"`
} }
if err := json.Unmarshal(cmd.Payload, &payload); err != nil { if err := json.Unmarshal(cmd.Payload, &payload); err != nil {
_ = a.client.SubmitCommandResult(ctx, cmd.ID, false, "invalid payload: "+err.Error(), nil) a.submitCommandResult(ctx, cmd.ID, false, "invalid payload: "+err.Error(), nil)
return return
} }
databases, err := backup.DiscoverDatabases(ctx, backup.NewOSCommandExecutor(), backup.DiscoverRequest{ databases, err := backup.DiscoverDatabases(ctx, backup.NewOSCommandExecutor(), backup.DiscoverRequest{
@@ -242,10 +257,10 @@ func (a *Agent) handleDiscoverDB(ctx context.Context, cmd *CommandPayload) {
Password: payload.Password, Password: payload.Password,
}) })
if err != nil { if err != nil {
_ = a.client.SubmitCommandResult(ctx, cmd.ID, false, err.Error(), nil) a.submitCommandResult(ctx, cmd.ID, false, err.Error(), nil)
return return
} }
_ = a.client.SubmitCommandResult(ctx, cmd.ID, true, "", map[string]any{"databases": databases}) a.submitCommandResult(ctx, cmd.ID, true, "", map[string]any{"databases": databases})
} }
// handleListDir 处理 list_dir 命令(阶段四实现) // handleListDir 处理 list_dir 命令(阶段四实现)
@@ -254,15 +269,44 @@ func (a *Agent) handleListDir(ctx context.Context, cmd *CommandPayload) {
Path string `json:"path"` Path string `json:"path"`
} }
if err := json.Unmarshal(cmd.Payload, &payload); err != nil { if err := json.Unmarshal(cmd.Payload, &payload); err != nil {
_ = a.client.SubmitCommandResult(ctx, cmd.ID, false, "invalid payload: "+err.Error(), nil) a.submitCommandResult(ctx, cmd.ID, false, "invalid payload: "+err.Error(), nil)
return return
} }
entries, err := listLocalDir(payload.Path) entries, err := listLocalDir(payload.Path)
if err != nil { if err != nil {
_ = a.client.SubmitCommandResult(ctx, cmd.ID, false, err.Error(), nil) a.submitCommandResult(ctx, cmd.ID, false, err.Error(), nil)
return return
} }
_ = a.client.SubmitCommandResult(ctx, cmd.ID, true, "", map[string]any{"entries": entries}) a.submitCommandResult(ctx, cmd.ID, true, "", map[string]any{"entries": entries})
}
func (a *Agent) submitCommandResult(ctx context.Context, commandID uint, success bool, message string, data any) {
reportCtx, cancel := agentFinalizationContext(ctx)
defer cancel()
var err error
attempts := 0
retryLoop:
for attempts < 3 {
attempts++
err = a.client.SubmitCommandResult(reportCtx, commandID, success, message, data)
if err == nil {
return
}
if attempts < 3 {
timer := time.NewTimer(time.Duration(attempts) * 100 * time.Millisecond)
select {
case <-reportCtx.Done():
timer.Stop()
break retryLoop
case <-timer.C:
}
}
}
a.logger.Error("submit agent command result failed",
zap.Uint("command_id", commandID),
zap.Bool("success", success),
zap.Int("attempts", attempts),
zap.Error(err))
} }
// 辅助函数 // 辅助函数
+36
View File
@@ -0,0 +1,36 @@
package agent
import (
"context"
"net/http"
"net/http/httptest"
"sync/atomic"
"testing"
"go.uber.org/zap"
)
func TestSubmitCommandResultRetriesWithCanceledCommandContext(t *testing.T) {
var requests atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
if requests.Add(1) < 3 {
http.Error(w, "temporarily unavailable", http.StatusServiceUnavailable)
return
}
w.WriteHeader(http.StatusOK)
}))
defer server.Close()
agent := &Agent{
client: NewMasterClient(server.URL, "token", false),
logger: zap.NewNop(),
}
ctx, cancel := context.WithCancel(context.Background())
cancel()
agent.submitCommandResult(ctx, 17, false, "backup failed", nil)
if got := requests.Load(); got != 3 {
t.Fatalf("submit requests = %d, want 3", got)
}
}
+79 -38
View File
@@ -17,6 +17,7 @@ import (
"backupx/server/internal/storage" "backupx/server/internal/storage"
storageRclone "backupx/server/internal/storage/rclone" storageRclone "backupx/server/internal/storage/rclone"
"backupx/server/pkg/compress" "backupx/server/pkg/compress"
"go.uber.org/zap"
) )
// Executor 负责在 Agent 本地执行命令。 // Executor 负责在 Agent 本地执行命令。
@@ -25,35 +26,26 @@ type Executor struct {
tempDir string tempDir string
backupRegistry *backup.Registry backupRegistry *backup.Registry
storageRegistry *storage.Registry storageRegistry *storage.Registry
logger *zap.Logger
} }
// NewExecutor 构造执行器。预先初始化 backup runner 与 storage registry。 // NewExecutor 构造执行器。预先初始化 backup runner 与 storage registry。
func NewExecutor(client *MasterClient, tempDir string) *Executor { func NewExecutor(client *MasterClient, tempDir string) *Executor {
backupRegistry := backup.NewRegistry( backupRegistry := backup.NewDefaultRegistry()
backup.NewFileRunner(), storageRegistry := storageRclone.NewDefaultRegistry()
backup.NewSQLiteRunner(),
backup.NewMySQLRunner(nil),
backup.NewPostgreSQLRunner(nil),
backup.NewSAPHANARunner(nil),
backup.NewMongoDBRunner(nil),
)
storageRegistry := storage.NewRegistry(
storageRclone.NewLocalDiskFactory(),
storageRclone.NewS3Factory(),
storageRclone.NewWebDAVFactory(),
storageRclone.NewGoogleDriveFactory(),
storageRclone.NewAliyunOSSFactory(),
storageRclone.NewTencentCOSFactory(),
storageRclone.NewQiniuKodoFactory(),
storageRclone.NewFTPFactory(),
storageRclone.NewRcloneFactory(),
)
storageRclone.RegisterAllBackends(storageRegistry)
return &Executor{ return &Executor{
client: client, client: client,
tempDir: tempDir, tempDir: tempDir,
backupRegistry: backupRegistry, backupRegistry: backupRegistry,
storageRegistry: storageRegistry, storageRegistry: storageRegistry,
logger: zap.NewNop(),
}
}
// SetLogger attaches the Agent process logger to execution and reporting paths.
func (e *Executor) SetLogger(logger *zap.Logger) {
if logger != nil {
e.logger = logger
} }
} }
@@ -90,7 +82,7 @@ func (e *Executor) ExecuteRunTask(ctx context.Context, taskID, recordID uint) er
} }
// 3) 运行 runner // 3) 运行 runner
logger := newRecordLogger(ctx, e.client, recordID) logger := newRecordLogger(ctx, e.client, e.logger, recordID)
result, err := runner.Run(ctx, backupSpec, logger) result, err := runner.Run(ctx, backupSpec, logger)
if err != nil { if err != nil {
e.reportRecordFailure(ctx, recordID, err.Error()) e.reportRecordFailure(ctx, recordID, err.Error())
@@ -177,7 +169,9 @@ func (e *Executor) ExecuteRunTask(ctx context.Context, taskID, recordID uint) er
} }
// 6) 上报最终成功 // 6) 上报最终成功
return e.client.UpdateRecord(ctx, recordID, RecordUpdate{ reportCtx, cancel := agentFinalizationContext(ctx)
defer cancel()
if err := e.client.UpdateRecord(reportCtx, recordID, RecordUpdate{
Status: "success", Status: "success",
FileName: fileName, FileName: fileName,
FileSize: fileSize, FileSize: fileSize,
@@ -187,7 +181,10 @@ func (e *Executor) ExecuteRunTask(ctx context.Context, taskID, recordID uint) er
StorageTransferMode: selectedStorageTransferMode, StorageTransferMode: selectedStorageTransferMode,
StorageUploadResults: uploadResults, StorageUploadResults: uploadResults,
LogAppend: fmt.Sprintf("[agent] 任务完成,总计 %d 字节\n", fileSize), LogAppend: fmt.Sprintf("[agent] 任务完成,总计 %d 字节\n", fileSize),
}) }); err != nil {
return fmt.Errorf("report backup success to master: %w", err)
}
return nil
} }
// uploadToTarget 上传单个目标。为保持简化不做上传级重试(rclone 本身已有 low-level 重试)。 // uploadToTarget 上传单个目标。为保持简化不做上传级重试(rclone 本身已有 low-level 重试)。
@@ -222,7 +219,11 @@ func (e *Executor) uploadToTarget(ctx context.Context, recordID uint, target Sto
// appendLog 追加日志到 Master 记录(尽力而为,失败不中断主流程) // appendLog 追加日志到 Master 记录(尽力而为,失败不中断主流程)
func (e *Executor) appendLog(ctx context.Context, recordID uint, line string) { func (e *Executor) appendLog(ctx context.Context, recordID uint, line string) {
_ = e.client.UpdateRecord(ctx, recordID, RecordUpdate{LogAppend: line}) if err := e.client.UpdateRecord(ctx, recordID, RecordUpdate{LogAppend: line}); err != nil {
e.logger.Warn("append backup record log to master failed",
zap.Uint("record_id", recordID),
zap.Error(err))
}
} }
// reportRecordFailure 上报失败状态 // reportRecordFailure 上报失败状态
@@ -231,12 +232,27 @@ func (e *Executor) reportRecordFailure(ctx context.Context, recordID uint, msg s
} }
func (e *Executor) reportRecordFailureWithUploadResults(ctx context.Context, recordID uint, msg string, uploadResults []StorageResultItem) { func (e *Executor) reportRecordFailureWithUploadResults(ctx context.Context, recordID uint, msg string, uploadResults []StorageResultItem) {
_ = e.client.UpdateRecord(ctx, recordID, RecordUpdate{ reportCtx, cancel := agentFinalizationContext(ctx)
defer cancel()
if err := e.client.UpdateRecord(reportCtx, recordID, RecordUpdate{
Status: "failed", Status: "failed",
ErrorMessage: msg, ErrorMessage: msg,
StorageUploadResults: uploadResults, StorageUploadResults: uploadResults,
LogAppend: fmt.Sprintf("[agent] 错误: %s\n", msg), LogAppend: fmt.Sprintf("[agent] 错误: %s\n", msg),
}) }); err != nil {
e.logger.Error("report backup failure to master failed",
zap.Uint("record_id", recordID),
zap.Error(err))
}
}
// agentFinalizationContext lets terminal state reach the Master even when the
// command context was canceled, while bounding shutdown/network delays.
func agentFinalizationContext(ctx context.Context) (context.Context, context.CancelFunc) {
if ctx == nil {
ctx = context.Background()
}
return context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second)
} }
// buildBackupTaskSpec 把 AgentTaskSpec 转换为 backup.TaskSpec。 // buildBackupTaskSpec 把 AgentTaskSpec 转换为 backup.TaskSpec。
@@ -301,30 +317,40 @@ func compactStringList(items []string) []string {
type recordLogger struct { type recordLogger struct {
ctx context.Context ctx context.Context
client *MasterClient client *MasterClient
logger *zap.Logger
recordID uint recordID uint
} }
func newRecordLogger(ctx context.Context, client *MasterClient, recordID uint) *recordLogger { func newRecordLogger(ctx context.Context, client *MasterClient, logger *zap.Logger, recordID uint) *recordLogger {
return &recordLogger{ctx: ctx, client: client, recordID: recordID} return &recordLogger{ctx: ctx, client: client, logger: logger, recordID: recordID}
} }
func (l *recordLogger) WriteLine(message string) { func (l *recordLogger) WriteLine(message string) {
_ = l.client.UpdateRecord(l.ctx, l.recordID, RecordUpdate{LogAppend: message + "\n"}) if err := l.client.UpdateRecord(l.ctx, l.recordID, RecordUpdate{LogAppend: message + "\n"}); err != nil {
l.logger.Warn("append backup runner log to master failed",
zap.Uint("record_id", l.recordID),
zap.Error(err))
}
} }
// restoreLogger 把 runner 日志回传到 Master 恢复记录。 // restoreLogger 把 runner 日志回传到 Master 恢复记录。
type restoreLogger struct { type restoreLogger struct {
ctx context.Context ctx context.Context
client *MasterClient client *MasterClient
logger *zap.Logger
restoreID uint restoreID uint
} }
func newRestoreLogger(ctx context.Context, client *MasterClient, restoreID uint) *restoreLogger { func newRestoreLogger(ctx context.Context, client *MasterClient, logger *zap.Logger, restoreID uint) *restoreLogger {
return &restoreLogger{ctx: ctx, client: client, restoreID: restoreID} return &restoreLogger{ctx: ctx, client: client, logger: logger, restoreID: restoreID}
} }
func (l *restoreLogger) WriteLine(message string) { func (l *restoreLogger) WriteLine(message string) {
_ = l.client.UpdateRestore(l.ctx, l.restoreID, RestoreUpdate{LogAppend: message + "\n"}) if err := l.client.UpdateRestore(l.ctx, l.restoreID, RestoreUpdate{LogAppend: message + "\n"}); err != nil {
l.logger.Warn("append restore runner log to master failed",
zap.Uint("restore_record_id", l.restoreID),
zap.Error(err))
}
} }
// DeleteStorageObject 在 Agent 本机上删除指定存储对象(供跨节点清理调用)。 // DeleteStorageObject 在 Agent 本机上删除指定存储对象(供跨节点清理调用)。
@@ -453,29 +479,44 @@ func (e *Executor) ExecuteRestore(ctx context.Context, restoreRecordID uint) err
e.reportRestoreFailure(ctx, restoreRecordID, fmt.Sprintf("不支持的备份类型: %v", err)) e.reportRestoreFailure(ctx, restoreRecordID, fmt.Sprintf("不支持的备份类型: %v", err))
return err return err
} }
logger := newRestoreLogger(ctx, e.client, restoreRecordID) logger := newRestoreLogger(ctx, e.client, e.logger, restoreRecordID)
if err := runner.Restore(ctx, taskSpec, preparedPath, logger); err != nil { if err := runner.Restore(ctx, taskSpec, preparedPath, logger); err != nil {
e.reportRestoreFailure(ctx, restoreRecordID, err.Error()) e.reportRestoreFailure(ctx, restoreRecordID, err.Error())
return err return err
} }
// 5) 上报成功 // 5) 上报成功
return e.client.UpdateRestore(ctx, restoreRecordID, RestoreUpdate{ reportCtx, cancel := agentFinalizationContext(ctx)
defer cancel()
if err := e.client.UpdateRestore(reportCtx, restoreRecordID, RestoreUpdate{
Status: "success", Status: "success",
LogAppend: "[agent] 恢复执行完成\n", LogAppend: "[agent] 恢复执行完成\n",
}) }); err != nil {
return fmt.Errorf("report restore success to master: %w", err)
}
return nil
} }
func (e *Executor) appendRestoreLog(ctx context.Context, restoreID uint, line string) { func (e *Executor) appendRestoreLog(ctx context.Context, restoreID uint, line string) {
_ = e.client.UpdateRestore(ctx, restoreID, RestoreUpdate{LogAppend: line}) if err := e.client.UpdateRestore(ctx, restoreID, RestoreUpdate{LogAppend: line}); err != nil {
e.logger.Warn("append restore log to master failed",
zap.Uint("restore_record_id", restoreID),
zap.Error(err))
}
} }
func (e *Executor) reportRestoreFailure(ctx context.Context, restoreID uint, msg string) { func (e *Executor) reportRestoreFailure(ctx context.Context, restoreID uint, msg string) {
_ = e.client.UpdateRestore(ctx, restoreID, RestoreUpdate{ reportCtx, cancel := agentFinalizationContext(ctx)
defer cancel()
if err := e.client.UpdateRestore(reportCtx, restoreID, RestoreUpdate{
Status: "failed", Status: "failed",
ErrorMessage: msg, ErrorMessage: msg,
LogAppend: fmt.Sprintf("[agent] 错误: %s\n", msg), LogAppend: fmt.Sprintf("[agent] 错误: %s\n", msg),
}) }); err != nil {
e.logger.Error("report restore failure to master failed",
zap.Uint("restore_record_id", restoreID),
zap.Error(err))
}
} }
// buildRestoreBackupTaskSpec 把 RestoreSpec 转成 backup.TaskSpec。 // buildRestoreBackupTaskSpec 把 RestoreSpec 转成 backup.TaskSpec。
+29
View File
@@ -18,8 +18,35 @@ import (
"time" "time"
"backupx/server/internal/storage" "backupx/server/internal/storage"
"go.uber.org/zap"
"go.uber.org/zap/zapcore"
"go.uber.org/zap/zaptest/observer"
) )
func TestReportRecordFailureUsesFinalizationContextAndLogsUpdateError(t *testing.T) {
requestCount := 0
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
requestCount++
http.Error(w, "master unavailable", http.StatusServiceUnavailable)
}))
defer server.Close()
core, observed := observer.New(zapcore.ErrorLevel)
executor := NewExecutor(NewMasterClient(server.URL, "token", false), t.TempDir())
executor.SetLogger(zap.New(core))
ctx, cancel := context.WithCancel(context.Background())
cancel()
executor.reportRecordFailure(ctx, 42, "backup failed")
if requestCount != 1 {
t.Fatalf("terminal update requests = %d, want 1 despite canceled command context", requestCount)
}
if observed.Len() != 1 || observed.All()[0].Message != "report backup failure to master failed" {
t.Fatalf("observed logs = %#v", observed.All())
}
}
func TestBuildBackupTaskSpecParsesJSONSourcePaths(t *testing.T) { func TestBuildBackupTaskSpecParsesJSONSourcePaths(t *testing.T) {
spec := &TaskSpec{ spec := &TaskSpec{
TaskID: 7, TaskID: 7,
@@ -323,6 +350,8 @@ func (f *agentTestStorageFactory) Type() storage.ProviderType {
return "agent_test_storage" return "agent_test_storage"
} }
func (f *agentTestStorageFactory) SensitiveFields() []string { return nil }
func (f *agentTestStorageFactory) New(_ context.Context, config map[string]any) (storage.StorageProvider, error) { func (f *agentTestStorageFactory) New(_ context.Context, config map[string]any) (storage.StorageProvider, error) {
name, _ := config["name"].(string) name, _ := config["name"].(string)
provider := f.providers[name] provider := f.providers[name]
+13 -1
View File
@@ -20,7 +20,7 @@ type DirEntry struct {
func listLocalDir(path string) ([]DirEntry, error) { func listLocalDir(path string) ([]DirEntry, error) {
cleaned := filepath.Clean(strings.TrimSpace(path)) cleaned := filepath.Clean(strings.TrimSpace(path))
if strings.TrimSpace(path) == "" || cleaned == "." { if strings.TrimSpace(path) == "" || cleaned == "." {
cleaned = "/" cleaned = localFilesystemRoot()
} }
entries, err := os.ReadDir(cleaned) entries, err := os.ReadDir(cleaned)
if err != nil { if err != nil {
@@ -48,3 +48,15 @@ func listLocalDir(path string) ([]DirEntry, error) {
}) })
return result, nil return result, nil
} }
func localFilesystemRoot() string {
root := string(os.PathSeparator)
workingDir, err := os.Getwd()
if err != nil {
return root
}
if volume := filepath.VolumeName(workingDir); volume != "" {
return volume + root
}
return root
}
+103 -33
View File
@@ -5,6 +5,7 @@ import (
"errors" "errors"
"fmt" "fmt"
stdhttp "net/http" stdhttp "net/http"
"sync"
"time" "time"
"backupx/server/internal/backup" "backupx/server/internal/backup"
@@ -12,6 +13,7 @@ import (
"backupx/server/internal/config" "backupx/server/internal/config"
"backupx/server/internal/database" "backupx/server/internal/database"
aphttp "backupx/server/internal/http" aphttp "backupx/server/internal/http"
"backupx/server/internal/lifecycle"
"backupx/server/internal/logger" "backupx/server/internal/logger"
"backupx/server/internal/metrics" "backupx/server/internal/metrics"
"backupx/server/internal/notify" "backupx/server/internal/notify"
@@ -19,7 +21,6 @@ import (
"backupx/server/internal/scheduler" "backupx/server/internal/scheduler"
"backupx/server/internal/security" "backupx/server/internal/security"
"backupx/server/internal/service" "backupx/server/internal/service"
"backupx/server/internal/storage"
"backupx/server/internal/storage/codec" "backupx/server/internal/storage/codec"
storageRclone "backupx/server/internal/storage/rclone" storageRclone "backupx/server/internal/storage/rclone"
"go.uber.org/zap" "go.uber.org/zap"
@@ -33,6 +34,10 @@ type Application struct {
db *gorm.DB db *gorm.DB
httpServer *stdhttp.Server httpServer *stdhttp.Server
scheduler *scheduler.Service scheduler *scheduler.Service
background *lifecycle.Supervisor
shutdownMu sync.Mutex
closeOnce sync.Once
} }
func New(ctx context.Context, cfg config.Config, version string) (*Application, error) { func New(ctx context.Context, cfg config.Config, version string) (*Application, error) {
@@ -55,34 +60,30 @@ func New(ctx context.Context, cfg config.Config, version string) (*Application,
oauthSessionRepo := repository.NewOAuthSessionRepository(db) oauthSessionRepo := repository.NewOAuthSessionRepository(db)
resolvedSecurity, err := service.ResolveSecurity(ctx, cfg.Security, systemConfigRepo) resolvedSecurity, err := service.ResolveSecurity(ctx, cfg.Security, systemConfigRepo)
if err != nil { if err != nil {
return nil, fmt.Errorf("resolve security config: %w", err) resolveErr := fmt.Errorf("resolve security config: %w", err)
if sqlDB, handleErr := db.DB(); handleErr != nil {
resolveErr = errors.Join(resolveErr, fmt.Errorf("get database handle for cleanup: %w", handleErr))
} else if closeErr := sqlDB.Close(); closeErr != nil {
resolveErr = errors.Join(resolveErr, fmt.Errorf("close database after bootstrap failure: %w", closeErr))
}
return nil, resolveErr
} }
background := lifecycle.NewSupervisor(ctx)
jwtManager := security.NewJWTManager(resolvedSecurity.JWTSecret, config.MustJWTDuration(cfg.Security)) jwtManager := security.NewJWTManager(resolvedSecurity.JWTSecret, config.MustJWTDuration(cfg.Security))
rateLimiter := security.NewLoginRateLimiter(5, time.Minute) rateLimiter := security.NewLoginRateLimiter(5, time.Minute)
configCipher := codec.NewConfigCipher(resolvedSecurity.EncryptionKey) configCipher := codec.NewConfigCipher(resolvedSecurity.EncryptionKey)
authService := service.NewAuthService(userRepo, systemConfigRepo, jwtManager, rateLimiter, configCipher) authService := service.NewAuthService(userRepo, systemConfigRepo, jwtManager, rateLimiter, configCipher)
systemService := service.NewSystemService(cfg, version, time.Now().UTC()) systemService := service.NewSystemService(cfg, version, time.Now().UTC())
storageRegistry := storage.NewRegistry( storageRegistry := storageRclone.NewDefaultRegistry()
storageRclone.NewLocalDiskFactory(),
storageRclone.NewS3Factory(),
storageRclone.NewWebDAVFactory(),
storageRclone.NewGoogleDriveFactory(),
storageRclone.NewAliyunOSSFactory(),
storageRclone.NewTencentCOSFactory(),
storageRclone.NewQiniuKodoFactory(),
storageRclone.NewFTPFactory(),
storageRclone.NewRcloneFactory(),
)
// 将全部 rclone 后端注册为独立存储类型(sftp、azureblob、dropbox 等与 s3、ftp 完全平级)
storageRclone.RegisterAllBackends(storageRegistry)
storageTargetService := service.NewStorageTargetService(storageTargetRepo, oauthSessionRepo, storageRegistry, configCipher) storageTargetService := service.NewStorageTargetService(storageTargetRepo, oauthSessionRepo, storageRegistry, configCipher)
storageTargetService.SetBackgroundRunner(background)
storageTargetService.SetBackupTaskRepository(backupTaskRepo) storageTargetService.SetBackupTaskRepository(backupTaskRepo)
storageTargetService.SetBackupRecordRepository(backupRecordRepo) storageTargetService.SetBackupRecordRepository(backupRecordRepo)
backupTaskService := service.NewBackupTaskService(backupTaskRepo, storageTargetRepo, configCipher) backupTaskService := service.NewBackupTaskService(backupTaskRepo, storageTargetRepo, configCipher)
backupTaskService.SetRecordsAndStorage(backupRecordRepo, storageRegistry) backupTaskService.SetRecordsAndStorage(backupRecordRepo, storageRegistry)
// nodeRepo 在下方 Cluster 节点管理区块才实例化,这里延后注入 // nodeRepo 在下方 Cluster 节点管理区块才实例化,这里延后注入
backupRunnerRegistry := backup.NewRegistry(backup.NewFileRunner(), backup.NewSQLiteRunner(), backup.NewMySQLRunner(nil), backup.NewPostgreSQLRunner(nil), backup.NewSAPHANARunner(nil), backup.NewMongoDBRunner(nil)) backupRunnerRegistry := backup.NewDefaultRegistry()
logHub := backup.NewLogHub() logHub := backup.NewLogHub()
retentionService := backupretention.NewService(backupRecordRepo, configCipher.Key()) retentionService := backupretention.NewService(backupRecordRepo, configCipher.Key())
notifyRegistry := notify.NewRegistry(notify.NewEmailNotifier(), notify.NewWebhookNotifier(), notify.NewTelegramNotifier()) notifyRegistry := notify.NewRegistry(notify.NewEmailNotifier(), notify.NewWebhookNotifier(), notify.NewTelegramNotifier())
@@ -96,6 +97,7 @@ func New(ctx context.Context, cfg config.Config, version string) (*Application,
storageRclone.StartAccounting(rcloneCtx) storageRclone.StartAccounting(rcloneCtx)
backupExecutionService := service.NewBackupExecutionService(backupTaskRepo, backupRecordRepo, storageTargetRepo, storageRegistry, backupRunnerRegistry, logHub, retentionService, configCipher, notificationService, cfg.Backup.TempDir, cfg.Backup.MaxConcurrent, cfg.Backup.Retries, cfg.Backup.BandwidthLimit) backupExecutionService := service.NewBackupExecutionService(backupTaskRepo, backupRecordRepo, storageTargetRepo, storageRegistry, backupRunnerRegistry, logHub, retentionService, configCipher, notificationService, cfg.Backup.TempDir, cfg.Backup.MaxConcurrent, cfg.Backup.Retries, cfg.Backup.BandwidthLimit)
backupExecutionService.SetBackgroundRunner(background)
schedulerService := scheduler.NewService(backupTaskRepo, backupExecutionService, appLogger) schedulerService := scheduler.NewService(backupTaskRepo, backupExecutionService, appLogger)
backupTaskService.SetScheduler(schedulerService) backupTaskService.SetScheduler(schedulerService)
// 审计日志注入延迟到 auditService 创建后(见下方) // 审计日志注入延迟到 auditService 创建后(见下方)
@@ -104,12 +106,15 @@ func New(ctx context.Context, cfg config.Config, version string) (*Application,
restoreRecordRepo := repository.NewRestoreRecordRepository(db) restoreRecordRepo := repository.NewRestoreRecordRepository(db)
restoreLogHub := backup.NewLogHub() restoreLogHub := backup.NewLogHub()
dashboardService := service.NewDashboardService(backupTaskRepo, backupRecordRepo, storageTargetRepo) dashboardService := service.NewDashboardService(backupTaskRepo, backupRecordRepo, storageTargetRepo)
dashboardService.SetBackgroundRunner(background)
reportService := service.NewReportService(backupTaskRepo, backupRecordRepo) reportService := service.NewReportService(backupTaskRepo, backupRecordRepo)
settingsService := service.NewSettingsService(systemConfigRepo) settingsService := service.NewSettingsService(systemConfigRepo)
// Audit // Audit
auditLogRepo := repository.NewAuditLogRepository(db) auditLogRepo := repository.NewAuditLogRepository(db)
auditService := service.NewAuditService(auditLogRepo) auditService := service.NewAuditService(auditLogRepo)
auditService.SetLogger(appLogger)
auditService.SetBackgroundRunner(background)
authService.SetAuditService(auditService) authService.SetAuditService(auditService)
schedulerService.SetAuditRecorder(auditService) schedulerService.SetAuditRecorder(auditService)
// 审计日志外输:启动时用当前 settings 初始化 webhook,后续前端修改立即生效 // 审计日志外输:启动时用当前 settings 初始化 webhook,后续前端修改立即生效
@@ -125,6 +130,7 @@ func New(ctx context.Context, cfg config.Config, version string) (*Application,
backupTaskService.SetNodeRepository(nodeRepo) backupTaskService.SetNodeRepository(nodeRepo)
schedulerService.SetNodeRepository(nodeRepo) schedulerService.SetNodeRepository(nodeRepo)
nodeService := service.NewNodeService(nodeRepo, version) nodeService := service.NewNodeService(nodeRepo, version)
nodeService.SetBackgroundRunner(background)
nodeService.SetTaskRepository(backupTaskRepo) nodeService.SetTaskRepository(backupTaskRepo)
if err := nodeService.EnsureLocalNode(ctx); err != nil { if err := nodeService.EnsureLocalNode(ctx); err != nil {
appLogger.Warn("failed to ensure local node", zap.Error(err)) appLogger.Warn("failed to ensure local node", zap.Error(err))
@@ -136,12 +142,15 @@ func New(ctx context.Context, cfg config.Config, version string) (*Application,
agentCmdRepo := repository.NewAgentCommandRepository(db) agentCmdRepo := repository.NewAgentCommandRepository(db)
nodeService.SetAgentCommandRepository(agentCmdRepo) nodeService.SetAgentCommandRepository(agentCmdRepo)
agentService := service.NewAgentService(nodeRepo, backupTaskRepo, backupRecordRepo, storageTargetRepo, agentCmdRepo, configCipher, storageRegistry) agentService := service.NewAgentService(nodeRepo, backupTaskRepo, backupRecordRepo, storageTargetRepo, agentCmdRepo, configCipher, storageRegistry)
agentService.SetLogger(appLogger)
agentService.SetBackgroundRunner(background)
agentService.SetRestoreRepository(restoreRecordRepo) agentService.SetRestoreRepository(restoreRecordRepo)
agentService.StartCommandTimeoutMonitor(ctx, 30*time.Second, 10*time.Minute) agentService.StartCommandTimeoutMonitor(ctx, 30*time.Second, 10*time.Minute)
// 一键部署:install token service + 后台 GC // 一键部署:install token service + 后台 GC
installTokenRepo := repository.NewAgentInstallTokenRepository(db) installTokenRepo := repository.NewAgentInstallTokenRepository(db)
installTokenService := service.NewInstallTokenService(installTokenRepo, nodeRepo) installTokenService := service.NewInstallTokenService(installTokenRepo, nodeRepo)
installTokenService.SetBackgroundRunner(background)
installTokenService.StartGC(ctx, time.Hour) installTokenService.StartGC(ctx, time.Hour)
// 把 Agent 下发能力注入到备份执行服务,实现多节点路由 // 把 Agent 下发能力注入到备份执行服务,实现多节点路由
@@ -166,6 +175,7 @@ func New(ctx context.Context, cfg config.Config, version string) (*Application,
cfg.Backup.TempDir, cfg.Backup.TempDir,
cfg.Backup.MaxConcurrent, cfg.Backup.MaxConcurrent,
) )
restoreService.SetBackgroundRunner(background)
// 验证服务:定期校验备份可恢复性(企业合规刚需) // 验证服务:定期校验备份可恢复性(企业合规刚需)
verificationRecordRepo := repository.NewVerificationRecordRepository(db) verificationRecordRepo := repository.NewVerificationRecordRepository(db)
@@ -182,6 +192,7 @@ func New(ctx context.Context, cfg config.Config, version string) (*Application,
cfg.Backup.TempDir, cfg.Backup.TempDir,
cfg.Backup.MaxConcurrent, cfg.Backup.MaxConcurrent,
) )
verificationService.SetBackgroundRunner(background)
// 验证失败通知:通过 NotificationService 的事件总线派发 verify_failed // 验证失败通知:通过 NotificationService 的事件总线派发 verify_failed
verificationService.SetNotifier(service.NewVerificationEventNotifier(notificationService)) verificationService.SetNotifier(service.NewVerificationEventNotifier(notificationService))
// 恢复完成/失败事件派发(restore_success / restore_failed // 恢复完成/失败事件派发(restore_success / restore_failed
@@ -206,6 +217,8 @@ func New(ctx context.Context, cfg config.Config, version string) (*Application,
nodeRepo, storageRegistry, configCipher, nodeRepo, storageRegistry, configCipher,
cfg.Backup.TempDir, cfg.Backup.MaxConcurrent, cfg.Backup.TempDir, cfg.Backup.MaxConcurrent,
) )
replicationService.SetLogger(appLogger)
replicationService.SetBackgroundRunner(background)
replicationService.SetEventDispatcher(notificationService) replicationService.SetEventDispatcher(notificationService)
backupExecutionService.SetReplicationTrigger(replicationService) backupExecutionService.SetReplicationTrigger(replicationService)
// 备份成功后触发下游依赖任务(任务依赖链工作流) // 备份成功后触发下游依赖任务(任务依赖链工作流)
@@ -229,6 +242,7 @@ func New(ctx context.Context, cfg config.Config, version string) (*Application,
// 集群版本监控:每 30 分钟扫描,节点 24 小时内只告警一次 // 集群版本监控:每 30 分钟扫描,节点 24 小时内只告警一次
clusterVersionMonitor := service.NewClusterVersionMonitor(nodeRepo, version) clusterVersionMonitor := service.NewClusterVersionMonitor(nodeRepo, version)
clusterVersionMonitor.SetEventDispatcher(notificationService) clusterVersionMonitor.SetEventDispatcher(notificationService)
clusterVersionMonitor.SetBackgroundRunner(background)
clusterVersionMonitor.Start(ctx, 30*time.Minute, 24*time.Hour) clusterVersionMonitor.Start(ctx, 30*time.Minute, 24*time.Hour)
// Dashboard 集群概览依赖注入 // Dashboard 集群概览依赖注入
@@ -247,6 +261,7 @@ func New(ctx context.Context, cfg config.Config, version string) (*Application,
metrics.NewRepoSource(storageTargetRepo, backupRecordRepo, nodeRepo, backupTaskRepo, agentCmdRepo), metrics.NewRepoSource(storageTargetRepo, backupRecordRepo, nodeRepo, backupTaskRepo, agentCmdRepo),
30*time.Second, 30*time.Second,
) )
metricsCollector.SetBackgroundRunner(background)
metricsCollector.Start(ctx) metricsCollector.Start(ctx)
router := aphttp.NewRouter(aphttp.RouterDependencies{ router := aphttp.NewRouter(aphttp.RouterDependencies{
@@ -299,13 +314,21 @@ func New(ctx context.Context, cfg config.Config, version string) (*Application,
db: db, db: db,
httpServer: httpServer, httpServer: httpServer,
scheduler: schedulerService, scheduler: schedulerService,
background: background,
}, nil }, nil
} }
func (a *Application) Run(ctx context.Context) error { func (a *Application) Run(ctx context.Context) error {
if a.scheduler != nil { if a.scheduler != nil {
if err := a.scheduler.Start(context.Background()); err != nil { runCtx := ctx
return fmt.Errorf("start scheduler: %w", err) if a.background != nil {
runCtx = a.background.Context()
}
if err := a.scheduler.Start(runCtx); err != nil {
shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
shutdownErr := a.Shutdown(shutdownCtx)
return errors.Join(fmt.Errorf("start scheduler: %w", err), shutdownErr)
} }
} }
errCh := make(chan error, 1) errCh := make(chan error, 1)
@@ -320,30 +343,77 @@ func (a *Application) Run(ctx context.Context) error {
select { select {
case <-ctx.Done(): case <-ctx.Done():
shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) shutdownCtx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
defer cancel() defer cancel()
a.logger.Info("shutdown signal received") a.logger.Info("shutdown signal received")
if err := a.httpServer.Shutdown(shutdownCtx); err != nil { return a.Shutdown(shutdownCtx)
return fmt.Errorf("shutdown http server: %w", err)
}
if a.scheduler != nil {
if err := a.scheduler.Stop(context.Background()); err != nil {
return fmt.Errorf("stop scheduler: %w", err)
}
}
return nil
case err := <-errCh: case err := <-errCh:
shutdownCtx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
defer cancel()
shutdownErr := a.Shutdown(shutdownCtx)
if err != nil { if err != nil {
return fmt.Errorf("serve http: %w", err) return errors.Join(fmt.Errorf("serve http: %w", err), shutdownErr)
} }
return nil return shutdownErr
} }
} }
func (a *Application) Close() { // Shutdown stops new scheduled and HTTP work before canceling and waiting for
if a.logger != nil { // application-owned background tasks. Every phase is attempted even if an
_ = a.logger.Sync() // earlier phase fails, so a timeout cannot leave workers detached.
func (a *Application) Shutdown(ctx context.Context) error {
if a == nil {
return nil
} }
if ctx == nil {
ctx = context.Background()
}
a.shutdownMu.Lock()
defer a.shutdownMu.Unlock()
var shutdownErrors []error
if a.scheduler != nil {
if err := a.scheduler.Stop(ctx); err != nil {
shutdownErrors = append(shutdownErrors, fmt.Errorf("stop scheduler: %w", err))
}
}
if a.httpServer != nil {
if err := a.httpServer.Shutdown(ctx); err != nil {
shutdownErrors = append(shutdownErrors, fmt.Errorf("shutdown http server: %w", err))
}
}
if a.background != nil {
if err := a.background.Shutdown(ctx); err != nil {
shutdownErrors = append(shutdownErrors, fmt.Errorf("stop background tasks: %w", err))
}
}
return errors.Join(shutdownErrors...)
}
func (a *Application) Close() {
if a == nil {
return
}
a.closeOnce.Do(func() {
shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
if err := a.Shutdown(shutdownCtx); err != nil && a.logger != nil {
a.logger.Warn("application cleanup incomplete", zap.Error(err))
}
if a.db != nil {
if sqlDB, err := a.db.DB(); err != nil {
if a.logger != nil {
a.logger.Warn("get database handle for close failed", zap.Error(err))
}
} else if err := sqlDB.Close(); err != nil && a.logger != nil {
a.logger.Warn("close database failed", zap.Error(err))
}
}
if a.logger != nil {
if err := a.logger.Sync(); err != nil {
a.logger.Warn("flush logger failed", zap.Error(err))
}
}
})
} }
func (a *Application) Logger() *zap.Logger { func (a *Application) Logger() *zap.Logger {
+68
View File
@@ -0,0 +1,68 @@
package app
import (
"context"
"errors"
"path/filepath"
"testing"
"backupx/server/internal/lifecycle"
"github.com/glebarez/sqlite"
"go.uber.org/zap"
"gorm.io/gorm"
)
func TestApplicationCloseStopsBackgroundAndClosesDatabaseOnce(t *testing.T) {
db, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "app.db")), &gorm.Config{})
if err != nil {
t.Fatalf("open database: %v", err)
}
sqlDB, err := db.DB()
if err != nil {
t.Fatalf("get database handle: %v", err)
}
background := lifecycle.NewSupervisor(context.Background())
started := make(chan struct{})
finished := make(chan struct{})
if !background.Go(func(ctx context.Context) {
close(started)
<-ctx.Done()
close(finished)
}) {
t.Fatal("expected background task to be accepted")
}
<-started
application := &Application{db: db, logger: zap.NewNop(), background: background}
application.Close()
application.Close()
select {
case <-finished:
default:
t.Fatal("Close returned before the background task exited")
}
if err := sqlDB.Ping(); err == nil {
t.Fatal("database remained open after Close")
}
}
func TestApplicationShutdownCanRetryAfterWaitTimeout(t *testing.T) {
background := lifecycle.NewSupervisor(context.Background())
release := make(chan struct{})
if !background.Go(func(context.Context) { <-release }) {
t.Fatal("expected background task to be accepted")
}
application := &Application{logger: zap.NewNop(), background: background}
waitCtx, cancel := context.WithCancel(context.Background())
cancel()
if err := application.Shutdown(waitCtx); !errors.Is(err, context.Canceled) {
t.Fatalf("first Shutdown error = %v, want context.Canceled", err)
}
close(release)
if err := application.Shutdown(context.Background()); err != nil {
t.Fatalf("second Shutdown returned error: %v", err)
}
}
+1 -21
View File
@@ -30,7 +30,7 @@ type Agent struct {
// NewAgent 构造 Agent,初始化 storage provider 与 catalog。 // NewAgent 构造 Agent,初始化 storage provider 与 catalog。
func NewAgent(ctx context.Context, cfg *Config) (*Agent, error) { func NewAgent(ctx context.Context, cfg *Config) (*Agent, error) {
registry := buildStorageRegistry() registry := storageRclone.NewDefaultRegistry()
provider, err := registry.Create(ctx, cfg.StorageType, cfg.StorageConfig) provider, err := registry.Create(ctx, cfg.StorageType, cfg.StorageConfig)
if err != nil { if err != nil {
return nil, fmt.Errorf("create storage provider: %w", err) return nil, fmt.Errorf("create storage provider: %w", err)
@@ -337,23 +337,3 @@ func boolStr(b bool) string {
} }
return "false" return "false"
} }
// buildStorageRegistry 构造与主程序一致的 storage registry。
//
// Backint Agent 作为独立 CLI 进程运行,不依赖 BackupX HTTP 服务,
// 因此这里直接引用 storage/rclone 包注册所有后端。
func buildStorageRegistry() *storage.Registry {
registry := storage.NewRegistry(
storageRclone.NewLocalDiskFactory(),
storageRclone.NewS3Factory(),
storageRclone.NewWebDAVFactory(),
storageRclone.NewGoogleDriveFactory(),
storageRclone.NewAliyunOSSFactory(),
storageRclone.NewTencentCOSFactory(),
storageRclone.NewQiniuKodoFactory(),
storageRclone.NewFTPFactory(),
storageRclone.NewRcloneFactory(),
)
storageRclone.RegisterAllBackends(registry)
return registry
}
@@ -1,37 +0,0 @@
//go:build ignore
package backup
import (
"context"
"io"
"os"
"os/exec"
)
type CommandExecutor interface {
LookPath(file string) (string, error)
Run(ctx context.Context, name string, args []string, env map[string]string, stdin io.Reader, stdout io.Writer, stderr io.Writer) error
}
type OSCommandExecutor struct{}
func NewOSCommandExecutor() *OSCommandExecutor {
return &OSCommandExecutor{}
}
func (e *OSCommandExecutor) LookPath(file string) (string, error) {
return exec.LookPath(file)
}
func (e *OSCommandExecutor) Run(ctx context.Context, name string, args []string, env map[string]string, stdin io.Reader, stdout io.Writer, stderr io.Writer) error {
command := exec.CommandContext(ctx, name, args...)
command.Stdin = stdin
command.Stdout = stdout
command.Stderr = stderr
command.Env = os.Environ()
for key, value := range env {
command.Env = append(command.Env, key+"="+value)
}
return command.Run()
}
@@ -0,0 +1,145 @@
package backup
import (
"context"
"io"
"os"
"reflect"
"testing"
)
type capturingLogWriter struct {
lines []string
}
func (w *capturingLogWriter) WriteLine(message string) {
w.lines = append(w.lines, message)
}
func TestLogLineWriterHandlesFragmentedWrites(t *testing.T) {
log := &capturingLogWriter{}
w := newLogLineWriter(log, "tool")
assertWrite := func(chunk string) {
t.Helper()
n, err := w.Write([]byte(chunk))
if err != nil || n != len(chunk) {
t.Fatalf("Write(%q) = (%d, %v), want (%d, nil)", chunk, n, err, len(chunk))
}
}
assertWrite("fir")
if len(log.lines) != 0 || string(w.pending) != "fir" {
t.Fatalf("first fragment logged or buffered incorrectly: lines=%#v pending=%q", log.lines, w.pending)
}
assertWrite("st\nsec")
if !reflect.DeepEqual(log.lines, []string{"[tool] first"}) || string(w.pending) != "sec" {
t.Fatalf("second fragment handled incorrectly: lines=%#v pending=%q", log.lines, w.pending)
}
assertWrite("ond\n")
if !reflect.DeepEqual(log.lines, []string{"[tool] first", "[tool] second"}) || len(w.pending) != 0 {
t.Fatalf("completed fragments handled incorrectly: lines=%#v pending=%q", log.lines, w.pending)
}
if got := w.collected(); got != "first\nsecond" {
t.Fatalf("collected() = %q, want complete raw output", got)
}
}
func TestLogLineWriterEmitsMultipleCompleteLinesOnce(t *testing.T) {
log := &capturingLogWriter{}
w := newLogLineWriter(log, "tool")
n, err := w.Write([]byte("one\ntwo\n\n three \r\n"))
if err != nil || n != len("one\ntwo\n\n three \r\n") {
t.Fatalf("Write() = (%d, %v)", n, err)
}
want := []string{"[tool] one", "[tool] two", "[tool] three"}
if !reflect.DeepEqual(log.lines, want) {
t.Fatalf("lines = %#v, want %#v", log.lines, want)
}
if len(w.pending) != 0 {
t.Fatalf("complete input left pending bytes: %q", w.pending)
}
}
func TestLogLineWriterFlushIsIdempotentAndPreservesCollection(t *testing.T) {
log := &capturingLogWriter{}
w := newLogLineWriter(log, "tool")
_, _ = w.Write([]byte("complete\n tail "))
w.Flush()
w.Flush()
want := []string{"[tool] complete", "[tool] tail"}
if !reflect.DeepEqual(log.lines, want) {
t.Fatalf("lines after repeated Flush = %#v, want %#v", log.lines, want)
}
if len(w.pending) != 0 {
t.Fatalf("Flush left pending bytes: %q", w.pending)
}
if got := w.collected(); got != "complete\n tail" {
t.Fatalf("collected() = %q, want complete stderr independent of Flush", got)
}
}
func TestPostgreSQLRunnerFlushesEachCommandTail(t *testing.T) {
executor := &fakeCommandExecutor{runFunc: func(_ string, args []string, options CommandOptions) error {
name := args[len(args)-1]
_, _ = io.WriteString(options.Stdout, name)
_, _ = io.WriteString(options.Stderr, "warning "+name)
return nil
}}
log := &capturingLogWriter{}
runner := NewPostgreSQLRunner(executor)
result, err := runner.Run(context.Background(), TaskSpec{
Name: "pg-log-lines",
TempDir: t.TempDir(),
Database: DatabaseSpec{
Host: "127.0.0.1", Port: 5432, User: "postgres", Names: []string{"app", "audit"},
},
}, log)
if err != nil {
t.Fatalf("Run returned error: %v", err)
}
t.Cleanup(func() { _ = os.RemoveAll(result.TempDir) })
for _, expected := range []string{"[pg_dump] warning app", "[pg_dump] warning audit"} {
count := 0
for _, line := range log.lines {
if line == expected {
count++
}
}
if count != 1 {
t.Fatalf("line %q occurred %d times in %#v", expected, count, log.lines)
}
}
}
func TestMongoDBRunnerFlushesUnterminatedStderr(t *testing.T) {
executor := &fakeCommandExecutor{runFunc: func(_ string, _ []string, options CommandOptions) error {
_, _ = io.WriteString(options.Stdout, "archive")
_, _ = io.WriteString(options.Stderr, "tail warning")
return nil
}}
log := &capturingLogWriter{}
runner := NewMongoDBRunner(executor)
result, err := runner.Run(context.Background(), TaskSpec{
Name: "mongo-log-tail",
Database: DatabaseSpec{
Host: "127.0.0.1", Port: 27017, Names: []string{"app"},
},
}, log)
if err != nil {
t.Fatalf("Run returned error: %v", err)
}
t.Cleanup(func() { _ = os.RemoveAll(result.TempDir) })
count := 0
for _, line := range log.lines {
if line == "[mongodump] tail warning" {
count++
}
}
if count != 1 {
t.Fatalf("unterminated stderr line occurred %d times in %#v", count, log.lines)
}
}
+4 -2
View File
@@ -62,8 +62,10 @@ func (r *MongoDBRunner) Run(ctx context.Context, task TaskSpec, writer LogWriter
writer.WriteLine(fmt.Sprintf("连接到 MongoDB: %s:%d", task.Database.Host, task.Database.Port)) writer.WriteLine(fmt.Sprintf("连接到 MongoDB: %s:%d", task.Database.Host, task.Database.Port))
stderrWriter := newLogLineWriter(writer, "mongodump") stderrWriter := newLogLineWriter(writer, "mongodump")
writer.WriteLine("开始执行 mongodump") writer.WriteLine("开始执行 mongodump")
if err := r.executor.Run(ctx, "mongodump", args, CommandOptions{Stdout: file, Stderr: stderrWriter}); err != nil { runErr := r.executor.Run(ctx, "mongodump", args, CommandOptions{Stdout: file, Stderr: stderrWriter})
return nil, fmt.Errorf("run mongodump: %w: %s", err, stderrWriter.collected()) stderrWriter.Flush()
if runErr != nil {
return nil, fmt.Errorf("run mongodump: %w: %s", runErr, stderrWriter.collected())
} }
info, err := file.Stat() info, err := file.Stat()
if err != nil { if err != nil {
+39 -22
View File
@@ -1,7 +1,6 @@
package backup package backup
import ( import (
"bufio"
"bytes" "bytes"
"context" "context"
"fmt" "fmt"
@@ -70,8 +69,10 @@ func (r *MySQLRunner) Run(ctx context.Context, task TaskSpec, writer LogWriter)
stderrWriter := newLogLineWriter(writer, "mysqldump") stderrWriter := newLogLineWriter(writer, "mysqldump")
writer.WriteLine("开始执行 mysqldump") writer.WriteLine("开始执行 mysqldump")
if err := r.executor.Run(ctx, "mysqldump", args, CommandOptions{Stdout: file, Stderr: stderrWriter, Env: mysqlEnv(task.Database.Password)}); err != nil { runErr := r.executor.Run(ctx, "mysqldump", args, CommandOptions{Stdout: file, Stderr: stderrWriter, Env: mysqlEnv(task.Database.Password)})
return nil, fmt.Errorf("run mysqldump: %w: %s", err, stderrWriter.collected()) stderrWriter.Flush()
if runErr != nil {
return nil, fmt.Errorf("run mysqldump: %w: %s", runErr, stderrWriter.collected())
} }
info, err := file.Stat() info, err := file.Stat()
if err != nil { if err != nil {
@@ -109,9 +110,10 @@ func mysqlEnv(password string) []string {
// logLineWriter streams each line of output to a LogWriter in real-time. // logLineWriter streams each line of output to a LogWriter in real-time.
type logLineWriter struct { type logLineWriter struct {
writer LogWriter writer LogWriter
prefix string prefix string
buf bytes.Buffer pending []byte
output []byte
} }
func newLogLineWriter(w LogWriter, prefix string) *logLineWriter { func newLogLineWriter(w LogWriter, prefix string) *logLineWriter {
@@ -119,28 +121,43 @@ func newLogLineWriter(w LogWriter, prefix string) *logLineWriter {
} }
func (w *logLineWriter) Write(p []byte) (int, error) { func (w *logLineWriter) Write(p []byte) (int, error) {
n := len(p) w.output = append(w.output, p...)
w.buf.Write(p) w.pending = append(w.pending, p...)
scanner := bufio.NewScanner(strings.NewReader(w.buf.String())) consumed := 0
var remaining string for {
for scanner.Scan() { newline := bytes.IndexByte(w.pending[consumed:], '\n')
line := strings.TrimSpace(scanner.Text()) if newline < 0 {
if line != "" { break
w.writer.WriteLine(fmt.Sprintf("[%s] %s", w.prefix, line))
} }
end := consumed + newline
w.emit(w.pending[consumed:end])
consumed = end + 1
} }
// Keep any partial last line (no newline yet) if consumed > 0 {
lastNl := bytes.LastIndexByte(p, '\n') copy(w.pending, w.pending[consumed:])
if lastNl >= 0 { w.pending = w.pending[:len(w.pending)-consumed]
remaining = w.buf.String()[w.buf.Len()-(len(p)-lastNl-1):] }
w.buf.Reset() return len(p), nil
w.buf.WriteString(remaining) }
// Flush emits the final unterminated line. It is safe to call more than once.
func (w *logLineWriter) Flush() {
if len(w.pending) == 0 {
return
}
w.emit(w.pending)
w.pending = w.pending[:0]
}
func (w *logLineWriter) emit(raw []byte) {
line := strings.TrimSpace(string(raw))
if line != "" {
w.writer.WriteLine(fmt.Sprintf("[%s] %s", w.prefix, line))
} }
return n, nil
} }
func (w *logLineWriter) collected() string { func (w *logLineWriter) collected() string {
return strings.TrimSpace(w.buf.String()) return strings.TrimSpace(string(w.output))
} }
func formatFileSize(size int64) string { func formatFileSize(size int64) string {
-171
View File
@@ -1,171 +0,0 @@
//go:build ignore
package backup
import (
"bytes"
"context"
"fmt"
"os"
"path/filepath"
"strings"
)
type PostgreSQLRunner struct {
executor CommandExecutor
}
func NewPostgreSQLRunner(executor CommandExecutor) *PostgreSQLRunner {
if executor == nil {
executor = NewOSCommandExecutor()
}
return &PostgreSQLRunner{executor: executor}
}
func (r *PostgreSQLRunner) Type() string {
return "postgresql"
}
func (r *PostgreSQLRunner) Run(ctx context.Context, spec TaskSpec, logger LogSink) (*Result, error) {
if _, err := r.executor.LookPath("pg_dump"); err != nil {
return nil, fmt.Errorf("pg_dump is required: %w", err)
}
databases := splitDatabaseNames(spec.DBName)
if len(databases) == 0 {
return nil, fmt.Errorf("postgresql database name is required")
}
tempDir, err := CreateTaskTempDir(spec.TaskName, spec.StartedAt)
if err != nil {
return nil, err
}
if len(databases) == 1 {
return r.dumpSingleDatabase(ctx, spec, databases[0], tempDir, logger)
}
multiDumpDir := filepath.Join(tempDir, "postgres-dumps")
if err := os.MkdirAll(multiDumpDir, 0o755); err != nil {
return nil, fmt.Errorf("create postgres multi dump directory: %w", err)
}
for _, databaseName := range databases {
if _, err := r.dumpDatabaseToFile(ctx, spec, databaseName, filepath.Join(multiDumpDir, sanitizeDumpName(databaseName)+".sql"), logger); err != nil {
return nil, err
}
}
fileName := BuildArtifactName(spec.TaskName, spec.StartedAt, "tar.gz")
artifactPath := filepath.Join(tempDir, fileName)
size, err := CreateTarGz(ctx, multiDumpDir, nil, artifactPath, logger)
if err != nil {
return nil, err
}
return &Result{ArtifactPath: artifactPath, FileName: fileName, Size: size, StorageKey: BuildStorageKey("postgresql", spec.StartedAt, fileName)}, nil
}
func (r *PostgreSQLRunner) Restore(ctx context.Context, spec TaskSpec, artifactPath string, logger LogSink) error {
if _, err := r.executor.LookPath("psql"); err != nil {
return fmt.Errorf("psql is required: %w", err)
}
databases := splitDatabaseNames(spec.DBName)
if len(databases) == 0 {
return fmt.Errorf("postgresql database name is required")
}
if strings.HasSuffix(strings.ToLower(artifactPath), ".tar.gz") {
restoreDir, err := CreateTaskTempDir(spec.TaskName+"-restore", spec.StartedAt)
if err != nil {
return err
}
if err := ExtractTarGz(ctx, artifactPath, restoreDir, logger); err != nil {
return err
}
for _, databaseName := range databases {
filePath := filepath.Join(restoreDir, filepath.Base(restoreDir), sanitizeDumpName(databaseName)+".sql")
if _, err := os.Stat(filePath); err != nil {
fallback := filepath.Join(restoreDir, "postgres-dumps", sanitizeDumpName(databaseName)+".sql")
filePath = fallback
}
if err := r.restoreDatabaseFromFile(ctx, spec, databaseName, filePath, logger); err != nil {
return err
}
}
return nil
}
return r.restoreDatabaseFromFile(ctx, spec, databases[0], artifactPath, logger)
}
func (r *PostgreSQLRunner) dumpSingleDatabase(ctx context.Context, spec TaskSpec, databaseName string, tempDir string, logger LogSink) (*Result, error) {
fileName := BuildArtifactName(spec.TaskName, spec.StartedAt, "sql")
artifactPath := filepath.Join(tempDir, fileName)
size, err := r.dumpDatabaseToFile(ctx, spec, databaseName, artifactPath, logger)
if err != nil {
return nil, err
}
return &Result{ArtifactPath: artifactPath, FileName: fileName, Size: size, StorageKey: BuildStorageKey("postgresql", spec.StartedAt, fileName)}, nil
}
func (r *PostgreSQLRunner) dumpDatabaseToFile(ctx context.Context, spec TaskSpec, databaseName string, artifactPath string, logger LogSink) (int64, error) {
output, err := os.Create(filepath.Clean(artifactPath))
if err != nil {
return 0, fmt.Errorf("create postgres dump file: %w", err)
}
defer output.Close()
stderr := &bytes.Buffer{}
args := []string{"-h", spec.DBHost, "-p", fmt.Sprintf("%d", spec.DBPort), "-U", spec.DBUser, "-d", databaseName, "--no-owner", "--no-privileges"}
if logger != nil {
logger.Infof("开始执行 pg_dump%s", databaseName)
}
if err := r.executor.Run(ctx, "pg_dump", args, postgresEnv(spec.DBPassword), nil, output, stderr); err != nil {
return 0, fmt.Errorf("run pg_dump: %w: %s", err, strings.TrimSpace(stderr.String()))
}
info, err := output.Stat()
if err != nil {
return 0, fmt.Errorf("stat postgres dump file: %w", err)
}
return info.Size(), nil
}
func (r *PostgreSQLRunner) restoreDatabaseFromFile(ctx context.Context, spec TaskSpec, databaseName string, artifactPath string, logger LogSink) error {
input, err := os.Open(filepath.Clean(artifactPath))
if err != nil {
return fmt.Errorf("open postgres restore file: %w", err)
}
defer input.Close()
stderr := &bytes.Buffer{}
args := []string{"-h", spec.DBHost, "-p", fmt.Sprintf("%d", spec.DBPort), "-U", spec.DBUser, "-d", databaseName}
if logger != nil {
logger.Infof("开始执行 psql 恢复:%s", databaseName)
}
if err := r.executor.Run(ctx, "psql", args, postgresEnv(spec.DBPassword), input, nil, stderr); err != nil {
return fmt.Errorf("run psql restore: %w: %s", err, strings.TrimSpace(stderr.String()))
}
return nil
}
func postgresEnv(password string) map[string]string {
if strings.TrimSpace(password) == "" {
return nil
}
return map[string]string{"PGPASSWORD": password}
}
func splitDatabaseNames(value string) []string {
parts := strings.Split(value, ",")
result := make([]string, 0, len(parts))
for _, part := range parts {
trimmed := strings.TrimSpace(part)
if trimmed == "" {
continue
}
result = append(result, trimmed)
}
return result
}
func sanitizeDumpName(value string) string {
trimmed := strings.TrimSpace(strings.ToLower(value))
trimmed = strings.ReplaceAll(trimmed, " ", "-")
trimmed = strings.ReplaceAll(trimmed, "/", "-")
trimmed = strings.ReplaceAll(trimmed, "\\", "-")
trimmed = strings.Trim(trimmed, "-._")
if trimmed == "" {
return "database"
}
return trimmed
}
+5 -3
View File
@@ -43,12 +43,14 @@ func (r *PostgreSQLRunner) Run(ctx context.Context, task TaskSpec, writer LogWri
} }
writer.WriteLine(fmt.Sprintf("连接到 PostgreSQL: %s:%d", task.Database.Host, task.Database.Port)) writer.WriteLine(fmt.Sprintf("连接到 PostgreSQL: %s:%d", task.Database.Host, task.Database.Port))
writer.WriteLine(fmt.Sprintf("备份数据库: %s", strings.Join(dbNames, ", "))) writer.WriteLine(fmt.Sprintf("备份数据库: %s", strings.Join(dbNames, ", ")))
stderrWriter := newLogLineWriter(writer, "pg_dump")
for index, name := range dbNames { for index, name := range dbNames {
args := []string{"--clean", "--if-exists", "--create", "--format=plain", "-h", task.Database.Host, "-p", strconv.Itoa(task.Database.Port), "-U", task.Database.User, "--dbname", name} args := []string{"--clean", "--if-exists", "--create", "--format=plain", "-h", task.Database.Host, "-p", strconv.Itoa(task.Database.Port), "-U", task.Database.User, "--dbname", name}
writer.WriteLine(fmt.Sprintf("开始导出数据库 [%d/%d]: %s", index+1, len(dbNames), name)) writer.WriteLine(fmt.Sprintf("开始导出数据库 [%d/%d]: %s", index+1, len(dbNames), name))
if err := r.executor.Run(ctx, "pg_dump", args, CommandOptions{Stdout: file, Stderr: stderrWriter, Env: append(os.Environ(), "PGPASSWORD="+task.Database.Password)}); err != nil { stderrWriter := newLogLineWriter(writer, "pg_dump")
return nil, fmt.Errorf("run pg_dump for %s: %w", name, err) runErr := r.executor.Run(ctx, "pg_dump", args, CommandOptions{Stdout: file, Stderr: stderrWriter, Env: append(os.Environ(), "PGPASSWORD="+task.Database.Password)})
stderrWriter.Flush()
if runErr != nil {
return nil, fmt.Errorf("run pg_dump for %s: %w", name, runErr)
} }
writer.WriteLine(fmt.Sprintf("数据库 %s 导出完成", name)) writer.WriteLine(fmt.Sprintf("数据库 %s 导出完成", name))
if index < len(dbNames)-1 { if index < len(dbNames)-1 {
+12
View File
@@ -20,6 +20,18 @@ func NewRegistry(runners ...BackupRunner) *Registry {
return registry return registry
} }
// NewDefaultRegistry returns the runner set shared by Master and Agent.
func NewDefaultRegistry() *Registry {
return NewRegistry(
NewFileRunner(),
NewSQLiteRunner(),
NewMySQLRunner(nil),
NewPostgreSQLRunner(nil),
NewSAPHANARunner(nil),
NewMongoDBRunner(nil),
)
}
func (r *Registry) Register(runner BackupRunner) { func (r *Registry) Register(runner BackupRunner) {
if runner == nil { if runner == nil {
return return
+1
View File
@@ -304,6 +304,7 @@ func (r *SAPHANARunner) runHdbsqlWithRetry(ctx context.Context, name string, arg
} }
stderrWriter := newLogLineWriter(writer, "hdbsql") stderrWriter := newLogLineWriter(writer, "hdbsql")
err := r.executor.Run(ctx, name, args, CommandOptions{Stderr: stderrWriter}) err := r.executor.Run(ctx, name, args, CommandOptions{Stderr: stderrWriter})
stderrWriter.Flush()
if err == nil { if err == nil {
return nil return nil
} }
+113 -2
View File
@@ -1,10 +1,12 @@
package database package database
import ( import (
"encoding/json"
"fmt" "fmt"
"os" "os"
"path/filepath" "path/filepath"
"strings" "strings"
"time"
"backupx/server/internal/config" "backupx/server/internal/config"
"backupx/server/internal/model" "backupx/server/internal/model"
@@ -30,6 +32,22 @@ func Open(cfg config.DatabaseConfig, logger *zap.Logger) (*gorm.DB, error) {
if err != nil { if err != nil {
return nil, fmt.Errorf("open sqlite: %w", err) return nil, fmt.Errorf("open sqlite: %w", err)
} }
initialized := false
defer func() {
if initialized {
return
}
sqlDB, dbErr := db.DB()
if dbErr != nil {
if logger != nil {
logger.Warn("get database handle after initialization failure", zap.Error(dbErr))
}
return
}
if closeErr := sqlDB.Close(); closeErr != nil && logger != nil {
logger.Warn("close database after initialization failure", zap.Error(closeErr))
}
}()
if err := db.AutoMigrate(&model.User{}, &model.SystemConfig{}, &model.StorageTarget{}, &model.OAuthSession{}, &model.BackupTask{}, &model.BackupRecord{}, &model.Notification{}, &model.Node{}, &model.BackupTaskStorageTarget{}, &model.AuditLog{}, &model.AgentCommand{}, &model.AgentInstallToken{}, &model.RestoreRecord{}, &model.VerificationRecord{}, &model.ApiKey{}, &model.ReplicationRecord{}, &model.TaskTemplate{}); err != nil { if err := db.AutoMigrate(&model.User{}, &model.SystemConfig{}, &model.StorageTarget{}, &model.OAuthSession{}, &model.BackupTask{}, &model.BackupRecord{}, &model.Notification{}, &model.Node{}, &model.BackupTaskStorageTarget{}, &model.AuditLog{}, &model.AgentCommand{}, &model.AgentInstallToken{}, &model.RestoreRecord{}, &model.VerificationRecord{}, &model.ApiKey{}, &model.ReplicationRecord{}, &model.TaskTemplate{}); err != nil {
return nil, fmt.Errorf("migrate schema: %w", err) return nil, fmt.Errorf("migrate schema: %w", err)
@@ -37,11 +55,104 @@ func Open(cfg config.DatabaseConfig, logger *zap.Logger) (*gorm.DB, error) {
// 一次性数据迁移:从 backup_tasks.storage_target_id 回填到多对多中间表 // 一次性数据迁移:从 backup_tasks.storage_target_id 回填到多对多中间表
var count int64 var count int64
db.Model(&model.BackupTaskStorageTarget{}).Count(&count) if err := db.Model(&model.BackupTaskStorageTarget{}).Count(&count).Error; err != nil {
return nil, fmt.Errorf("count backup task storage target mappings: %w", err)
}
if count == 0 { if count == 0 {
db.Exec("INSERT INTO backup_task_storage_targets (backup_task_id, storage_target_id) SELECT id, storage_target_id FROM backup_tasks WHERE storage_target_id > 0") if err := db.Exec("INSERT INTO backup_task_storage_targets (backup_task_id, storage_target_id) SELECT id, storage_target_id FROM backup_tasks WHERE storage_target_id > 0").Error; err != nil {
return nil, fmt.Errorf("backfill backup task storage target mappings: %w", err)
}
}
reconciled, err := reconcileInterruptedOperations(db, time.Now().UTC())
if err != nil {
return nil, fmt.Errorf("reconcile interrupted operations: %w", err)
}
if reconciled > 0 {
logger.Warn("interrupted operations marked as failed", zap.Int64("records", reconciled))
} }
logger.Info("database initialized", zap.String("path", cfg.Path)) logger.Info("database initialized", zap.String("path", cfg.Path))
initialized = true
return db, nil return db, nil
} }
func reconcileInterruptedOperations(db *gorm.DB, completedAt time.Time) (int64, error) {
const message = "应用在任务完成前重启,执行状态已自动收敛为失败"
var reconciled int64
err := db.Transaction(func(tx *gorm.DB) error {
// Pending/dispatched Agent commands survive a Master restart. Their Agent
// may still be executing (or may claim the pending command after startup),
// so their linked records must not be mistaken for orphaned local work.
var activeCommands []model.AgentCommand
if err := tx.Where("status IN ? AND type IN ?",
[]string{model.AgentCommandStatusPending, model.AgentCommandStatusDispatched},
[]string{model.AgentCommandTypeRunTask, model.AgentCommandTypeRestoreRecord}).
Find(&activeCommands).Error; err != nil {
return fmt.Errorf("active agent commands: %w", err)
}
activeBackupRecordIDs := make([]uint, 0, len(activeCommands))
activeRestoreRecordIDs := make([]uint, 0, len(activeCommands))
for i := range activeCommands {
cmd := &activeCommands[i]
switch cmd.Type {
case model.AgentCommandTypeRunTask:
var payload struct {
RecordID uint `json:"recordId"`
}
if json.Unmarshal([]byte(cmd.Payload), &payload) == nil && payload.RecordID > 0 {
activeBackupRecordIDs = append(activeBackupRecordIDs, payload.RecordID)
}
case model.AgentCommandTypeRestoreRecord:
var payload struct {
RestoreRecordID uint `json:"restoreRecordId"`
}
if json.Unmarshal([]byte(cmd.Payload), &payload) == nil && payload.RestoreRecordID > 0 {
activeRestoreRecordIDs = append(activeRestoreRecordIDs, payload.RestoreRecordID)
}
}
}
markFailed := func(entity any, runningStatus, failedStatus string, activeAgentRecordIDs []uint) error {
query := tx.Model(entity).Where("status = ?", runningStatus)
if len(activeAgentRecordIDs) > 0 {
query = query.Where("id NOT IN ?", activeAgentRecordIDs)
}
result := query.
Updates(map[string]any{
"status": failedStatus,
"error_message": message,
"completed_at": completedAt,
"duration_seconds": gorm.Expr("CAST(MAX(0, (julianday(?) - julianday(started_at)) * 86400) AS INTEGER)", completedAt),
})
if result.Error != nil {
return result.Error
}
reconciled += result.RowsAffected
return nil
}
if err := markFailed(&model.BackupRecord{}, model.BackupRecordStatusRunning, model.BackupRecordStatusFailed, activeBackupRecordIDs); err != nil {
return fmt.Errorf("backup records: %w", err)
}
if err := markFailed(&model.RestoreRecord{}, model.RestoreRecordStatusRunning, model.RestoreRecordStatusFailed, activeRestoreRecordIDs); err != nil {
return fmt.Errorf("restore records: %w", err)
}
if err := markFailed(&model.VerificationRecord{}, model.VerificationRecordStatusRunning, model.VerificationRecordStatusFailed, nil); err != nil {
return fmt.Errorf("verification records: %w", err)
}
if err := markFailed(&model.ReplicationRecord{}, model.ReplicationStatusRunning, model.ReplicationStatusFailed, nil); err != nil {
return fmt.Errorf("replication records: %w", err)
}
result := tx.Model(&model.BackupTask{}).
Where("last_status = ? AND NOT EXISTS (SELECT 1 FROM backup_records WHERE backup_records.task_id = backup_tasks.id AND backup_records.status = ?)", model.BackupTaskStatusRunning, model.BackupRecordStatusRunning).
Update("last_status", model.BackupTaskStatusFailed)
if result.Error != nil {
return fmt.Errorf("backup tasks: %w", result.Error)
}
reconciled += result.RowsAffected
return nil
})
return reconciled, err
}
+178
View File
@@ -1,11 +1,14 @@
package database package database
import ( import (
"fmt"
"path/filepath" "path/filepath"
"testing" "testing"
"time"
"backupx/server/internal/config" "backupx/server/internal/config"
"backupx/server/internal/logger" "backupx/server/internal/logger"
"backupx/server/internal/model"
) )
func TestOpenConfiguresSQLiteForSingleMasterConcurrency(t *testing.T) { func TestOpenConfiguresSQLiteForSingleMasterConcurrency(t *testing.T) {
@@ -38,3 +41,178 @@ func TestOpenConfiguresSQLiteForSingleMasterConcurrency(t *testing.T) {
t.Fatalf("busy_timeout = %d, want 5000", busyTimeout) t.Fatalf("busy_timeout = %d, want 5000", busyTimeout)
} }
} }
func TestReconcileInterruptedOperations(t *testing.T) {
log, err := logger.New(config.LogConfig{Level: "error"})
if err != nil {
t.Fatal(err)
}
db, err := Open(config.DatabaseConfig{Path: filepath.Join(t.TempDir(), "reconcile.db")}, log)
if err != nil {
t.Fatal(err)
}
sqlDB, err := db.DB()
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
if err := sqlDB.Close(); err != nil {
t.Errorf("close database: %v", err)
}
})
startedAt := time.Now().UTC().Add(-time.Minute)
task := model.BackupTask{Name: "interrupted", Type: model.BackupTaskTypeFile, LastStatus: model.BackupTaskStatusRunning}
if err := db.Create(&task).Error; err != nil {
t.Fatal(err)
}
items := []any{
&model.BackupRecord{TaskID: task.ID, Status: model.BackupRecordStatusRunning, StartedAt: startedAt},
&model.RestoreRecord{TaskID: task.ID, Status: model.RestoreRecordStatusRunning, StartedAt: startedAt},
&model.VerificationRecord{TaskID: task.ID, Status: model.VerificationRecordStatusRunning, StartedAt: startedAt},
&model.ReplicationRecord{TaskID: task.ID, Status: model.ReplicationStatusRunning, StartedAt: startedAt},
}
for _, item := range items {
if err := db.Create(item).Error; err != nil {
t.Fatal(err)
}
}
completedAt := time.Now().UTC()
count, err := reconcileInterruptedOperations(db, completedAt)
if err != nil {
t.Fatal(err)
}
if count != 5 {
t.Fatalf("reconciled records = %d, want 5", count)
}
var runningRecords int64
for _, entity := range []any{&model.BackupRecord{}, &model.RestoreRecord{}, &model.VerificationRecord{}, &model.ReplicationRecord{}} {
if err := db.Model(entity).Where("status = ?", "running").Count(&runningRecords).Error; err != nil {
t.Fatal(err)
}
if runningRecords != 0 {
t.Fatalf("%T still has %d running records", entity, runningRecords)
}
}
if err := db.First(&task, task.ID).Error; err != nil {
t.Fatal(err)
}
if task.LastStatus != model.BackupTaskStatusFailed {
t.Fatalf("task last status = %q, want failed", task.LastStatus)
}
var backupRecord model.BackupRecord
if err := db.Where("task_id = ?", task.ID).First(&backupRecord).Error; err != nil {
t.Fatal(err)
}
if backupRecord.DurationSeconds <= 0 {
t.Fatalf("backup duration = %d, want positive duration", backupRecord.DurationSeconds)
}
}
func TestReconcileInterruptedOperationsPreservesRemoteAgentWork(t *testing.T) {
log, err := logger.New(config.LogConfig{Level: "error"})
if err != nil {
t.Fatal(err)
}
db, err := Open(config.DatabaseConfig{Path: filepath.Join(t.TempDir(), "remote-reconcile.db")}, log)
if err != nil {
t.Fatal(err)
}
sqlDB, err := db.DB()
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
if err := sqlDB.Close(); err != nil {
t.Errorf("close database: %v", err)
}
})
localNode := model.Node{Name: "local", Token: "local-token", Status: model.NodeStatusOnline, IsLocal: true}
remoteNode := model.Node{Name: "remote", Token: "remote-token", Status: model.NodeStatusOnline, IsLocal: false}
if err := db.Create(&localNode).Error; err != nil {
t.Fatal(err)
}
if err := db.Create(&remoteNode).Error; err != nil {
t.Fatal(err)
}
localTask := model.BackupTask{Name: "local-interrupted", Type: model.BackupTaskTypeFile, LastStatus: model.BackupTaskStatusRunning}
remoteTask := model.BackupTask{Name: "remote-still-running", Type: model.BackupTaskTypeFile, LastStatus: model.BackupTaskStatusRunning}
orphanRemoteTask := model.BackupTask{Name: "remote-without-command", Type: model.BackupTaskTypeFile, LastStatus: model.BackupTaskStatusRunning}
if err := db.Create(&localTask).Error; err != nil {
t.Fatal(err)
}
if err := db.Create(&remoteTask).Error; err != nil {
t.Fatal(err)
}
if err := db.Create(&orphanRemoteTask).Error; err != nil {
t.Fatal(err)
}
startedAt := time.Now().UTC().Add(-time.Minute)
localBackup := &model.BackupRecord{TaskID: localTask.ID, NodeID: localNode.ID, Status: model.BackupRecordStatusRunning, StartedAt: startedAt}
localRestore := &model.RestoreRecord{TaskID: localTask.ID, NodeID: localNode.ID, Status: model.RestoreRecordStatusRunning, StartedAt: startedAt}
remoteBackup := &model.BackupRecord{TaskID: remoteTask.ID, NodeID: remoteNode.ID, Status: model.BackupRecordStatusRunning, StartedAt: startedAt}
remoteRestore := &model.RestoreRecord{TaskID: remoteTask.ID, NodeID: remoteNode.ID, Status: model.RestoreRecordStatusRunning, StartedAt: startedAt}
orphanRemoteBackup := &model.BackupRecord{TaskID: orphanRemoteTask.ID, NodeID: remoteNode.ID, Status: model.BackupRecordStatusRunning, StartedAt: startedAt}
items := []any{localBackup, localRestore, remoteBackup, remoteRestore, orphanRemoteBackup}
for _, item := range items {
if err := db.Create(item).Error; err != nil {
t.Fatal(err)
}
}
activeCommands := []model.AgentCommand{
{NodeID: remoteNode.ID, Type: model.AgentCommandTypeRunTask, Status: model.AgentCommandStatusDispatched, Payload: fmt.Sprintf(`{"recordId":%d}`, remoteBackup.ID)},
{NodeID: remoteNode.ID, Type: model.AgentCommandTypeRestoreRecord, Status: model.AgentCommandStatusPending, Payload: fmt.Sprintf(`{"restoreRecordId":%d}`, remoteRestore.ID)},
}
if err := db.Create(&activeCommands).Error; err != nil {
t.Fatal(err)
}
count, err := reconcileInterruptedOperations(db, time.Now().UTC())
if err != nil {
t.Fatal(err)
}
if count != 5 {
t.Fatalf("reconciled rows = %d, want local and unlinked remote records plus their tasks", count)
}
for _, entity := range []any{&model.BackupRecord{}, &model.RestoreRecord{}} {
var remoteRunning int64
if err := db.Model(entity).Where("task_id = ? AND status = ?", remoteTask.ID, "running").Count(&remoteRunning).Error; err != nil {
t.Fatal(err)
}
if remoteRunning != 1 {
t.Fatalf("%T remote running records = %d, want 1", entity, remoteRunning)
}
var localRunning int64
if err := db.Model(entity).Where("task_id = ? AND status = ?", localTask.ID, "running").Count(&localRunning).Error; err != nil {
t.Fatal(err)
}
if localRunning != 0 {
t.Fatalf("%T local running records = %d, want 0", entity, localRunning)
}
}
if err := db.First(&localTask, localTask.ID).Error; err != nil {
t.Fatal(err)
}
if err := db.First(&remoteTask, remoteTask.ID).Error; err != nil {
t.Fatal(err)
}
if err := db.First(&orphanRemoteTask, orphanRemoteTask.ID).Error; err != nil {
t.Fatal(err)
}
if localTask.LastStatus != model.BackupTaskStatusFailed || remoteTask.LastStatus != model.BackupTaskStatusRunning {
t.Fatalf("task statuses = local:%q remote:%q", localTask.LastStatus, remoteTask.LastStatus)
}
if orphanRemoteTask.LastStatus != model.BackupTaskStatusFailed {
t.Fatalf("unlinked remote task status = %q, want failed", orphanRemoteTask.LastStatus)
}
if err := db.First(orphanRemoteBackup, orphanRemoteBackup.ID).Error; err != nil {
t.Fatal(err)
}
if orphanRemoteBackup.Status != model.BackupRecordStatusFailed {
t.Fatalf("unlinked remote backup status = %q, want failed", orphanRemoteBackup.Status)
}
}
+1 -3
View File
@@ -46,9 +46,7 @@ func (h *HealthHandler) Ready(c *gin.Context) {
checks["database"] = "error: " + err.Error() checks["database"] = "error: " + err.Error()
overallOK = false overallOK = false
} else { } else {
ctx, cancel := c.Request.Context(), func() {} if err := sqlDB.PingContext(c.Request.Context()); err != nil {
_ = cancel
if err := sqlDB.PingContext(ctx); err != nil {
checks["database"] = "ping failed: " + err.Error() checks["database"] = "ping failed: " + err.Error()
overallOK = false overallOK = false
} else { } else {
+1
View File
@@ -399,6 +399,7 @@ func NewRouter(deps RouterDependencies) *gin.Engine {
func requestLogger(logger *zap.Logger) gin.HandlerFunc { func requestLogger(logger *zap.Logger) gin.HandlerFunc {
return func(c *gin.Context) { return func(c *gin.Context) {
response.SetLogger(c, logger)
c.Next() c.Next()
logger.Info("http request", logger.Info("http request",
zap.String("method", c.Request.Method), zap.String("method", c.Request.Method),
-98
View File
@@ -1,98 +0,0 @@
//go:build ignore
package httpapi
import (
"net/http"
"backupx/server/internal/service"
"backupx/server/pkg/response"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
)
type authHandler struct {
service *service.AuthService
logger *zap.Logger
}
type setupRequest struct {
Username string `json:"username" binding:"required,min=3,max=64"`
Password string `json:"password" binding:"required,min=8,max=128"`
DisplayName string `json:"displayName" binding:"required,min=1,max=128"`
}
type loginRequest struct {
Username string `json:"username" binding:"required,min=3,max=64"`
Password string `json:"password" binding:"required,min=8,max=128"`
}
func newAuthHandler(service *service.AuthService, logger *zap.Logger) *authHandler {
return &authHandler{service: service, logger: logger}
}
func (h *authHandler) registerRoutes(router gin.IRouter, protected gin.IRouter) {
router.GET("/auth/setup/status", h.getSetupStatus)
router.POST("/auth/setup", h.setup)
router.POST("/auth/login", h.login)
protected.GET("/auth/profile", h.profile)
}
func (h *authHandler) getSetupStatus(c *gin.Context) {
initialized, err := h.service.GetSetupStatus(c.Request.Context())
if err != nil {
writeError(c, h.logger, err)
return
}
response.Success(c, gin.H{"initialized": initialized})
}
func (h *authHandler) setup(c *gin.Context) {
payload, err := bindJSON[setupRequest](c, h.logger)
if err != nil {
writeError(c, h.logger, err)
return
}
result, err := h.service.Setup(c.Request.Context(), service.SetupInput{
Username: payload.Username,
Password: payload.Password,
DisplayName: payload.DisplayName,
})
if err != nil {
writeError(c, h.logger, err)
return
}
c.JSON(http.StatusCreated, response.Envelope{Code: "OK", Message: "success", Data: result})
}
func (h *authHandler) login(c *gin.Context) {
payload, err := bindJSON[loginRequest](c, h.logger)
if err != nil {
writeError(c, h.logger, err)
return
}
result, err := h.service.Login(c.Request.Context(), service.LoginInput{
Username: payload.Username,
Password: payload.Password,
RemoteAddr: c.ClientIP(),
})
if err != nil {
writeError(c, h.logger, err)
return
}
response.Success(c, result)
}
func (h *authHandler) profile(c *gin.Context) {
userID, err := getUserID(c)
if err != nil {
response.Error(c, http.StatusUnauthorized, "AUTH_UNAUTHORIZED", "认证信息无效")
return
}
result, err := h.service.GetCurrentUser(c.Request.Context(), userID)
if err != nil {
writeError(c, h.logger, err)
return
}
response.Success(c, result)
}
-23
View File
@@ -1,23 +0,0 @@
//go:build ignore
package httpapi
import (
"fmt"
"github.com/gin-gonic/gin"
)
const claimsContextKey = "authClaims"
func getUserID(c *gin.Context) (uint, error) {
value, ok := c.Get(claimsContextKey)
if !ok {
return 0, fmt.Errorf("missing auth claims")
}
claims, ok := value.(AuthClaims)
if !ok {
return 0, fmt.Errorf("invalid auth claims")
}
return claims.UserID, nil
}
-92
View File
@@ -1,92 +0,0 @@
//go:build ignore
package httpapi
import (
"errors"
"fmt"
"net/http"
"strings"
"backupx/server/internal/apperror"
"backupx/server/internal/security"
"backupx/server/pkg/response"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
)
type AuthClaims struct {
UserID uint
Username string
Role string
}
func Recovery(logger *zap.Logger) gin.HandlerFunc {
return func(c *gin.Context) {
defer func() {
if recovered := recover(); recovered != nil {
logger.Error("panic recovered", zap.Any("panic", recovered), zap.String("path", c.Request.URL.Path))
response.Error(c, http.StatusInternalServerError, "INTERNAL_ERROR", "服务器内部错误")
c.Abort()
}
}()
c.Next()
}
}
func RequestLogger(logger *zap.Logger) gin.HandlerFunc {
return func(c *gin.Context) {
c.Next()
logger.Info("http request",
zap.String("method", c.Request.Method),
zap.String("path", c.Request.URL.Path),
zap.Int("status", c.Writer.Status()),
zap.String("client_ip", c.ClientIP()),
)
}
}
func AuthMiddleware(jwtManager *security.JWTManager) gin.HandlerFunc {
return func(c *gin.Context) {
authorization := strings.TrimSpace(c.GetHeader("Authorization"))
if authorization == "" || !strings.HasPrefix(strings.ToLower(authorization), "bearer ") {
response.Error(c, http.StatusUnauthorized, "AUTH_UNAUTHORIZED", "缺少有效的认证令牌")
c.Abort()
return
}
tokenValue := strings.TrimSpace(strings.TrimPrefix(authorization, "Bearer"))
if tokenValue == authorization {
tokenValue = strings.TrimSpace(strings.TrimPrefix(authorization, "bearer"))
}
claims, err := jwtManager.Parse(tokenValue)
if err != nil {
response.Error(c, http.StatusUnauthorized, "AUTH_UNAUTHORIZED", "认证令牌无效或已过期")
c.Abort()
return
}
c.Set(claimsContextKey, AuthClaims{UserID: claims.UserID, Username: claims.Username, Role: claims.Role})
c.Next()
}
}
func writeError(c *gin.Context, logger *zap.Logger, err error) {
var appErr *apperror.AppError
if errors.As(err, &appErr) {
if appErr.Err != nil {
logger.Warn("request failed", zap.String("code", appErr.Code), zap.Error(appErr.Err))
}
response.Error(c, appErr.Status, appErr.Code, appErr.Message)
return
}
logger.Error("unexpected error", zap.Error(err))
response.Error(c, http.StatusInternalServerError, "INTERNAL_ERROR", "服务器内部错误")
}
func bindJSON[T any](c *gin.Context, logger *zap.Logger) (*T, error) {
var payload T
if err := c.ShouldBindJSON(&payload); err != nil {
logger.Warn("bind json failed", zap.Error(err))
return nil, apperror.Wrap(http.StatusBadRequest, "INVALID_REQUEST", fmt.Sprintf("请求参数错误: %v", err), err)
}
return &payload, nil
}
-38
View File
@@ -1,38 +0,0 @@
//go:build ignore
package httpapi
import (
"backupx/server/internal/security"
"backupx/server/internal/service"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
)
type Dependencies struct {
Logger *zap.Logger
AuthService *service.AuthService
SystemService *service.SystemService
JWTManager *security.JWTManager
Mode string
}
func NewRouter(deps Dependencies) *gin.Engine {
gin.SetMode(deps.Mode)
router := gin.New()
router.Use(Recovery(deps.Logger), RequestLogger(deps.Logger))
api := router.Group("/api")
authHandler := newAuthHandler(deps.AuthService, deps.Logger)
systemHandler := newSystemHandler(deps.SystemService)
protected := api.Group("")
protected.Use(AuthMiddleware(deps.JWTManager))
authHandler.registerRoutes(api, protected)
systemHandler.registerRoutes(protected)
api.GET("/healthz", func(c *gin.Context) {
c.JSON(200, gin.H{"status": "ok"})
})
return router
}
-96
View File
@@ -1,96 +0,0 @@
//go:build ignore
package httpapi
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"path/filepath"
"testing"
"time"
"backupx/server/internal/config"
"backupx/server/internal/database"
"backupx/server/internal/logger"
"backupx/server/internal/repository"
"backupx/server/internal/security"
"backupx/server/internal/service"
)
func TestSetupLoginProfileAndSystemInfo(t *testing.T) {
tmpDir := t.TempDir()
cfg := config.Config{
Server: config.ServerConfig{Mode: "test"},
Database: config.DatabaseConfig{Path: filepath.Join(tmpDir, "backupx.db")},
Security: config.SecurityConfig{JWTSecret: "test-jwt-secret", JWTExpire: "1h", EncryptionKey: "test-encryption-key"},
Log: config.LogConfig{Level: "error"},
}
log, err := logger.New(cfg.Log)
if err != nil {
t.Fatalf("logger.New() error = %v", err)
}
db, err := database.Open(cfg.Database, log)
if err != nil {
t.Fatalf("database.Open() error = %v", err)
}
jwtManager := security.NewJWTManager(cfg.Security.JWTSecret, time.Hour)
authService := service.NewAuthService(repository.NewUserRepository(db), jwtManager, security.NewLoginLimiter(5, time.Minute))
systemService := service.NewSystemService(cfg, "test", time.Now().Add(-time.Minute))
router := NewRouter(Dependencies{Logger: log, AuthService: authService, SystemService: systemService, JWTManager: jwtManager, Mode: "test"})
setupBody := map[string]string{"username": "admin", "password": "super-secret", "displayName": "管理员"}
setupResp := performJSONRequest(t, router, http.MethodPost, "/api/auth/setup", setupBody, "")
if setupResp.Code != http.StatusCreated {
t.Fatalf("unexpected setup status: %d body=%s", setupResp.Code, setupResp.Body.String())
}
var setupPayload struct {
Code string `json:"code"`
Data struct {
Token string `json:"token"`
} `json:"data"`
}
if err := json.Unmarshal(setupResp.Body.Bytes(), &setupPayload); err != nil {
t.Fatalf("decode setup response: %v", err)
}
if setupPayload.Data.Token == "" {
t.Fatal("expected token in setup response")
}
profileResp := performJSONRequest(t, router, http.MethodGet, "/api/auth/profile", nil, setupPayload.Data.Token)
if profileResp.Code != http.StatusOK {
t.Fatalf("unexpected profile status: %d body=%s", profileResp.Code, profileResp.Body.String())
}
loginBody := map[string]string{"username": "admin", "password": "super-secret"}
loginResp := performJSONRequest(t, router, http.MethodPost, "/api/auth/login", loginBody, "")
if loginResp.Code != http.StatusOK {
t.Fatalf("unexpected login status: %d body=%s", loginResp.Code, loginResp.Body.String())
}
systemResp := performJSONRequest(t, router, http.MethodGet, "/api/system/info", nil, setupPayload.Data.Token)
if systemResp.Code != http.StatusOK {
t.Fatalf("unexpected system info status: %d body=%s", systemResp.Code, systemResp.Body.String())
}
}
func performJSONRequest(t *testing.T, handler http.Handler, method string, path string, payload any, token string) *httptest.ResponseRecorder {
t.Helper()
var body []byte
if payload != nil {
encoded, err := json.Marshal(payload)
if err != nil {
t.Fatalf("json.Marshal() error = %v", err)
}
body = encoded
}
request := httptest.NewRequest(method, path, bytes.NewReader(body))
request.Header.Set("Content-Type", "application/json")
if token != "" {
request.Header.Set("Authorization", "Bearer "+token)
}
response := httptest.NewRecorder()
handler.ServeHTTP(response, request)
return response
}
-25
View File
@@ -1,25 +0,0 @@
//go:build ignore
package httpapi
import (
"backupx/server/internal/service"
"backupx/server/pkg/response"
"github.com/gin-gonic/gin"
)
type systemHandler struct {
service *service.SystemService
}
func newSystemHandler(service *service.SystemService) *systemHandler {
return &systemHandler{service: service}
}
func (h *systemHandler) registerRoutes(protected gin.IRouter) {
protected.GET("/system/info", h.info)
}
func (h *systemHandler) info(c *gin.Context) {
response.Success(c, h.service.GetInfo())
}
+86
View File
@@ -0,0 +1,86 @@
package lifecycle
import (
"context"
"sync"
)
// Supervisor owns application background tasks. It rejects new work once
// shutdown starts, cancels the shared task context, and waits for accepted work.
type Supervisor struct {
ctx context.Context
cancel context.CancelFunc
mu sync.Mutex
stopping bool
wg sync.WaitGroup
done chan struct{}
stopOnce sync.Once
}
func NewSupervisor(parent context.Context) *Supervisor {
if parent == nil {
parent = context.Background()
}
ctx, cancel := context.WithCancel(parent)
return &Supervisor{
ctx: ctx,
cancel: cancel,
done: make(chan struct{}),
}
}
// Context is the root context passed to every accepted task.
func (s *Supervisor) Context() context.Context {
return s.ctx
}
// Go starts task unless shutdown has begun or the root context is already
// canceled. The lock makes Add and the transition to Wait mutually exclusive.
func (s *Supervisor) Go(task func(context.Context)) bool {
if s == nil || task == nil {
return false
}
s.mu.Lock()
if s.stopping || s.ctx.Err() != nil {
s.mu.Unlock()
return false
}
s.wg.Add(1)
s.mu.Unlock()
go func() {
defer s.wg.Done()
task(s.ctx)
}()
return true
}
// Shutdown is idempotent. Cancellation always happens, even when waitCtx has
// already expired; callers may call Shutdown again to wait for eventual exit.
func (s *Supervisor) Shutdown(waitCtx context.Context) error {
if s == nil {
return nil
}
if waitCtx == nil {
waitCtx = context.Background()
}
s.stopOnce.Do(func() {
s.mu.Lock()
s.stopping = true
s.cancel()
s.mu.Unlock()
go func() {
s.wg.Wait()
close(s.done)
}()
})
select {
case <-s.done:
return nil
case <-waitCtx.Done():
return waitCtx.Err()
}
}
@@ -0,0 +1,54 @@
package lifecycle
import (
"context"
"errors"
"testing"
"time"
)
func TestSupervisorShutdownCancelsAndWaits(t *testing.T) {
supervisor := NewSupervisor(context.Background())
started := make(chan struct{})
finished := make(chan struct{})
if !supervisor.Go(func(ctx context.Context) {
close(started)
<-ctx.Done()
close(finished)
}) {
t.Fatal("expected task to be accepted")
}
<-started
waitCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := supervisor.Shutdown(waitCtx); err != nil {
t.Fatalf("Shutdown returned error: %v", err)
}
select {
case <-finished:
default:
t.Fatal("Shutdown returned before the task finished")
}
if supervisor.Go(func(context.Context) {}) {
t.Fatal("expected task submitted after shutdown to be rejected")
}
}
func TestSupervisorShutdownHonorsWaitContext(t *testing.T) {
supervisor := NewSupervisor(context.Background())
release := make(chan struct{})
if !supervisor.Go(func(context.Context) { <-release }) {
t.Fatal("expected task to be accepted")
}
waitCtx, cancel := context.WithCancel(context.Background())
cancel()
if err := supervisor.Shutdown(waitCtx); !errors.Is(err, context.Canceled) {
t.Fatalf("Shutdown error = %v, want context.Canceled", err)
}
close(release)
if err := supervisor.Shutdown(context.Background()); err != nil {
t.Fatalf("second Shutdown returned error: %v", err)
}
}
+26 -8
View File
@@ -17,6 +17,11 @@ type SampleSource interface {
CountSLABreach(ctx context.Context) (int, error) CountSLABreach(ctx context.Context) (int, error)
} }
// BackgroundRunner is implemented by the application lifecycle supervisor.
type BackgroundRunner interface {
Go(func(context.Context)) bool
}
// repoSource 把 repository 适配到 SampleSource。 // repoSource 把 repository 适配到 SampleSource。
type repoSource struct { type repoSource struct {
targets repository.StorageTargetRepository targets repository.StorageTargetRepository
@@ -90,9 +95,10 @@ func (s *repoSource) CountSLABreach(ctx context.Context) (int, error) {
// Collector 周期性采集 gauge 类指标(存储用量、节点在线、SLA 违约)。 // Collector 周期性采集 gauge 类指标(存储用量、节点在线、SLA 违约)。
// 用后台 goroutine 驱动,避免在 /metrics 请求路径做慢 IO。 // 用后台 goroutine 驱动,避免在 /metrics 请求路径做慢 IO。
type Collector struct { type Collector struct {
metrics *Metrics metrics *Metrics
source SampleSource source SampleSource
interval time.Duration interval time.Duration
background BackgroundRunner
} }
// NewCollector 创建周期采集器。interval=0 走默认 30s。 // NewCollector 创建周期采集器。interval=0 走默认 30s。
@@ -103,25 +109,37 @@ func NewCollector(m *Metrics, source SampleSource, interval time.Duration) *Coll
return &Collector{metrics: m, source: source, interval: interval} return &Collector{metrics: m, source: source, interval: interval}
} }
func (c *Collector) SetBackgroundRunner(runner BackgroundRunner) {
c.background = runner
}
// Start 在后台运行采集循环;随 ctx 取消而终止。 // Start 在后台运行采集循环;随 ctx 取消而终止。
// 启动时立即采一次,之后按 interval 轮询。 // 启动时立即采一次,之后按 interval 轮询。
func (c *Collector) Start(ctx context.Context) { func (c *Collector) Start(ctx context.Context) {
if c == nil || c.metrics == nil || c.source == nil { if c == nil || c.metrics == nil || c.source == nil {
return return
} }
go func() { run := func(runCtx context.Context) {
c.collect(ctx) c.collect(runCtx)
ticker := time.NewTicker(c.interval) ticker := time.NewTicker(c.interval)
defer ticker.Stop() defer ticker.Stop()
for { for {
select { select {
case <-ctx.Done(): case <-runCtx.Done():
return return
case <-ticker.C: case <-ticker.C:
c.collect(ctx) c.collect(runCtx)
} }
} }
}() }
if c.background != nil {
c.background.Go(run)
return
}
if ctx == nil {
ctx = context.Background()
}
go run(ctx)
} }
// collect 执行一次采样;单轮失败不影响下次。 // collect 执行一次采样;单轮失败不影响下次。
@@ -0,0 +1,64 @@
package metrics
import (
"context"
"testing"
"time"
"backupx/server/internal/model"
"backupx/server/internal/repository"
)
type lifecycleSampleSource struct{}
func (lifecycleSampleSource) ListStorageTargets(context.Context) ([]model.StorageTarget, error) {
return nil, nil
}
func (lifecycleSampleSource) StorageUsage(context.Context) ([]repository.BackupStorageUsageItem, error) {
return nil, nil
}
func (lifecycleSampleSource) ListNodes(context.Context) ([]model.Node, error) {
return nil, nil
}
func (lifecycleSampleSource) AgentQueueSummaries(context.Context) (map[uint]repository.AgentCommandQueueSummary, error) {
return nil, nil
}
func (lifecycleSampleSource) CountSLABreach(context.Context) (int, error) {
return 0, nil
}
type capturingCollectorRunner struct {
task func(context.Context)
}
func (r *capturingCollectorRunner) Go(task func(context.Context)) bool {
r.task = task
return true
}
func TestCollectorUsesConfiguredBackgroundRunner(t *testing.T) {
runner := &capturingCollectorRunner{}
collector := NewCollector(New("test"), lifecycleSampleSource{}, time.Hour)
collector.SetBackgroundRunner(runner)
collector.Start(context.Background())
if runner.task == nil {
t.Fatal("collector did not register with the background runner")
}
runCtx, cancel := context.WithCancel(context.Background())
cancel()
done := make(chan struct{})
go func() {
runner.task(runCtx)
close(done)
}()
select {
case <-done:
case <-time.After(time.Second):
t.Fatal("collector did not stop when background runner context was canceled")
}
}
+13 -3
View File
@@ -38,6 +38,7 @@ type Service struct {
verifyRunner VerifyRunner verifyRunner VerifyRunner
logger *zap.Logger logger *zap.Logger
audit AuditRecorder audit AuditRecorder
runCtx context.Context
entries map[uint]cron.EntryID // 备份 cron 条目 entries map[uint]cron.EntryID // 备份 cron 条目
verifyEntries map[uint]cron.EntryID // 验证 cron 条目 verifyEntries map[uint]cron.EntryID // 验证 cron 条目
} }
@@ -49,6 +50,7 @@ func NewService(tasks repository.BackupTaskRepository, runner TaskRunner, logger
tasks: tasks, tasks: tasks,
runner: runner, runner: runner,
logger: logger, logger: logger,
runCtx: context.Background(),
entries: make(map[uint]cron.EntryID), entries: make(map[uint]cron.EntryID),
verifyEntries: make(map[uint]cron.EntryID), verifyEntries: make(map[uint]cron.EntryID),
} }
@@ -72,6 +74,12 @@ func (s *Service) SetNodeRepository(nodes repository.NodeRepository) {
} }
func (s *Service) Start(ctx context.Context) error { func (s *Service) Start(ctx context.Context) error {
if ctx == nil {
ctx = context.Background()
}
s.mu.Lock()
s.runCtx = ctx
s.mu.Unlock()
if err := s.Reload(ctx); err != nil { if err := s.Reload(ctx); err != nil {
return err return err
} }
@@ -163,10 +171,11 @@ func (s *Service) syncTaskLocked(task *model.BackupTask) error {
taskNodeID := task.NodeID taskNodeID := task.NodeID
cronExpr := task.CronExpr cronExpr := task.CronExpr
maintenanceWindows := task.MaintenanceWindows maintenanceWindows := task.MaintenanceWindows
runCtx := s.runCtx
entryID, err := s.cron.AddFunc(cronExpr, func() { entryID, err := s.cron.AddFunc(cronExpr, func() {
// 集群感知:若任务绑定了离线的远程节点,跳过本轮触发避免堆积 failed 记录 // 集群感知:若任务绑定了离线的远程节点,跳过本轮触发避免堆积 failed 记录
if taskNodeID > 0 && s.nodes != nil { if taskNodeID > 0 && s.nodes != nil {
node, err := s.nodes.FindByID(context.Background(), taskNodeID) node, err := s.nodes.FindByID(runCtx, taskNodeID)
// 用实时推导的状态判定,避免后台监控刷新前把任务下发给刚失联的节点。 // 用实时推导的状态判定,避免后台监控刷新前把任务下发给刚失联的节点。
if err == nil && node != nil && !node.IsLocal && node.EffectiveStatus(time.Now().UTC()) != model.NodeStatusOnline { if err == nil && node != nil && !node.IsLocal && node.EffectiveStatus(time.Now().UTC()) != model.NodeStatusOnline {
if s.logger != nil { if s.logger != nil {
@@ -213,7 +222,7 @@ func (s *Service) syncTaskLocked(task *model.BackupTask) error {
TargetName: taskName, Detail: fmt.Sprintf("定时调度触发备份任务: %s (cron: %s)", taskName, cronExpr), TargetName: taskName, Detail: fmt.Sprintf("定时调度触发备份任务: %s (cron: %s)", taskName, cronExpr),
}) })
} }
if _, runErr := s.runner.RunTaskByID(context.Background(), taskID); runErr != nil && s.logger != nil { if _, runErr := s.runner.RunTaskByID(runCtx, taskID); runErr != nil && s.logger != nil {
s.logger.Warn("scheduled backup run failed", zap.Uint("task_id", taskID), zap.Error(runErr)) s.logger.Warn("scheduled backup run failed", zap.Uint("task_id", taskID), zap.Error(runErr))
} }
}) })
@@ -245,6 +254,7 @@ func (s *Service) syncVerifyTaskLocked(task *model.BackupTask) error {
taskName := task.Name taskName := task.Name
mode := task.VerifyMode mode := task.VerifyMode
verifyCron := task.VerifyCronExpr verifyCron := task.VerifyCronExpr
runCtx := s.runCtx
entryID, err := s.cron.AddFunc(verifyCron, func() { entryID, err := s.cron.AddFunc(verifyCron, func() {
if s.audit != nil { if s.audit != nil {
s.audit.Record(servicepkg.AuditEntry{ s.audit.Record(servicepkg.AuditEntry{
@@ -253,7 +263,7 @@ func (s *Service) syncVerifyTaskLocked(task *model.BackupTask) error {
TargetName: taskName, Detail: fmt.Sprintf("定时验证演练: %s (cron: %s, mode: %s)", taskName, verifyCron, mode), TargetName: taskName, Detail: fmt.Sprintf("定时验证演练: %s (cron: %s, mode: %s)", taskName, verifyCron, mode),
}) })
} }
if _, runErr := s.verifyRunner.StartByTask(context.Background(), taskID, mode, "system"); runErr != nil && s.logger != nil { if _, runErr := s.verifyRunner.StartByTask(runCtx, taskID, mode, "system"); runErr != nil && s.logger != nil {
s.logger.Warn("scheduled verify run failed", zap.Uint("task_id", taskID), zap.Error(runErr)) s.logger.Warn("scheduled verify run failed", zap.Uint("task_id", taskID), zap.Error(runErr))
} }
}) })
-60
View File
@@ -1,60 +0,0 @@
//go:build ignore
package security
import (
"fmt"
"time"
"backupx/server/internal/model"
"github.com/golang-jwt/jwt/v5"
)
type Claims struct {
UserID uint `json:"userId"`
Username string `json:"username"`
Role string `json:"role"`
jwt.RegisteredClaims
}
type JWTManager struct {
secret []byte
duration time.Duration
}
func NewJWTManager(secret string, duration time.Duration) *JWTManager {
return &JWTManager{secret: []byte(secret), duration: duration}
}
func (m *JWTManager) IssueToken(user *model.User) (string, error) {
now := time.Now().UTC()
claims := Claims{
UserID: user.ID,
Username: user.Username,
Role: user.Role,
RegisteredClaims: jwt.RegisteredClaims{
Subject: fmt.Sprintf("%d", user.ID),
IssuedAt: jwt.NewNumericDate(now),
ExpiresAt: jwt.NewNumericDate(now.Add(m.duration)),
},
}
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
return token.SignedString(m.secret)
}
func (m *JWTManager) Parse(tokenValue string) (*Claims, error) {
token, err := jwt.ParseWithClaims(tokenValue, &Claims{}, func(token *jwt.Token) (any, error) {
if token.Method != jwt.SigningMethodHS256 {
return nil, fmt.Errorf("unexpected signing method")
}
return m.secret, nil
})
if err != nil {
return nil, err
}
claims, ok := token.Claims.(*Claims)
if !ok || !token.Valid {
return nil, fmt.Errorf("invalid token")
}
return claims, nil
}
-25
View File
@@ -1,25 +0,0 @@
//go:build ignore
package security
import (
"testing"
"time"
"backupx/server/internal/model"
)
func TestJWTManagerIssueAndParse(t *testing.T) {
manager := NewJWTManager("test-secret", time.Hour)
token, err := manager.IssueToken(&model.User{ID: 7, Username: "admin", Role: "admin"})
if err != nil {
t.Fatalf("IssueToken() error = %v", err)
}
claims, err := manager.Parse(token)
if err != nil {
t.Fatalf("Parse() error = %v", err)
}
if claims.UserID != 7 || claims.Username != "admin" {
t.Fatalf("unexpected claims: %+v", claims)
}
}
-54
View File
@@ -1,54 +0,0 @@
//go:build ignore
package security
import (
"sync"
"time"
)
type limiterEntry struct {
Count int
ResetAt time.Time
}
type LoginLimiter struct {
mu sync.Mutex
window time.Duration
max int
records map[string]limiterEntry
}
func NewLoginLimiter(max int, window time.Duration) *LoginLimiter {
return &LoginLimiter{window: window, max: max, records: make(map[string]limiterEntry)}
}
func (l *LoginLimiter) Allow(key string) bool {
l.mu.Lock()
defer l.mu.Unlock()
entry, ok := l.records[key]
if !ok || time.Now().After(entry.ResetAt) {
delete(l.records, key)
return true
}
return entry.Count < l.max
}
func (l *LoginLimiter) RegisterFailure(key string) {
l.mu.Lock()
defer l.mu.Unlock()
now := time.Now()
entry, ok := l.records[key]
if !ok || now.After(entry.ResetAt) {
l.records[key] = limiterEntry{Count: 1, ResetAt: now.Add(l.window)}
return
}
entry.Count++
l.records[key] = entry
}
func (l *LoginLimiter) Reset(key string) {
l.mu.Lock()
defer l.mu.Unlock()
delete(l.records, key)
}
-93
View File
@@ -1,93 +0,0 @@
//go:build ignore
package security
import (
"crypto/rand"
"encoding/hex"
"encoding/json"
"fmt"
"os"
"path/filepath"
"backupx/server/internal/config"
)
type PersistedSecrets struct {
JWTSecret string `json:"jwtSecret"`
EncryptionKey string `json:"encryptionKey"`
}
func EnsureSecrets(cfg *config.Config) error {
if cfg.Security.JWTSecret != "" && cfg.Security.EncryptionKey != "" {
return nil
}
storePath := filepath.Join(filepath.Dir(cfg.Database.Path), "backupx.secrets.json")
current, err := loadSecrets(storePath)
if err != nil {
return err
}
if current == nil {
current = &PersistedSecrets{}
}
if current.JWTSecret == "" {
current.JWTSecret, err = randomHex(32)
if err != nil {
return err
}
}
if current.EncryptionKey == "" {
current.EncryptionKey, err = randomHex(32)
if err != nil {
return err
}
}
if err := saveSecrets(storePath, current); err != nil {
return err
}
if cfg.Security.JWTSecret == "" {
cfg.Security.JWTSecret = current.JWTSecret
}
if cfg.Security.EncryptionKey == "" {
cfg.Security.EncryptionKey = current.EncryptionKey
}
return nil
}
func loadSecrets(path string) (*PersistedSecrets, error) {
content, err := os.ReadFile(path)
if err != nil {
if os.IsNotExist(err) {
return nil, nil
}
return nil, fmt.Errorf("read secrets: %w", err)
}
var secrets PersistedSecrets
if err := json.Unmarshal(content, &secrets); err != nil {
return nil, fmt.Errorf("decode secrets: %w", err)
}
return &secrets, nil
}
func saveSecrets(path string, secrets *PersistedSecrets) error {
if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil {
return fmt.Errorf("create secrets dir: %w", err)
}
content, err := json.MarshalIndent(secrets, "", " ")
if err != nil {
return fmt.Errorf("encode secrets: %w", err)
}
if err := os.WriteFile(path, content, 0o600); err != nil {
return fmt.Errorf("write secrets: %w", err)
}
return nil
}
func randomHex(size int) (string, error) {
bytes := make([]byte, size)
if _, err := rand.Read(bytes); err != nil {
return "", fmt.Errorf("generate random secret: %w", err)
}
return hex.EncodeToString(bytes), nil
}
+153 -44
View File
@@ -19,6 +19,7 @@ import (
"backupx/server/internal/repository" "backupx/server/internal/repository"
"backupx/server/internal/storage" "backupx/server/internal/storage"
"backupx/server/internal/storage/codec" "backupx/server/internal/storage/codec"
"go.uber.org/zap"
) )
// AgentService 实现 Master 端 Agent 协议,提供给远程 Agent 通过 HTTP 调用。 // AgentService 实现 Master 端 Agent 协议,提供给远程 Agent 通过 HTTP 调用。
@@ -32,6 +33,8 @@ type AgentService struct {
restoreRepo repository.RestoreRecordRepository restoreRepo repository.RestoreRecordRepository
registry *storage.Registry registry *storage.Registry
cipher *codec.ConfigCipher cipher *codec.ConfigCipher
logger *zap.Logger
background BackgroundRunner
} }
func NewAgentService( func NewAgentService(
@@ -51,9 +54,24 @@ func NewAgentService(
cmdRepo: cmdRepo, cmdRepo: cmdRepo,
registry: registry, registry: registry,
cipher: cipher, cipher: cipher,
logger: zap.NewNop(),
} }
} }
// SetLogger attaches the application logger used by background command
// reconciliation. The no-op default keeps the service safe in tests.
func (s *AgentService) SetLogger(logger *zap.Logger) {
if logger != nil {
s.logger = logger
}
}
// SetBackgroundRunner makes the command timeout monitor part of the
// application lifecycle so shutdown waits for an in-flight reconciliation.
func (s *AgentService) SetBackgroundRunner(background BackgroundRunner) {
s.background = background
}
// SetRestoreRepository 注入恢复记录仓储,用于命令超时时联动 restore_record 状态。 // SetRestoreRepository 注入恢复记录仓储,用于命令超时时联动 restore_record 状态。
// 可选注入:未注入时恢复命令超时仅标记命令 timeout,记录需另行查验。 // 可选注入:未注入时恢复命令超时仅标记命令 timeout,记录需另行查验。
func (s *AgentService) SetRestoreRepository(repo repository.RestoreRecordRepository) { func (s *AgentService) SetRestoreRepository(repo repository.RestoreRecordRepository) {
@@ -117,6 +135,16 @@ func (s *AgentService) SubmitCommandResult(ctx context.Context, node *model.Node
if cmd.NodeID != node.ID { if cmd.NodeID != node.ID {
return apperror.Unauthorized("AGENT_COMMAND_FORBIDDEN", "命令不属于当前节点", nil) return apperror.Unauthorized("AGENT_COMMAND_FORBIDDEN", "命令不属于当前节点", nil)
} }
// A failed terminal report may be retried after the command row was already
// completed but before its linked business record was updated. Re-run that
// idempotent convergence step so a transient database error cannot leave a
// backup or restore record stuck in running forever.
if cmd.Status == model.AgentCommandStatusFailed {
if result.Success {
return nil
}
return s.failLinkedRecord(ctx, cmd, agentCommandFailureMessage(result.ErrorMessage))
}
now := time.Now().UTC() now := time.Now().UTC()
if result.Success { if result.Success {
cmd.Status = model.AgentCommandStatusSucceeded cmd.Status = model.AgentCommandStatusSucceeded
@@ -128,8 +156,21 @@ func (s *AgentService) SubmitCommandResult(ctx context.Context, node *model.Node
cmd.Result = string(result.Result) cmd.Result = string(result.Result)
} }
cmd.CompletedAt = &now cmd.CompletedAt = &now
_, err = s.cmdRepo.CompleteDispatched(ctx, cmd) completed, err := s.cmdRepo.CompleteDispatched(ctx, cmd)
return err if err != nil || !completed || result.Success {
return err
}
persistCtx, cancel := finalizationContext(ctx)
defer cancel()
return s.failLinkedRecord(persistCtx, cmd, agentCommandFailureMessage(result.ErrorMessage))
}
func agentCommandFailureMessage(message string) string {
message = strings.TrimSpace(message)
if message == "" {
return "Agent 命令执行失败"
}
return "Agent 命令执行失败:" + message
} }
// AgentTaskSpec 给 Agent 返回的任务规格,包含解密后的存储配置,供 Agent 直接执行。 // AgentTaskSpec 给 Agent 返回的任务规格,包含解密后的存储配置,供 Agent 直接执行。
@@ -616,19 +657,29 @@ func (s *AgentService) StartCommandTimeoutMonitor(ctx context.Context, interval
if timeout <= 0 { if timeout <= 0 {
timeout = 10 * time.Minute timeout = 10 * time.Minute
} }
ticker := time.NewTicker(interval) monitor := func(runCtx context.Context) {
go func() { ticker := time.NewTicker(interval)
defer ticker.Stop() defer ticker.Stop()
for { for {
select { select {
case <-ctx.Done(): case <-runCtx.Done():
return return
case <-ticker.C: case <-ticker.C:
threshold := time.Now().UTC().Add(-timeout) threshold := time.Now().UTC().Add(-timeout)
s.processStaleCommands(ctx, threshold) s.processStaleCommands(runCtx, threshold)
} }
} }
}() }
if s.background != nil {
if !s.background.Go(monitor) {
s.logger.Warn("agent command timeout monitor not started: application is shutting down")
}
return
}
if ctx == nil {
ctx = context.Background()
}
go monitor(ctx)
} }
// processStaleCommands 扫描已超时的 pending/dispatched 命令并联动关联记录。 // processStaleCommands 扫描已超时的 pending/dispatched 命令并联动关联记录。
@@ -636,12 +687,24 @@ func (s *AgentService) StartCommandTimeoutMonitor(ctx context.Context, interval
// 单条失败不影响后续处理。 // 单条失败不影响后续处理。
func (s *AgentService) processStaleCommands(ctx context.Context, threshold time.Time) { func (s *AgentService) processStaleCommands(ctx context.Context, threshold time.Time) {
commands, err := s.cmdRepo.ListStaleActive(ctx, threshold) commands, err := s.cmdRepo.ListStaleActive(ctx, threshold)
if err != nil || len(commands) == 0 { if err != nil {
s.logger.Error("list stale agent commands failed", zap.Error(err))
return
}
if len(commands) == 0 {
return return
} }
for i := range commands { for i := range commands {
cmd := commands[i] cmd := commands[i]
if s.commandStillActive(ctx, &cmd, threshold) { stillActive, activeErr := s.commandStillActive(ctx, &cmd, threshold)
if activeErr != nil {
s.logger.Warn("check stale agent command activity failed",
zap.Uint("command_id", cmd.ID),
zap.String("command_type", cmd.Type),
zap.Error(activeErr))
continue
}
if stillActive {
continue continue
} }
now := time.Now().UTC() now := time.Now().UTC()
@@ -649,18 +712,33 @@ func (s *AgentService) processStaleCommands(ctx context.Context, threshold time.
cmd.ErrorMessage = "agent did not report result before timeout" cmd.ErrorMessage = "agent did not report result before timeout"
cmd.CompletedAt = &now cmd.CompletedAt = &now
timedOut, err := s.cmdRepo.TimeoutActive(ctx, &cmd) timedOut, err := s.cmdRepo.TimeoutActive(ctx, &cmd)
if err != nil || !timedOut { if err != nil {
s.logger.Error("mark agent command timed out failed",
zap.Uint("command_id", cmd.ID),
zap.String("command_type", cmd.Type),
zap.Error(err))
continue continue
} }
s.failLinkedRecord(ctx, &cmd) if !timedOut {
continue
}
persistCtx, cancel := finalizationContext(ctx)
failErr := s.failLinkedRecord(persistCtx, &cmd)
cancel()
if failErr != nil {
s.logger.Error("mark timed-out agent command record failed",
zap.Uint("command_id", cmd.ID),
zap.String("command_type", cmd.Type),
zap.Error(failErr))
}
} }
} }
// commandStillActive 用关联记录状态、记录更新时间和节点心跳作为长任务续租信号。 // commandStillActive 用关联记录状态、记录更新时间和节点心跳作为长任务续租信号。
// 仅 run_task / restore_record 允许续租,避免短 RPC 命令被在线节点长期保留。 // 仅 run_task / restore_record 允许续租,避免短 RPC 命令被在线节点长期保留。
func (s *AgentService) commandStillActive(ctx context.Context, cmd *model.AgentCommand, threshold time.Time) bool { func (s *AgentService) commandStillActive(ctx context.Context, cmd *model.AgentCommand, threshold time.Time) (bool, error) {
if cmd.Status != model.AgentCommandStatusDispatched { if cmd.Status != model.AgentCommandStatusDispatched {
return false return false, nil
} }
switch cmd.Type { switch cmd.Type {
case model.AgentCommandTypeRunTask: case model.AgentCommandTypeRunTask:
@@ -668,90 +746,121 @@ func (s *AgentService) commandStillActive(ctx context.Context, cmd *model.AgentC
RecordID uint `json:"recordId"` RecordID uint `json:"recordId"`
} }
if err := json.Unmarshal([]byte(cmd.Payload), &payload); err != nil || payload.RecordID == 0 { if err := json.Unmarshal([]byte(cmd.Payload), &payload); err != nil || payload.RecordID == 0 {
return false return false, nil
} }
record, err := s.recordRepo.FindByID(ctx, payload.RecordID) record, err := s.recordRepo.FindByID(ctx, payload.RecordID)
if err != nil || record == nil || record.Status != model.BackupRecordStatusRunning { if err != nil {
return false return false, fmt.Errorf("find backup record %d: %w", payload.RecordID, err)
} }
if s.nodeRecentlySeen(ctx, cmd.NodeID, threshold) { if record == nil || record.Status != model.BackupRecordStatusRunning {
return true return false, nil
} }
return record.UpdatedAt.After(threshold) nodeActive, err := s.nodeRecentlySeen(ctx, cmd.NodeID, threshold)
if err != nil {
return false, err
}
return nodeActive || record.UpdatedAt.After(threshold), nil
case model.AgentCommandTypeRestoreRecord: case model.AgentCommandTypeRestoreRecord:
if s.restoreRepo == nil { if s.restoreRepo == nil {
return false return false, nil
} }
var payload struct { var payload struct {
RestoreRecordID uint `json:"restoreRecordId"` RestoreRecordID uint `json:"restoreRecordId"`
} }
if err := json.Unmarshal([]byte(cmd.Payload), &payload); err != nil || payload.RestoreRecordID == 0 { if err := json.Unmarshal([]byte(cmd.Payload), &payload); err != nil || payload.RestoreRecordID == 0 {
return false return false, nil
} }
restore, err := s.restoreRepo.FindByID(ctx, payload.RestoreRecordID) restore, err := s.restoreRepo.FindByID(ctx, payload.RestoreRecordID)
if err != nil || restore == nil || restore.Status != model.RestoreRecordStatusRunning { if err != nil {
return false return false, fmt.Errorf("find restore record %d: %w", payload.RestoreRecordID, err)
} }
if s.nodeRecentlySeen(ctx, cmd.NodeID, threshold) { if restore == nil || restore.Status != model.RestoreRecordStatusRunning {
return true return false, nil
} }
return restore.UpdatedAt.After(threshold) nodeActive, err := s.nodeRecentlySeen(ctx, cmd.NodeID, threshold)
if err != nil {
return false, err
}
return nodeActive || restore.UpdatedAt.After(threshold), nil
default: default:
return false return false, nil
} }
} }
func (s *AgentService) nodeRecentlySeen(ctx context.Context, nodeID uint, threshold time.Time) bool { func (s *AgentService) nodeRecentlySeen(ctx context.Context, nodeID uint, threshold time.Time) (bool, error) {
node, err := s.nodeRepo.FindByID(ctx, nodeID) node, err := s.nodeRepo.FindByID(ctx, nodeID)
if err != nil || node == nil { if err != nil {
return false return false, fmt.Errorf("find agent node %d: %w", nodeID, err)
} }
return node.Status == model.NodeStatusOnline && node.LastSeen.After(threshold) if node == nil {
return false, nil
}
return node.Status == model.NodeStatusOnline && node.LastSeen.After(threshold), nil
} }
// failLinkedRecord 根据命令类型把关联记录标记为 failed。 // failLinkedRecord 根据命令类型把关联记录标记为 failed。
// 只对仍然处于 running 状态的记录生效,避免覆盖已完成的结果。 // 只对仍然处于 running 状态的记录生效,避免覆盖已完成的结果。
func (s *AgentService) failLinkedRecord(ctx context.Context, cmd *model.AgentCommand) { func (s *AgentService) failLinkedRecord(ctx context.Context, cmd *model.AgentCommand, messages ...string) error {
const failureMessage = "Agent 未在超时前回传状态(节点可能已离线或崩溃)" failureMessage := "Agent 未在超时前回传状态(节点可能已离线或崩溃)"
if len(messages) > 0 && strings.TrimSpace(messages[0]) != "" {
failureMessage = strings.TrimSpace(messages[0])
}
switch cmd.Type { switch cmd.Type {
case model.AgentCommandTypeRunTask: case model.AgentCommandTypeRunTask:
var payload struct { var payload struct {
RecordID uint `json:"recordId"` RecordID uint `json:"recordId"`
} }
if err := json.Unmarshal([]byte(cmd.Payload), &payload); err != nil || payload.RecordID == 0 { if err := json.Unmarshal([]byte(cmd.Payload), &payload); err != nil {
return return fmt.Errorf("decode run_task payload: %w", err)
}
if payload.RecordID == 0 {
return errors.New("run_task payload has no recordId")
} }
record, err := s.recordRepo.FindByID(ctx, payload.RecordID) record, err := s.recordRepo.FindByID(ctx, payload.RecordID)
if err != nil || record == nil || record.Status != model.BackupRecordStatusRunning { if err != nil {
return return fmt.Errorf("find backup record %d: %w", payload.RecordID, err)
}
if record == nil || record.Status != model.BackupRecordStatusRunning {
return nil
} }
completedAt := time.Now().UTC() completedAt := time.Now().UTC()
record.Status = model.BackupRecordStatusFailed record.Status = model.BackupRecordStatusFailed
record.ErrorMessage = failureMessage record.ErrorMessage = failureMessage
record.CompletedAt = &completedAt record.CompletedAt = &completedAt
record.DurationSeconds = int(completedAt.Sub(record.StartedAt).Seconds()) record.DurationSeconds = int(completedAt.Sub(record.StartedAt).Seconds())
_ = s.recordRepo.Update(ctx, record) if err := s.recordRepo.Update(ctx, record); err != nil {
return fmt.Errorf("update backup record %d: %w", record.ID, err)
}
case model.AgentCommandTypeRestoreRecord: case model.AgentCommandTypeRestoreRecord:
if s.restoreRepo == nil { if s.restoreRepo == nil {
return return errors.New("restore record repository is not configured")
} }
var payload struct { var payload struct {
RestoreRecordID uint `json:"restoreRecordId"` RestoreRecordID uint `json:"restoreRecordId"`
} }
if err := json.Unmarshal([]byte(cmd.Payload), &payload); err != nil || payload.RestoreRecordID == 0 { if err := json.Unmarshal([]byte(cmd.Payload), &payload); err != nil {
return return fmt.Errorf("decode restore_record payload: %w", err)
}
if payload.RestoreRecordID == 0 {
return errors.New("restore_record payload has no restoreRecordId")
} }
restore, err := s.restoreRepo.FindByID(ctx, payload.RestoreRecordID) restore, err := s.restoreRepo.FindByID(ctx, payload.RestoreRecordID)
if err != nil || restore == nil || restore.Status != model.RestoreRecordStatusRunning { if err != nil {
return return fmt.Errorf("find restore record %d: %w", payload.RestoreRecordID, err)
}
if restore == nil || restore.Status != model.RestoreRecordStatusRunning {
return nil
} }
completedAt := time.Now().UTC() completedAt := time.Now().UTC()
restore.Status = model.RestoreRecordStatusFailed restore.Status = model.RestoreRecordStatusFailed
restore.ErrorMessage = failureMessage restore.ErrorMessage = failureMessage
restore.CompletedAt = &completedAt restore.CompletedAt = &completedAt
restore.DurationSeconds = int(completedAt.Sub(restore.StartedAt).Seconds()) restore.DurationSeconds = int(completedAt.Sub(restore.StartedAt).Seconds())
_ = s.restoreRepo.Update(ctx, restore) if err := s.restoreRepo.Update(ctx, restore); err != nil {
return fmt.Errorf("update restore record %d: %w", restore.ID, err)
}
} }
return nil
} }
// AgentSelfStatus 是 /api/v1/agent/self 端点返回给 Agent 的轻量状态摘要。 // AgentSelfStatus 是 /api/v1/agent/self 端点返回给 Agent 的轻量状态摘要。
@@ -24,6 +24,15 @@ import (
"gorm.io/gorm" "gorm.io/gorm"
) )
type failingUpdateBackupRecordRepository struct {
repository.BackupRecordRepository
updateErr error
}
func (r *failingUpdateBackupRecordRepository) Update(context.Context, *model.BackupRecord) error {
return r.updateErr
}
func newAgentServicePoolTestHarness(t *testing.T) (*AgentService, *gorm.DB, repository.BackupRecordRepository, repository.AgentCommandRepository, *model.Node, *model.Node) { func newAgentServicePoolTestHarness(t *testing.T) (*AgentService, *gorm.DB, repository.BackupRecordRepository, repository.AgentCommandRepository, *model.Node, *model.Node) {
t.Helper() t.Helper()
log, err := logger.New(config.LogConfig{Level: "error"}) log, err := logger.New(config.LogConfig{Level: "error"})
@@ -34,6 +43,7 @@ func newAgentServicePoolTestHarness(t *testing.T) (*AgentService, *gorm.DB, repo
if err != nil { if err != nil {
t.Fatalf("database.Open returned error: %v", err) t.Fatalf("database.Open returned error: %v", err)
} }
closeTestDatabase(t, db)
cipher := codec.NewConfigCipher("agent-service-secret") cipher := codec.NewConfigCipher("agent-service-secret")
nodeRepo := repository.NewNodeRepository(db) nodeRepo := repository.NewNodeRepository(db)
taskRepo := repository.NewBackupTaskRepository(db) taskRepo := repository.NewBackupTaskRepository(db)
@@ -87,6 +97,23 @@ func newAgentServicePoolTestHarness(t *testing.T) (*AgentService, *gorm.DB, repo
return NewAgentService(nodeRepo, taskRepo, recordRepo, storageRepo, cmdRepo, cipher, storageRegistry), db, recordRepo, cmdRepo, owner, other return NewAgentService(nodeRepo, taskRepo, recordRepo, storageRepo, cmdRepo, cipher, storageRegistry), db, recordRepo, cmdRepo, owner, other
} }
func TestAgentServiceFailLinkedRecordPropagatesTerminalUpdateError(t *testing.T) {
svc, _, records, _, _, _ := newAgentServicePoolTestHarness(t)
wantErr := errors.New("record update failed")
svc.recordRepo = &failingUpdateBackupRecordRepository{
BackupRecordRepository: records,
updateErr: wantErr,
}
err := svc.failLinkedRecord(context.Background(), &model.AgentCommand{
Type: model.AgentCommandTypeRunTask,
Payload: `{"recordId":1}`,
})
if err == nil || !errors.Is(err, wantErr) {
t.Fatalf("failLinkedRecord error = %v, want wrapped update error", err)
}
}
func TestAgentServicePooledTaskUsesRecordNodeForSpecAndRecordUpdates(t *testing.T) { func TestAgentServicePooledTaskUsesRecordNodeForSpecAndRecordUpdates(t *testing.T) {
svc, _, records, _, owner, other := newAgentServicePoolTestHarness(t) svc, _, records, _, owner, other := newAgentServicePoolTestHarness(t)
ctx := context.Background() ctx := context.Background()
@@ -752,6 +779,44 @@ func TestAgentServiceSubmitCommandResultDoesNotOverwriteTerminalCommand(t *testi
} }
} }
func TestAgentServiceSubmitFailedCommandConvergesLinkedRecord(t *testing.T) {
svc, _, records, commands, owner, _ := newAgentServicePoolTestHarness(t)
ctx := context.Background()
dispatchedAt := time.Now().UTC()
command := &model.AgentCommand{
NodeID: owner.ID,
Type: model.AgentCommandTypeRunTask,
Status: model.AgentCommandStatusDispatched,
Payload: `{"recordId":1}`,
DispatchedAt: &dispatchedAt,
}
if err := commands.Create(ctx, command); err != nil {
t.Fatalf("Create command returned error: %v", err)
}
if err := svc.SubmitCommandResult(ctx, owner, command.ID, AgentCommandResult{
Success: false,
ErrorMessage: "terminal update could not reach Master",
}); err != nil {
t.Fatalf("SubmitCommandResult returned error: %v", err)
}
record, err := records.FindByID(ctx, 1)
if err != nil {
t.Fatalf("FindByID record returned error: %v", err)
}
if record.Status != model.BackupRecordStatusFailed || !strings.Contains(record.ErrorMessage, "terminal update") {
t.Fatalf("linked record did not converge: %#v", record)
}
updatedCommand, err := commands.FindByID(ctx, command.ID)
if err != nil {
t.Fatalf("FindByID command returned error: %v", err)
}
if updatedCommand.Status != model.AgentCommandStatusFailed {
t.Fatalf("command status = %q, want failed", updatedCommand.Status)
}
}
func TestAgentServiceUpdateRecordDoesNotOverwriteTerminalRecord(t *testing.T) { func TestAgentServiceUpdateRecordDoesNotOverwriteTerminalRecord(t *testing.T) {
svc, _, records, _, owner, _ := newAgentServicePoolTestHarness(t) svc, _, records, _, owner, _ := newAgentServicePoolTestHarness(t)
ctx := context.Background() ctx := context.Background()
@@ -24,6 +24,7 @@ func newApiKeyTestService(t *testing.T) *ApiKeyService {
if err != nil { if err != nil {
t.Fatalf("database.Open: %v", err) t.Fatalf("database.Open: %v", err)
} }
closeTestDatabase(t, db)
return NewApiKeyService(repository.NewApiKeyRepository(db)) return NewApiKeyService(repository.NewApiKeyRepository(db))
} }
@@ -23,6 +23,11 @@ func TestAuditRetention(t *testing.T) {
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
sqlDB, err := db.DB()
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = sqlDB.Close() })
auditRepo := repository.NewAuditLogRepository(db) auditRepo := repository.NewAuditLogRepository(db)
configRepo := repository.NewSystemConfigRepository(db) configRepo := repository.NewSystemConfigRepository(db)
svc := NewAuditService(auditRepo) svc := NewAuditService(auditRepo)
+83 -28
View File
@@ -8,7 +8,6 @@ import (
"encoding/hex" "encoding/hex"
"encoding/json" "encoding/json"
"fmt" "fmt"
"log"
"net/http" "net/http"
"strconv" "strconv"
"strings" "strings"
@@ -18,6 +17,7 @@ import (
"backupx/server/internal/apperror" "backupx/server/internal/apperror"
"backupx/server/internal/model" "backupx/server/internal/model"
"backupx/server/internal/repository" "backupx/server/internal/repository"
"go.uber.org/zap"
) )
// AuditEntry 是记录审计日志的输入结构 // AuditEntry 是记录审计日志的输入结构
@@ -41,14 +41,35 @@ type AuditService struct {
webhookURL string webhookURL string
webhookSecret string webhookSecret string
httpClient *http.Client httpClient *http.Client
async func(func(context.Context)) bool
logger *zap.Logger
inFlight chan struct{}
} }
const maxAuditInFlight = 64
func NewAuditService(repo repository.AuditLogRepository) *AuditService { func NewAuditService(repo repository.AuditLogRepository) *AuditService {
return &AuditService{ return &AuditService{
repo: repo, repo: repo,
httpClient: &http.Client{ httpClient: &http.Client{
Timeout: 3 * time.Second, // 短超时:审计 webhook 不应拖慢业务 Timeout: 3 * time.Second, // 短超时:审计 webhook 不应拖慢业务
}, },
async: runDetached,
logger: zap.NewNop(),
inFlight: make(chan struct{}, maxAuditInFlight),
}
}
func (s *AuditService) SetLogger(logger *zap.Logger) {
if logger != nil {
s.logger = logger
}
}
// SetBackgroundRunner binds audit persistence and webhook delivery to the application lifecycle.
func (s *AuditService) SetBackgroundRunner(runner BackgroundRunner) {
if runner != nil {
s.async = runner.Go
} }
} }
@@ -71,37 +92,55 @@ func (s *AuditService) StartRetentionMonitor(ctx context.Context, configs reposi
if interval <= 0 { if interval <= 0 {
interval = 6 * time.Hour interval = 6 * time.Hour
} }
go func() { if ctx == nil {
ctx = context.Background()
}
accepted := s.async(func(workerCtx context.Context) {
monitorCtx, cancel := context.WithCancel(workerCtx)
defer cancel()
stopLink := context.AfterFunc(ctx, cancel)
defer stopLink()
ticker := time.NewTicker(interval) ticker := time.NewTicker(interval)
defer ticker.Stop() defer ticker.Stop()
s.runRetentionOnce(ctx, configs) // 启动后立即跑一次 s.runRetentionOnce(monitorCtx, configs) // 启动后立即跑一次
for { for {
select { select {
case <-ctx.Done(): case <-monitorCtx.Done():
return return
case <-ticker.C: case <-ticker.C:
s.runRetentionOnce(ctx, configs) s.runRetentionOnce(monitorCtx, configs)
} }
} }
}() })
if !accepted {
s.logger.Warn("audit retention monitor not started: application is shutting down")
}
} }
func (s *AuditService) runRetentionOnce(ctx context.Context, configs repository.SystemConfigRepository) { func (s *AuditService) runRetentionOnce(ctx context.Context, configs repository.SystemConfigRepository) {
cfg, err := configs.GetByKey(ctx, SettingKeyAuditRetentionDays) cfg, err := configs.GetByKey(ctx, SettingKeyAuditRetentionDays)
if err != nil || cfg == nil { if err != nil {
s.logger.Warn("read audit retention setting failed", zap.Error(err))
return
}
if cfg == nil {
return return
} }
days, err := strconv.Atoi(strings.TrimSpace(cfg.Value)) days, err := strconv.Atoi(strings.TrimSpace(cfg.Value))
if err != nil || days <= 0 { if err != nil {
s.logger.Warn("invalid audit retention setting", zap.String("value", cfg.Value), zap.Error(err))
return
}
if days <= 0 {
return return
} }
deleted, err := s.PurgeOlderThan(ctx, days) deleted, err := s.PurgeOlderThan(ctx, days)
if err != nil { if err != nil {
log.Printf("[audit] retention purge failed: %v", err) s.logger.Warn("audit retention purge failed", zap.Error(err))
return return
} }
if deleted > 0 { if deleted > 0 {
log.Printf("[audit] retention purge: deleted %d logs older than %d days", deleted, days) s.logger.Info("audit retention purge completed", zap.Int64("deleted", deleted), zap.Int("retention_days", days))
} }
} }
@@ -123,12 +162,24 @@ func (s *AuditService) SetWebhook(url, secret string) {
s.webhookSecret = strings.TrimSpace(secret) s.webhookSecret = strings.TrimSpace(secret)
} }
// Record 异步 fire-and-forget 写入审计日志,不阻塞业务逻辑 // Record asynchronously persists an audit event without blocking the request.
func (s *AuditService) Record(entry AuditEntry) { func (s *AuditService) Record(entry AuditEntry) {
if s == nil || s.repo == nil { if s == nil || s.repo == nil {
return return
} }
go func() { select {
case s.inFlight <- struct{}{}:
default:
s.logger.Error("audit event rejected: in-flight limit reached",
zap.Int("limit", cap(s.inFlight)),
zap.String("category", entry.Category),
zap.String("action", entry.Action))
return
}
accepted := s.async(func(workerCtx context.Context) {
defer func() { <-s.inFlight }()
persistCtx, cancel := finalizationContext(workerCtx)
defer cancel()
record := &model.AuditLog{ record := &model.AuditLog{
UserID: entry.UserID, UserID: entry.UserID,
Username: entry.Username, Username: entry.Username,
@@ -140,24 +191,30 @@ func (s *AuditService) Record(entry AuditEntry) {
Detail: entry.Detail, Detail: entry.Detail,
ClientIP: entry.ClientIP, ClientIP: entry.ClientIP,
} }
if err := s.repo.Create(context.Background(), record); err != nil { if err := s.repo.Create(persistCtx, record); err != nil {
log.Printf("[audit] failed to write audit log: %v", err) s.logger.Error("failed to write audit log", zap.String("category", entry.Category), zap.String("action", entry.Action), zap.Error(err))
} }
s.fireWebhook(record) if err := s.fireWebhook(persistCtx, record); err != nil {
}() s.logger.Warn("audit webhook delivery failed", zap.String("category", entry.Category), zap.String("action", entry.Action), zap.Error(err))
}
})
if !accepted {
<-s.inFlight
s.logger.Warn("audit event rejected: application is shutting down", zap.String("category", entry.Category), zap.String("action", entry.Action))
}
} }
// fireWebhook 异步向外部系统转发审计事件。失败降级到本地日志,永不影响主流程。 // fireWebhook forwards an audit event. The caller owns asynchronous execution.
func (s *AuditService) fireWebhook(record *model.AuditLog) { func (s *AuditService) fireWebhook(ctx context.Context, record *model.AuditLog) error {
if s == nil { if s == nil {
return return nil
} }
s.webhookMu.RLock() s.webhookMu.RLock()
url := s.webhookURL url := s.webhookURL
secret := s.webhookSecret secret := s.webhookSecret
s.webhookMu.RUnlock() s.webhookMu.RUnlock()
if url == "" { if url == "" {
return return nil
} }
payload := map[string]any{ payload := map[string]any{
"eventType": "audit.log", "eventType": "audit.log",
@@ -176,13 +233,11 @@ func (s *AuditService) fireWebhook(record *model.AuditLog) {
} }
body, err := json.Marshal(payload) body, err := json.Marshal(payload)
if err != nil { if err != nil {
log.Printf("[audit] webhook marshal failed: %v", err) return fmt.Errorf("marshal audit webhook: %w", err)
return
} }
req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, url, bytes.NewReader(body)) req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body))
if err != nil { if err != nil {
log.Printf("[audit] webhook build request failed: %v", err) return fmt.Errorf("build audit webhook request: %w", err)
return
} }
req.Header.Set("Content-Type", "application/json") req.Header.Set("Content-Type", "application/json")
req.Header.Set("User-Agent", "BackupX-Audit/1.0") req.Header.Set("User-Agent", "BackupX-Audit/1.0")
@@ -193,13 +248,13 @@ func (s *AuditService) fireWebhook(record *model.AuditLog) {
} }
resp, err := s.httpClient.Do(req) resp, err := s.httpClient.Do(req)
if err != nil { if err != nil {
log.Printf("[audit] webhook POST failed: %v", err) return fmt.Errorf("post audit webhook: %w", err)
return
} }
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode >= 400 { if resp.StatusCode >= 400 {
log.Printf("[audit] webhook returned status %d", resp.StatusCode) return fmt.Errorf("audit webhook returned status %d", resp.StatusCode)
} }
return nil
} }
// List 分页查询审计日志 // List 分页查询审计日志
@@ -14,6 +14,7 @@ import (
"testing" "testing"
"time" "time"
"backupx/server/internal/lifecycle"
"backupx/server/internal/model" "backupx/server/internal/model"
"backupx/server/internal/repository" "backupx/server/internal/repository"
) )
@@ -131,3 +132,40 @@ func TestAuditService_WebhookDisabledWhenURLEmpty(t *testing.T) {
time.Sleep(100 * time.Millisecond) time.Sleep(100 * time.Millisecond)
// 无显式断言:能不 panic 即算通过 // 无显式断言:能不 panic 即算通过
} }
func TestAuditServiceSupervisorShutdownFlushesAcceptedRecord(t *testing.T) {
repo := newFakeAuditRepo()
supervisor := lifecycle.NewSupervisor(context.Background())
svc := NewAuditService(repo)
svc.SetBackgroundRunner(supervisor)
svc.Record(AuditEntry{Username: "alice", Category: "auth", Action: "logout"})
waitCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := supervisor.Shutdown(waitCtx); err != nil {
t.Fatalf("Shutdown: %v", err)
}
select {
case <-repo.created:
default:
t.Fatal("accepted audit record was not flushed during shutdown")
}
}
func TestAuditServiceBoundsInFlightWork(t *testing.T) {
repo := newFakeAuditRepo()
svc := NewAuditService(repo)
svc.inFlight = make(chan struct{}, 1)
accepted := 0
svc.async = func(func(context.Context)) bool {
accepted++
return true
}
svc.Record(AuditEntry{Category: "auth", Action: "first"})
svc.Record(AuditEntry{Category: "auth", Action: "second"})
if accepted != 1 {
t.Fatalf("accepted tasks = %d, want 1", accepted)
}
}
@@ -0,0 +1,84 @@
package service
import (
"context"
"testing"
"time"
)
type capturingMonitorRunner struct {
tasks []func(context.Context)
}
func (r *capturingMonitorRunner) Go(task func(context.Context)) bool {
r.tasks = append(r.tasks, task)
return true
}
func TestLongRunningMonitorsUseConfiguredBackgroundRunner(t *testing.T) {
runner := &capturingMonitorRunner{}
nodes := NewNodeService(nil, "test")
nodes.SetBackgroundRunner(runner)
nodes.StartOfflineMonitor(context.Background(), time.Hour)
installTokens := NewInstallTokenService(nil, nil)
installTokens.SetBackgroundRunner(runner)
installTokens.StartGC(context.Background(), time.Hour)
dashboard := NewDashboardService(nil, nil, nil)
dashboard.SetBackgroundRunner(runner)
dashboard.StartSLAMonitor(context.Background(), nil, time.Hour, time.Hour)
versions := NewClusterVersionMonitor(nil, "test")
versions.SetBackgroundRunner(runner)
versions.Start(context.Background(), time.Hour, time.Hour)
storageTargets := NewStorageTargetService(nil, nil, nil, nil)
storageTargets.SetBackgroundRunner(runner)
storageTargets.StartHealthMonitor(context.Background(), nil, time.Hour)
if len(runner.tasks) != 5 {
t.Fatalf("background runner received %d tasks, want 5", len(runner.tasks))
}
// The monitor must listen to the supervisor-provided context, not retain
// the context passed to Start. Running one captured task is sufficient to
// lock this ownership contract for the shared helper.
runCtx, cancel := context.WithCancel(context.Background())
done := make(chan struct{})
go func() {
runner.tasks[0](runCtx)
close(done)
}()
cancel()
select {
case <-done:
case <-time.After(time.Second):
t.Fatal("monitor did not stop when background runner context was canceled")
}
}
func TestBackgroundMonitorFallbackUsesCallerContext(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
started := make(chan struct{})
done := make(chan struct{})
if !startBackgroundMonitor(nil, ctx, func(runCtx context.Context) {
close(started)
<-runCtx.Done()
close(done)
}) {
t.Fatal("fallback monitor was rejected")
}
select {
case <-started:
case <-time.After(time.Second):
t.Fatal("fallback monitor did not start")
}
cancel()
select {
case <-done:
case <-time.After(time.Second):
t.Fatal("fallback monitor did not stop when caller context was canceled")
}
}
@@ -0,0 +1,62 @@
package service
import (
"context"
"net/http"
"time"
"backupx/server/internal/apperror"
)
// BackgroundRunner is the narrow lifecycle dependency used by asynchronous
// services. lifecycle.Supervisor implements it at the application boundary.
type BackgroundRunner interface {
Go(func(context.Context)) bool
}
func runDetached(task func(context.Context)) bool {
if task == nil {
return false
}
go task(context.Background())
return true
}
// startBackgroundMonitor keeps the legacy caller-owned context when no
// lifecycle runner is configured, while allowing the application supervisor
// to own and wait for long-running monitors in production.
func startBackgroundMonitor(runner BackgroundRunner, fallbackCtx context.Context, task func(context.Context)) bool {
if task == nil {
return false
}
if runner != nil {
return runner.Go(task)
}
if fallbackCtx == nil {
fallbackCtx = context.Background()
}
go task(fallbackCtx)
return true
}
func backgroundTaskUnavailable(code string) *apperror.AppError {
return apperror.New(http.StatusServiceUnavailable, code, "服务正在关闭,无法启动新的后台任务", context.Canceled)
}
// finalizationContext lets a canceled task persist its terminal state. It is
// intentionally short-lived so shutdown cannot wait forever on cleanup I/O.
func finalizationContext(ctx context.Context) (context.Context, context.CancelFunc) {
if ctx == nil {
ctx = context.Background()
}
return context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second)
}
func acquireBackgroundSlot(ctx context.Context, semaphore chan struct{}) bool {
select {
case semaphore <- struct{}{}:
return true
case <-ctx.Done():
return false
}
}
@@ -105,7 +105,7 @@ type BackupExecutionService struct {
agentDispatcher AgentDispatcher agentDispatcher AgentDispatcher
replicationHook ReplicationTrigger replicationHook ReplicationTrigger
dependentsResolver DependentsResolver dependentsResolver DependentsResolver
async func(func()) async func(func(context.Context)) bool
now func() time.Time now func() time.Time
tempDir string tempDir string
semaphore chan struct{} semaphore chan struct{}
@@ -127,6 +127,13 @@ func (s *BackupExecutionService) SetMetrics(m *metrics.Metrics) {
s.metrics = m s.metrics = m
} }
// SetBackgroundRunner binds local asynchronous executions to the application lifecycle.
func (s *BackupExecutionService) SetBackgroundRunner(runner BackgroundRunner) {
if runner != nil {
s.async = runner.Go
}
}
// ReplicationTrigger 抽象备份成功后的副本派发(实现者:ReplicationService)。 // ReplicationTrigger 抽象备份成功后的副本派发(实现者:ReplicationService)。
type ReplicationTrigger interface { type ReplicationTrigger interface {
TriggerAutoReplication(ctx context.Context, task *model.BackupTask, record *model.BackupRecord) TriggerAutoReplication(ctx context.Context, task *model.BackupTask, record *model.BackupRecord)
@@ -194,14 +201,12 @@ func NewBackupExecutionService(
retention: retention, retention: retention,
cipher: cipher, cipher: cipher,
notifier: notifier, notifier: notifier,
async: func(job func()) { async: runDetached,
go job() now: func() time.Time { return time.Now().UTC() },
}, tempDir: tempDir,
now: func() time.Time { return time.Now().UTC() }, semaphore: make(chan struct{}, maxConcurrent),
tempDir: tempDir, retries: retries,
semaphore: make(chan struct{}, maxConcurrent), bandwidthLimit: bandwidthLimit,
retries: retries,
bandwidthLimit: bandwidthLimit,
} }
} }
@@ -259,62 +264,6 @@ func (s *BackupExecutionService) DownloadRecord(ctx context.Context, recordID ui
return &DownloadedArtifact{FileName: fileName, Reader: reader}, nil return &DownloadedArtifact{FileName: fileName, Reader: reader}, nil
} }
func (s *BackupExecutionService) RestoreRecord(ctx context.Context, recordID uint) error {
record, provider, err := s.loadRecordProvider(ctx, recordID)
if err != nil {
return err
}
task, err := s.tasks.FindByID(ctx, record.TaskID)
if err != nil {
return apperror.Internal("BACKUP_TASK_GET_FAILED", "无法获取关联备份任务", err)
}
if task == nil {
return apperror.New(404, "BACKUP_TASK_NOT_FOUND", "关联的备份任务不存在,无法执行恢复", fmt.Errorf("backup task %d not found", record.TaskID))
}
if record.BackupKind == model.BackupKindRepository {
spec, specErr := s.buildTaskSpec(task, record.StartedAt)
if specErr != nil {
return specErr
}
if err := backup.NewRepositoryStore(s.cipher.Key()).Restore(ctx, provider, record.StoragePath, record.Checksum, spec, backup.NopLogWriter{}); err != nil {
return apperror.Internal("BACKUP_RECORD_RESTORE_FAILED", "从 CDC 仓库恢复备份失败", err)
}
return nil
}
tempDir, err := os.MkdirTemp("", "backupx-restore-*")
if err != nil {
return apperror.Internal("BACKUP_RECORD_RESTORE_FAILED", "无法创建恢复目录", err)
}
defer os.RemoveAll(tempDir)
artifactPath := filepath.Join(tempDir, filepath.Base(record.FileName))
if strings.TrimSpace(filepath.Base(record.FileName)) == "" {
artifactPath = filepath.Join(tempDir, filepath.Base(record.StoragePath))
}
reader, err := provider.Download(ctx, record.StoragePath)
if err != nil {
return apperror.Internal("BACKUP_RECORD_RESTORE_FAILED", "无法下载备份文件", err)
}
if err := writeReaderToFile(artifactPath, reader); err != nil {
return apperror.Internal("BACKUP_RECORD_RESTORE_FAILED", "无法写入恢复文件", err)
}
preparedPath, err := s.prepareArtifactForRestore(artifactPath)
if err != nil {
return apperror.Internal("BACKUP_RECORD_RESTORE_FAILED", "无法准备恢复文件", err)
}
spec, err := s.buildTaskSpec(task, record.StartedAt)
if err != nil {
return err
}
runner, err := s.runnerRegistry.Runner(spec.Type)
if err != nil {
return apperror.BadRequest("BACKUP_TASK_INVALID", "不支持的备份任务类型", err)
}
if err := runner.Restore(ctx, spec, preparedPath, backup.NopLogWriter{}); err != nil {
return apperror.Internal("BACKUP_RECORD_RESTORE_FAILED", "恢复备份失败", err)
}
return nil
}
func (s *BackupExecutionService) DeleteRecord(ctx context.Context, recordID uint) error { func (s *BackupExecutionService) DeleteRecord(ctx context.Context, recordID uint) error {
record, err := s.records.FindByID(ctx, recordID) record, err := s.records.FindByID(ctx, recordID)
if err != nil { if err != nil {
@@ -504,6 +453,11 @@ func (s *BackupExecutionService) startTask(ctx context.Context, id uint, async b
task.LastRunAt = &startedAt task.LastRunAt = &startedAt
task.LastStatus = "running" task.LastStatus = "running"
if err := s.tasks.Update(ctx, task); err != nil { if err := s.tasks.Update(ctx, task); err != nil {
finalizeErr := s.finalizeRecord(ctx, &runTask, record.ID, startedAt, model.BackupRecordStatusFailed,
"无法更新任务状态: "+err.Error(), "", "", 0, "", "", primaryTargetID)
if finalizeErr != nil {
err = errors.Join(err, finalizeErr)
}
return nil, apperror.Internal("BACKUP_TASK_UPDATE_FAILED", "无法更新任务状态", err) return nil, apperror.Internal("BACKUP_TASK_UPDATE_FAILED", "无法更新任务状态", err)
} }
// 多节点路由:task.NodeID 指向远程节点时,把执行任务入队给 Agent; // 多节点路由:task.NodeID 指向远程节点时,把执行任务入队给 Agent;
@@ -512,8 +466,10 @@ func (s *BackupExecutionService) startTask(ctx context.Context, id uint, async b
// 节点离线 → 立即把刚创建的 running 记录标记 failed,返回明确错误 // 节点离线 → 立即把刚创建的 running 记录标记 failed,返回明确错误
if remoteNode.Status != model.NodeStatusOnline { if remoteNode.Status != model.NodeStatusOnline {
offlineMsg := fmt.Sprintf("节点 %s 当前离线,无法执行备份任务", remoteNode.Name) offlineMsg := fmt.Sprintf("节点 %s 当前离线,无法执行备份任务", remoteNode.Name)
_ = s.finalizeRecord(ctx, &runTask, record.ID, startedAt, model.BackupRecordStatusFailed, if finalizeErr := s.finalizeRecord(ctx, &runTask, record.ID, startedAt, model.BackupRecordStatusFailed,
offlineMsg, "", "", 0, "", "", primaryTargetID) offlineMsg, "", "", 0, "", "", primaryTargetID); finalizeErr != nil {
return nil, apperror.Internal("BACKUP_RECORD_FINALIZE_FAILED", "无法写回备份失败状态", finalizeErr)
}
return nil, apperror.BadRequest("NODE_OFFLINE", offlineMsg, nil) return nil, apperror.BadRequest("NODE_OFFLINE", offlineMsg, nil)
} }
if _, enqueueErr := s.agentDispatcher.EnqueueCommand(ctx, resolvedNodeID, model.AgentCommandTypeRunTask, map[string]any{ if _, enqueueErr := s.agentDispatcher.EnqueueCommand(ctx, resolvedNodeID, model.AgentCommandTypeRunTask, map[string]any{
@@ -521,19 +477,28 @@ func (s *BackupExecutionService) startTask(ctx context.Context, id uint, async b
"recordId": record.ID, "recordId": record.ID,
}); enqueueErr != nil { }); enqueueErr != nil {
// 入队失败 → 在记录中标记失败,继续返回详情 // 入队失败 → 在记录中标记失败,继续返回详情
_ = s.finalizeRecord(ctx, &runTask, record.ID, startedAt, model.BackupRecordStatusFailed, if finalizeErr := s.finalizeRecord(ctx, &runTask, record.ID, startedAt, model.BackupRecordStatusFailed,
"无法下发任务到远程节点: "+enqueueErr.Error(), "", "", 0, "", "", primaryTargetID) "无法下发任务到远程节点: "+enqueueErr.Error(), "", "", 0, "", "", primaryTargetID); finalizeErr != nil {
enqueueErr = errors.Join(enqueueErr, finalizeErr)
}
return nil, apperror.Internal("AGENT_COMMAND_ENQUEUE_FAILED", "无法下发任务到远程节点", enqueueErr) return nil, apperror.Internal("AGENT_COMMAND_ENQUEUE_FAILED", "无法下发任务到远程节点", enqueueErr)
} }
return s.getRecordDetail(ctx, record.ID) return s.getRecordDetail(ctx, record.ID)
} }
run := func() { run := func(runCtx context.Context) {
s.executeTask(context.Background(), &runTask, record.ID, startedAt) s.executeTask(runCtx, &runTask, record.ID, startedAt)
} }
if async { if async {
s.async(run) if !s.async(run) {
message := "服务正在关闭,备份任务未启动"
if finalizeErr := s.finalizeRecord(ctx, &runTask, record.ID, startedAt, model.BackupRecordStatusFailed,
message, "", "", 0, "", "", primaryTargetID); finalizeErr != nil {
return nil, apperror.Internal("BACKUP_RECORD_FINALIZE_FAILED", "无法写回备份失败状态", finalizeErr)
}
return nil, backgroundTaskUnavailable("BACKUP_SERVICE_SHUTTING_DOWN")
}
} else { } else {
run() run(ctx)
} }
return s.getRecordDetail(ctx, record.ID) return s.getRecordDetail(ctx, record.ID)
} }
@@ -844,20 +809,23 @@ func (s *BackupExecutionService) executeRepositoryTask(ctx context.Context, task
logger.Warnf("部分存储目标 CDC 仓库同步失败:%s", strings.Join(failures, "; ")) logger.Warnf("部分存储目标 CDC 仓库同步失败:%s", strings.Join(failures, "; "))
} }
if s.dependentsResolver != nil { if s.dependentsResolver != nil {
go func(upstreamID uint, upstreamName string) { accepted := s.async(func(runCtx context.Context) {
dependents, resolveErr := s.dependentsResolver.TriggerDependents(context.Background(), upstreamID) dependents, resolveErr := s.dependentsResolver.TriggerDependents(runCtx, task.ID)
if resolveErr != nil { if resolveErr != nil {
logger.Warnf("解析任务 %s 的下游依赖失败:%v", upstreamName, resolveErr) logger.Warnf("解析任务 %s 的下游依赖失败:%v", task.Name, resolveErr)
return return
} }
for _, dependentID := range dependents { for _, dependentID := range dependents {
if _, runErr := s.RunTaskByID(context.Background(), dependentID); runErr != nil { if _, runErr := s.RunTaskByID(runCtx, dependentID); runErr != nil {
logger.Warnf("触发下游任务 #%d 失败(上游: %s):%v", dependentID, upstreamName, runErr) logger.Warnf("触发下游任务 #%d 失败(上游: %s):%v", dependentID, task.Name, runErr)
} else { } else {
logger.Infof("已触发下游任务 #%d(上游: %s", dependentID, upstreamName) logger.Infof("已触发下游任务 #%d(上游: %s", dependentID, task.Name)
} }
} }
}(task.ID, task.Name) })
if !accepted {
logger.Warnf("服务正在关闭,跳过触发任务 %s 的下游依赖", task.Name)
}
} }
return result, nil return result, nil
} }
@@ -871,20 +839,6 @@ func (s *BackupExecutionService) acquireRepositoryLock(targetID uint) func() {
} }
func (s *BackupExecutionService) executeTask(ctx context.Context, task *model.BackupTask, recordID uint, startedAt time.Time) { func (s *BackupExecutionService) executeTask(ctx context.Context, task *model.BackupTask, recordID uint, startedAt time.Time) {
// 节点级并发限流:当任务绑定节点且节点配置了 MaxConcurrent>0
// 该节点上所有任务共享一个节点专属 semaphore,互相排队
nodeSem := s.acquireNodeSemaphore(ctx, task.NodeID)
if nodeSem != nil {
nodeSem <- struct{}{}
defer func() { <-nodeSem }()
}
s.semaphore <- struct{}{}
defer func() { <-s.semaphore }()
// Prometheus: running gauge + 完成时 observe 耗时/字节/状态
s.metrics.IncTaskRunning()
defer s.metrics.DecTaskRunning()
logger := backup.NewExecutionLogger(recordID, s.logHub) logger := backup.NewExecutionLogger(recordID, s.logHub)
status := "failed" status := "failed"
errMessage := "" errMessage := ""
@@ -900,8 +854,10 @@ func (s *BackupExecutionService) executeTask(ctx context.Context, task *model.Ba
var manifestJSON string var manifestJSON string
var repositoryProviders map[uint]storage.StorageProvider var repositoryProviders map[uint]storage.StorageProvider
completeRecord := func() { completeRecord := func() {
persistCtx, cancel := finalizationContext(ctx)
defer cancel()
readyForRepositoryRetention := status == model.BackupRecordStatusSuccess readyForRepositoryRetention := status == model.BackupRecordStatusSuccess
if finalizeErr := s.finalizeRecord(ctx, task, recordID, startedAt, status, errMessage, logger.String(), fileName, fileSize, checksum, storagePath, selectedStorageTargetID); finalizeErr != nil { if finalizeErr := s.finalizeRecord(persistCtx, task, recordID, startedAt, status, errMessage, logger.String(), fileName, fileSize, checksum, storagePath, selectedStorageTargetID); finalizeErr != nil {
logger.Errorf("写回备份记录失败:%v", finalizeErr) logger.Errorf("写回备份记录失败:%v", finalizeErr)
readyForRepositoryRetention = false readyForRepositoryRetention = false
} }
@@ -913,7 +869,7 @@ func (s *BackupExecutionService) executeTask(ctx context.Context, task *model.Ba
if marshalErr != nil { if marshalErr != nil {
logger.Warnf("序列化多目标上传结果失败:%v", marshalErr) logger.Warnf("序列化多目标上传结果失败:%v", marshalErr)
readyForRepositoryRetention = false readyForRepositoryRetention = false
} else if record, findErr := s.records.FindByID(ctx, recordID); findErr != nil || record == nil { } else if record, findErr := s.records.FindByID(persistCtx, recordID); findErr != nil || record == nil {
if findErr != nil { if findErr != nil {
logger.Warnf("读取备份记录以写回多目标结果失败:%v", findErr) logger.Warnf("读取备份记录以写回多目标结果失败:%v", findErr)
} else { } else {
@@ -922,7 +878,7 @@ func (s *BackupExecutionService) executeTask(ctx context.Context, task *model.Ba
readyForRepositoryRetention = false readyForRepositoryRetention = false
} else { } else {
record.StorageUploadResults = string(resultsJSON) record.StorageUploadResults = string(resultsJSON)
if updateErr := s.records.Update(ctx, record); updateErr != nil { if updateErr := s.records.Update(persistCtx, record); updateErr != nil {
logger.Warnf("写回多目标上传结果失败:%v", updateErr) logger.Warnf("写回多目标上传结果失败:%v", updateErr)
readyForRepositoryRetention = false readyForRepositoryRetention = false
} }
@@ -930,11 +886,11 @@ func (s *BackupExecutionService) executeTask(ctx context.Context, task *model.Ba
} }
// 持久化差异链信息:全量记录其清单(供后续差异比对),差异记录其基线全量 ID。 // 持久化差异链信息:全量记录其清单(供后续差异比对),差异记录其基线全量 ID。
if status == model.BackupRecordStatusSuccess && (backupKind != model.BackupKindFull || baseRecordID != 0 || manifestJSON != "") { if status == model.BackupRecordStatusSuccess && (backupKind != model.BackupKindFull || baseRecordID != 0 || manifestJSON != "") {
if record, findErr := s.records.FindByID(ctx, recordID); findErr == nil && record != nil { if record, findErr := s.records.FindByID(persistCtx, recordID); findErr == nil && record != nil {
record.BackupKind = backupKind record.BackupKind = backupKind
record.BaseRecordID = baseRecordID record.BaseRecordID = baseRecordID
record.Manifest = manifestJSON record.Manifest = manifestJSON
if updErr := s.records.Update(ctx, record); updErr != nil { if updErr := s.records.Update(persistCtx, record); updErr != nil {
logger.Warnf("写回差异链信息失败:%v", updErr) logger.Warnf("写回差异链信息失败:%v", updErr)
readyForRepositoryRetention = false readyForRepositoryRetention = false
} }
@@ -947,7 +903,7 @@ func (s *BackupExecutionService) executeTask(ctx context.Context, task *model.Ba
readyForRepositoryRetention = false readyForRepositoryRetention = false
} }
} }
if readyForRepositoryRetention && backupKind == model.BackupKindRepository && s.retention != nil && len(repositoryProviders) > 0 { if ctx.Err() == nil && readyForRepositoryRetention && backupKind == model.BackupKindRepository && s.retention != nil && len(repositoryProviders) > 0 {
targetIDs := make([]uint, 0, len(repositoryProviders)) targetIDs := make([]uint, 0, len(repositoryProviders))
for targetID := range repositoryProviders { for targetID := range repositoryProviders {
targetIDs = append(targetIDs, targetID) targetIDs = append(targetIDs, targetID)
@@ -969,8 +925,8 @@ func (s *BackupExecutionService) executeTask(ctx context.Context, task *model.Ba
} }
} }
} }
if s.shouldNotify(ctx, task, status) { if s.shouldNotify(persistCtx, task, status) {
if err := s.notifier.NotifyBackupResult(ctx, BackupExecutionNotification{Task: task, Record: &model.BackupRecord{ID: recordID, TaskID: task.ID, Status: status, FileName: fileName, FileSize: fileSize, StoragePath: storagePath, ErrorMessage: errMessage, StartedAt: startedAt}, Error: buildOptionalError(errMessage)}); err != nil { if err := s.notifier.NotifyBackupResult(persistCtx, BackupExecutionNotification{Task: task, Record: &model.BackupRecord{ID: recordID, TaskID: task.ID, Status: status, FileName: fileName, FileSize: fileSize, StoragePath: storagePath, ErrorMessage: errMessage, StartedAt: startedAt}, Error: buildOptionalError(errMessage)}); err != nil {
logger.Warnf("发送备份通知失败:%v", err) logger.Warnf("发送备份通知失败:%v", err)
} }
} else { } else {
@@ -980,6 +936,28 @@ func (s *BackupExecutionService) executeTask(ctx context.Context, task *model.Ba
} }
defer completeRecord() defer completeRecord()
// 节点级并发限流:当任务绑定节点且节点配置了 MaxConcurrent>0
// 该节点上所有任务共享一个节点专属 semaphore,互相排队。
nodeSem := s.acquireNodeSemaphore(ctx, task.NodeID)
if nodeSem != nil {
if !acquireBackgroundSlot(ctx, nodeSem) {
errMessage = ctx.Err().Error()
logger.Warnf("等待节点执行槽时任务被取消:%v", ctx.Err())
return
}
defer func() { <-nodeSem }()
}
if !acquireBackgroundSlot(ctx, s.semaphore) {
errMessage = ctx.Err().Error()
logger.Warnf("等待全局执行槽时任务被取消:%v", ctx.Err())
return
}
defer func() { <-s.semaphore }()
// Prometheus: running gauge + 完成时 observe 耗时/字节/状态
s.metrics.IncTaskRunning()
defer s.metrics.DecTaskRunning()
spec, err := s.buildTaskSpec(task, startedAt) spec, err := s.buildTaskSpec(task, startedAt)
if err != nil { if err != nil {
errMessage = err.Error() errMessage = err.Error()
@@ -1217,20 +1195,24 @@ func (s *BackupExecutionService) executeTask(ctx context.Context, task *model.Ba
// 自动派发复制(3-2-1):任务配置 ReplicationTargetIDs 且本次有任意目标成功时生效 // 自动派发复制(3-2-1):任务配置 ReplicationTargetIDs 且本次有任意目标成功时生效
// 触发下游依赖任务(best-effort,失败仅 warn // 触发下游依赖任务(best-effort,失败仅 warn
if s.dependentsResolver != nil { if s.dependentsResolver != nil {
go func(upstreamID uint, upstreamName string) { accepted := s.async(func(runCtx context.Context) {
dependents, err := s.dependentsResolver.TriggerDependents(context.Background(), upstreamID) dependents, err := s.dependentsResolver.TriggerDependents(runCtx, task.ID)
if err != nil { if err != nil {
logger.Warnf("解析任务 %s 的下游依赖失败:%v", task.Name, err)
return return
} }
for _, depID := range dependents { for _, depID := range dependents {
_, runErr := s.RunTaskByID(context.Background(), depID) _, runErr := s.RunTaskByID(runCtx, depID)
if runErr != nil { if runErr != nil {
logger.Warnf("触发下游任务 #%d 失败(上游: %s: %v", depID, upstreamName, runErr) logger.Warnf("触发下游任务 #%d 失败(上游: %s: %v", depID, task.Name, runErr)
} else { } else {
logger.Infof("已触发下游任务 #%d(上游: %s", depID, upstreamName) logger.Infof("已触发下游任务 #%d(上游: %s", depID, task.Name)
} }
} }
}(task.ID, task.Name) })
if !accepted {
logger.Warnf("服务正在关闭,跳过触发任务 %s 的下游依赖", task.Name)
}
} }
if s.replicationHook != nil && strings.TrimSpace(task.ReplicationTargetIDs) != "" { if s.replicationHook != nil && strings.TrimSpace(task.ReplicationTargetIDs) != "" {
record := &model.BackupRecord{ record := &model.BackupRecord{
@@ -1253,7 +1235,7 @@ func (s *BackupExecutionService) executeTask(ctx context.Context, task *model.Ba
} }
} }
logger.Infof("触发自动复制(3-2-1 规则):%s", task.ReplicationTargetIDs) logger.Infof("触发自动复制(3-2-1 规则):%s", task.ReplicationTargetIDs)
s.replicationHook.TriggerAutoReplication(context.Background(), task, record) s.replicationHook.TriggerAutoReplication(ctx, task, record)
} }
} else { } else {
errMessage = strings.Join(failedMessages, "; ") errMessage = strings.Join(failedMessages, "; ")
@@ -1356,28 +1338,6 @@ func applyHANAExtraConfig(spec *backup.DatabaseSpec, extra map[string]any) {
} }
} }
func (s *BackupExecutionService) loadRecordProvider(ctx context.Context, recordID uint) (*model.BackupRecord, storage.StorageProvider, error) {
record, err := s.records.FindByID(ctx, recordID)
if err != nil {
return nil, nil, apperror.Internal("BACKUP_RECORD_GET_FAILED", "无法获取备份记录详情", err)
}
if record == nil {
return nil, nil, apperror.New(404, "BACKUP_RECORD_NOT_FOUND", "备份记录不存在", fmt.Errorf("backup record %d not found", recordID))
}
if err := s.validateClusterAccessible(ctx, record); err != nil {
return nil, nil, err
}
provider, err := s.resolveProvider(ctx, record.StorageTargetID)
if err != nil {
return nil, nil, err
}
return record, provider, nil
}
func (s *BackupExecutionService) prepareArtifactForRestore(artifactPath string) (string, error) {
return prepareBackupArtifact(s.cipher, artifactPath, nil)
}
func (s *BackupExecutionService) getRecordDetail(ctx context.Context, recordID uint) (*BackupRecordDetail, error) { func (s *BackupExecutionService) getRecordDetail(ctx context.Context, recordID uint) (*BackupRecordDetail, error) {
record, err := s.records.FindByID(ctx, recordID) record, err := s.records.FindByID(ctx, recordID)
if err != nil { if err != nil {
@@ -1410,25 +1370,6 @@ func buildOptionalError(message string) error {
return fmt.Errorf("%s", message) return fmt.Errorf("%s", message)
} }
func buildStorageProviderFromRepos(ctx context.Context, storageTargetID uint, storageTargets repository.StorageTargetRepository, storageRegistry *storage.Registry, cipher *codec.ConfigCipher) (storage.StorageProvider, *model.StorageTarget, error) {
target, err := storageTargets.FindByID(ctx, storageTargetID)
if err != nil {
return nil, nil, apperror.Internal("BACKUP_STORAGE_TARGET_LOOKUP_FAILED", "无法读取存储目标", err)
}
if target == nil {
return nil, nil, apperror.BadRequest("BACKUP_STORAGE_TARGET_INVALID", "存储目标不存在", nil)
}
var configMap map[string]any
if err := cipher.DecryptJSON(target.ConfigCiphertext, &configMap); err != nil {
return nil, nil, apperror.Internal("BACKUP_STORAGE_TARGET_DECRYPT_FAILED", "无法解密存储目标配置", err)
}
provider, err := storageRegistry.Create(ctx, storage.ParseProviderType(target.Type), configMap)
if err != nil {
return nil, nil, err
}
return provider, target, nil
}
// hashingReader 在上传过程中同步计算字节数和 SHA-256,零额外 I/O // hashingReader 在上传过程中同步计算字节数和 SHA-256,零额外 I/O
type hashingReader struct { type hashingReader struct {
reader io.Reader reader io.Reader
@@ -33,6 +33,8 @@ func (f *testStorageFactory) Type() storage.ProviderType {
return "test_storage" return "test_storage"
} }
func (f *testStorageFactory) SensitiveFields() []string { return nil }
func (f *testStorageFactory) New(_ context.Context, config map[string]any) (storage.StorageProvider, error) { func (f *testStorageFactory) New(_ context.Context, config map[string]any) (storage.StorageProvider, error) {
name, _ := config["name"].(string) name, _ := config["name"].(string)
provider := f.providers[name] provider := f.providers[name]
@@ -250,20 +252,6 @@ func TestBackupExecutionServiceRepositoryModeRoundTrip(t *testing.T) {
t.Fatalf("unexpected repository export: name=%s size=%d", download.FileName, len(exported)) t.Fatalf("unexpected repository export: name=%s size=%d", download.FileName, len(exported))
} }
if err := os.WriteFile(largePath, bytes.Repeat([]byte{0}, len(large)), 0o640); err != nil {
t.Fatalf("damage source before restore: %v", err)
}
if err := executionService.RestoreRecord(ctx, second.ID); err != nil {
t.Fatalf("restore repository record returned error: %v", err)
}
restored, err := os.ReadFile(largePath)
if err != nil {
t.Fatalf("read restored source: %v", err)
}
if !bytes.Equal(restored, large) {
t.Fatalf("repository restore did not reproduce the source")
}
if err := recordService.Delete(ctx, first.ID); err != nil { if err := recordService.Delete(ctx, first.ID); err != nil {
t.Fatalf("delete first repository record: %v", err) t.Fatalf("delete first repository record: %v", err)
} }
@@ -428,40 +416,6 @@ func TestBackupExecutionServiceDeleteRecordDispatchesRemoteLocalDiskCleanup(t *t
} }
} }
func TestBackupExecutionServiceRestoreRecordRejectsRemoteLocalDisk(t *testing.T) {
executionService, _, tasks, _, records, _, _ := newExecutionTestServices(t)
ctx := context.Background()
executionService.SetClusterDependencies(&nodeRepoStub{nodes: []model.Node{
{ID: 10, Name: "edge-a", Token: "edge-a-token", Status: model.NodeStatusOnline},
}}, &fakeDispatcher{})
task, err := tasks.FindByID(ctx, 1)
if err != nil {
t.Fatalf("FindByID task returned error: %v", err)
}
completedAt := time.Now().UTC()
record := &model.BackupRecord{
TaskID: task.ID,
StorageTargetID: task.StorageTargetID,
NodeID: 10,
Status: model.BackupRecordStatusSuccess,
FileName: "remote.tar.gz",
StoragePath: "file/2026/05/09/remote.tar.gz",
StartedAt: completedAt.Add(-time.Second),
CompletedAt: &completedAt,
}
if err := records.Create(ctx, record); err != nil {
t.Fatalf("Create record returned error: %v", err)
}
err = executionService.RestoreRecord(ctx, record.ID)
if err == nil {
t.Fatal("expected remote local_disk restore to be rejected")
}
if !strings.Contains(err.Error(), "Master 无法跨节点访问") {
t.Fatalf("expected cross-node local_disk error, got %v", err)
}
}
func TestBackupExecutionServiceDownloadsMasterRelayedLocalDiskRecord(t *testing.T) { func TestBackupExecutionServiceDownloadsMasterRelayedLocalDiskRecord(t *testing.T) {
executionService, _, tasks, _, records, _, storageDir := newExecutionTestServices(t) executionService, _, tasks, _, records, _, storageDir := newExecutionTestServices(t)
ctx := context.Background() ctx := context.Background()
@@ -712,27 +666,6 @@ func TestBackupExecutionServiceContinuesWhenStorageUsageSnapshotFails(t *testing
} }
} }
func TestBackupRecordServiceRestore(t *testing.T) {
executionService, recordService, _, _, _, sourceDir, _ := newExecutionTestServices(t)
detail, err := executionService.RunTaskByIDSync(context.Background(), 1)
if err != nil {
t.Fatalf("RunTaskByIDSync returned error: %v", err)
}
if err := os.RemoveAll(sourceDir); err != nil {
t.Fatalf("RemoveAll returned error: %v", err)
}
if err := recordService.Restore(context.Background(), detail.ID); err != nil {
t.Fatalf("Restore returned error: %v", err)
}
content, err := os.ReadFile(filepath.Join(sourceDir, "index.html"))
if err != nil {
t.Fatalf("ReadFile returned error: %v", err)
}
if string(content) != "hello" {
t.Fatalf("unexpected restored content: %s", string(content))
}
}
type storageUsageCountingRecordRepo struct { type storageUsageCountingRecordRepo struct {
repository.BackupRecordRepository repository.BackupRecordRepository
mu sync.Mutex mu sync.Mutex
@@ -37,6 +37,7 @@ func newLockTestHarness(t *testing.T) (*BackupRecordService, *BackupExecutionSer
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
closeTestDatabase(t, db)
cipher := codec.NewConfigCipher("lock-secret") cipher := codec.NewConfigCipher("lock-secret")
targets := repository.NewStorageTargetRepository(db) targets := repository.NewStorageTargetRepository(db)
tasks := repository.NewBackupTaskRepository(db) tasks := repository.NewBackupTaskRepository(db)
@@ -156,10 +156,6 @@ func (s *BackupRecordService) Download(ctx context.Context, id uint) (*Downloade
return s.execution.DownloadRecord(ctx, id) return s.execution.DownloadRecord(ctx, id)
} }
func (s *BackupRecordService) Restore(ctx context.Context, id uint) error {
return s.execution.RestoreRecord(ctx, id)
}
func (s *BackupRecordService) Delete(ctx context.Context, id uint) error { func (s *BackupRecordService) Delete(ctx context.Context, id uint) error {
return s.execution.DeleteRecord(ctx, id) return s.execution.DeleteRecord(ctx, id)
} }
+18 -1
View File
@@ -3,6 +3,7 @@ package service
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"net/http" "net/http"
"strconv" "strconv"
@@ -389,7 +390,15 @@ func (s *BackupTaskService) Delete(ctx context.Context, id uint) (*DeleteResult,
return nil, apperror.New(http.StatusNotFound, "BACKUP_TASK_NOT_FOUND", "备份任务不存在", fmt.Errorf("backup task %d not found", id)) return nil, apperror.New(http.StatusNotFound, "BACKUP_TASK_NOT_FOUND", "备份任务不存在", fmt.Errorf("backup task %d not found", id))
} }
if s.scheduler != nil { if s.scheduler != nil {
_ = s.scheduler.RemoveTask(ctx, id) if err := s.scheduler.RemoveTask(ctx, id); err != nil {
rollbackCtx, cancel := finalizationContext(ctx)
rollbackErr := s.scheduler.SyncTask(rollbackCtx, existing)
cancel()
if rollbackErr != nil {
err = errors.Join(err, fmt.Errorf("restore task schedule: %w", rollbackErr))
}
return nil, apperror.Internal("BACKUP_TASK_UNSCHEDULE_FAILED", "无法移除备份任务调度", err)
}
} }
// 清理远端存储文件(尽力而为,不阻止删除) // 清理远端存储文件(尽力而为,不阻止删除)
@@ -397,6 +406,14 @@ func (s *BackupTaskService) Delete(ctx context.Context, id uint) (*DeleteResult,
result.RecordCount, result.CleanedFiles = s.cleanupRemoteFiles(ctx, id) result.RecordCount, result.CleanedFiles = s.cleanupRemoteFiles(ctx, id)
if err := s.tasks.Delete(ctx, id); err != nil { if err := s.tasks.Delete(ctx, id); err != nil {
if s.scheduler != nil {
rollbackCtx, cancel := finalizationContext(ctx)
rollbackErr := s.scheduler.SyncTask(rollbackCtx, existing)
cancel()
if rollbackErr != nil {
err = errors.Join(err, fmt.Errorf("restore task schedule: %w", rollbackErr))
}
}
return nil, apperror.Internal("BACKUP_TASK_DELETE_FAILED", "无法删除备份任务", err) return nil, apperror.Internal("BACKUP_TASK_DELETE_FAILED", "无法删除备份任务", err)
} }
return result, nil return result, nil
@@ -2,10 +2,12 @@ package service
import ( import (
"context" "context"
"errors"
"path/filepath" "path/filepath"
"strings" "strings"
"testing" "testing"
"backupx/server/internal/apperror"
"backupx/server/internal/config" "backupx/server/internal/config"
"backupx/server/internal/database" "backupx/server/internal/database"
"backupx/server/internal/logger" "backupx/server/internal/logger"
@@ -14,6 +16,34 @@ import (
"backupx/server/internal/storage/codec" "backupx/server/internal/storage/codec"
) )
type backupTaskSchedulerStub struct {
removeErr error
syncErr error
removedIDs []uint
syncedTasks []model.BackupTask
}
func (s *backupTaskSchedulerStub) SyncTask(_ context.Context, task *model.BackupTask) error {
if task != nil {
s.syncedTasks = append(s.syncedTasks, *task)
}
return s.syncErr
}
func (s *backupTaskSchedulerStub) RemoveTask(_ context.Context, taskID uint) error {
s.removedIDs = append(s.removedIDs, taskID)
return s.removeErr
}
type failingDeleteBackupTaskRepository struct {
repository.BackupTaskRepository
deleteErr error
}
func (r *failingDeleteBackupTaskRepository) Delete(context.Context, uint) error {
return r.deleteErr
}
func newBackupTaskServiceForTest(t *testing.T) (*BackupTaskService, repository.StorageTargetRepository, repository.BackupTaskRepository) { func newBackupTaskServiceForTest(t *testing.T) (*BackupTaskService, repository.StorageTargetRepository, repository.BackupTaskRepository) {
t.Helper() t.Helper()
log, err := logger.New(config.LogConfig{Level: "error"}) log, err := logger.New(config.LogConfig{Level: "error"})
@@ -24,6 +54,7 @@ func newBackupTaskServiceForTest(t *testing.T) (*BackupTaskService, repository.S
if err != nil { if err != nil {
t.Fatalf("database.Open returned error: %v", err) t.Fatalf("database.Open returned error: %v", err)
} }
closeTestDatabase(t, db)
targets := repository.NewStorageTargetRepository(db) targets := repository.NewStorageTargetRepository(db)
tasks := repository.NewBackupTaskRepository(db) tasks := repository.NewBackupTaskRepository(db)
service := NewBackupTaskService(tasks, targets, codec.NewConfigCipher("task-service-secret")) service := NewBackupTaskService(tasks, targets, codec.NewConfigCipher("task-service-secret"))
@@ -138,6 +169,76 @@ func TestBackupTaskServiceCreateAndGet(t *testing.T) {
} }
} }
func TestBackupTaskServiceDeleteKeepsTaskWhenUnscheduleFails(t *testing.T) {
ctx := context.Background()
service, targets, tasks := newBackupTaskServiceForTest(t)
if err := targets.Create(ctx, &model.StorageTarget{Name: "local", Type: "local_disk", Enabled: true, ConfigCiphertext: "ciphertext", ConfigVersion: 1, LastTestStatus: "unknown"}); err != nil {
t.Fatal(err)
}
created, err := service.Create(ctx, BackupTaskUpsertInput{
Name: "unschedule-failure", Type: "file", Enabled: true, SourcePath: "/srv/data",
StorageTargetID: 1, RetentionDays: 7, Compression: "gzip", MaxBackups: 3,
})
if err != nil {
t.Fatal(err)
}
scheduler := &backupTaskSchedulerStub{removeErr: errors.New("remove failed")}
service.SetScheduler(scheduler)
if _, err := service.Delete(ctx, created.ID); err == nil {
t.Fatal("Delete should fail when the scheduler cannot remove the task")
} else {
var appErr *apperror.AppError
if !errors.As(err, &appErr) || appErr.Code != "BACKUP_TASK_UNSCHEDULE_FAILED" {
t.Fatalf("Delete error = %#v", err)
}
}
stored, err := tasks.FindByID(ctx, created.ID)
if err != nil || stored == nil {
t.Fatalf("task should remain after unschedule failure: task=%#v err=%v", stored, err)
}
if len(scheduler.removedIDs) != 1 || len(scheduler.syncedTasks) != 1 {
t.Fatalf("scheduler calls = remove:%v sync:%d", scheduler.removedIDs, len(scheduler.syncedTasks))
}
}
func TestBackupTaskServiceDeleteRestoresScheduleWhenPersistenceFails(t *testing.T) {
ctx := context.Background()
service, targets, tasks := newBackupTaskServiceForTest(t)
if err := targets.Create(ctx, &model.StorageTarget{Name: "local", Type: "local_disk", Enabled: true, ConfigCiphertext: "ciphertext", ConfigVersion: 1, LastTestStatus: "unknown"}); err != nil {
t.Fatal(err)
}
created, err := service.Create(ctx, BackupTaskUpsertInput{
Name: "delete-failure", Type: "file", Enabled: true, SourcePath: "/srv/data",
StorageTargetID: 1, RetentionDays: 7, Compression: "gzip", MaxBackups: 3,
})
if err != nil {
t.Fatal(err)
}
service.tasks = &failingDeleteBackupTaskRepository{
BackupTaskRepository: tasks,
deleteErr: errors.New("database delete failed"),
}
scheduler := &backupTaskSchedulerStub{}
service.SetScheduler(scheduler)
if _, err := service.Delete(ctx, created.ID); err == nil {
t.Fatal("Delete should fail when persistence fails")
} else {
var appErr *apperror.AppError
if !errors.As(err, &appErr) || appErr.Code != "BACKUP_TASK_DELETE_FAILED" {
t.Fatalf("Delete error = %#v", err)
}
}
if len(scheduler.removedIDs) != 1 || len(scheduler.syncedTasks) != 1 || scheduler.syncedTasks[0].ID != created.ID {
t.Fatalf("scheduler rollback calls = remove:%v sync:%#v", scheduler.removedIDs, scheduler.syncedTasks)
}
stored, err := tasks.FindByID(ctx, created.ID)
if err != nil || stored == nil {
t.Fatalf("task should remain after persistence failure: task=%#v err=%v", stored, err)
}
}
func TestBackupTaskServiceKeepsMaskedPasswordOnUpdate(t *testing.T) { func TestBackupTaskServiceKeepsMaskedPasswordOnUpdate(t *testing.T) {
ctx := context.Background() ctx := context.Background()
service, targets, tasks := newBackupTaskServiceForTest(t) service, targets, tasks := newBackupTaskServiceForTest(t)
+10 -5
View File
@@ -21,6 +21,7 @@ type ClusterVersionMonitor struct {
nodeRepo repository.NodeRepository nodeRepo repository.NodeRepository
eventDispatcher EventDispatcher eventDispatcher EventDispatcher
masterVersion string masterVersion string
background BackgroundRunner
mu sync.Mutex mu sync.Mutex
notified map[uint]time.Time notified map[uint]time.Time
} }
@@ -37,6 +38,10 @@ func (m *ClusterVersionMonitor) SetEventDispatcher(dispatcher EventDispatcher) {
m.eventDispatcher = dispatcher m.eventDispatcher = dispatcher
} }
func (m *ClusterVersionMonitor) SetBackgroundRunner(runner BackgroundRunner) {
m.background = runner
}
// Start 启动后台扫描。ctx 取消时退出。 // Start 启动后台扫描。ctx 取消时退出。
// scanInterval 建议 30 分钟;resetInterval 建议 24 小时。 // scanInterval 建议 30 分钟;resetInterval 建议 24 小时。
func (m *ClusterVersionMonitor) Start(ctx context.Context, scanInterval, resetInterval time.Duration) { func (m *ClusterVersionMonitor) Start(ctx context.Context, scanInterval, resetInterval time.Duration) {
@@ -47,19 +52,19 @@ func (m *ClusterVersionMonitor) Start(ctx context.Context, scanInterval, resetIn
resetInterval = 24 * time.Hour resetInterval = 24 * time.Hour
} }
// 启动立即跑一次,让控制台尽快看到 // 启动立即跑一次,让控制台尽快看到
go func() { startBackgroundMonitor(m.background, ctx, func(runCtx context.Context) {
m.scan(ctx, resetInterval) m.scan(runCtx, resetInterval)
ticker := time.NewTicker(scanInterval) ticker := time.NewTicker(scanInterval)
defer ticker.Stop() defer ticker.Stop()
for { for {
select { select {
case <-ctx.Done(): case <-runCtx.Done():
return return
case <-ticker.C: case <-ticker.C:
m.scan(ctx, resetInterval) m.scan(runCtx, resetInterval)
} }
} }
}() })
} }
func (m *ClusterVersionMonitor) scan(ctx context.Context, resetInterval time.Duration) { func (m *ClusterVersionMonitor) scan(ctx context.Context, resetInterval time.Duration) {
@@ -45,6 +45,7 @@ func newDashboardNotificationTestDeps(t *testing.T) (*DashboardService, *Notific
if err != nil { if err != nil {
t.Fatalf("database.Open returned error: %v", err) t.Fatalf("database.Open returned error: %v", err)
} }
closeTestDatabase(t, db)
tasks := repository.NewBackupTaskRepository(db) tasks := repository.NewBackupTaskRepository(db)
records := repository.NewBackupRecordRepository(db) records := repository.NewBackupRecordRepository(db)
targets := repository.NewStorageTargetRepository(db) targets := repository.NewStorageTargetRepository(db)
+10 -5
View File
@@ -34,6 +34,7 @@ type DashboardService struct {
targets repository.StorageTargetRepository targets repository.StorageTargetRepository
nodes repository.NodeRepository nodes repository.NodeRepository
masterVersion string masterVersion string
background BackgroundRunner
// slaMonitor 内部跟踪已告警的违约任务,避免每次扫描重复派发事件 // slaMonitor 内部跟踪已告警的违约任务,避免每次扫描重复派发事件
slaNotified map[uint]time.Time slaNotified map[uint]time.Time
slaMu sync.Mutex slaMu sync.Mutex
@@ -43,6 +44,10 @@ func NewDashboardService(tasks repository.BackupTaskRepository, records reposito
return &DashboardService{tasks: tasks, records: records, targets: targets, slaNotified: map[uint]time.Time{}} return &DashboardService{tasks: tasks, records: records, targets: targets, slaNotified: map[uint]time.Time{}}
} }
func (s *DashboardService) SetBackgroundRunner(runner BackgroundRunner) {
s.background = runner
}
// SetClusterDependencies 注入节点仓储与 Master 版本,启用集群概览。 // SetClusterDependencies 注入节点仓储与 Master 版本,启用集群概览。
func (s *DashboardService) SetClusterDependencies(nodes repository.NodeRepository, masterVersion string) { func (s *DashboardService) SetClusterDependencies(nodes repository.NodeRepository, masterVersion string) {
s.nodes = nodes s.nodes = nodes
@@ -561,18 +566,18 @@ func (s *DashboardService) StartSLAMonitor(ctx context.Context, dispatcher Event
if resetInterval <= 0 { if resetInterval <= 0 {
resetInterval = 6 * time.Hour resetInterval = 6 * time.Hour
} }
ticker := time.NewTicker(scanInterval) startBackgroundMonitor(s.background, ctx, func(runCtx context.Context) {
go func() { ticker := time.NewTicker(scanInterval)
defer ticker.Stop() defer ticker.Stop()
for { for {
select { select {
case <-ctx.Done(): case <-runCtx.Done():
return return
case <-ticker.C: case <-ticker.C:
s.scanAndDispatchSLA(ctx, dispatcher, resetInterval) s.scanAndDispatchSLA(runCtx, dispatcher, resetInterval)
} }
} }
}() })
} }
// scanAndDispatchSLA 执行一次 SLA 违约扫描并按需派发事件。 // scanAndDispatchSLA 执行一次 SLA 违约扫描并按需派发事件。
@@ -27,6 +27,7 @@ func TestGoogleDriveOAuthServiceStartAndComplete(t *testing.T) {
if err != nil { if err != nil {
t.Fatalf("database.Open returned error: %v", err) t.Fatalf("database.Open returned error: %v", err)
} }
closeTestDatabase(t, db)
sessions := repository.NewOAuthSessionRepository(db) sessions := repository.NewOAuthSessionRepository(db)
service := NewGoogleDriveOAuthService(sessions, codec.New("encryption-secret")) service := NewGoogleDriveOAuthService(sessions, codec.New("encryption-secret"))
service.now = func() time.Time { return time.Date(2026, 3, 7, 0, 0, 0, 0, time.UTC) } service.now = func() time.Time { return time.Date(2026, 3, 7, 0, 0, 0, 0, time.UTC) }
@@ -18,14 +18,19 @@ import (
// InstallTokenService 负责一次性安装令牌的创建/消费/校验。 // InstallTokenService 负责一次性安装令牌的创建/消费/校验。
type InstallTokenService struct { type InstallTokenService struct {
repo repository.AgentInstallTokenRepository repo repository.AgentInstallTokenRepository
nodeRepo repository.NodeRepository nodeRepo repository.NodeRepository
background BackgroundRunner
} }
func NewInstallTokenService(repo repository.AgentInstallTokenRepository, nodeRepo repository.NodeRepository) *InstallTokenService { func NewInstallTokenService(repo repository.AgentInstallTokenRepository, nodeRepo repository.NodeRepository) *InstallTokenService {
return &InstallTokenService{repo: repo, nodeRepo: nodeRepo} return &InstallTokenService{repo: repo, nodeRepo: nodeRepo}
} }
func (s *InstallTokenService) SetBackgroundRunner(runner BackgroundRunner) {
s.background = runner
}
// InstallTokenInput 生成一次性安装令牌的输入。 // InstallTokenInput 生成一次性安装令牌的输入。
type InstallTokenInput struct { type InstallTokenInput struct {
NodeID uint NodeID uint
@@ -247,18 +252,18 @@ func (s *InstallTokenService) StartGC(ctx context.Context, interval time.Duratio
if interval <= 0 { if interval <= 0 {
interval = time.Hour interval = time.Hour
} }
go func() { startBackgroundMonitor(s.background, ctx, func(runCtx context.Context) {
ticker := time.NewTicker(interval) ticker := time.NewTicker(interval)
defer ticker.Stop() defer ticker.Stop()
for { for {
select { select {
case <-ctx.Done(): case <-runCtx.Done():
return return
case <-ticker.C: case <-ticker.C:
_, _ = s.repo.DeleteExpiredBefore(ctx, time.Now().UTC().Add(-7*24*time.Hour)) _, _ = s.repo.DeleteExpiredBefore(runCtx, time.Now().UTC().Add(-7*24*time.Hour))
} }
} }
}() })
} }
func (s *InstallTokenService) validate(in InstallTokenInput) error { func (s *InstallTokenService) validate(in InstallTokenInput) error {
@@ -22,6 +22,7 @@ func openInstallTokenTestDB(t *testing.T) *gorm.DB {
if err != nil { if err != nil {
t.Fatalf("open: %v", err) t.Fatalf("open: %v", err)
} }
closeTestDatabase(t, db)
if err := db.AutoMigrate(&model.AgentInstallToken{}, &model.Node{}); err != nil { if err := db.AutoMigrate(&model.AgentInstallToken{}, &model.Node{}); err != nil {
t.Fatalf("migrate: %v", err) t.Fatalf("migrate: %v", err)
} }
+15 -10
View File
@@ -67,11 +67,12 @@ type NodeUpdateInput struct {
// NodeService manages the cluster nodes. // NodeService manages the cluster nodes.
type NodeService struct { type NodeService struct {
repo repository.NodeRepository repo repository.NodeRepository
taskRepo repository.BackupTaskRepository taskRepo repository.BackupTaskRepository
agentRPC NodeAgentRPC agentRPC NodeAgentRPC
cmdRepo repository.AgentCommandRepository cmdRepo repository.AgentCommandRepository
version string version string
background BackgroundRunner
} }
// NodeAgentRPC 抽象 Agent 远程调用能力(避免 service 内循环依赖)。 // NodeAgentRPC 抽象 Agent 远程调用能力(避免 service 内循环依赖)。
@@ -85,6 +86,10 @@ func NewNodeService(repo repository.NodeRepository, version string) *NodeService
return &NodeService{repo: repo, version: version} return &NodeService{repo: repo, version: version}
} }
func (s *NodeService) SetBackgroundRunner(runner BackgroundRunner) {
s.background = runner
}
// SetTaskRepository 注入任务仓储以支持删除前引用检查。可选注入,便于测试。 // SetTaskRepository 注入任务仓储以支持删除前引用检查。可选注入,便于测试。
func (s *NodeService) SetTaskRepository(taskRepo repository.BackupTaskRepository) { func (s *NodeService) SetTaskRepository(taskRepo repository.BackupTaskRepository) {
s.taskRepo = taskRepo s.taskRepo = taskRepo
@@ -315,19 +320,19 @@ func (s *NodeService) StartOfflineMonitor(ctx context.Context, interval time.Dur
if interval <= 0 { if interval <= 0 {
interval = 15 * time.Second interval = 15 * time.Second
} }
ticker := time.NewTicker(interval) startBackgroundMonitor(s.background, ctx, func(runCtx context.Context) {
go func() { ticker := time.NewTicker(interval)
defer ticker.Stop() defer ticker.Stop()
for { for {
select { select {
case <-ctx.Done(): case <-runCtx.Done():
return return
case <-ticker.C: case <-ticker.C:
threshold := time.Now().UTC().Add(-OfflineThreshold) threshold := time.Now().UTC().Add(-OfflineThreshold)
_, _ = s.repo.MarkStaleOffline(ctx, threshold) _, _ = s.repo.MarkStaleOffline(runCtx, threshold)
} }
} }
}() })
} }
// Heartbeat updates the node status when an agent reports in. // Heartbeat updates the node status when an agent reports in.
@@ -20,6 +20,7 @@ func openNodeServiceDB(t *testing.T) *gorm.DB {
if err != nil { if err != nil {
t.Fatalf("open: %v", err) t.Fatalf("open: %v", err)
} }
closeTestDatabase(t, db)
if err := db.AutoMigrate(&model.Node{}); err != nil { if err := db.AutoMigrate(&model.Node{}); err != nil {
t.Fatalf("migrate: %v", err) t.Fatalf("migrate: %v", err)
} }
@@ -172,7 +172,7 @@ func (s *NotificationService) NotifyBackupResult(ctx context.Context, event Back
if success { if success {
eventType = model.NotificationEventBackupSuccess eventType = model.NotificationEventBackupSuccess
} }
items, err := s.collectSubscribers(ctx, eventType, success) items, err := s.collectSubscribers(ctx, eventType)
if err != nil { if err != nil {
return err return err
} }
@@ -194,9 +194,7 @@ func (s *NotificationService) DispatchEvent(ctx context.Context, eventType strin
if s.broadcaster != nil { if s.broadcaster != nil {
_ = s.broadcaster.Publish(ctx, eventType, title, body, fields) _ = s.broadcaster.Publish(ctx, eventType, title, body, fields)
} }
// 将 fallback 布尔用于旧语义场景(backup_success / backup_failed)。 items, err := s.collectSubscribers(ctx, eventType)
fallbackSuccess := eventType == model.NotificationEventBackupSuccess
items, err := s.collectSubscribers(ctx, eventType, fallbackSuccess)
if err != nil { if err != nil {
return err return err
} }
@@ -254,7 +252,7 @@ func (s *NotificationService) sendFirstByType(ctx context.Context, notificationT
// collectSubscribers 按事件类型收集启用的订阅者。 // collectSubscribers 按事件类型收集启用的订阅者。
// 列出启用通知后按事件类型再过滤(避免引入新 repository 方法)。 // 列出启用通知后按事件类型再过滤(避免引入新 repository 方法)。
func (s *NotificationService) collectSubscribers(ctx context.Context, eventType string, fallbackSuccess bool) ([]model.Notification, error) { func (s *NotificationService) collectSubscribers(ctx context.Context, eventType string) ([]model.Notification, error) {
all, err := s.notifications.List(ctx) all, err := s.notifications.List(ctx)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -284,8 +282,6 @@ func (s *NotificationService) collectSubscribers(ctx context.Context, eventType
// 其他事件类型必须显式订阅才推送 // 其他事件类型必须显式订阅才推送
continue continue
} }
// 额外校验 fallbackSuccess 参数,保持历史行为一致
_ = fallbackSuccess
} }
matched = append(matched, item) matched = append(matched, item)
} }
+62 -23
View File
@@ -15,6 +15,7 @@ import (
"backupx/server/internal/repository" "backupx/server/internal/repository"
"backupx/server/internal/storage" "backupx/server/internal/storage"
"backupx/server/internal/storage/codec" "backupx/server/internal/storage/codec"
"go.uber.org/zap"
) )
// ReplicationService 实现备份复制(3-2-1 规则核心)。 // ReplicationService 实现备份复制(3-2-1 规则核心)。
@@ -36,9 +37,10 @@ type ReplicationService struct {
eventDispatcher EventDispatcher eventDispatcher EventDispatcher
tempDir string tempDir string
semaphore chan struct{} semaphore chan struct{}
async func(func()) async func(func(context.Context)) bool
now func() time.Time now func() time.Time
metrics *metrics.Metrics metrics *metrics.Metrics
logger *zap.Logger
} }
// SetMetrics 注入 Prometheus 采集器。 // SetMetrics 注入 Prometheus 采集器。
@@ -46,6 +48,19 @@ func (s *ReplicationService) SetMetrics(m *metrics.Metrics) {
s.metrics = m s.metrics = m
} }
func (s *ReplicationService) SetLogger(logger *zap.Logger) {
if logger != nil {
s.logger = logger
}
}
// SetBackgroundRunner binds replication work to the application lifecycle.
func (s *ReplicationService) SetBackgroundRunner(runner BackgroundRunner) {
if runner != nil {
s.async = runner.Go
}
}
func NewReplicationService( func NewReplicationService(
replications repository.ReplicationRecordRepository, replications repository.ReplicationRecordRepository,
records repository.BackupRecordRepository, records repository.BackupRecordRepository,
@@ -71,8 +86,9 @@ func NewReplicationService(
cipher: cipher, cipher: cipher,
tempDir: tempDir, tempDir: tempDir,
semaphore: make(chan struct{}, maxConcurrent), semaphore: make(chan struct{}, maxConcurrent),
async: func(job func()) { go job() }, async: runDetached,
now: func() time.Time { return time.Now().UTC() }, now: func() time.Time { return time.Now().UTC() },
logger: zap.NewNop(),
} }
} }
@@ -123,13 +139,16 @@ func (s *ReplicationService) TriggerAutoReplication(ctx context.Context, task *m
} }
// 跨节点 local_disk 场景保护:Master 无法访问远程节点本地文件 // 跨节点 local_disk 场景保护:Master 无法访问远程节点本地文件
if err := s.validateClusterAccessible(ctx, record); err != nil { if err := s.validateClusterAccessible(ctx, record); err != nil {
s.logger.Warn("automatic replication skipped: source is not accessible", zap.Uint("backup_record_id", record.ID), zap.Error(err))
return return
} }
for _, destID := range destIDs { for _, destID := range destIDs {
if destID == record.StorageTargetID { if destID == record.StorageTargetID {
continue // 源与目标相同,跳过 continue // 源与目标相同,跳过
} }
_, _ = s.Start(ctx, record.ID, destID, "system") if _, err := s.Start(ctx, record.ID, destID, "system"); err != nil {
s.logger.Warn("automatic replication start failed", zap.Uint("backup_record_id", record.ID), zap.Uint("dest_target_id", destID), zap.Error(err))
}
} }
} }
@@ -172,20 +191,24 @@ func (s *ReplicationService) Start(ctx context.Context, backupRecordID, destTarg
if err := s.replications.Create(ctx, rep); err != nil { if err := s.replications.Create(ctx, rep); err != nil {
return nil, apperror.Internal("REPLICATION_CREATE_FAILED", "无法创建复制记录", err) return nil, apperror.Internal("REPLICATION_CREATE_FAILED", "无法创建复制记录", err)
} }
s.async(func() { repForRun := *rep
s.executeReplication(context.Background(), rep.ID) if !s.async(func(runCtx context.Context) {
}) s.executeReplication(runCtx, &repForRun)
}) {
message := "服务正在关闭,复制任务未启动"
if finalizeErr := s.finalizeReplication(ctx, rep, model.ReplicationStatusFailed, message, 0); finalizeErr != nil {
return nil, apperror.Internal("REPLICATION_FINALIZE_FAILED", "无法写回复制失败状态", finalizeErr)
}
return nil, backgroundTaskUnavailable("REPLICATION_SERVICE_SHUTTING_DOWN")
}
summary := s.toSummary(rep, "", dest.Name) summary := s.toSummary(rep, "", dest.Name)
return &summary, nil return &summary, nil
} }
// executeReplication 实际执行:下载源对象到本地临时文件 → 上传到目标存储。 // executeReplication 实际执行:下载源对象到本地临时文件 → 上传到目标存储。
func (s *ReplicationService) executeReplication(ctx context.Context, repID uint) { func (s *ReplicationService) executeReplication(ctx context.Context, rep *model.ReplicationRecord) {
s.semaphore <- struct{}{} if rep == nil {
defer func() { <-s.semaphore }() s.logger.Error("replication record is nil")
rep, err := s.replications.FindByID(ctx, repID)
if err != nil || rep == nil {
return return
} }
status := model.ReplicationStatusFailed status := model.ReplicationStatusFailed
@@ -193,19 +216,25 @@ func (s *ReplicationService) executeReplication(ctx context.Context, repID uint)
fileSize := int64(0) fileSize := int64(0)
defer func() { defer func() {
completedAt := s.now() persistCtx, cancel := finalizationContext(ctx)
rep.Status = status defer cancel()
rep.FileSize = fileSize if finalizeErr := s.finalizeReplication(persistCtx, rep, status, errMessage, fileSize); finalizeErr != nil {
rep.ErrorMessage = strings.TrimSpace(errMessage) s.logger.Error("finalize replication record failed", zap.Uint("replication_id", rep.ID), zap.Error(finalizeErr))
rep.DurationSeconds = int(completedAt.Sub(rep.StartedAt).Seconds()) }
rep.CompletedAt = &completedAt
_ = s.replications.Update(ctx, rep)
s.metrics.ObserveReplication(status) s.metrics.ObserveReplication(status)
if status == model.ReplicationStatusFailed { if status == model.ReplicationStatusFailed {
s.dispatchFailed(ctx, rep, errMessage) if dispatchErr := s.dispatchFailed(persistCtx, rep, errMessage); dispatchErr != nil {
s.logger.Warn("dispatch replication failure event failed", zap.Uint("replication_id", rep.ID), zap.Error(dispatchErr))
}
} }
}() }()
if !acquireBackgroundSlot(ctx, s.semaphore) {
errMessage = ctx.Err().Error()
return
}
defer func() { <-s.semaphore }()
sourceProvider, err := s.resolveProvider(ctx, rep.SourceTargetID) sourceProvider, err := s.resolveProvider(ctx, rep.SourceTargetID)
if err != nil { if err != nil {
errMessage = err.Error() errMessage = err.Error()
@@ -271,9 +300,19 @@ func (s *ReplicationService) validateClusterAccessible(ctx context.Context, reco
"REPLICATION_CROSS_NODE_LOCAL_DISK", "复制。请改用云存储作为主备份") "REPLICATION_CROSS_NODE_LOCAL_DISK", "复制。请改用云存储作为主备份")
} }
func (s *ReplicationService) dispatchFailed(ctx context.Context, rep *model.ReplicationRecord, message string) { func (s *ReplicationService) finalizeReplication(ctx context.Context, rep *model.ReplicationRecord, status, message string, fileSize int64) error {
completedAt := s.now()
rep.Status = status
rep.FileSize = fileSize
rep.ErrorMessage = strings.TrimSpace(message)
rep.DurationSeconds = int(completedAt.Sub(rep.StartedAt).Seconds())
rep.CompletedAt = &completedAt
return s.replications.Update(ctx, rep)
}
func (s *ReplicationService) dispatchFailed(ctx context.Context, rep *model.ReplicationRecord, message string) error {
if s.eventDispatcher == nil || rep == nil { if s.eventDispatcher == nil || rep == nil {
return return nil
} }
title := "BackupX 备份复制失败" title := "BackupX 备份复制失败"
body := fmt.Sprintf("备份记录:#%d\n源 → 目标:#%d → #%d\n错误:%s", rep.BackupRecordID, rep.SourceTargetID, rep.DestTargetID, message) body := fmt.Sprintf("备份记录:#%d\n源 → 目标:#%d → #%d\n错误:%s", rep.BackupRecordID, rep.SourceTargetID, rep.DestTargetID, message)
@@ -285,7 +324,7 @@ func (s *ReplicationService) dispatchFailed(ctx context.Context, rep *model.Repl
"destTargetId": rep.DestTargetID, "destTargetId": rep.DestTargetID,
"error": message, "error": message,
} }
_ = s.eventDispatcher.DispatchEvent(ctx, model.NotificationEventReplicationFailed, title, body, fields) return s.eventDispatcher.DispatchEvent(ctx, model.NotificationEventReplicationFailed, title, body, fields)
} }
// List / Get / toSummary // List / Get / toSummary
@@ -47,6 +47,11 @@ func newReplicationTestHarness(t *testing.T) *replicationTestHarness {
if err != nil { if err != nil {
t.Fatalf("database.Open: %v", err) t.Fatalf("database.Open: %v", err)
} }
sqlDB, err := db.DB()
if err != nil {
t.Fatalf("db.DB: %v", err)
}
t.Cleanup(func() { _ = sqlDB.Close() })
cipher := codec.NewConfigCipher("replicate-secret") cipher := codec.NewConfigCipher("replicate-secret")
targets := repository.NewStorageTargetRepository(db) targets := repository.NewStorageTargetRepository(db)
tasks := repository.NewBackupTaskRepository(db) tasks := repository.NewBackupTaskRepository(db)
@@ -109,8 +114,9 @@ func TestReplicationService_MirrorsToDestTarget(t *testing.T) {
} }
done := make(chan struct{}) done := make(chan struct{})
h.repl.async = func(job func()) { h.repl.async = func(job func(context.Context)) bool {
go func() { job(); close(done) }() go func() { job(context.Background()); close(done) }()
return true
} }
summary, err := h.repl.Start(ctx, backupDetail.ID, 2, "tester") summary, err := h.repl.Start(ctx, backupDetail.ID, 2, "tester")
if err != nil { if err != nil {
@@ -36,6 +36,7 @@ func newReportTestHarness(t *testing.T) (*ReportService, *BackupExecutionService
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
closeTestDatabase(t, db)
cipher := codec.NewConfigCipher("report-secret") cipher := codec.NewConfigCipher("report-secret")
targets := repository.NewStorageTargetRepository(db) targets := repository.NewStorageTargetRepository(db)
tasks := repository.NewBackupTaskRepository(db) tasks := repository.NewBackupTaskRepository(db)
+47 -17
View File
@@ -3,6 +3,7 @@ package service
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"io" "io"
"os" "os"
@@ -39,7 +40,7 @@ type RestoreService struct {
eventDispatcher EventDispatcher eventDispatcher EventDispatcher
tempDir string tempDir string
semaphore chan struct{} semaphore chan struct{}
async func(func()) async func(func(context.Context)) bool
now func() time.Time now func() time.Time
metrics *metrics.Metrics metrics *metrics.Metrics
} }
@@ -49,6 +50,13 @@ func (s *RestoreService) SetMetrics(m *metrics.Metrics) {
s.metrics = m s.metrics = m
} }
// SetBackgroundRunner binds local restore work to the application lifecycle.
func (s *RestoreService) SetBackgroundRunner(runner BackgroundRunner) {
if runner != nil {
s.async = runner.Go
}
}
// NewRestoreService 构造恢复服务。maxConcurrent 控制本地并发恢复数。 // NewRestoreService 构造恢复服务。maxConcurrent 控制本地并发恢复数。
func NewRestoreService( func NewRestoreService(
restores repository.RestoreRecordRepository, restores repository.RestoreRecordRepository,
@@ -83,7 +91,7 @@ func NewRestoreService(
dispatcher: dispatcher, dispatcher: dispatcher,
tempDir: tempDir, tempDir: tempDir,
semaphore: make(chan struct{}, maxConcurrent), semaphore: make(chan struct{}, maxConcurrent),
async: func(job func()) { go job() }, async: runDetached,
now: func() time.Time { return time.Now().UTC() }, now: func() time.Time { return time.Now().UTC() },
} }
} }
@@ -187,12 +195,18 @@ func (s *RestoreService) StartSelective(ctx context.Context, backupRecordID uint
// 远程节点路由 // 远程节点路由
if remoteNode := s.resolveRemoteNode(ctx, restoreNodeID); remoteNode != nil { if remoteNode := s.resolveRemoteNode(ctx, restoreNodeID); remoteNode != nil {
if s.dispatcher == nil { if s.dispatcher == nil {
message := "Agent 下发通道未就绪"
if finalizeErr := s.finalize(ctx, restore.ID, model.RestoreRecordStatusFailed, message); finalizeErr != nil {
return nil, apperror.Internal("RESTORE_FINALIZE_FAILED", "无法写回恢复失败状态", finalizeErr)
}
return nil, apperror.Internal("RESTORE_DISPATCH_UNAVAILABLE", "Agent 下发通道未就绪", nil) return nil, apperror.Internal("RESTORE_DISPATCH_UNAVAILABLE", "Agent 下发通道未就绪", nil)
} }
// 节点离线 → 立即标记 failed,避免记录永远卡在 running // 节点离线 → 立即标记 failed,避免记录永远卡在 running
if remoteNode.Status != model.NodeStatusOnline { if remoteNode.Status != model.NodeStatusOnline {
offlineMsg := fmt.Sprintf("节点 %s 当前离线,无法执行恢复", remoteNode.Name) offlineMsg := fmt.Sprintf("节点 %s 当前离线,无法执行恢复", remoteNode.Name)
_ = s.finalize(ctx, restore.ID, model.RestoreRecordStatusFailed, offlineMsg) if finalizeErr := s.finalize(ctx, restore.ID, model.RestoreRecordStatusFailed, offlineMsg); finalizeErr != nil {
return nil, apperror.Internal("RESTORE_FINALIZE_FAILED", "无法写回恢复失败状态", finalizeErr)
}
s.logHub.Append(restore.ID, "error", offlineMsg) s.logHub.Append(restore.ID, "error", offlineMsg)
s.logHub.Complete(restore.ID, model.RestoreRecordStatusFailed) s.logHub.Complete(restore.ID, model.RestoreRecordStatusFailed)
return nil, apperror.BadRequest("NODE_OFFLINE", offlineMsg, nil) return nil, apperror.BadRequest("NODE_OFFLINE", offlineMsg, nil)
@@ -200,8 +214,10 @@ func (s *RestoreService) StartSelective(ctx context.Context, backupRecordID uint
if _, dispatchErr := s.dispatcher.EnqueueCommand(ctx, restoreNodeID, model.AgentCommandTypeRestoreRecord, map[string]any{ if _, dispatchErr := s.dispatcher.EnqueueCommand(ctx, restoreNodeID, model.AgentCommandTypeRestoreRecord, map[string]any{
"restoreRecordId": restore.ID, "restoreRecordId": restore.ID,
}); dispatchErr != nil { }); dispatchErr != nil {
_ = s.finalize(ctx, restore.ID, model.RestoreRecordStatusFailed, if finalizeErr := s.finalize(ctx, restore.ID, model.RestoreRecordStatusFailed,
"下发恢复任务到远程节点失败: "+dispatchErr.Error()) "下发恢复任务到远程节点失败: "+dispatchErr.Error()); finalizeErr != nil {
dispatchErr = errors.Join(dispatchErr, finalizeErr)
}
return nil, apperror.Internal("AGENT_COMMAND_ENQUEUE_FAILED", "无法下发恢复任务到远程节点", dispatchErr) return nil, apperror.Internal("AGENT_COMMAND_ENQUEUE_FAILED", "无法下发恢复任务到远程节点", dispatchErr)
} }
s.logHub.Append(restore.ID, "info", fmt.Sprintf("已下发恢复任务到节点 %s#%d),等待 Agent 执行", remoteNode.Name, restoreNodeID)) s.logHub.Append(restore.ID, "info", fmt.Sprintf("已下发恢复任务到节点 %s#%d),等待 Agent 执行", remoteNode.Name, restoreNodeID))
@@ -209,10 +225,16 @@ func (s *RestoreService) StartSelective(ctx context.Context, backupRecordID uint
} }
// 本地节点:异步执行 // 本地节点:异步执行
run := func() { run := func(runCtx context.Context) {
s.executeLocally(context.Background(), restore.ID, task, record, selectedPaths, targetPath) s.executeLocally(runCtx, restore.ID, task, record, selectedPaths, targetPath)
}
if !s.async(run) {
message := "服务正在关闭,恢复任务未启动"
if finalizeErr := s.finalize(ctx, restore.ID, model.RestoreRecordStatusFailed, message); finalizeErr != nil {
return nil, apperror.Internal("RESTORE_FINALIZE_FAILED", "无法写回恢复失败状态", finalizeErr)
}
return nil, backgroundTaskUnavailable("RESTORE_SERVICE_SHUTTING_DOWN")
} }
s.async(run)
return s.getDetail(ctx, restore.ID) return s.getDetail(ctx, restore.ID)
} }
@@ -238,22 +260,30 @@ func (s *RestoreService) resolveRemoteNode(ctx context.Context, nodeID uint) *mo
// executeLocally 在 Master 本地执行恢复。 // executeLocally 在 Master 本地执行恢复。
func (s *RestoreService) executeLocally(ctx context.Context, restoreID uint, task *model.BackupTask, backupRecord *model.BackupRecord, selectedPaths []string, targetPath string) { func (s *RestoreService) executeLocally(ctx context.Context, restoreID uint, task *model.BackupTask, backupRecord *model.BackupRecord, selectedPaths []string, targetPath string) {
s.semaphore <- struct{}{}
defer func() { <-s.semaphore }()
logger := backup.NewExecutionLogger(restoreID, s.logHub) logger := backup.NewExecutionLogger(restoreID, s.logHub)
status := model.RestoreRecordStatusFailed status := model.RestoreRecordStatusFailed
errMessage := "" errMessage := ""
defer func() { defer func() {
finalizeErr := s.finalizeWithLog(ctx, restoreID, status, errMessage, logger.String()) persistCtx, cancel := finalizationContext(ctx)
defer cancel()
finalizeErr := s.finalizeWithLog(persistCtx, restoreID, status, errMessage, logger.String())
if finalizeErr != nil { if finalizeErr != nil {
logger.Errorf("写回恢复记录失败:%v", finalizeErr) logger.Errorf("写回恢复记录失败:%v", finalizeErr)
} }
s.logHub.Complete(restoreID, status) s.logHub.Complete(restoreID, status)
s.dispatchRestoreEvent(ctx, restoreID, status, errMessage, task) if dispatchErr := s.dispatchRestoreEvent(persistCtx, restoreID, status, errMessage, task); dispatchErr != nil {
logger.Warnf("派发恢复结果事件失败:%v", dispatchErr)
}
}() }()
if !acquireBackgroundSlot(ctx, s.semaphore) {
errMessage = ctx.Err().Error()
logger.Warnf("等待恢复执行槽时任务被取消:%v", ctx.Err())
return
}
defer func() { <-s.semaphore }()
logger.Infof("开始在本地执行恢复(备份记录 #%d)", backupRecord.ID) logger.Infof("开始在本地执行恢复(备份记录 #%d)", backupRecord.ID)
spec, specErr := s.buildTaskSpec(task, backupRecord.StartedAt) spec, specErr := s.buildTaskSpec(task, backupRecord.StartedAt)
@@ -387,9 +417,9 @@ func backupKindLabel(kind string) string {
// dispatchRestoreEvent 按终态向事件总线派发 restore_success 或 restore_failed。 // dispatchRestoreEvent 按终态向事件总线派发 restore_success 或 restore_failed。
// eventDispatcher 未注入时静默忽略,保持向后兼容。 // eventDispatcher 未注入时静默忽略,保持向后兼容。
func (s *RestoreService) dispatchRestoreEvent(ctx context.Context, restoreID uint, status, errMessage string, task *model.BackupTask) { func (s *RestoreService) dispatchRestoreEvent(ctx context.Context, restoreID uint, status, errMessage string, task *model.BackupTask) error {
if s.eventDispatcher == nil { if s.eventDispatcher == nil {
return return nil
} }
var eventType, title string var eventType, title string
switch status { switch status {
@@ -400,7 +430,7 @@ func (s *RestoreService) dispatchRestoreEvent(ctx context.Context, restoreID uin
eventType = model.NotificationEventRestoreFailed eventType = model.NotificationEventRestoreFailed
title = "BackupX 恢复失败" title = "BackupX 恢复失败"
default: default:
return return nil
} }
taskName := "未知任务" taskName := "未知任务"
if task != nil { if task != nil {
@@ -419,7 +449,7 @@ func (s *RestoreService) dispatchRestoreEvent(ctx context.Context, restoreID uin
if task != nil { if task != nil {
fields["taskId"] = task.ID fields["taskId"] = task.ID
} }
_ = s.eventDispatcher.DispatchEvent(ctx, eventType, title, body, fields) return s.eventDispatcher.DispatchEvent(ctx, eventType, title, body, fields)
} }
// resolveProvider 解密存储目标配置并创建 provider(共享实现)。 // resolveProvider 解密存储目标配置并创建 provider(共享实现)。
@@ -84,6 +84,11 @@ func newRestoreTestHarness(t *testing.T, remoteNode bool) *restoreTestHarness {
if err != nil { if err != nil {
t.Fatalf("database.Open: %v", err) t.Fatalf("database.Open: %v", err)
} }
sqlDB, err := db.DB()
if err != nil {
t.Fatalf("db.DB: %v", err)
}
t.Cleanup(func() { _ = sqlDB.Close() })
cipher := codec.NewConfigCipher("restore-secret") cipher := codec.NewConfigCipher("restore-secret")
targets := repository.NewStorageTargetRepository(db) targets := repository.NewStorageTargetRepository(db)
tasks := repository.NewBackupTaskRepository(db) tasks := repository.NewBackupTaskRepository(db)
@@ -159,11 +164,12 @@ func TestRestoreServiceStart_LocalNodeExecutesInline(t *testing.T) {
// 用同步 async 让测试可等待 // 用同步 async 让测试可等待
done := make(chan struct{}) done := make(chan struct{})
h.service.async = func(job func()) { h.service.async = func(job func(context.Context)) bool {
go func() { go func() {
job() job(context.Background())
close(done) close(done)
}() }()
return true
} }
detail, err := h.service.Start(ctx, backupDetail.ID, "tester") detail, err := h.service.Start(ctx, backupDetail.ID, "tester")
if err != nil { if err != nil {
@@ -200,6 +206,62 @@ func TestRestoreServiceStart_LocalNodeExecutesInline(t *testing.T) {
} }
} }
func TestRestoreServiceStart_RepositoryRecord(t *testing.T) {
h := newRestoreTestHarness(t, false)
ctx := context.Background()
task, err := h.tasks.FindByID(ctx, 1)
if err != nil {
t.Fatalf("FindByID task: %v", err)
}
task.BackupMode = model.BackupModeRepository
if err := h.tasks.Update(ctx, task); err != nil {
t.Fatalf("Update repository task: %v", err)
}
backupDetail, err := h.execution.RunTaskByIDSync(ctx, task.ID)
if err != nil {
t.Fatalf("RunTaskByIDSync repository: %v", err)
}
if backupDetail.BackupKind != model.BackupKindRepository {
t.Fatalf("expected repository backup, got %#v", backupDetail)
}
if err := os.RemoveAll(h.sourceDir); err != nil {
t.Fatalf("remove source: %v", err)
}
done := make(chan struct{})
h.service.async = func(job func(context.Context)) bool {
go func() {
job(context.Background())
close(done)
}()
return true
}
detail, err := h.service.Start(ctx, backupDetail.ID, "repository-test")
if err != nil {
t.Fatalf("Start repository restore: %v", err)
}
select {
case <-done:
case <-time.After(15 * time.Second):
t.Fatal("repository restore did not complete in time")
}
final, err := h.service.Get(ctx, detail.ID)
if err != nil {
t.Fatalf("Get repository restore: %v", err)
}
if final.Status != model.RestoreRecordStatusSuccess {
t.Fatalf("expected repository restore success, got %s (err=%s)", final.Status, final.ErrorMessage)
}
content, err := os.ReadFile(filepath.Join(h.sourceDir, "index.html"))
if err != nil {
t.Fatalf("read repository-restored file: %v", err)
}
if string(content) != "hello-restore" {
t.Fatalf("unexpected repository-restored content: %q", content)
}
}
// TestRestoreServiceStart_RejectsCorruptedBackup 验证恢复在还原前做 SHA-256 完整性 // TestRestoreServiceStart_RejectsCorruptedBackup 验证恢复在还原前做 SHA-256 完整性
// 校验:若已存储的备份对象被损坏/篡改,恢复必须失败且不触碰源数据。 // 校验:若已存储的备份对象被损坏/篡改,恢复必须失败且不触碰源数据。
func TestRestoreServiceStart_RejectsCorruptedBackup(t *testing.T) { func TestRestoreServiceStart_RejectsCorruptedBackup(t *testing.T) {
@@ -242,8 +304,9 @@ func TestRestoreServiceStart_RejectsCorruptedBackup(t *testing.T) {
} }
done := make(chan struct{}) done := make(chan struct{})
h.service.async = func(job func()) { h.service.async = func(job func(context.Context)) bool {
go func() { job(); close(done) }() go func() { job(context.Background()); close(done) }()
return true
} }
detail, err := h.service.Start(ctx, backupDetail.ID, "tester") detail, err := h.service.Start(ctx, backupDetail.ID, "tester")
if err != nil { if err != nil {
@@ -293,11 +356,12 @@ func TestRestoreServiceStart_RestoresToAlternatePath(t *testing.T) {
} }
done := make(chan struct{}) done := make(chan struct{})
h.service.async = func(job func()) { h.service.async = func(job func(context.Context)) bool {
go func() { go func() {
job() job(context.Background())
close(done) close(done)
}() }()
return true
} }
detail, err := h.service.StartSelective(ctx, backupDetail.ID, nil, altDir, "tester") detail, err := h.service.StartSelective(ctx, backupDetail.ID, nil, altDir, "tester")
if err != nil { if err != nil {
@@ -1,71 +0,0 @@
package service
import (
"context"
"time"
"backupx/server/internal/model"
"backupx/server/internal/repository"
"backupx/server/internal/storage"
"backupx/server/internal/storage/codec"
)
type RetentionService struct {
records repository.BackupRecordRepository
storageTargets repository.StorageTargetRepository
storageRegistry *storage.Registry
cipher *codec.ConfigCipher
}
func NewRetentionService(records repository.BackupRecordRepository, storageTargets repository.StorageTargetRepository, storageRegistry *storage.Registry, cipher *codec.ConfigCipher) *RetentionService {
return &RetentionService{records: records, storageTargets: storageTargets, storageRegistry: storageRegistry, cipher: cipher}
}
func (s *RetentionService) Apply(ctx context.Context, task *model.BackupTask) error {
if task == nil || (task.RetentionDays <= 0 && task.MaxBackups <= 0) {
return nil
}
items, err := s.records.ListSuccessfulByTask(ctx, task.ID)
if err != nil {
return err
}
removeSet := make(map[uint]model.BackupRecord)
if task.RetentionDays > 0 {
cutoff := time.Now().UTC().AddDate(0, 0, -task.RetentionDays)
for _, item := range items {
if item.CompletedAt != nil && item.CompletedAt.Before(cutoff) {
removeSet[item.ID] = item
}
}
}
if task.MaxBackups > 0 {
kept := 0
for _, item := range items {
if _, marked := removeSet[item.ID]; marked {
continue
}
kept++
if kept > task.MaxBackups {
removeSet[item.ID] = item
}
}
}
if len(removeSet) == 0 {
return nil
}
provider, _, err := buildStorageProviderFromRepos(ctx, task.StorageTargetID, s.storageTargets, s.storageRegistry, s.cipher)
if err != nil {
return err
}
for _, item := range removeSet {
if item.StoragePath != "" {
if err := provider.Delete(ctx, item.StoragePath); err != nil {
return err
}
}
if err := s.records.Delete(ctx, item.ID); err != nil {
return err
}
}
return nil
}
@@ -2,6 +2,7 @@ package service
import ( import (
"context" "context"
"errors"
"fmt" "fmt"
"net/http" "net/http"
"strings" "strings"
@@ -95,6 +96,7 @@ type StorageTargetService struct {
records repository.BackupRecordRepository records repository.BackupRecordRepository
registry *storage.Registry registry *storage.Registry
cipher *codec.ConfigCipher cipher *codec.ConfigCipher
background BackgroundRunner
} }
func NewStorageTargetService( func NewStorageTargetService(
@@ -114,6 +116,10 @@ func (s *StorageTargetService) SetBackupRecordRepository(records repository.Back
s.records = records s.records = records
} }
func (s *StorageTargetService) SetBackgroundRunner(runner BackgroundRunner) {
s.background = runner
}
func (s *StorageTargetService) List(ctx context.Context) ([]StorageTargetSummary, error) { func (s *StorageTargetService) List(ctx context.Context) ([]StorageTargetSummary, error) {
items, err := s.targets.List(ctx) items, err := s.targets.List(ctx)
if err != nil { if err != nil {
@@ -254,7 +260,12 @@ func (s *StorageTargetService) TestConnection(ctx context.Context, input Storage
item.LastTestMessage = "连接成功" item.LastTestMessage = "连接成功"
} }
if item.ID != 0 { if item.ID != 0 {
_ = s.targets.Update(ctx, item) if updateErr := s.targets.Update(ctx, item); updateErr != nil {
if testErr != nil {
return apperror.BadRequest("STORAGE_TARGET_TEST_FAILED", sanitizeMessage(testErr.Error()), errors.Join(testErr, fmt.Errorf("save connection test result: %w", updateErr)))
}
return apperror.Internal("STORAGE_TARGET_TEST_RESULT_SAVE_FAILED", "连接成功,但无法保存测试结果", updateErr)
}
} }
if testErr != nil { if testErr != nil {
return apperror.BadRequest("STORAGE_TARGET_TEST_FAILED", sanitizeMessage(testErr.Error()), testErr) return apperror.BadRequest("STORAGE_TARGET_TEST_FAILED", sanitizeMessage(testErr.Error()), testErr)
@@ -269,23 +280,23 @@ func (s *StorageTargetService) StartHealthMonitor(ctx context.Context, dispatche
if interval <= 0 { if interval <= 0 {
interval = 5 * time.Minute interval = 5 * time.Minute
} }
ticker := time.NewTicker(interval)
// notified 跟踪已告警的目标,避免每轮重复 // notified 跟踪已告警的目标,避免每轮重复
notified := map[uint]bool{} notified := map[uint]bool{}
capacityNotified := map[uint]bool{} capacityNotified := map[uint]bool{}
var mu sync.Mutex var mu sync.Mutex
go func() { startBackgroundMonitor(s.background, ctx, func(runCtx context.Context) {
ticker := time.NewTicker(interval)
defer ticker.Stop() defer ticker.Stop()
for { for {
select { select {
case <-ctx.Done(): case <-runCtx.Done():
return return
case <-ticker.C: case <-ticker.C:
s.runHealthCheckOnce(ctx, dispatcher, &mu, notified) s.runHealthCheckOnce(runCtx, dispatcher, &mu, notified)
s.runCapacityCheckOnce(ctx, dispatcher, &mu, capacityNotified) s.runCapacityCheckOnce(runCtx, dispatcher, &mu, capacityNotified)
} }
} }
}() })
} }
// StorageCapacityWarningThreshold 存储使用率告警阈值(85%)。 // StorageCapacityWarningThreshold 存储使用率告警阈值(85%)。
@@ -371,7 +382,6 @@ func (s *StorageTargetService) runHealthCheckOnce(ctx context.Context, dispatche
if !target.Enabled { if !target.Enabled {
continue continue
} }
previousStatus := target.LastTestStatus
configMap := map[string]any{} configMap := map[string]any{}
if err := s.cipher.DecryptJSON(target.ConfigCiphertext, &configMap); err != nil { if err := s.cipher.DecryptJSON(target.ConfigCiphertext, &configMap); err != nil {
continue continue
@@ -380,13 +390,13 @@ func (s *StorageTargetService) runHealthCheckOnce(ctx context.Context, dispatche
now := time.Now().UTC() now := time.Now().UTC()
if err != nil { if err != nil {
s.applyHealthResult(ctx, &target, now, false, err.Error()) s.applyHealthResult(ctx, &target, now, false, err.Error())
s.notifyUnhealthyTransition(ctx, dispatcher, mu, notified, &target, previousStatus, err.Error()) s.notifyUnhealthyTransition(ctx, dispatcher, mu, notified, &target, err.Error())
continue continue
} }
testErr := provider.TestConnection(ctx) testErr := provider.TestConnection(ctx)
if testErr != nil { if testErr != nil {
s.applyHealthResult(ctx, &target, now, false, testErr.Error()) s.applyHealthResult(ctx, &target, now, false, testErr.Error())
s.notifyUnhealthyTransition(ctx, dispatcher, mu, notified, &target, previousStatus, testErr.Error()) s.notifyUnhealthyTransition(ctx, dispatcher, mu, notified, &target, testErr.Error())
continue continue
} }
s.applyHealthResult(ctx, &target, now, true, "连接成功") s.applyHealthResult(ctx, &target, now, true, "连接成功")
@@ -408,7 +418,7 @@ func (s *StorageTargetService) applyHealthResult(ctx context.Context, target *mo
_ = s.targets.Update(ctx, target) _ = s.targets.Update(ctx, target)
} }
func (s *StorageTargetService) notifyUnhealthyTransition(ctx context.Context, dispatcher EventDispatcher, mu *sync.Mutex, notified map[uint]bool, target *model.StorageTarget, previousStatus string, message string) { func (s *StorageTargetService) notifyUnhealthyTransition(ctx context.Context, dispatcher EventDispatcher, mu *sync.Mutex, notified map[uint]bool, target *model.StorageTarget, message string) {
if dispatcher == nil { if dispatcher == nil {
return return
} }
@@ -423,7 +433,6 @@ func (s *StorageTargetService) notifyUnhealthyTransition(ctx context.Context, di
if already { if already {
return return
} }
_ = previousStatus // 保留参数便于未来扩展:区分"从未测试"与"从 success 掉线"
title := "BackupX 存储目标连接失败" title := "BackupX 存储目标连接失败"
body := fmt.Sprintf("存储目标:%s (类型: %s)\n错误:%s", target.Name, target.Type, message) body := fmt.Sprintf("存储目标:%s (类型: %s)\n错误:%s", target.Name, target.Type, message)
fields := map[string]any{ fields := map[string]any{
@@ -473,7 +482,9 @@ func (s *StorageTargetService) CompleteGoogleDriveOAuth(ctx context.Context, inp
// Mark used immediately to prevent duplicate requests (e.g. React StrictMode double invocation) // Mark used immediately to prevent duplicate requests (e.g. React StrictMode double invocation)
now := time.Now().UTC() now := time.Now().UTC()
session.UsedAt = &now session.UsedAt = &now
_ = s.oauthSessions.Update(ctx, session) if err := s.oauthSessions.Update(ctx, session); err != nil {
return nil, apperror.Internal("STORAGE_GOOGLE_OAUTH_SESSION_FAILED", "无法锁定 Google Drive 授权会话", err)
}
var draft googleDriveOAuthDraft var draft googleDriveOAuthDraft
if err := s.cipher.DecryptJSON(session.PayloadCiphertext, &draft); err != nil { if err := s.cipher.DecryptJSON(session.PayloadCiphertext, &draft); err != nil {
@@ -180,7 +180,7 @@ func (s *TaskExportService) Import(ctx context.Context, payload ExportPayload) (
results = append(results, ImportResult{Name: t.Name, TaskID: detail.ID, Success: true}) results = append(results, ImportResult{Name: t.Name, TaskID: detail.ID, Success: true})
} }
// 第二阶段:依赖链接(上游任务名 → 新 ID) // 第二阶段:依赖链接(上游任务名 → 新 ID)
for i, t := range payload.Tasks { for _, t := range payload.Tasks {
if len(t.DependsOnTaskNames) == 0 { if len(t.DependsOnTaskNames) == 0 {
continue continue
} }
@@ -206,7 +206,6 @@ func (s *TaskExportService) Import(ctx context.Context, payload ExportPayload) (
break break
} }
} }
_ = i
} }
} }
return results, nil return results, nil
@@ -3,7 +3,6 @@ package service
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"fmt"
"strings" "strings"
"backupx/server/internal/apperror" "backupx/server/internal/apperror"
@@ -235,6 +234,3 @@ func toTemplateSummary(item *model.TaskTemplate) TaskTemplateSummary {
UpdatedAt: item.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"), UpdatedAt: item.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"),
} }
} }
// 确保未使用告警
var _ = fmt.Sprintf
@@ -0,0 +1,22 @@
package service
import (
"testing"
"gorm.io/gorm"
)
// closeTestDatabase releases SQLite file handles before testing.TempDir cleanup.
// Windows does not permit removal of an open database file.
func closeTestDatabase(t *testing.T, db *gorm.DB) {
t.Helper()
sqlDB, err := db.DB()
if err != nil {
t.Fatalf("get test database handle: %v", err)
}
t.Cleanup(func() {
if err := sqlDB.Close(); err != nil {
t.Errorf("close test database: %v", err)
}
})
}
+36 -11
View File
@@ -39,7 +39,7 @@ type VerificationService struct {
notifier VerificationNotifier notifier VerificationNotifier
tempDir string tempDir string
semaphore chan struct{} semaphore chan struct{}
async func(func()) async func(func(context.Context)) bool
now func() time.Time now func() time.Time
metrics *metrics.Metrics metrics *metrics.Metrics
} }
@@ -49,6 +49,13 @@ func (s *VerificationService) SetMetrics(m *metrics.Metrics) {
s.metrics = m s.metrics = m
} }
// SetBackgroundRunner binds local verification work to the application lifecycle.
func (s *VerificationService) SetBackgroundRunner(runner BackgroundRunner) {
if runner != nil {
s.async = runner.Go
}
}
// VerificationNotifier 给用户推送验证完成/失败通知。 // VerificationNotifier 给用户推送验证完成/失败通知。
// 可选注入:未注入时仅写记录。 // 可选注入:未注入时仅写记录。
type VerificationNotifier interface { type VerificationNotifier interface {
@@ -129,7 +136,7 @@ func NewVerificationService(
notifier: noopVerificationNotifier{}, notifier: noopVerificationNotifier{},
tempDir: tempDir, tempDir: tempDir,
semaphore: make(chan struct{}, maxConcurrent), semaphore: make(chan struct{}, maxConcurrent),
async: func(job func()) { go job() }, async: runDetached,
now: func() time.Time { return time.Now().UTC() }, now: func() time.Time { return time.Now().UTC() },
} }
} }
@@ -231,10 +238,16 @@ func (s *VerificationService) Start(ctx context.Context, backupRecordID uint, mo
if err := s.verifications.Create(ctx, verification); err != nil { if err := s.verifications.Create(ctx, verification); err != nil {
return nil, apperror.Internal("VERIFY_RECORD_CREATE_FAILED", "无法创建验证记录", err) return nil, apperror.Internal("VERIFY_RECORD_CREATE_FAILED", "无法创建验证记录", err)
} }
run := func() { run := func(runCtx context.Context) {
s.executeLocally(context.Background(), verification.ID, task, record) s.executeLocally(runCtx, verification.ID, task, record)
}
if !s.async(run) {
message := "服务正在关闭,验证任务未启动"
if finalizeErr := s.finalize(ctx, verification.ID, model.VerificationRecordStatusFailed, message, "", ""); finalizeErr != nil {
return nil, apperror.Internal("VERIFY_FINALIZE_FAILED", "无法写回验证失败状态", finalizeErr)
}
return nil, backgroundTaskUnavailable("VERIFY_SERVICE_SHUTTING_DOWN")
} }
s.async(run)
return s.getDetail(ctx, verification.ID) return s.getDetail(ctx, verification.ID)
} }
@@ -247,25 +260,37 @@ func (s *VerificationService) validateClusterAccessible(ctx context.Context, rec
// executeLocally 异步执行验证:下载 → 解密 → 解压 → 按类型校验。 // executeLocally 异步执行验证:下载 → 解密 → 解压 → 按类型校验。
func (s *VerificationService) executeLocally(ctx context.Context, verID uint, task *model.BackupTask, backupRecord *model.BackupRecord) { func (s *VerificationService) executeLocally(ctx context.Context, verID uint, task *model.BackupTask, backupRecord *model.BackupRecord) {
s.semaphore <- struct{}{}
defer func() { <-s.semaphore }()
logger := backup.NewExecutionLogger(verID, s.logHub) logger := backup.NewExecutionLogger(verID, s.logHub)
status := model.VerificationRecordStatusFailed status := model.VerificationRecordStatusFailed
errMessage := "" errMessage := ""
summary := "" summary := ""
defer func() { defer func() {
_ = s.finalize(ctx, verID, status, errMessage, summary, logger.String()) persistCtx, cancel := finalizationContext(ctx)
defer cancel()
if finalizeErr := s.finalize(persistCtx, verID, status, errMessage, summary, logger.String()); finalizeErr != nil {
logger.Errorf("写回验证记录失败:%v", finalizeErr)
}
s.logHub.Complete(verID, status) s.logHub.Complete(verID, status)
// 失败时推送通知(best-effort // 失败时推送通知(best-effort
if status == model.VerificationRecordStatusFailed && s.notifier != nil { if status == model.VerificationRecordStatusFailed && s.notifier != nil {
if record, err := s.verifications.FindByID(ctx, verID); err == nil && record != nil { if record, findErr := s.verifications.FindByID(persistCtx, verID); findErr != nil {
_ = s.notifier.NotifyVerificationResult(ctx, task, record) logger.Warnf("读取验证记录以发送通知失败:%v", findErr)
} else if record != nil {
if notifyErr := s.notifier.NotifyVerificationResult(persistCtx, task, record); notifyErr != nil {
logger.Warnf("发送验证失败通知失败:%v", notifyErr)
}
} }
} }
}() }()
if !acquireBackgroundSlot(ctx, s.semaphore) {
errMessage = ctx.Err().Error()
logger.Warnf("等待验证执行槽时任务被取消:%v", ctx.Err())
return
}
defer func() { <-s.semaphore }()
logger.Infof("开始验证备份记录 #%d(模式:%s)", backupRecord.ID, model.VerificationModeQuick) logger.Infof("开始验证备份记录 #%d(模式:%s)", backupRecord.ID, model.VerificationModeQuick)
if err := os.MkdirAll(s.tempDir, 0o755); err != nil { if err := os.MkdirAll(s.tempDir, 0o755); err != nil {
@@ -11,6 +11,7 @@ import (
"backupx/server/internal/backup" "backupx/server/internal/backup"
"backupx/server/internal/config" "backupx/server/internal/config"
"backupx/server/internal/database" "backupx/server/internal/database"
"backupx/server/internal/lifecycle"
"backupx/server/internal/logger" "backupx/server/internal/logger"
"backupx/server/internal/model" "backupx/server/internal/model"
"backupx/server/internal/repository" "backupx/server/internal/repository"
@@ -45,6 +46,11 @@ func newVerifyTestHarness(t *testing.T) *verifyTestHarness {
if err != nil { if err != nil {
t.Fatalf("database.Open: %v", err) t.Fatalf("database.Open: %v", err)
} }
sqlDB, err := db.DB()
if err != nil {
t.Fatalf("db.DB: %v", err)
}
t.Cleanup(func() { _ = sqlDB.Close() })
cipher := codec.NewConfigCipher("verify-secret") cipher := codec.NewConfigCipher("verify-secret")
targets := repository.NewStorageTargetRepository(db) targets := repository.NewStorageTargetRepository(db)
tasks := repository.NewBackupTaskRepository(db) tasks := repository.NewBackupTaskRepository(db)
@@ -77,8 +83,9 @@ func (h *verifyTestHarness) runVerify(t *testing.T, backupRecordID uint) *Verifi
t.Helper() t.Helper()
ctx := context.Background() ctx := context.Background()
done := make(chan struct{}) done := make(chan struct{})
h.verify.async = func(job func()) { h.verify.async = func(job func(context.Context)) bool {
go func() { job(); close(done) }() go func() { job(context.Background()); close(done) }()
return true
} }
detail, err := h.verify.Start(ctx, backupRecordID, "quick", "tester") detail, err := h.verify.Start(ctx, backupRecordID, "quick", "tester")
if err != nil { if err != nil {
@@ -96,6 +103,48 @@ func (h *verifyTestHarness) runVerify(t *testing.T, backupRecordID uint) *Verifi
return final return final
} }
func TestVerificationServiceSupervisorCancellationFinalizesRecord(t *testing.T) {
h := newVerifyTestHarness(t)
backupDetail, err := h.execution.RunTaskByIDSync(context.Background(), 1)
if err != nil {
t.Fatalf("RunTaskByIDSync: %v", err)
}
supervisor := lifecycle.NewSupervisor(context.Background())
h.verify.SetBackgroundRunner(supervisor)
// Occupy the only available execution path before cancellation so the test
// deterministically exercises cancellation while queued.
for i := 0; i < cap(h.verify.semaphore); i++ {
h.verify.semaphore <- struct{}{}
}
detail, err := h.verify.Start(context.Background(), backupDetail.ID, model.VerificationModeQuick, "tester")
if err != nil {
t.Fatalf("Start: %v", err)
}
waitCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := supervisor.Shutdown(waitCtx); err != nil {
t.Fatalf("Shutdown: %v", err)
}
for i := 0; i < cap(h.verify.semaphore); i++ {
<-h.verify.semaphore
}
final, err := h.verify.Get(context.Background(), detail.ID)
if err != nil {
t.Fatalf("Get: %v", err)
}
if final.Status != model.VerificationRecordStatusFailed {
t.Fatalf("status = %q, want failed", final.Status)
}
if final.CompletedAt == nil {
t.Fatal("canceled verification was not finalized")
}
if !strings.Contains(strings.ToLower(final.ErrorMessage), "canceled") {
t.Fatalf("error message = %q, want cancellation", final.ErrorMessage)
}
}
// TestVerificationService_Success 覆盖正常路径:对一个有效(gzip 压缩)的备份做验证应通过。 // TestVerificationService_Success 覆盖正常路径:对一个有效(gzip 压缩)的备份做验证应通过。
// 同时回归保护 #77——新增的 SHA-256 校验不得误伤合法的压缩备份。 // 同时回归保护 #77——新增的 SHA-256 校验不得误伤合法的压缩备份。
func TestVerificationService_Success(t *testing.T) { func TestVerificationService_Success(t *testing.T) {
@@ -0,0 +1,21 @@
package rclone
import "backupx/server/internal/storage"
// NewDefaultRegistry returns the storage factory set shared by Master, Agent,
// and the standalone Backint process.
func NewDefaultRegistry() *storage.Registry {
registry := storage.NewRegistry(
NewLocalDiskFactory(),
NewS3Factory(),
NewWebDAVFactory(),
NewGoogleDriveFactory(),
NewAliyunOSSFactory(),
NewTencentCOSFactory(),
NewQiniuKodoFactory(),
NewFTPFactory(),
NewRcloneFactory(),
)
RegisterAllBackends(registry)
return registry
}
+8 -124
View File
@@ -2,7 +2,6 @@ package storage
import ( import (
"context" "context"
"encoding/json"
"fmt" "fmt"
"sort" "sort"
"sync" "sync"
@@ -10,26 +9,6 @@ import (
"backupx/server/internal/apperror" "backupx/server/internal/apperror"
) )
type providerFactoryWithNew interface {
New(context.Context, map[string]any) (StorageProvider, error)
}
type providerFactoryWithCreate interface {
Create(context.Context, json.RawMessage) (StorageProvider, error)
}
type providerFactoryWithSensitiveFields interface {
SensitiveFields() []string
}
type providerFactoryWithSensitiveKeys interface {
SensitiveKeys() []string
}
type providerFactoryWithValidate interface {
Validate(json.RawMessage) error
}
type Registry struct { type Registry struct {
mu sync.RWMutex mu sync.RWMutex
factories map[ProviderType]ProviderFactory factories map[ProviderType]ProviderFactory
@@ -78,115 +57,20 @@ func (r *Registry) SensitiveFields(providerType string) []string {
if !ok { if !ok {
return nil return nil
} }
if typed, ok := factory.(providerFactoryWithSensitiveFields); ok { return factory.SensitiveFields()
return typed.SensitiveFields()
}
if typed, ok := factory.(providerFactoryWithSensitiveKeys); ok {
return typed.SensitiveKeys()
}
return nil
} }
func (r *Registry) SensitiveKeys(providerType string) []string { func (r *Registry) Create(ctx context.Context, providerType string, config map[string]any) (StorageProvider, error) {
return r.SensitiveFields(providerType)
}
func (r *Registry) Validate(providerType string, raw json.RawMessage) error {
factory, ok := r.Factory(providerType)
if !ok {
return apperror.BadRequest("STORAGE_PROVIDER_UNSUPPORTED", "不支持的存储类型", fmt.Errorf("unsupported storage provider type: %s", providerType))
}
if typed, ok := factory.(providerFactoryWithValidate); ok {
if err := typed.Validate(raw); err != nil {
return apperror.BadRequest("STORAGE_TARGET_INVALID_CONFIG", "存储目标配置不合法", err)
}
return nil
}
configMap, err := decodeConfigMap(raw)
if err != nil {
return apperror.BadRequest("STORAGE_TARGET_INVALID_CONFIG", "存储目标配置不合法", err)
}
if typed, ok := factory.(providerFactoryWithNew); ok {
if _, err := typed.New(context.Background(), configMap); err != nil {
return apperror.BadRequest("STORAGE_TARGET_INVALID_CONFIG", "存储目标配置不合法", err)
}
return nil
}
return apperror.BadRequest("STORAGE_TARGET_INVALID_CONFIG", "存储目标配置不合法", fmt.Errorf("provider %s has no validator", providerType))
}
func (r *Registry) Create(ctx context.Context, providerType string, rawConfig any) (StorageProvider, error) {
factory, ok := r.Factory(providerType) factory, ok := r.Factory(providerType)
if !ok { if !ok {
return nil, apperror.BadRequest("STORAGE_PROVIDER_UNSUPPORTED", "不支持的存储类型", fmt.Errorf("unsupported storage provider type: %s", providerType)) return nil, apperror.BadRequest("STORAGE_PROVIDER_UNSUPPORTED", "不支持的存储类型", fmt.Errorf("unsupported storage provider type: %s", providerType))
} }
raw, configMap, err := normalizeConfig(rawConfig) if config == nil {
config = map[string]any{}
}
provider, err := factory.New(ctx, config)
if err != nil { if err != nil {
return nil, apperror.BadRequest("STORAGE_TARGET_INVALID_CONFIG", "存储目标配置不合法", err) return nil, apperror.BadRequest("STORAGE_TARGET_INVALID_CONFIG", "无法创建存储客户端", err)
} }
if typed, ok := factory.(providerFactoryWithNew); ok { return provider, nil
provider, err := typed.New(ctx, configMap)
if err != nil {
return nil, apperror.BadRequest("STORAGE_TARGET_INVALID_CONFIG", "无法创建存储客户端", err)
}
return provider, nil
}
if typed, ok := factory.(providerFactoryWithCreate); ok {
provider, err := typed.Create(ctx, raw)
if err != nil {
return nil, apperror.BadRequest("STORAGE_TARGET_INVALID_CONFIG", "无法创建存储客户端", err)
}
return provider, nil
}
return nil, apperror.BadRequest("STORAGE_TARGET_INVALID_CONFIG", "无法创建存储客户端", fmt.Errorf("provider %s has no constructor", providerType))
}
func normalizeConfig(rawConfig any) (json.RawMessage, map[string]any, error) {
switch value := rawConfig.(type) {
case nil:
return json.RawMessage("{}"), map[string]any{}, nil
case map[string]any:
raw, err := json.Marshal(value)
if err != nil {
return nil, nil, fmt.Errorf("marshal config: %w", err)
}
return raw, value, nil
case json.RawMessage:
configMap, err := decodeConfigMap(value)
if err != nil {
return nil, nil, err
}
return value, configMap, nil
case []byte:
raw := json.RawMessage(value)
configMap, err := decodeConfigMap(raw)
if err != nil {
return nil, nil, err
}
return raw, configMap, nil
default:
raw, err := json.Marshal(value)
if err != nil {
return nil, nil, fmt.Errorf("marshal config: %w", err)
}
configMap, err := decodeConfigMap(raw)
if err != nil {
return nil, nil, err
}
return raw, configMap, nil
}
}
func decodeConfigMap(raw json.RawMessage) (map[string]any, error) {
if len(raw) == 0 {
return map[string]any{}, nil
}
var configMap map[string]any
if err := json.Unmarshal(raw, &configMap); err != nil {
return nil, fmt.Errorf("decode config: %w", err)
}
if configMap == nil {
return map[string]any{}, nil
}
return configMap, nil
} }
+2
View File
@@ -66,6 +66,8 @@ type StorageRangeDownloader interface {
type ProviderFactory interface { type ProviderFactory interface {
Type() ProviderType Type() ProviderType
SensitiveFields() []string
New(context.Context, map[string]any) (StorageProvider, error)
} }
// StorageAbout 是可选能力接口,支持查询远端存储空间。 // StorageAbout 是可选能力接口,支持查询远端存储空间。
+42 -2
View File
@@ -2,13 +2,15 @@ package response
import ( import (
"errors" "errors"
"fmt"
"net/http" "net/http"
"backupx/server/internal/apperror" "backupx/server/internal/apperror"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"go.uber.org/zap"
) )
const loggerContextKey = "backupx.response.logger"
type Envelope struct { type Envelope struct {
Code string `json:"code"` Code string `json:"code"`
Message string `json:"message"` Message string `json:"message"`
@@ -19,12 +21,50 @@ func Success(c *gin.Context, data any) {
c.JSON(http.StatusOK, Envelope{Code: "OK", Message: "success", Data: data}) c.JSON(http.StatusOK, Envelope{Code: "OK", Message: "success", Data: data})
} }
// SetLogger attaches the application logger to a request so response helpers
// can report failures without relying on a process-global logger.
func SetLogger(c *gin.Context, logger *zap.Logger) {
if logger != nil {
c.Set(loggerContextKey, logger)
}
}
func Error(c *gin.Context, err error) { func Error(c *gin.Context, err error) {
fmt.Printf("HTTP Error: %v\n", err)
var appErr *apperror.AppError var appErr *apperror.AppError
if errors.As(err, &appErr) { if errors.As(err, &appErr) {
logError(c, appErr.Status, appErr.Code, err)
c.JSON(appErr.Status, Envelope{Code: appErr.Code, Message: appErr.Message}) c.JSON(appErr.Status, Envelope{Code: appErr.Code, Message: appErr.Message})
return return
} }
logError(c, http.StatusInternalServerError, "INTERNAL_ERROR", err)
c.JSON(http.StatusInternalServerError, Envelope{Code: "INTERNAL_ERROR", Message: "服务器内部错误"}) c.JSON(http.StatusInternalServerError, Envelope{Code: "INTERNAL_ERROR", Message: "服务器内部错误"})
} }
func logError(c *gin.Context, status int, code string, err error) {
value, exists := c.Get(loggerContextKey)
if !exists {
return
}
logger, ok := value.(*zap.Logger)
if !ok || logger == nil {
return
}
method := ""
path := ""
if c.Request != nil {
method = c.Request.Method
path = c.Request.URL.Path
}
fields := []zap.Field{
zap.Int("status", status),
zap.String("code", code),
zap.String("method", method),
zap.String("path", path),
zap.Error(err),
}
if status >= http.StatusInternalServerError {
logger.Error("http request failed", fields...)
return
}
logger.Warn("http request rejected", fields...)
}
+39
View File
@@ -0,0 +1,39 @@
package response
import (
"errors"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
"go.uber.org/zap/zaptest/observer"
)
func TestErrorLogsInternalFailureAndHidesDetail(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
ctx.Request = httptest.NewRequest(http.MethodGet, "/api/test", nil)
core, observed := observer.New(zap.WarnLevel)
SetLogger(ctx, zap.New(core))
Error(ctx, errors.New("database password should stay private"))
if recorder.Code != http.StatusInternalServerError {
t.Fatalf("status = %d, want %d", recorder.Code, http.StatusInternalServerError)
}
if strings.Contains(recorder.Body.String(), "database password") {
t.Fatalf("internal detail leaked in response: %s", recorder.Body.String())
}
entries := observed.All()
if len(entries) != 1 || entries[0].Level != zap.ErrorLevel {
t.Fatalf("observed logs = %#v, want one error", entries)
}
context := entries[0].ContextMap()
if context["code"] != "INTERNAL_ERROR" || context["path"] != "/api/test" {
t.Fatalf("unexpected log context: %#v", context)
}
}
+4
View File
@@ -0,0 +1,4 @@
coverage
dist
node_modules
package-lock.json
+6
View File
@@ -0,0 +1,6 @@
{
"semi": false,
"singleQuote": true,
"trailingComma": "all",
"printWidth": 100
}
+38
View File
@@ -0,0 +1,38 @@
import js from '@eslint/js'
import eslintConfigPrettier from 'eslint-config-prettier'
import reactHooks from 'eslint-plugin-react-hooks'
import reactRefresh from 'eslint-plugin-react-refresh'
import tseslint from 'typescript-eslint'
export default tseslint.config(
{
ignores: ['coverage', 'dist'],
},
{
files: ['**/*.{js,mjs,cjs}'],
...js.configs.recommended,
},
{
files: ['**/*.{ts,tsx}'],
extends: [...tseslint.configs.recommended],
plugins: {
'react-hooks': reactHooks,
'react-refresh': reactRefresh,
},
rules: {
'@typescript-eslint/no-explicit-any': 'off',
'@typescript-eslint/no-unused-vars': [
'error',
{
argsIgnorePattern: '^_',
caughtErrorsIgnorePattern: '^_',
varsIgnorePattern: '^_',
},
],
'react-hooks/exhaustive-deps': 'error',
'react-hooks/rules-of-hooks': 'error',
'react-refresh/only-export-components': 'off',
},
},
eslintConfigPrettier,
)
+1235 -22
View File
File diff suppressed because it is too large Load Diff
+11 -1
View File
@@ -8,6 +8,9 @@
}, },
"scripts": { "scripts": {
"dev": "vite", "dev": "vite",
"lint": "eslint . --max-warnings 0",
"format": "prettier --write .",
"format:check": "prettier --check .",
"typecheck": "tsc --noEmit -p tsconfig.json && tsc --noEmit -p tsconfig.node.json", "typecheck": "tsc --noEmit -p tsconfig.json && tsc --noEmit -p tsconfig.node.json",
"build": "npm run typecheck && vite build", "build": "npm run typecheck && vite build",
"preview": "vite preview", "preview": "vite preview",
@@ -22,10 +25,11 @@
"react": "^18.3.1", "react": "^18.3.1",
"react-dom": "^18.3.1", "react-dom": "^18.3.1",
"react-i18next": "^16.5.6", "react-i18next": "^16.5.6",
"react-router-dom": "^7.0.0", "react-router-dom": "^7.18.2",
"zustand": "^5.0.15" "zustand": "^5.0.15"
}, },
"devDependencies": { "devDependencies": {
"@eslint/js": "^10.0.1",
"@testing-library/jest-dom": "^6.6.3", "@testing-library/jest-dom": "^6.6.3",
"@testing-library/react": "^16.2.0", "@testing-library/react": "^16.2.0",
"@testing-library/user-event": "^14.6.4", "@testing-library/user-event": "^14.6.4",
@@ -33,8 +37,14 @@
"@types/react": "^18.3.20", "@types/react": "^18.3.20",
"@types/react-dom": "^18.3.6", "@types/react-dom": "^18.3.6",
"@vitejs/plugin-react": "^4.3.4", "@vitejs/plugin-react": "^4.3.4",
"eslint": "^10.9.1",
"eslint-config-prettier": "^10.1.8",
"eslint-plugin-react-hooks": "^7.1.1",
"eslint-plugin-react-refresh": "^0.5.4",
"jsdom": "^26.0.0", "jsdom": "^26.0.0",
"prettier": "^3.9.6",
"typescript": "^5.7.3", "typescript": "^5.7.3",
"typescript-eslint": "^8.68.0",
"vite": "^6.4.3", "vite": "^6.4.3",
"vitest": "^4.1.10" "vitest": "^4.1.10"
} }
-13
View File
@@ -1,13 +0,0 @@
import { BrowserRouter } from 'react-router-dom'
import { RouterView } from './router'
import { AuthBootstrap } from './components/AuthBootstrap'
export function App() {
return (
<BrowserRouter>
<AuthBootstrap>
<RouterView />
</AuthBootstrap>
</BrowserRouter>
)
}
+40 -12
View File
@@ -1,4 +1,4 @@
import { Button, Divider, Input, Select, Space, Switch, Typography } from '@arco-design/web-react' import { Button, Input, Select, Space, Switch, Typography } from '@arco-design/web-react'
import { useEffect, useMemo, useState } from 'react' import { useEffect, useMemo, useState } from 'react'
export interface CronInputProps { export interface CronInputProps {
@@ -72,8 +72,20 @@ function describeCron(expr: string): string {
// 每周某天 // 每周某天
if (day === '*' && week !== '*') { if (day === '*' && week !== '*') {
const weekNames: Record<string, string> = { '0': '日', '1': '一', '2': '二', '3': '三', '4': '四', '5': '五', '6': '六', '7': '日' } const weekNames: Record<string, string> = {
const days = week.split(',').map((w) => `${weekNames[w] || w}`).join('') '0': '',
'1': '一',
'2': '二',
'3': '三',
'4': '四',
'5': '五',
'6': '六',
'7': '日',
}
const days = week
.split(',')
.map((w) => `${weekNames[w] || w}`)
.join('、')
return `${days} ${time} 执行` return `${days} ${time} 执行`
} }
// 每月某日 // 每月某日
@@ -103,7 +115,7 @@ export function CronInput({ value, onChange }: CronInputProps) {
// 从 prop 同步 // 从 prop 同步
useEffect(() => { useEffect(() => {
if (value !== undefined && value !== cronExpr) { if (value !== undefined) {
setCronExpr(value || DEFAULT_CRON) setCronExpr(value || DEFAULT_CRON)
} }
}, [value]) }, [value])
@@ -205,12 +217,12 @@ export function CronInput({ value, onChange }: CronInputProps) {
if (isAdvanced) emit(val) if (isAdvanced) emit(val)
}} }}
/> />
{description && ( {description && <Typography.Text type="secondary">{description}</Typography.Text>}
<Typography.Text type="secondary">{description}</Typography.Text>
)}
<div style={{ marginLeft: 'auto' }}> <div style={{ marginLeft: 'auto' }}>
<Space size="mini"> <Space size="mini">
<Typography.Text type="secondary" style={{ fontSize: 12 }}></Typography.Text> <Typography.Text type="secondary" style={{ fontSize: 12 }}>
</Typography.Text>
<Switch <Switch
size="small" size="small"
checked={isAdvanced} checked={isAdvanced}
@@ -230,16 +242,32 @@ export function CronInput({ value, onChange }: CronInputProps) {
{showCustom && !isAdvanced && ( {showCustom && !isAdvanced && (
<div style={{ padding: '12px 16px', background: 'var(--color-fill-1)', borderRadius: 6 }}> <div style={{ padding: '12px 16px', background: 'var(--color-fill-1)', borderRadius: 6 }}>
<Space size="large" style={{ marginBottom: 12 }}> <Space size="large" style={{ marginBottom: 12 }}>
<Button size="small" type={mode === 'daily' ? 'primary' : 'text'} onClick={() => handleCustomChange({ mode: 'daily' })}> <Button
size="small"
type={mode === 'daily' ? 'primary' : 'text'}
onClick={() => handleCustomChange({ mode: 'daily' })}
>
</Button> </Button>
<Button size="small" type={mode === 'weekly' ? 'primary' : 'text'} onClick={() => handleCustomChange({ mode: 'weekly' })}> <Button
size="small"
type={mode === 'weekly' ? 'primary' : 'text'}
onClick={() => handleCustomChange({ mode: 'weekly' })}
>
</Button> </Button>
<Button size="small" type={mode === 'monthly' ? 'primary' : 'text'} onClick={() => handleCustomChange({ mode: 'monthly' })}> <Button
size="small"
type={mode === 'monthly' ? 'primary' : 'text'}
onClick={() => handleCustomChange({ mode: 'monthly' })}
>
</Button> </Button>
<Button size="small" type={mode === 'interval' ? 'primary' : 'text'} onClick={() => handleCustomChange({ mode: 'interval' })}> <Button
size="small"
type={mode === 'interval' ? 'primary' : 'text'}
onClick={() => handleCustomChange({ mode: 'interval' })}
>
</Button> </Button>
</Space> </Space>
@@ -0,0 +1,69 @@
import { Avatar, Button, Dropdown, Menu } from '@arco-design/web-react'
import { useState } from 'react'
import { IconDown, IconLock, IconPoweroff, IconSafe } from '../icons'
import type { UserInfo } from '../../services/auth'
import { useAuthStore } from '../../stores/auth'
import { roleLabel } from '../../utils/permissions'
import { ChangePasswordModal } from './ChangePasswordModal'
import { MfaSettingsModal } from './MfaSettingsModal'
interface AccountControlsProps {
user: UserInfo | null
}
export function AccountControls({ user }: AccountControlsProps) {
const [passwordVisible, setPasswordVisible] = useState(false)
const [securityVisible, setSecurityVisible] = useState(false)
const logout = useAuthStore((state) => state.logout)
const droplist = (
<Menu
onClickMenuItem={(key) => {
if (key === 'password') {
setPasswordVisible(true)
} else if (key === 'two-factor') {
setSecurityVisible(true)
} else if (key === 'logout') {
logout()
}
}}
>
<Menu.Item key="password">
<IconLock style={{ marginRight: 8 }} />
</Menu.Item>
<Menu.Item key="two-factor">
<IconSafe style={{ marginRight: 8 }} />
</Menu.Item>
<Menu.Item key="logout">
<IconPoweroff style={{ marginRight: 8 }} />
退
</Menu.Item>
</Menu>
)
return (
<>
<Dropdown droplist={droplist} position="br">
<Button type="text" style={{ display: 'flex', alignItems: 'center', gap: 6 }}>
<Avatar size={28} style={{ backgroundColor: 'var(--color-primary-6)' }}>
{(user?.displayName ?? user?.username ?? '管')[0]}
</Avatar>
<span>{user?.displayName ?? user?.username ?? '管理员'}</span>
<span style={{ color: 'var(--color-text-3)', fontSize: 12 }}>
[{roleLabel(user?.role)}]
</span>
<IconDown />
</Button>
</Dropdown>
{passwordVisible ? (
<ChangePasswordModal username={user?.username} onClose={() => setPasswordVisible(false)} />
) : null}
{securityVisible ? (
<MfaSettingsModal user={user} onClose={() => setSecurityVisible(false)} />
) : null}
</>
)
}
@@ -0,0 +1,71 @@
import { Form, Input, Message, Modal } from '@arco-design/web-react'
import { useState } from 'react'
import {
changePassword,
clearTrustedDeviceToken,
type ChangePasswordPayload,
} from '../../services/auth'
import { resolveErrorMessage } from '../../utils/error'
interface ChangePasswordModalProps {
username?: string
onClose: () => void
}
export function ChangePasswordModal({ username, onClose }: ChangePasswordModalProps) {
const [loading, setLoading] = useState(false)
const [form] = Form.useForm<ChangePasswordPayload & { confirmPassword: string }>()
function close() {
form.resetFields()
onClose()
}
async function handleChangePassword() {
try {
const values = await form.validate()
if (values.newPassword !== values.confirmPassword) {
Message.error('两次输入的新密码不一致')
return
}
setLoading(true)
await changePassword({ oldPassword: values.oldPassword, newPassword: values.newPassword })
clearTrustedDeviceToken(username)
Message.success('密码修改成功')
close()
} catch (error) {
if (error) {
Message.error(resolveErrorMessage(error, '密码修改失败'))
}
} finally {
setLoading(false)
}
}
return (
<Modal
title="修改密码"
visible
onCancel={close}
onOk={handleChangePassword}
confirmLoading={loading}
unmountOnExit
>
<Form form={form} layout="vertical">
<Form.Item field="oldPassword" label="当前密码" rules={[{ required: true, minLength: 8 }]}>
<Input.Password placeholder="请输入当前密码" />
</Form.Item>
<Form.Item field="newPassword" label="新密码" rules={[{ required: true, minLength: 8 }]}>
<Input.Password placeholder="请输入新密码(至少 8 位)" />
</Form.Item>
<Form.Item
field="confirmPassword"
label="确认新密码"
rules={[{ required: true, minLength: 8 }]}
>
<Input.Password placeholder="请再次输入新密码" />
</Form.Item>
</Form>
</Modal>
)
}
@@ -0,0 +1,501 @@
import {
Alert,
Button,
Divider,
Form,
Input,
Message,
Modal,
Space,
Tag,
Typography,
} from '@arco-design/web-react'
import { useCallback, useEffect, useRef, useState } from 'react'
import {
beginWebAuthnRegistration,
clearTrustedDeviceToken,
configureOtp,
deleteWebAuthnCredential,
disableTwoFactor,
enableTwoFactor,
finishWebAuthnRegistration,
listTrustedDevices,
listWebAuthnCredentials,
prepareTwoFactor,
regenerateRecoveryCodes,
revokeTrustedDevice,
type TrustedDevice,
type TwoFactorSetupResult,
type UserInfo,
type WebAuthnCredential,
} from '../../services/auth'
import { useAuthStore } from '../../stores/auth'
import { resolveErrorMessage } from '../../utils/error'
import { createWebAuthnCredential } from '../../utils/webauthn'
interface MfaSettingsModalProps {
user: UserInfo | null
onClose: () => void
}
export function MfaSettingsModal({ user, onClose }: MfaSettingsModalProps) {
const [loading, setLoading] = useState(false)
const [setup, setSetup] = useState<TwoFactorSetupResult | null>(null)
const [recoveryCodes, setRecoveryCodes] = useState<string[]>([])
const [webAuthnCredentials, setWebAuthnCredentials] = useState<WebAuthnCredential[]>([])
const [trustedDevices, setTrustedDevices] = useState<TrustedDevice[]>([])
const [detailsLoading, setDetailsLoading] = useState(false)
const [form] = Form.useForm<{
currentPassword: string
code: string
email: string
phone: string
}>()
const setUser = useAuthStore((state) => state.setUser)
const initializedRef = useRef(false)
const loadSecurityDetails = useCallback(async () => {
setDetailsLoading(true)
try {
const [credentials, devices] = await Promise.all([
listWebAuthnCredentials(),
listTrustedDevices(),
])
setWebAuthnCredentials(credentials)
setTrustedDevices(devices)
} catch (error) {
Message.error(resolveErrorMessage(error, '加载安全配置失败'))
} finally {
setDetailsLoading(false)
}
}, [])
useEffect(() => {
if (initializedRef.current) {
return
}
initializedRef.current = true
form.setFieldValue('email', user?.email ?? '')
form.setFieldValue('phone', user?.phone ?? '')
void loadSecurityDetails()
}, [form, loadSecurityDetails, user?.email, user?.phone])
function applySecurityUserUpdate(updated: UserInfo) {
setUser(updated)
if (!updated.mfaEnabled) {
clearTrustedDeviceToken(updated.username)
}
}
async function copyRecoveryCodes() {
if (recoveryCodes.length === 0) return
try {
await navigator.clipboard.writeText(recoveryCodes.join('\n'))
Message.success('已复制到剪贴板')
} catch {
Message.info('请手动选择文本复制')
}
}
async function handleTwoFactorSetupAction() {
try {
const values = await form.validate()
setLoading(true)
if (!setup) {
const result = await prepareTwoFactor({ currentPassword: values.currentPassword })
setSetup(result)
Message.success('TOTP 密钥已生成')
return
}
const result = await enableTwoFactor({ code: values.code })
setUser(result.user)
setRecoveryCodes(result.recoveryCodes)
Message.success('TOTP 已启用')
} catch (error) {
if (error) {
Message.error(resolveErrorMessage(error, 'TOTP 操作失败'))
}
} finally {
setLoading(false)
}
}
async function handleRegenerateRecoveryCodes() {
try {
const values = await form.validate()
setLoading(true)
const result = await regenerateRecoveryCodes({
currentPassword: values.currentPassword,
code: values.code,
})
setUser(result.user)
setRecoveryCodes(result.recoveryCodes)
form.resetFields()
Message.success('恢复码已重新生成')
} catch (error) {
if (error) {
Message.error(resolveErrorMessage(error, '恢复码生成失败'))
}
} finally {
setLoading(false)
}
}
async function handleDisableTwoFactor() {
try {
const values = await form.validate()
setLoading(true)
const updated = await disableTwoFactor({
currentPassword: values.currentPassword,
code: values.code,
})
applySecurityUserUpdate(updated)
Message.success('TOTP 已关闭')
onClose()
} catch (error) {
if (error) {
Message.error(resolveErrorMessage(error, '关闭 TOTP 失败'))
}
} finally {
setLoading(false)
}
}
function readCurrentPassword() {
const currentPassword = String(form.getFieldValue('currentPassword') ?? '')
if (currentPassword.trim().length < 8) {
Message.error('请输入当前密码')
return ''
}
return currentPassword
}
async function handleRegisterWebAuthn() {
const currentPassword = readCurrentPassword()
if (!currentPassword) return
try {
setLoading(true)
const options = await beginWebAuthnRegistration({ currentPassword })
const credential = await createWebAuthnCredential(options)
const updated = await finishWebAuthnRegistration({
name: navigator.userAgent.slice(0, 120),
credential,
})
applySecurityUserUpdate(updated)
await loadSecurityDetails()
Message.success('通行密钥已注册')
} catch (error) {
Message.error(resolveErrorMessage(error, '通行密钥注册失败'))
} finally {
setLoading(false)
}
}
async function handleDeleteWebAuthnCredential(id: string) {
const currentPassword = readCurrentPassword()
if (!currentPassword) return
try {
setLoading(true)
const updated = await deleteWebAuthnCredential(id, { currentPassword })
applySecurityUserUpdate(updated)
await loadSecurityDetails()
Message.success('通行密钥已删除')
} catch (error) {
Message.error(resolveErrorMessage(error, '删除通行密钥失败'))
} finally {
setLoading(false)
}
}
async function handleConfigureOtp(channel: 'email' | 'sms', enabled: boolean) {
const currentPassword = readCurrentPassword()
if (!currentPassword) return
const email = String(form.getFieldValue('email') ?? '')
const phone = String(form.getFieldValue('phone') ?? '')
try {
setLoading(true)
const updated = await configureOtp({ currentPassword, channel, enabled, email, phone })
applySecurityUserUpdate(updated)
form.setFieldValue('email', updated.email ?? '')
form.setFieldValue('phone', updated.phone ?? '')
Message.success(enabled ? 'OTP 已启用' : 'OTP 已关闭')
} catch (error) {
Message.error(resolveErrorMessage(error, 'OTP 配置失败'))
} finally {
setLoading(false)
}
}
async function handleRevokeTrustedDevice(id: string) {
const currentPassword = readCurrentPassword()
if (!currentPassword) return
try {
setLoading(true)
await revokeTrustedDevice(id, { currentPassword })
clearTrustedDeviceToken(user?.username)
await loadSecurityDetails()
Message.success('可信设备已移除')
} catch (error) {
Message.error(resolveErrorMessage(error, '移除可信设备失败'))
} finally {
setLoading(false)
}
}
function renderFooter() {
if (recoveryCodes.length > 0) {
return (
<Space>
<Button onClick={() => void copyRecoveryCodes()}></Button>
<Button type="primary" onClick={onClose}>
</Button>
</Space>
)
}
if (user?.twoFactorEnabled) {
return (
<Space>
<Button onClick={onClose}></Button>
<Button loading={loading} onClick={() => void handleRegenerateRecoveryCodes()}>
</Button>
<Button status="danger" loading={loading} onClick={() => void handleDisableTwoFactor()}>
TOTP
</Button>
</Space>
)
}
return (
<Space>
<Button onClick={onClose}></Button>
<Button type="primary" loading={loading} onClick={() => void handleTwoFactorSetupAction()}>
{setup ? '启用 TOTP' : '生成 TOTP 二维码'}
</Button>
</Space>
)
}
return (
<Modal title="多因素认证" visible onCancel={onClose} footer={renderFooter()} unmountOnExit>
{recoveryCodes.length > 0 ? (
<Space direction="vertical" size="medium" style={{ width: '100%' }}>
<Alert
type="warning"
content="恢复码只会显示一次。请立即保存;每个恢复码只能使用一次。"
/>
<Input.TextArea value={recoveryCodes.join('\n')} autoSize readOnly />
</Space>
) : (
<Form form={form} layout="vertical">
{user?.twoFactorEnabled ? (
<>
<Alert
type="success"
content={`当前账号已启用 TOTP,恢复码剩余 ${user.twoFactorRecoveryCodesRemaining ?? 0} 个。`}
style={{ marginBottom: 16 }}
/>
<Form.Item
field="currentPassword"
label="当前密码"
rules={[{ required: true, minLength: 8 }]}
>
<Input.Password placeholder="请输入当前密码" />
</Form.Item>
<Form.Item
field="code"
label="TOTP 验证码"
rules={[{ required: true, minLength: 6, maxLength: 10 }]}
>
<Input placeholder="请输入 6 位验证码" maxLength={10} />
</Form.Item>
</>
) : (
<>
{!setup ? (
<>
<Alert
type="info"
content="启用前需要验证当前密码。"
style={{ marginBottom: 16 }}
/>
<Form.Item
field="currentPassword"
label="当前密码"
rules={[{ required: true, minLength: 8 }]}
>
<Input.Password placeholder="请输入当前密码" />
</Form.Item>
</>
) : (
<>
<Alert
type="warning"
content="密钥仅在本次启用流程中显示。启用后会生成一次性恢复码。"
style={{ marginBottom: 16 }}
/>
<div style={{ display: 'flex', gap: 20, alignItems: 'center', marginBottom: 16 }}>
<img
src={setup.qrCodeDataUrl}
alt="TOTP 二维码"
style={{
width: 160,
height: 160,
border: '1px solid var(--color-border)',
borderRadius: 8,
}}
/>
<Space direction="vertical" size={8} style={{ flex: 1, minWidth: 0 }}>
<Typography.Text type="secondary"></Typography.Text>
<Input value={setup.secret} readOnly />
</Space>
</div>
<Form.Item
field="code"
label="TOTP 验证码"
rules={[{ required: true, minLength: 6, maxLength: 10 }]}
>
<Input placeholder="请输入 6 位验证码" maxLength={10} />
</Form.Item>
</>
)}
</>
)}
<Divider />
<Space direction="vertical" size="medium" style={{ width: '100%' }}>
<Space style={{ justifyContent: 'space-between', width: '100%' }}>
<Typography.Title heading={6} style={{ margin: 0 }}>
</Typography.Title>
<Tag color={webAuthnCredentials.length > 0 ? 'green' : 'gray'} bordered>
{webAuthnCredentials.length > 0 ? `${webAuthnCredentials.length}` : '未注册'}
</Tag>
</Space>
<Typography.Paragraph type="secondary" style={{ margin: 0 }}>
Passkey
</Typography.Paragraph>
<Button loading={loading} onClick={() => void handleRegisterWebAuthn()}>
</Button>
<Space direction="vertical" size={8} style={{ width: '100%' }}>
{detailsLoading ? (
<Typography.Text type="secondary">...</Typography.Text>
) : null}
{webAuthnCredentials.map((item) => (
<div
key={item.id}
style={{
display: 'flex',
justifyContent: 'space-between',
gap: 12,
alignItems: 'center',
padding: '8px 0',
borderTop: '1px solid var(--color-border)',
}}
>
<Space direction="vertical" size={2}>
<Typography.Text>{item.name}</Typography.Text>
<Typography.Text type="secondary" style={{ fontSize: 12 }}>
{item.lastUsedAt ? `最近使用 ${item.lastUsedAt}` : `创建于 ${item.createdAt}`}
</Typography.Text>
</Space>
<Button
size="small"
status="danger"
onClick={() => void handleDeleteWebAuthnCredential(item.id)}
>
</Button>
</div>
))}
</Space>
</Space>
<Divider />
<Space direction="vertical" size="medium" style={{ width: '100%' }}>
<Typography.Title heading={6} style={{ margin: 0 }}>
/ OTP
</Typography.Title>
<Alert
type="info"
content="邮件 OTP 使用已启用的 Email 通知配置发送;短信 OTP 使用 Webhook 通知配置发送,payload 会包含 phone/code/purpose 字段。"
/>
<Space wrap>
<Tag color={user?.emailOtpEnabled ? 'green' : 'gray'} bordered>
OTP {user?.emailOtpEnabled ? '已启用' : '未启用'}
</Tag>
<Tag color={user?.smsOtpEnabled ? 'green' : 'gray'} bordered>
OTP {user?.smsOtpEnabled ? '已启用' : '未启用'}
</Tag>
</Space>
<Form.Item field="email" label="邮箱">
<Input placeholder="启用邮件 OTP 时填写" />
</Form.Item>
<Form.Item field="phone" label="手机号">
<Input placeholder="启用短信 OTP 时填写" />
</Form.Item>
<Space wrap>
<Button
loading={loading}
onClick={() => void handleConfigureOtp('email', !user?.emailOtpEnabled)}
>
{user?.emailOtpEnabled ? '关闭邮件 OTP' : '启用邮件 OTP'}
</Button>
<Button
loading={loading}
onClick={() => void handleConfigureOtp('sms', !user?.smsOtpEnabled)}
>
{user?.smsOtpEnabled ? '关闭短信 OTP' : '启用短信 OTP'}
</Button>
</Space>
</Space>
<Divider />
<Space direction="vertical" size="medium" style={{ width: '100%' }}>
<Space style={{ justifyContent: 'space-between', width: '100%' }}>
<Typography.Title heading={6} style={{ margin: 0 }}>
</Typography.Title>
<Tag color={trustedDevices.length > 0 ? 'green' : 'gray'} bordered>
{trustedDevices.length}
</Tag>
</Space>
<Typography.Paragraph type="secondary" style={{ margin: 0 }}>
30
</Typography.Paragraph>
<Space direction="vertical" size={8} style={{ width: '100%' }}>
{trustedDevices.map((item) => (
<div
key={item.id}
style={{
display: 'flex',
justifyContent: 'space-between',
gap: 12,
alignItems: 'center',
padding: '8px 0',
borderTop: '1px solid var(--color-border)',
}}
>
<Space direction="vertical" size={2}>
<Typography.Text>{item.name}</Typography.Text>
<Typography.Text type="secondary" style={{ fontSize: 12 }}>
使 {item.lastUsedAt || '-'} {item.expiresAt}
</Typography.Text>
</Space>
<Button
size="small"
status="danger"
onClick={() => void handleRevokeTrustedDevice(item.id)}
>
</Button>
</div>
))}
{!detailsLoading && trustedDevices.length === 0 ? (
<Typography.Text type="secondary"></Typography.Text>
) : null}
</Space>
</Space>
</Form>
)}
</Modal>
)
}
@@ -16,7 +16,14 @@ interface AdminDataSectionProps {
children: ReactNode children: ReactNode
} }
export function AdminDataSection({ title, description, actions, metrics, toolbar, children }: AdminDataSectionProps) { export function AdminDataSection({
title,
description,
actions,
metrics,
toolbar,
children,
}: AdminDataSectionProps) {
return ( return (
<section className="admin-section" aria-labelledby="admin-section-title"> <section className="admin-section" aria-labelledby="admin-section-title">
<header className="admin-section__header"> <header className="admin-section__header">
-58
View File
@@ -1,58 +0,0 @@
import { render, screen } from '@testing-library/react';
import { MemoryRouter, Route, Routes } from 'react-router-dom';
import { AuthGuard } from './auth-guard';
import { useAuthStore } from '../stores/auth';
function renderWithRoutes(initialEntry: string) {
return render(
<MemoryRouter initialEntries={[initialEntry]}>
<Routes>
<Route path="/login" element={<div>login-page</div>} />
<Route
path="/"
element={
<AuthGuard>
<div>protected-page</div>
</AuthGuard>
}
/>
</Routes>
</MemoryRouter>,
);
}
describe('AuthGuard', () => {
beforeEach(() => {
useAuthStore.setState({
token: null,
user: null,
hydrated: true,
status: 'anonymous',
});
});
it('redirects anonymous users to login page', async () => {
renderWithRoutes('/');
expect(await screen.findByText('login-page')).toBeInTheDocument();
});
it('renders children for authenticated users', async () => {
useAuthStore.setState({
token: 'token',
user: {
id: 1,
username: 'admin',
displayName: '管理员',
role: 'admin',
},
hydrated: true,
status: 'authenticated',
});
renderWithRoutes('/');
expect(await screen.findByText('protected-page')).toBeInTheDocument();
});
});
-25
View File
@@ -1,25 +0,0 @@
import { Spin } from '@arco-design/web-react';
import type { PropsWithChildren } from 'react';
import { Navigate, useLocation } from 'react-router-dom';
import { useAuthStore } from '../stores/auth';
export function AuthGuard({ children }: PropsWithChildren) {
const hydrated = useAuthStore((state) => state.hydrated);
const status = useAuthStore((state) => state.status);
const location = useLocation();
if (!hydrated || status === 'bootstrapping' || status === 'idle') {
return (
<div className="fullscreen-center">
<Spin tip="正在加载登录状态..." />
</div>
);
}
if (status !== 'authenticated') {
return <Navigate to="/login" replace state={{ from: location }} />;
}
return <>{children}</>;
}
@@ -15,7 +15,12 @@ interface BackupRecordContentsModalProps {
// BackupRecordContentsModal 浏览某次备份捕获的文件清单(只读)。 // BackupRecordContentsModal 浏览某次备份捕获的文件清单(只读)。
// 数据来源于全量备份记录的清单,无需下载归档,秒级展示并支持按路径筛选。 // 数据来源于全量备份记录的清单,无需下载归档,秒级展示并支持按路径筛选。
export function BackupRecordContentsModal({ visible, recordId, onClose, onRestoreSelected }: BackupRecordContentsModalProps) { export function BackupRecordContentsModal({
visible,
recordId,
onClose,
onRestoreSelected,
}: BackupRecordContentsModalProps) {
const [loading, setLoading] = useState(false) const [loading, setLoading] = useState(false)
const [error, setError] = useState('') const [error, setError] = useState('')
const [contents, setContents] = useState<BackupRecordContents | null>(null) const [contents, setContents] = useState<BackupRecordContents | null>(null)
@@ -63,7 +68,14 @@ export function BackupRecordContentsModal({ visible, recordId, onClose, onRestor
}, [contents, keyword]) }, [contents, keyword])
return ( return (
<Modal visible={visible} title="备份内容" footer={null} onCancel={onClose} unmountOnExit style={{ width: 760 }}> <Modal
visible={visible}
title="备份内容"
footer={null}
onCancel={onClose}
unmountOnExit
style={{ width: 760 }}
>
{loading ? ( {loading ? (
<Spin style={{ display: 'block', textAlign: 'center', padding: 40 }} /> <Spin style={{ display: 'block', textAlign: 'center', padding: 40 }} />
) : error ? ( ) : error ? (
@@ -76,9 +88,20 @@ export function BackupRecordContentsModal({ visible, recordId, onClose, onRestor
{contents.basedOnFull ? `;差异备份,清单取自基线全量 #${contents.basedOnFull}` : ''} {contents.basedOnFull ? `;差异备份,清单取自基线全量 #${contents.basedOnFull}` : ''}
</Typography.Text> </Typography.Text>
<div style={{ display: 'flex', gap: 8, alignItems: 'center', margin: '8px 0' }}> <div style={{ display: 'flex', gap: 8, alignItems: 'center', margin: '8px 0' }}>
<Input.Search allowClear placeholder="按路径筛选" value={keyword} onChange={setKeyword} style={{ flex: 1 }} /> <Input.Search
allowClear
placeholder="按路径筛选"
value={keyword}
onChange={setKeyword}
style={{ flex: 1 }}
/>
{onRestoreSelected && ( {onRestoreSelected && (
<Button type="primary" status="warning" disabled={selectedKeys.length === 0} onClick={() => onRestoreSelected(selectedKeys)}> <Button
type="primary"
status="warning"
disabled={selectedKeys.length === 0}
onClick={() => onRestoreSelected(selectedKeys)}
>
{selectedKeys.length} {selectedKeys.length}
</Button> </Button>
)} )}
@@ -89,7 +112,11 @@ export function BackupRecordContentsModal({ visible, recordId, onClose, onRestor
data={filtered} data={filtered}
rowSelection={ rowSelection={
onRestoreSelected onRestoreSelected
? { type: 'checkbox', selectedRowKeys: selectedKeys, onChange: (keys) => setSelectedKeys(keys as string[]) } ? {
type: 'checkbox',
selectedRowKeys: selectedKeys,
onChange: (keys) => setSelectedKeys(keys as string[]),
}
: undefined : undefined
} }
pagination={{ pageSize: 50, sizeCanChange: false }} pagination={{ pageSize: 50, sizeCanChange: false }}
@@ -114,7 +141,8 @@ export function BackupRecordContentsModal({ visible, recordId, onClose, onRestor
dataIndex: 'size', dataIndex: 'size',
width: 120, width: 120,
align: 'right', align: 'right',
render: (_: unknown, row: BackupRecordContentEntry) => (row.isDir ? '-' : formatBytes(row.size)), render: (_: unknown, row: BackupRecordContentEntry) =>
row.isDir ? '-' : formatBytes(row.size),
}, },
]} ]}
/> />
@@ -1,13 +1,33 @@
import { Alert, Button, Descriptions, Drawer, Message, Space, Spin, Tag, Typography } from '@arco-design/web-react' import {
Alert,
Button,
Descriptions,
Drawer,
Message,
Space,
Spin,
Tag,
Typography,
} from '@arco-design/web-react'
import { useEffect, useMemo, useState } from 'react' import { useEffect, useMemo, useState } from 'react'
import { useNavigate } from 'react-router-dom' import { useNavigate } from 'react-router-dom'
import { deleteBackupRecord, downloadBackupRecord, getBackupRecord, streamBackupRecordLogs } from '../../services/backup-records' import {
deleteBackupRecord,
downloadBackupRecord,
getBackupRecord,
streamBackupRecordLogs,
} from '../../services/backup-records'
import { getBackupTask } from '../../services/backup-tasks' import { getBackupTask } from '../../services/backup-tasks'
import { startRestoreFromBackup } from '../../services/restore-records' import { startRestoreFromBackup } from '../../services/restore-records'
import { startVerifyByRecord } from '../../services/verification-records' import { startVerifyByRecord } from '../../services/verification-records'
import { useAuthStore } from '../../stores/auth' import { useAuthStore } from '../../stores/auth'
import { canWrite } from '../../utils/permissions' import { canWrite } from '../../utils/permissions'
import type { BackupLogEvent, BackupRecordDetail, BackupRecordStatus, StorageUploadResultItem } from '../../types/backup-records' import type {
BackupLogEvent,
BackupRecordDetail,
BackupRecordStatus,
StorageUploadResultItem,
} from '../../types/backup-records'
import type { BackupTaskDetail } from '../../types/backup-tasks' import type { BackupTaskDetail } from '../../types/backup-tasks'
import { resolveErrorMessage } from '../../utils/error' import { resolveErrorMessage } from '../../utils/error'
import { formatBytes, formatDateTime, formatDuration } from '../../utils/format' import { formatBytes, formatDateTime, formatDuration } from '../../utils/format'
@@ -39,7 +59,12 @@ function buildLogText(record: BackupRecordDetail | null, events: BackupLogEvent[
return record?.logContent ?? '' return record?.logContent ?? ''
} }
export function BackupRecordLogDrawer({ visible, recordId, onCancel, onChanged }: BackupRecordLogDrawerProps) { export function BackupRecordLogDrawer({
visible,
recordId,
onCancel,
onChanged,
}: BackupRecordLogDrawerProps) {
const navigate = useNavigate() const navigate = useNavigate()
const currentUser = useAuthStore((state) => state.user) const currentUser = useAuthStore((state) => state.user)
const writable = canWrite(currentUser) const writable = canWrite(currentUser)
@@ -90,7 +115,9 @@ export function BackupRecordLogDrawer({ visible, recordId, onCancel, onChanged }
return [...current, event] return [...current, event]
}) })
if (event.completed) { if (event.completed) {
setRecord((current) => (current ? { ...current, status: event.status as BackupRecordStatus } : current)) setRecord((current) =>
current ? { ...current, status: event.status as BackupRecordStatus } : current,
)
} }
}, },
onDone: () => { onDone: () => {
@@ -218,7 +245,11 @@ export function BackupRecordLogDrawer({ visible, recordId, onCancel, onChanged }
if (!recordId || paths.length === 0) { if (!recordId || paths.length === 0) {
return return
} }
if (!window.confirm(`确定将选中的 ${paths.length} 项恢复到原位置吗?这会覆盖目标位置的现有文件,不可撤销。`)) { if (
!window.confirm(
`确定将选中的 ${paths.length} 项恢复到原位置吗?这会覆盖目标位置的现有文件,不可撤销。`,
)
) {
return return
} }
try { try {
@@ -268,10 +299,20 @@ export function BackupRecordLogDrawer({ visible, recordId, onCancel, onChanged }
<Space> <Space>
{record.status && ( {record.status && (
<Tag color={getStatusColor(record.status)} bordered> <Tag color={getStatusColor(record.status)} bordered>
{record.status === 'success' ? '成功' : record.status === 'failed' ? '失败' : record.status === 'running' ? '执行中' : record.status} {record.status === 'success'
? '成功'
: record.status === 'failed'
? '失败'
: record.status === 'running'
? '执行中'
: record.status}
</Tag>
)}
{record.storageTargetName && (
<Tag color="arcoblue" bordered>
{record.storageTargetName}
</Tag> </Tag>
)} )}
{record.storageTargetName && <Tag color="arcoblue" bordered>{record.storageTargetName}</Tag>}
</Space> </Space>
</div> </div>
<Descriptions <Descriptions
@@ -280,10 +321,17 @@ export function BackupRecordLogDrawer({ visible, recordId, onCancel, onChanged }
{ label: '文件名', value: record.fileName || '-' }, { label: '文件名', value: record.fileName || '-' },
{ label: '文件大小', value: formatBytes(record.fileSize) }, { label: '文件大小', value: formatBytes(record.fileSize) },
{ label: '存储路径', value: record.storagePath || '-' }, { label: '存储路径', value: record.storagePath || '-' },
...(record.storageTransferMode ? [{ ...(record.storageTransferMode
label: '传输路径', ? [
value: record.storageTransferMode === 'master_relay' ? 'Master 流式中转' : 'Agent 直传', {
}] : []), label: '传输路径',
value:
record.storageTransferMode === 'master_relay'
? 'Master 流式中转'
: 'Agent 直传',
},
]
: []),
{ label: '开始时间', value: formatDateTime(record.startedAt) }, { label: '开始时间', value: formatDateTime(record.startedAt) },
{ label: '完成时间', value: formatDateTime(record.completedAt) }, { label: '完成时间', value: formatDateTime(record.completedAt) },
{ label: '耗时', value: formatDuration(record.durationSeconds) }, { label: '耗时', value: formatDuration(record.durationSeconds) },
@@ -320,20 +368,23 @@ export function BackupRecordLogDrawer({ visible, recordId, onCancel, onChanged }
</Button> </Button>
)} )}
</Space> </Space>
{record.storageUploadResults && (record.storageUploadResults.length > 1 || record.storageUploadResults.some((result) => result.transferMode)) && ( {record.storageUploadResults &&
<div> (record.storageUploadResults.length > 1 ||
<Typography.Title heading={6}></Typography.Title> record.storageUploadResults.some((result) => result.transferMode)) && (
<Descriptions <div>
column={1} <Typography.Title heading={6}></Typography.Title>
data={record.storageUploadResults.map((r: StorageUploadResultItem) => ({ <Descriptions
label: r.storageTargetName, column={1}
value: r.status === 'success' data={record.storageUploadResults.map((r: StorageUploadResultItem) => ({
? `上传成功${r.transferMode === 'master_relay' ? ' · Master 流式中转' : r.transferMode === 'direct' ? ' · Agent 直传' : ''}` label: r.storageTargetName,
: `上传失败: ${r.error || '未知错误'}`, value:
}))} r.status === 'success'
/> ? `上传成功${r.transferMode === 'master_relay' ? ' · Master 流式中转' : r.transferMode === 'direct' ? ' · Agent 直传' : ''}`
</div> : `上传失败: ${r.error || '未知错误'}`,
)} }))}
/>
</div>
)}
<div> <div>
<Typography.Title heading={6}></Typography.Title> <Typography.Title heading={6}></Typography.Title>
@@ -357,7 +408,11 @@ export function BackupRecordLogDrawer({ visible, recordId, onCancel, onChanged }
visible={contentsVisible} visible={contentsVisible}
recordId={recordId} recordId={recordId}
onClose={() => setContentsVisible(false)} onClose={() => setContentsVisible(false)}
onRestoreSelected={writable && record?.status === 'success' ? (paths) => void handleSelectiveRestore(paths) : undefined} onRestoreSelected={
writable && record?.status === 'success'
? (paths) => void handleSelectiveRestore(paths)
: undefined
}
/> />
</Drawer> </Drawer>
) )

Some files were not shown because too many files have changed in this diff Show More