mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-03 09:06:36 +08:00
fix(backend): include forward_port in backup export/import
This commit is contained in:
@@ -2012,19 +2012,25 @@ type ChainTunnelBackup struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type ForwardBackup struct {
|
type ForwardBackup struct {
|
||||||
ID int64 `json:"id"`
|
ID int64 `json:"id"`
|
||||||
UserID int64 `json:"userId"`
|
UserID int64 `json:"userId"`
|
||||||
UserName string `json:"userName"`
|
UserName string `json:"userName"`
|
||||||
Name string `json:"name"`
|
Name string `json:"name"`
|
||||||
TunnelID int64 `json:"tunnelId"`
|
TunnelID int64 `json:"tunnelId"`
|
||||||
RemoteAddr string `json:"remoteAddr"`
|
RemoteAddr string `json:"remoteAddr"`
|
||||||
Strategy string `json:"strategy"`
|
Strategy string `json:"strategy"`
|
||||||
InFlow int64 `json:"inFlow"`
|
InFlow int64 `json:"inFlow"`
|
||||||
OutFlow int64 `json:"outFlow"`
|
OutFlow int64 `json:"outFlow"`
|
||||||
CreatedTime int64 `json:"createdTime"`
|
CreatedTime int64 `json:"createdTime"`
|
||||||
UpdatedTime int64 `json:"updatedTime"`
|
UpdatedTime int64 `json:"updatedTime"`
|
||||||
Status int `json:"status"`
|
Status int `json:"status"`
|
||||||
Inx int `json:"inx"`
|
Inx int `json:"inx"`
|
||||||
|
ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type ForwardPortBackup struct {
|
||||||
|
NodeID int64 `json:"nodeId"`
|
||||||
|
Port int `json:"port"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type UserTunnelBackup struct {
|
type UserTunnelBackup struct {
|
||||||
@@ -2426,11 +2432,47 @@ func (r *Repository) exportForwards() ([]ForwardBackup, error) {
|
|||||||
if inx.Valid {
|
if inx.Valid {
|
||||||
f.Inx = int(inx.Int64)
|
f.Inx = int(inx.Int64)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
forwardPorts, err := r.exportForwardPorts(f.ID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
portsCopy := append([]ForwardPortBackup(nil), forwardPorts...)
|
||||||
|
f.ForwardPorts = &portsCopy
|
||||||
|
|
||||||
forwards = append(forwards, f)
|
forwards = append(forwards, f)
|
||||||
}
|
}
|
||||||
return forwards, rows.Err()
|
return forwards, rows.Err()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r *Repository) exportForwardPorts(forwardID int64) ([]ForwardPortBackup, error) {
|
||||||
|
rows, err := r.db.Query(`
|
||||||
|
SELECT node_id, port
|
||||||
|
FROM forward_port
|
||||||
|
WHERE forward_id = ?
|
||||||
|
ORDER BY id ASC
|
||||||
|
`, forwardID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
|
||||||
|
ports := make([]ForwardPortBackup, 0)
|
||||||
|
for rows.Next() {
|
||||||
|
var fp ForwardPortBackup
|
||||||
|
if err := rows.Scan(&fp.NodeID, &fp.Port); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
ports = append(ports, fp)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := rows.Err(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return ports, nil
|
||||||
|
}
|
||||||
|
|
||||||
func (r *Repository) exportUserTunnels() ([]UserTunnelBackup, error) {
|
func (r *Repository) exportUserTunnels() ([]UserTunnelBackup, error) {
|
||||||
rows, err := r.db.Query(`
|
rows, err := r.db.Query(`
|
||||||
SELECT id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status
|
SELECT id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status
|
||||||
@@ -2845,6 +2887,18 @@ func (r *Repository) importForwards(db Execer, forwards []ForwardBackup, now int
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return count, err
|
return count, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if f.ForwardPorts != nil {
|
||||||
|
if _, err := db.Exec(`DELETE FROM forward_port WHERE forward_id = ?`, f.ID); err != nil {
|
||||||
|
return count, err
|
||||||
|
}
|
||||||
|
for _, fp := range *f.ForwardPorts {
|
||||||
|
if _, err := db.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, f.ID, fp.NodeID, fp.Port); err != nil {
|
||||||
|
return count, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
count++
|
count++
|
||||||
}
|
}
|
||||||
return count, nil
|
return count, nil
|
||||||
|
|||||||
@@ -311,6 +311,171 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
|
t.Run("backup export and import preserve forward ports", func(t *testing.T) {
|
||||||
|
now := time.Now().UnixMilli()
|
||||||
|
|
||||||
|
tunnelRes, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||||
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||||
|
`, "backup-forward-tunnel", 1.0, 1, "tls", 0, now, now, 1, "", 88)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("seed tunnel for forward backup: %v", err)
|
||||||
|
}
|
||||||
|
tunnelID, err := tunnelRes.LastInsertId()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read tunnel id for forward backup: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
forwardRes, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||||
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||||
|
`, 1, "admin_user", "backup-forward", tunnelID, "127.0.0.1:9000", "fifo", 0, 0, now, now, 1, 88)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("seed forward for backup: %v", err)
|
||||||
|
}
|
||||||
|
forwardID, err := forwardRes.LastInsertId()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read forward id for backup: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
expected := map[int64]int{
|
||||||
|
2001: 21001,
|
||||||
|
2002: 21002,
|
||||||
|
}
|
||||||
|
for nodeID, port := range expected {
|
||||||
|
if _, err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port); err != nil {
|
||||||
|
t.Fatalf("seed forward_port %d:%d: %v", nodeID, port, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
exportReq := httptest.NewRequest(http.MethodPost, "/api/v1/backup/export", bytes.NewBufferString(`{"types":["forwards"]}`))
|
||||||
|
exportReq.Header.Set("Authorization", adminToken)
|
||||||
|
exportReq.Header.Set("Content-Type", "application/json")
|
||||||
|
exportResp := httptest.NewRecorder()
|
||||||
|
router.ServeHTTP(exportResp, exportReq)
|
||||||
|
|
||||||
|
if exportResp.Code != http.StatusOK {
|
||||||
|
t.Fatalf("expected export status 200, got %d", exportResp.Code)
|
||||||
|
}
|
||||||
|
|
||||||
|
exportBody, err := io.ReadAll(exportResp.Body)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read forwards backup body: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var payload map[string]interface{}
|
||||||
|
if err := json.Unmarshal(exportBody, &payload); err != nil {
|
||||||
|
t.Fatalf("decode forwards backup payload: %v", err)
|
||||||
|
}
|
||||||
|
version, _ := payload["version"].(string)
|
||||||
|
if strings.TrimSpace(version) == "" {
|
||||||
|
t.Fatalf("expected backup payload version, body=%s", string(exportBody))
|
||||||
|
}
|
||||||
|
|
||||||
|
forwardsRaw, ok := payload["forwards"].([]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("expected forwards array in payload, body=%s", string(exportBody))
|
||||||
|
}
|
||||||
|
|
||||||
|
foundForward := false
|
||||||
|
foundPorts := map[int64]int{}
|
||||||
|
for _, item := range forwardsRaw {
|
||||||
|
forwardMap, ok := item.(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
idValue, ok := forwardMap["id"].(float64)
|
||||||
|
if !ok || int64(idValue) != forwardID {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
foundForward = true
|
||||||
|
|
||||||
|
portsRaw, ok := forwardMap["forwardPorts"].([]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("expected forwardPorts for forward %d in payload", forwardID)
|
||||||
|
}
|
||||||
|
for _, p := range portsRaw {
|
||||||
|
portMap, ok := p.(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
nodeID, nodeOK := portMap["nodeId"].(float64)
|
||||||
|
port, portOK := portMap["port"].(float64)
|
||||||
|
if nodeOK && portOK {
|
||||||
|
foundPorts[int64(nodeID)] = int(port)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
if !foundForward {
|
||||||
|
t.Fatalf("expected forward %d in exported forwards payload", forwardID)
|
||||||
|
}
|
||||||
|
if len(foundPorts) != len(expected) {
|
||||||
|
t.Fatalf("expected %d exported forward ports, got %d", len(expected), len(foundPorts))
|
||||||
|
}
|
||||||
|
for nodeID, port := range expected {
|
||||||
|
if got, ok := foundPorts[nodeID]; !ok || got != port {
|
||||||
|
t.Fatalf("expected exported forward port node=%d port=%d, got %v", nodeID, port, foundPorts)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := repo.DB().Exec(`DELETE FROM forward_port WHERE forward_id = ?`, forwardID); err != nil {
|
||||||
|
t.Fatalf("clear forward_port before import: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, 9999, 39999); err != nil {
|
||||||
|
t.Fatalf("seed wrong forward_port before import: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
payload["types"] = []string{"forwards"}
|
||||||
|
importBody, err := json.Marshal(payload)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("marshal forwards import payload: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
importReq := httptest.NewRequest(http.MethodPost, "/api/v1/backup/import", bytes.NewReader(importBody))
|
||||||
|
importReq.Header.Set("Authorization", adminToken)
|
||||||
|
importReq.Header.Set("Content-Type", "application/json")
|
||||||
|
importResp := httptest.NewRecorder()
|
||||||
|
router.ServeHTTP(importResp, importReq)
|
||||||
|
|
||||||
|
var out response.R
|
||||||
|
if err := json.NewDecoder(importResp.Body).Decode(&out); err != nil {
|
||||||
|
t.Fatalf("decode forwards import response: %v", err)
|
||||||
|
}
|
||||||
|
if out.Code != 0 {
|
||||||
|
t.Fatalf("expected forwards import code 0, got %d (%s)", out.Code, out.Msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
rows, err := repo.DB().Query(`SELECT node_id, port FROM forward_port WHERE forward_id = ? ORDER BY id ASC`, forwardID)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("query forward ports after import: %v", err)
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
|
||||||
|
after := make(map[int64]int)
|
||||||
|
for rows.Next() {
|
||||||
|
var nodeID int64
|
||||||
|
var port int
|
||||||
|
if err := rows.Scan(&nodeID, &port); err != nil {
|
||||||
|
t.Fatalf("scan forward_port row: %v", err)
|
||||||
|
}
|
||||||
|
after[nodeID] = port
|
||||||
|
}
|
||||||
|
if err := rows.Err(); err != nil {
|
||||||
|
t.Fatalf("iterate forward_port rows: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(after) != len(expected) {
|
||||||
|
t.Fatalf("expected %d forward ports after import, got %d (%v)", len(expected), len(after), after)
|
||||||
|
}
|
||||||
|
for nodeID, port := range expected {
|
||||||
|
if got, ok := after[nodeID]; !ok || got != port {
|
||||||
|
t.Fatalf("expected forward_port node=%d port=%d after import, got %v", nodeID, port, after)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
t.Run("backup export tolerates nullable legacy tunnel chain fields", func(t *testing.T) {
|
t.Run("backup export tolerates nullable legacy tunnel chain fields", func(t *testing.T) {
|
||||||
now := time.Now().UnixMilli()
|
now := time.Now().UnixMilli()
|
||||||
res, err := repo.DB().Exec(`
|
res, err := repo.DB().Exec(`
|
||||||
|
|||||||
Reference in New Issue
Block a user