From 5a9715eb2639a8e219189148975de7d480c1f155 Mon Sep 17 00:00:00 2001 From: sagit Date: Fri, 13 Feb 2026 08:17:34 +0000 Subject: [PATCH 1/2] fix(backup): restore backup export/import APIs and route compatibility --- go-backend/internal/http/handler/backup.go | 368 ++++++++++++++++++ go-backend/internal/http/handler/handler.go | 6 + go-backend/internal/http/middleware/auth.go | 8 + .../tests/contract/migration_contract_test.go | 136 +++++++ vite-frontend/src/api/index.ts | 4 + 5 files changed, 522 insertions(+) create mode 100644 go-backend/internal/http/handler/backup.go diff --git a/go-backend/internal/http/handler/backup.go b/go-backend/internal/http/handler/backup.go new file mode 100644 index 0000000..935df7d --- /dev/null +++ b/go-backend/internal/http/handler/backup.go @@ -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 +} diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index 204613f..a8069ef 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -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) diff --git a/go-backend/internal/http/middleware/auth.go b/go-backend/internal/http/middleware/auth.go index bf3aebf..178ed71 100644 --- a/go-backend/internal/http/middleware/auth.go +++ b/go-backend/internal/http/middleware/auth.go @@ -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 diff --git a/go-backend/tests/contract/migration_contract_test.go b/go-backend/tests/contract/migration_contract_test.go index 89780cc..033296f 100644 --- a/go-backend/tests/contract/migration_contract_test.go +++ b/go-backend/tests/contract/migration_contract_test.go @@ -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") diff --git a/vite-frontend/src/api/index.ts b/vite-frontend/src/api/index.ts index 846283d..d09e257 100644 --- a/vite-frontend/src/api/index.ts +++ b/vite-frontend/src/api/index.ts @@ -130,6 +130,10 @@ export const updateConfigs = (configMap: Record) => 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`); From c049ceaacf7da57e5bafaf5afa9f1e0cd370a212 Mon Sep 17 00:00:00 2001 From: sagit Date: Fri, 13 Feb 2026 08:31:33 +0000 Subject: [PATCH 2/2] fix(backend): resolve backup handler build conflict after main merge --- go-backend/internal/http/handler/backup.go | 368 ------------------ go-backend/internal/http/handler/handler.go | 3 - .../internal/store/sqlite/repository.go | 5 +- .../tests/contract/migration_contract_test.go | 80 ++-- 4 files changed, 52 insertions(+), 404 deletions(-) delete mode 100644 go-backend/internal/http/handler/backup.go diff --git a/go-backend/internal/http/handler/backup.go b/go-backend/internal/http/handler/backup.go deleted file mode 100644 index 935df7d..0000000 --- a/go-backend/internal/http/handler/backup.go +++ /dev/null @@ -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 -} diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index 2e4033d..767dce3 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -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) diff --git a/go-backend/internal/store/sqlite/repository.go b/go-backend/internal/store/sqlite/repository.go index d72a6ec..c205df3 100644 --- a/go-backend/internal/store/sqlite/repository.go +++ b/go-backend/internal/store/sqlite/repository.go @@ -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 { diff --git a/go-backend/tests/contract/migration_contract_test.go b/go-backend/tests/contract/migration_contract_test.go index 033296f..4bc80a3 100644 --- a/go-backend/tests/contract/migration_contract_test.go +++ b/go-backend/tests/contract/migration_contract_test.go @@ -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")