mirror of
https://github.com/Awuqing/BackupX.git
synced 2026-09-07 16:36:42 +08:00
Merge pull request #137 from Awuqing/codex/architecture-simplification
refactor: simplify architecture and harden lifecycle
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
+7
-7
@@ -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,并通过仓库中的 ESLint 和 Prettier 检查
|
||||||
- **Commit 粒度** — 每个 commit 一件事,不要把顺手的小修改和功能代码混在一起
|
- **Commit 粒度** — 每个 commit 一件事,不要把顺手的小修改和功能代码混在一起
|
||||||
|
|||||||
Vendored
BIN
Binary file not shown.
+3
-8
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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))
|
||||||
}
|
}
|
||||||
|
|
||||||
// 辅助函数
|
// 辅助函数
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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。
|
||||||
|
|||||||
@@ -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]
|
||||||
|
|||||||
@@ -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
@@ -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 {
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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
|
|
||||||
}
|
|
||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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),
|
||||||
|
|||||||
@@ -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)
|
|
||||||
}
|
|
||||||
@@ -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
|
|
||||||
}
|
|
||||||
@@ -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
|
|
||||||
}
|
|
||||||
@@ -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
|
|
||||||
}
|
|
||||||
@@ -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
|
|
||||||
}
|
|
||||||
@@ -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())
|
|
||||||
}
|
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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))
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -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
|
|
||||||
}
|
|
||||||
@@ -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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -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)
|
|
||||||
}
|
|
||||||
@@ -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
|
|
||||||
}
|
|
||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 是可选能力接口,支持查询远端存储空间。
|
||||||
|
|||||||
@@ -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...)
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,4 @@
|
|||||||
|
coverage
|
||||||
|
dist
|
||||||
|
node_modules
|
||||||
|
package-lock.json
|
||||||
@@ -0,0 +1,6 @@
|
|||||||
|
{
|
||||||
|
"semi": false,
|
||||||
|
"singleQuote": true,
|
||||||
|
"trailingComma": "all",
|
||||||
|
"printWidth": 100
|
||||||
|
}
|
||||||
@@ -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,
|
||||||
|
)
|
||||||
Generated
+1235
-22
File diff suppressed because it is too large
Load Diff
+11
-1
@@ -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"
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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>
|
|
||||||
)
|
|
||||||
}
|
|
||||||
@@ -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">
|
||||||
|
|||||||
@@ -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();
|
|
||||||
});
|
|
||||||
});
|
|
||||||
@@ -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
Reference in New Issue
Block a user