fix(backend): resolve backup handler build conflict after main merge

This commit is contained in:
sagit
2026-02-13 08:31:33 +00:00
parent 3424221176
commit c049ceaacf
4 changed files with 52 additions and 404 deletions
-368
View File
@@ -1,368 +0,0 @@
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
}
@@ -178,9 +178,6 @@ func (h *Handler) Register(mux *http.ServeMux) {
mux.HandleFunc("/api/v1/federation/runtime/command", h.authPeer(h.federationRuntimeCommand))
mux.HandleFunc("/api/v1/federation/node/import", h.nodeImport)
mux.HandleFunc("/api/v1/backup/export", h.backupExport)
mux.HandleFunc("/api/v1/backup/import", h.backupImport)
mux.HandleFunc("/flow/test", h.flowTest)
mux.HandleFunc("/flow/config", h.flowConfig)
mux.HandleFunc("/flow/upload", h.flowUpload)
@@ -2484,7 +2484,7 @@ func (r *Repository) exportUserGroups() ([]UserGroupBackup, error) {
func (r *Repository) exportPermissions() ([]PermissionBackup, error) {
rows, err := r.db.Query(`
SELECT id, user_group_id, tunnel_group_id, created_time, created_by_group
SELECT id, user_group_id, tunnel_group_id, created_time
FROM group_permission ORDER BY id ASC
`)
if err != nil {
@@ -2495,9 +2495,10 @@ func (r *Repository) exportPermissions() ([]PermissionBackup, error) {
var permissions []PermissionBackup
for rows.Next() {
var p PermissionBackup
if err := rows.Scan(&p.ID, &p.UserGroupID, &p.TunnelGroupID, &p.CreatedTime, &p.CreatedByGroup); err != nil {
if err := rows.Scan(&p.ID, &p.UserGroupID, &p.TunnelGroupID, &p.CreatedTime); err != nil {
return nil, err
}
p.CreatedByGroup = 0
// Get grants for this permission
grantRows, err := r.db.Query(`SELECT id, user_group_id, tunnel_group_id, user_tunnel_id, created_time, created_by_group FROM group_permission_grant WHERE user_group_id = ? AND tunnel_group_id = ?`, p.UserGroupID, p.TunnelGroupID)
if err != nil {
@@ -236,23 +236,23 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
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 len(payloadA.Configs) == 0 {
t.Fatalf("expected exported configs, got none")
}
if _, ok := payloadA.Tables["vite_config"]; !ok {
t.Fatalf("expected vite_config in exported tables")
if _, ok := payloadA.Configs[key]; !ok {
t.Fatalf("expected %q in exported configs", key)
}
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")
if len(payloadB.Configs) == 0 {
t.Fatalf("expected exported configs 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)
payload.Configs[key] = "v2"
raw, err := json.Marshal(backupImportPayload{Types: []string{"configs"}, backupExportPayload: payload})
if err != nil {
t.Fatalf("marshal import payload: %v", err)
}
@@ -263,7 +263,13 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assertCode(t, resp, 0)
var out response.R
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
t.Fatalf("decode import response: %v", err)
}
if out.Code != 0 {
t.Fatalf("expected import code 0, got %d (%s)", out.Code, out.Msg)
}
cfg, err := repo.GetConfigByName(key)
if err != nil {
@@ -276,8 +282,8 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
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)
payload.Configs[key] = "v3"
raw, err := json.Marshal(backupImportPayload{Types: []string{"configs"}, backupExportPayload: payload})
if err != nil {
t.Fatalf("marshal restore payload: %v", err)
}
@@ -288,7 +294,13 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assertCode(t, resp, 0)
var out response.R
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
t.Fatalf("decode restore response: %v", err)
}
if out.Code != 0 {
t.Fatalf("expected restore code 0, got %d (%s)", out.Code, out.Msg)
}
cfg, err := repo.GetConfigByName(key)
if err != nil {
@@ -301,16 +313,21 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
}
type backupExportPayload struct {
Version int `json:"version"`
ExportedAt int64 `json:"exportedAt"`
Dialect string `json:"dialect"`
Tables map[string][]map[string]any `json:"tables"`
Version string `json:"version"`
ExportedAt int64 `json:"exportedAt"`
Configs map[string]string `json:"configs"`
}
type backupImportPayload struct {
Types []string `json:"types"`
backupExportPayload
}
func exportBackupPayload(t *testing.T, router http.Handler, path, token string) backupExportPayload {
t.Helper()
req := httptest.NewRequest(http.MethodPost, path, nil)
req := httptest.NewRequest(http.MethodPost, path, bytes.NewBufferString(`{"types":["configs"]}`))
req.Header.Set("Authorization", token)
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
@@ -318,27 +335,28 @@ func exportBackupPayload(t *testing.T, router http.Handler, path, token string)
t.Fatalf("expected status 200 on %s, got %d", path, resp.Code)
}
body, err := io.ReadAll(resp.Body)
if err != nil {
t.Fatalf("read backup payload from %s: %v", path, err)
}
var payload backupExportPayload
if err := json.NewDecoder(resp.Body).Decode(&payload); err != nil {
if err := json.Unmarshal(body, &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)
if strings.TrimSpace(payload.Version) == "" {
var out response.R
if err := json.Unmarshal(body, &out); err == nil {
t.Fatalf("expected backup payload on %s, got envelope code=%d msg=%q", path, out.Code, out.Msg)
}
t.Fatalf("expected non-empty backup payload version on %s, body=%s", path, string(body))
}
if payload.Configs == nil {
t.Fatalf("expected configs map in backup payload on %s", path)
}
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")