mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-03 01:06:36 +08:00
fix(backup): restore backup export/import APIs and route compatibility
This commit is contained in:
@@ -0,0 +1,368 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/store"
|
||||
)
|
||||
|
||||
type backupPayload struct {
|
||||
Version int `json:"version"`
|
||||
ExportedAt int64 `json:"exportedAt"`
|
||||
Dialect string `json:"dialect"`
|
||||
Tables map[string][]map[string]any `json:"tables"`
|
||||
}
|
||||
|
||||
func (h *Handler) backupExport(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if h == nil || h.repo == nil || h.repo.DB() == nil {
|
||||
response.WriteJSON(w, response.Err(-2, "database unavailable"))
|
||||
return
|
||||
}
|
||||
|
||||
db := h.repo.DB()
|
||||
tableNames, err := listBackupTables(db)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
tables := make(map[string][]map[string]any, len(tableNames))
|
||||
for _, tableName := range tableNames {
|
||||
rows, err := dumpTableRows(db, tableName)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
tables[tableName] = rows
|
||||
}
|
||||
|
||||
payload := backupPayload{
|
||||
Version: 1,
|
||||
ExportedAt: time.Now().UnixMilli(),
|
||||
Dialect: db.Dialect().String(),
|
||||
Tables: tables,
|
||||
}
|
||||
|
||||
data, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, "备份导出失败"))
|
||||
return
|
||||
}
|
||||
|
||||
fileName := fmt.Sprintf("flvx-backup-%s.json", time.Now().Format("20060102-150405"))
|
||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
w.Header().Set("Content-Disposition", fmt.Sprintf("attachment; filename=%q", fileName))
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write(data)
|
||||
}
|
||||
|
||||
func (h *Handler) backupImport(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if h == nil || h.repo == nil || h.repo.DB() == nil {
|
||||
response.WriteJSON(w, response.Err(-2, "database unavailable"))
|
||||
return
|
||||
}
|
||||
|
||||
raw, err := readBackupImportBody(r)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
var payload backupPayload
|
||||
if err := json.Unmarshal(raw, &payload); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("备份文件格式错误"))
|
||||
return
|
||||
}
|
||||
if len(payload.Tables) == 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("备份数据为空"))
|
||||
return
|
||||
}
|
||||
|
||||
db := h.repo.DB()
|
||||
tableNames, err := listBackupTables(db)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
tableSet := make(map[string]struct{}, len(tableNames))
|
||||
for _, tableName := range tableNames {
|
||||
tableSet[tableName] = struct{}{}
|
||||
}
|
||||
for tableName := range payload.Tables {
|
||||
if _, ok := tableSet[tableName]; !ok {
|
||||
response.WriteJSON(w, response.ErrDefault("备份文件包含未知数据表"))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
tx, err := db.Begin()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
defer tx.Rollback()
|
||||
|
||||
for i := len(tableNames) - 1; i >= 0; i-- {
|
||||
tableName := tableNames[i]
|
||||
if _, err := tx.Exec(fmt.Sprintf("DELETE FROM %s", quoteIdentifier(tableName))); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
inserted := 0
|
||||
for _, tableName := range tableNames {
|
||||
rows := payload.Tables[tableName]
|
||||
if len(rows) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
columns, err := tableColumns(db, tableName)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
for _, row := range rows {
|
||||
insertCols := make([]string, 0, len(columns))
|
||||
insertVals := make([]any, 0, len(columns))
|
||||
for _, col := range columns {
|
||||
val, ok := row[col]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
insertCols = append(insertCols, quoteIdentifier(col))
|
||||
insertVals = append(insertVals, val)
|
||||
}
|
||||
if len(insertCols) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
placeholders := make([]string, len(insertCols))
|
||||
for i := range placeholders {
|
||||
placeholders[i] = "?"
|
||||
}
|
||||
|
||||
query := fmt.Sprintf(
|
||||
"INSERT INTO %s (%s) VALUES (%s)",
|
||||
quoteIdentifier(tableName),
|
||||
strings.Join(insertCols, ","),
|
||||
strings.Join(placeholders, ","),
|
||||
)
|
||||
|
||||
if _, err := tx.Exec(query, insertVals...); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
inserted++
|
||||
}
|
||||
}
|
||||
|
||||
if err := tx.Commit(); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(map[string]any{
|
||||
"tables": len(tableNames),
|
||||
"inserted": inserted,
|
||||
}))
|
||||
}
|
||||
|
||||
func readBackupImportBody(r *http.Request) ([]byte, error) {
|
||||
contentType := strings.ToLower(strings.TrimSpace(r.Header.Get("Content-Type")))
|
||||
if strings.HasPrefix(contentType, "multipart/form-data") {
|
||||
if err := r.ParseMultipartForm(32 << 20); err != nil {
|
||||
return nil, fmt.Errorf("读取上传文件失败")
|
||||
}
|
||||
|
||||
if r.MultipartForm != nil {
|
||||
preferredFields := []string{"file", "backup", "data"}
|
||||
for _, field := range preferredFields {
|
||||
if files := r.MultipartForm.File[field]; len(files) > 0 {
|
||||
return readMultipartFile(files[0])
|
||||
}
|
||||
}
|
||||
for _, files := range r.MultipartForm.File {
|
||||
if len(files) > 0 {
|
||||
return readMultipartFile(files[0])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if data := strings.TrimSpace(r.FormValue("data")); data != "" {
|
||||
return []byte(data), nil
|
||||
}
|
||||
return nil, fmt.Errorf("未找到备份文件")
|
||||
}
|
||||
|
||||
defer r.Body.Close()
|
||||
body, err := io.ReadAll(io.LimitReader(r.Body, 32<<20))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("读取请求数据失败")
|
||||
}
|
||||
body = bytes.TrimSpace(body)
|
||||
if len(body) == 0 {
|
||||
return nil, fmt.Errorf("备份数据不能为空")
|
||||
}
|
||||
return body, nil
|
||||
}
|
||||
|
||||
func readMultipartFile(header *multipart.FileHeader) ([]byte, error) {
|
||||
if header == nil {
|
||||
return nil, fmt.Errorf("未找到备份文件")
|
||||
}
|
||||
file, err := header.Open()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("读取上传文件失败")
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
data, err := io.ReadAll(io.LimitReader(file, 32<<20))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("读取上传文件失败")
|
||||
}
|
||||
data = bytes.TrimSpace(data)
|
||||
if len(data) == 0 {
|
||||
return nil, fmt.Errorf("备份数据不能为空")
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
|
||||
func listBackupTables(db *store.DB) ([]string, error) {
|
||||
query := `
|
||||
SELECT name
|
||||
FROM sqlite_master
|
||||
WHERE type = 'table' AND name NOT LIKE 'sqlite_%'
|
||||
ORDER BY name
|
||||
`
|
||||
if db.Dialect() == store.DialectPostgres {
|
||||
query = `
|
||||
SELECT table_name
|
||||
FROM information_schema.tables
|
||||
WHERE table_schema = 'public' AND table_type = 'BASE TABLE'
|
||||
ORDER BY table_name
|
||||
`
|
||||
}
|
||||
|
||||
rows, err := db.Query(query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
out := make([]string, 0)
|
||||
for rows.Next() {
|
||||
var tableName string
|
||||
if err := rows.Scan(&tableName); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if isSafeIdentifier(tableName) {
|
||||
out = append(out, tableName)
|
||||
}
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func dumpTableRows(db *store.DB, tableName string) ([]map[string]any, error) {
|
||||
rows, err := db.Query(fmt.Sprintf("SELECT * FROM %s", quoteIdentifier(tableName)))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
columns, err := rows.Columns()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
items := make([]map[string]any, 0)
|
||||
for rows.Next() {
|
||||
vals := make([]any, len(columns))
|
||||
ptrs := make([]any, len(columns))
|
||||
for i := range vals {
|
||||
ptrs[i] = &vals[i]
|
||||
}
|
||||
if err := rows.Scan(ptrs...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
row := make(map[string]any, len(columns))
|
||||
for i, col := range columns {
|
||||
row[col] = normalizeExportedValue(vals[i])
|
||||
}
|
||||
items = append(items, row)
|
||||
}
|
||||
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func tableColumns(db *store.DB, tableName string) ([]string, error) {
|
||||
rows, err := db.Query(fmt.Sprintf("SELECT * FROM %s LIMIT 0", quoteIdentifier(tableName)))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
columns, err := rows.Columns()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sort.Strings(columns)
|
||||
return columns, nil
|
||||
}
|
||||
|
||||
func normalizeExportedValue(v any) any {
|
||||
switch t := v.(type) {
|
||||
case nil:
|
||||
return nil
|
||||
case []byte:
|
||||
return string(t)
|
||||
case time.Time:
|
||||
return t.UTC().Format(time.RFC3339Nano)
|
||||
default:
|
||||
return t
|
||||
}
|
||||
}
|
||||
|
||||
func quoteIdentifier(name string) string {
|
||||
if !isSafeIdentifier(name) {
|
||||
return "\"\""
|
||||
}
|
||||
return fmt.Sprintf("\"%s\"", name)
|
||||
}
|
||||
|
||||
func isSafeIdentifier(name string) bool {
|
||||
if name == "" {
|
||||
return false
|
||||
}
|
||||
for i := 0; i < len(name); i++ {
|
||||
ch := name[i]
|
||||
if (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z') || (ch >= '0' && ch <= '9') || ch == '_' {
|
||||
continue
|
||||
}
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -93,6 +93,12 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/config/list", h.getConfigs)
|
||||
mux.HandleFunc("/api/v1/config/update", h.updateConfigs)
|
||||
mux.HandleFunc("/api/v1/config/update-single", h.updateSingleConfig)
|
||||
mux.HandleFunc("/api/v1/backup/export", h.backupExport)
|
||||
mux.HandleFunc("/api/v1/backup/import", h.backupImport)
|
||||
mux.HandleFunc("/api/v1/backup/restore", h.backupImport)
|
||||
mux.HandleFunc("/api/v1/api/v1/backup/export", h.backupExport)
|
||||
mux.HandleFunc("/api/v1/api/v1/backup/import", h.backupImport)
|
||||
mux.HandleFunc("/api/v1/api/v1/backup/restore", h.backupImport)
|
||||
mux.HandleFunc("/api/v1/captcha/check", h.checkCaptcha)
|
||||
mux.HandleFunc("/api/v1/captcha/verify", h.captchaVerify)
|
||||
mux.HandleFunc("/api/v1/user/package", h.userPackage)
|
||||
|
||||
@@ -115,6 +115,14 @@ func requiresAdmin(path string) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
if strings.HasPrefix(path, "/api/v1/backup/") {
|
||||
return true
|
||||
}
|
||||
|
||||
if strings.HasPrefix(path, "/api/v1/api/v1/backup/") {
|
||||
return true
|
||||
}
|
||||
|
||||
if strings.HasPrefix(path, "/api/v1/tunnel/") {
|
||||
if strings.HasPrefix(path, "/api/v1/tunnel/user/tunnel") {
|
||||
return false
|
||||
|
||||
@@ -203,6 +203,142 @@ func TestSpeedLimitTunnelsRouteAlias(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestBackupExportImportRestoreContracts(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
userToken, err := auth.GenerateToken(2, "normal_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate user token: %v", err)
|
||||
}
|
||||
|
||||
key := "backup_contract_key"
|
||||
if _, err := repo.DB().Exec(`
|
||||
INSERT INTO vite_config(name, value, time)
|
||||
VALUES(?, ?, ?)
|
||||
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
|
||||
`, key, "v1", time.Now().UnixMilli()); err != nil {
|
||||
t.Fatalf("seed config for backup contract: %v", err)
|
||||
}
|
||||
|
||||
t.Run("non-admin is blocked on backup export", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/backup/export", nil)
|
||||
req.Header.Set("Authorization", userToken)
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
assertCodeMsg(t, resp, 403, "权限不足,仅管理员可操作")
|
||||
})
|
||||
|
||||
t.Run("standard and duplicate export routes both work", func(t *testing.T) {
|
||||
payloadA := exportBackupPayload(t, router, "/api/v1/backup/export", adminToken)
|
||||
if len(payloadA.Tables) == 0 {
|
||||
t.Fatalf("expected exported tables, got none")
|
||||
}
|
||||
if _, ok := payloadA.Tables["vite_config"]; !ok {
|
||||
t.Fatalf("expected vite_config in exported tables")
|
||||
}
|
||||
|
||||
payloadB := exportBackupPayload(t, router, "/api/v1/api/v1/backup/export", adminToken)
|
||||
if len(payloadB.Tables) == 0 {
|
||||
t.Fatalf("expected exported tables from duplicate-prefix route, got none")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("backup import applies exported data", func(t *testing.T) {
|
||||
payload := exportBackupPayload(t, router, "/api/v1/backup/export", adminToken)
|
||||
setBackupConfigValue(t, payload.Tables["vite_config"], key, "v2")
|
||||
raw, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal import payload: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/backup/import", bytes.NewReader(raw))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
assertCode(t, resp, 0)
|
||||
|
||||
cfg, err := repo.GetConfigByName(key)
|
||||
if err != nil {
|
||||
t.Fatalf("query imported config: %v", err)
|
||||
}
|
||||
if cfg == nil || cfg.Value != "v2" {
|
||||
t.Fatalf("expected imported config value v2, got %+v", cfg)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("backup restore alias applies exported data", func(t *testing.T) {
|
||||
payload := exportBackupPayload(t, router, "/api/v1/backup/export", adminToken)
|
||||
setBackupConfigValue(t, payload.Tables["vite_config"], key, "v3")
|
||||
raw, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal restore payload: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/backup/restore", bytes.NewReader(raw))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
assertCode(t, resp, 0)
|
||||
|
||||
cfg, err := repo.GetConfigByName(key)
|
||||
if err != nil {
|
||||
t.Fatalf("query restored config: %v", err)
|
||||
}
|
||||
if cfg == nil || cfg.Value != "v3" {
|
||||
t.Fatalf("expected restored config value v3, got %+v", cfg)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
type backupExportPayload struct {
|
||||
Version int `json:"version"`
|
||||
ExportedAt int64 `json:"exportedAt"`
|
||||
Dialect string `json:"dialect"`
|
||||
Tables map[string][]map[string]any `json:"tables"`
|
||||
}
|
||||
|
||||
func exportBackupPayload(t *testing.T, router http.Handler, path, token string) backupExportPayload {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodPost, path, nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
if resp.Code != http.StatusOK {
|
||||
t.Fatalf("expected status 200 on %s, got %d", path, resp.Code)
|
||||
}
|
||||
|
||||
var payload backupExportPayload
|
||||
if err := json.NewDecoder(resp.Body).Decode(&payload); err != nil {
|
||||
t.Fatalf("decode backup payload from %s: %v", path, err)
|
||||
}
|
||||
if payload.Version != 1 {
|
||||
t.Fatalf("expected backup payload version 1, got %d", payload.Version)
|
||||
}
|
||||
return payload
|
||||
}
|
||||
|
||||
func setBackupConfigValue(t *testing.T, rows []map[string]any, key, value string) {
|
||||
t.Helper()
|
||||
for _, row := range rows {
|
||||
if strings.TrimSpace(valueAsString(row["name"])) == key {
|
||||
row["value"] = value
|
||||
return
|
||||
}
|
||||
}
|
||||
t.Fatalf("did not find config row %q in backup payload", key)
|
||||
}
|
||||
|
||||
func setupContractRouter(t *testing.T, jwtSecret string) (http.Handler, *sqlite.Repository) {
|
||||
t.Helper()
|
||||
dbPath := filepath.Join(t.TempDir(), "contract.db")
|
||||
|
||||
@@ -130,6 +130,10 @@ export const updateConfigs = (configMap: Record<string, string>) =>
|
||||
export const updateConfig = (name: string, value: string) =>
|
||||
Network.post("/config/update-single", { name, value });
|
||||
|
||||
export const exportBackupData = () => Network.post("/backup/export");
|
||||
export const importBackupData = (data: any) => Network.post("/backup/import", data);
|
||||
export const restoreBackupData = (data: any) => Network.post("/backup/restore", data);
|
||||
|
||||
// 验证码相关接口
|
||||
export const checkCaptcha = () => Network.post("/captcha/check");
|
||||
export const generateCaptcha = () => Network.post(`/captcha/generate`);
|
||||
|
||||
Reference in New Issue
Block a user