diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index 03530d5..767dce3 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) @@ -172,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/http/middleware/auth.go b/go-backend/internal/http/middleware/auth.go index 7e60afc..00eb5a4 100644 --- a/go-backend/internal/http/middleware/auth.go +++ b/go-backend/internal/http/middleware/auth.go @@ -117,6 +117,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/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 89780cc..4bc80a3 100644 --- a/go-backend/tests/contract/migration_contract_test.go +++ b/go-backend/tests/contract/migration_contract_test.go @@ -203,6 +203,160 @@ 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.Configs) == 0 { + t.Fatalf("expected exported configs, got none") + } + 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.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) + payload.Configs[key] = "v2" + raw, err := json.Marshal(backupImportPayload{Types: []string{"configs"}, backupExportPayload: 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) + 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 { + 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) + payload.Configs[key] = "v3" + raw, err := json.Marshal(backupImportPayload{Types: []string{"configs"}, backupExportPayload: 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) + 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 { + 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 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, bytes.NewBufferString(`{"types":["configs"]}`)) + req.Header.Set("Authorization", token) + req.Header.Set("Content-Type", "application/json") + resp := httptest.NewRecorder() + + router.ServeHTTP(resp, req) + if resp.Code != http.StatusOK { + 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.Unmarshal(body, &payload); err != nil { + t.Fatalf("decode backup payload from %s: %v", path, err) + } + 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 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 7abb7f7..d183fd2 100644 --- a/vite-frontend/src/api/index.ts +++ b/vite-frontend/src/api/index.ts @@ -140,6 +140,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`);