mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-08 10:46:37 +08:00
fix(backup): add transaction support and auto-backup before import
- Add transaction support for import operations with rollback on failure - Add auto-backup before import to allow recovery on failure - Convert user import to use INSERT ON CONFLICT pattern - Add Execer interface to support both DB and Tx in import functions
This commit is contained in:
@@ -1203,6 +1203,12 @@ func (h *Handler) backupImport(w http.ResponseWriter, r *http.Request) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
autoBackup, err := h.repo.ExportAll()
|
||||||
|
if err != nil {
|
||||||
|
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("导入前自动备份失败: %v", err)))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
var backup sqlite.BackupData
|
var backup sqlite.BackupData
|
||||||
if err := decodeJSON(r.Body, &backup); err != nil {
|
if err := decodeJSON(r.Body, &backup); err != nil {
|
||||||
response.WriteJSON(w, response.Err(500, "备份数据格式错误"))
|
response.WriteJSON(w, response.Err(500, "备份数据格式错误"))
|
||||||
@@ -1211,9 +1217,10 @@ func (h *Handler) backupImport(w http.ResponseWriter, r *http.Request) {
|
|||||||
|
|
||||||
result, err := h.repo.Import(&backup, req.Types)
|
result, err := h.repo.Import(&backup, req.Types)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("导入失败: %v", err)))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
result.AutoBackup = autoBackup
|
||||||
response.WriteJSON(w, response.OK(result))
|
response.WriteJSON(w, response.OK(result))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -24,6 +24,14 @@ var embeddedSchema string
|
|||||||
//go:embed sql/data.sql
|
//go:embed sql/data.sql
|
||||||
var embeddedSeedData string
|
var embeddedSeedData string
|
||||||
|
|
||||||
|
// Execer is an interface that both *store.DB and *store.Tx satisfy.
|
||||||
|
// Used to allow import functions to work with both regular DB and transactions.
|
||||||
|
type Execer interface {
|
||||||
|
Exec(query string, args ...any) (sql.Result, error)
|
||||||
|
Query(query string, args ...any) (*sql.Rows, error)
|
||||||
|
QueryRow(query string, args ...any) *sql.Row
|
||||||
|
}
|
||||||
|
|
||||||
type Repository struct {
|
type Repository struct {
|
||||||
db *store.DB
|
db *store.DB
|
||||||
}
|
}
|
||||||
@@ -2513,19 +2521,20 @@ func (r *Repository) exportPermissions() ([]PermissionBackup, error) {
|
|||||||
|
|
||||||
// ImportResult contains the result of an import operation
|
// ImportResult contains the result of an import operation
|
||||||
type ImportResult struct {
|
type ImportResult struct {
|
||||||
UsersImported int `json:"usersImported"`
|
UsersImported int `json:"usersImported"`
|
||||||
NodesImported int `json:"nodesImported"`
|
NodesImported int `json:"nodesImported"`
|
||||||
TunnelsImported int `json:"tunnelsImported"`
|
TunnelsImported int `json:"tunnelsImported"`
|
||||||
ForwardsImported int `json:"forwardsImported"`
|
ForwardsImported int `json:"forwardsImported"`
|
||||||
UserTunnelsImported int `json:"userTunnelsImported"`
|
UserTunnelsImported int `json:"userTunnelsImported"`
|
||||||
SpeedLimitsImported int `json:"speedLimitsImported"`
|
SpeedLimitsImported int `json:"speedLimitsImported"`
|
||||||
TunnelGroupsImported int `json:"tunnelGroupsImported"`
|
TunnelGroupsImported int `json:"tunnelGroupsImported"`
|
||||||
UserGroupsImported int `json:"userGroupsImported"`
|
UserGroupsImported int `json:"userGroupsImported"`
|
||||||
PermissionsImported int `json:"permissionsImported"`
|
PermissionsImported int `json:"permissionsImported"`
|
||||||
ConfigsImported int `json:"configsImported"`
|
ConfigsImported int `json:"configsImported"`
|
||||||
|
AutoBackup *BackupData `json:"autoBackup,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// Import imports data from BackupData
|
// Import imports data from BackupData with transaction support
|
||||||
func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult, error) {
|
func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult, error) {
|
||||||
result := &ImportResult{}
|
result := &ImportResult{}
|
||||||
|
|
||||||
@@ -2534,10 +2543,16 @@ func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult,
|
|||||||
typeSet[t] = true
|
typeSet[t] = true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
tx, err := r.db.Begin()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to begin transaction: %w", err)
|
||||||
|
}
|
||||||
|
defer func() { _ = tx.Rollback() }()
|
||||||
|
|
||||||
now := unixMilliNow()
|
now := unixMilliNow()
|
||||||
|
|
||||||
if typeSet["users"] && len(backup.Users) > 0 {
|
if typeSet["users"] && len(backup.Users) > 0 {
|
||||||
count, err := r.importUsers(backup.Users, now)
|
count, err := r.importUsers(tx, backup.Users, now)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("import users failed: %w", err)
|
return nil, fmt.Errorf("import users failed: %w", err)
|
||||||
}
|
}
|
||||||
@@ -2545,7 +2560,7 @@ func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult,
|
|||||||
}
|
}
|
||||||
|
|
||||||
if typeSet["nodes"] && len(backup.Nodes) > 0 {
|
if typeSet["nodes"] && len(backup.Nodes) > 0 {
|
||||||
count, err := r.importNodes(backup.Nodes, now)
|
count, err := r.importNodes(tx, backup.Nodes, now)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("import nodes failed: %w", err)
|
return nil, fmt.Errorf("import nodes failed: %w", err)
|
||||||
}
|
}
|
||||||
@@ -2553,7 +2568,7 @@ func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult,
|
|||||||
}
|
}
|
||||||
|
|
||||||
if typeSet["tunnels"] && len(backup.Tunnels) > 0 {
|
if typeSet["tunnels"] && len(backup.Tunnels) > 0 {
|
||||||
count, err := r.importTunnels(backup.Tunnels, now)
|
count, err := r.importTunnels(tx, backup.Tunnels, now)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("import tunnels failed: %w", err)
|
return nil, fmt.Errorf("import tunnels failed: %w", err)
|
||||||
}
|
}
|
||||||
@@ -2561,7 +2576,7 @@ func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult,
|
|||||||
}
|
}
|
||||||
|
|
||||||
if typeSet["forwards"] && len(backup.Forwards) > 0 {
|
if typeSet["forwards"] && len(backup.Forwards) > 0 {
|
||||||
count, err := r.importForwards(backup.Forwards, now)
|
count, err := r.importForwards(tx, backup.Forwards, now)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("import forwards failed: %w", err)
|
return nil, fmt.Errorf("import forwards failed: %w", err)
|
||||||
}
|
}
|
||||||
@@ -2569,7 +2584,7 @@ func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult,
|
|||||||
}
|
}
|
||||||
|
|
||||||
if typeSet["userTunnels"] && len(backup.UserTunnels) > 0 {
|
if typeSet["userTunnels"] && len(backup.UserTunnels) > 0 {
|
||||||
count, err := r.importUserTunnels(backup.UserTunnels, now)
|
count, err := r.importUserTunnels(tx, backup.UserTunnels, now)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("import user tunnels failed: %w", err)
|
return nil, fmt.Errorf("import user tunnels failed: %w", err)
|
||||||
}
|
}
|
||||||
@@ -2577,7 +2592,7 @@ func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult,
|
|||||||
}
|
}
|
||||||
|
|
||||||
if typeSet["speedLimits"] && len(backup.SpeedLimits) > 0 {
|
if typeSet["speedLimits"] && len(backup.SpeedLimits) > 0 {
|
||||||
count, err := r.importSpeedLimits(backup.SpeedLimits, now)
|
count, err := r.importSpeedLimits(tx, backup.SpeedLimits, now)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("import speed limits failed: %w", err)
|
return nil, fmt.Errorf("import speed limits failed: %w", err)
|
||||||
}
|
}
|
||||||
@@ -2585,7 +2600,7 @@ func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult,
|
|||||||
}
|
}
|
||||||
|
|
||||||
if typeSet["tunnelGroups"] && len(backup.TunnelGroups) > 0 {
|
if typeSet["tunnelGroups"] && len(backup.TunnelGroups) > 0 {
|
||||||
count, err := r.importTunnelGroups(backup.TunnelGroups, now)
|
count, err := r.importTunnelGroups(tx, backup.TunnelGroups, now)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("import tunnel groups failed: %w", err)
|
return nil, fmt.Errorf("import tunnel groups failed: %w", err)
|
||||||
}
|
}
|
||||||
@@ -2593,7 +2608,7 @@ func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult,
|
|||||||
}
|
}
|
||||||
|
|
||||||
if typeSet["userGroups"] && len(backup.UserGroups) > 0 {
|
if typeSet["userGroups"] && len(backup.UserGroups) > 0 {
|
||||||
count, err := r.importUserGroups(backup.UserGroups, now)
|
count, err := r.importUserGroups(tx, backup.UserGroups, now)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("import user groups failed: %w", err)
|
return nil, fmt.Errorf("import user groups failed: %w", err)
|
||||||
}
|
}
|
||||||
@@ -2601,7 +2616,7 @@ func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult,
|
|||||||
}
|
}
|
||||||
|
|
||||||
if typeSet["permissions"] && len(backup.Permissions) > 0 {
|
if typeSet["permissions"] && len(backup.Permissions) > 0 {
|
||||||
count, err := r.importPermissions(backup.Permissions, now)
|
count, err := r.importPermissions(tx, backup.Permissions, now)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("import permissions failed: %w", err)
|
return nil, fmt.Errorf("import permissions failed: %w", err)
|
||||||
}
|
}
|
||||||
@@ -2609,43 +2624,42 @@ func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult,
|
|||||||
}
|
}
|
||||||
|
|
||||||
if typeSet["configs"] && len(backup.Configs) > 0 {
|
if typeSet["configs"] && len(backup.Configs) > 0 {
|
||||||
count, err := r.importConfigs(backup.Configs, now)
|
count, err := r.importConfigs(tx, backup.Configs, now)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("import configs failed: %w", err)
|
return nil, fmt.Errorf("import configs failed: %w", err)
|
||||||
}
|
}
|
||||||
result.ConfigsImported = count
|
result.ConfigsImported = count
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if err := tx.Commit(); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to commit transaction: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
return result, nil
|
return result, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Repository) importUsers(users []UserBackup, now int64) (int, error) {
|
func (r *Repository) importUsers(db Execer, users []UserBackup, now int64) (int, error) {
|
||||||
count := 0
|
count := 0
|
||||||
for _, u := range users {
|
for _, u := range users {
|
||||||
// Check if user exists
|
_, err := db.Exec(`
|
||||||
exists, err := r.UsernameExists(u.User)
|
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||||
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||||
|
ON CONFLICT(id) DO UPDATE SET
|
||||||
|
user = excluded.user,
|
||||||
|
pwd = excluded.pwd,
|
||||||
|
role_id = excluded.role_id,
|
||||||
|
exp_time = excluded.exp_time,
|
||||||
|
flow = excluded.flow,
|
||||||
|
in_flow = excluded.in_flow,
|
||||||
|
out_flow = excluded.out_flow,
|
||||||
|
flow_reset_time = excluded.flow_reset_time,
|
||||||
|
num = excluded.num,
|
||||||
|
updated_time = excluded.updated_time,
|
||||||
|
status = excluded.status
|
||||||
|
`, u.ID, u.User, u.Pwd, u.RoleID, u.ExpTime, u.Flow, u.InFlow, u.OutFlow, u.FlowResetTime, u.Num, u.CreatedTime, now, u.Status)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return count, err
|
return count, err
|
||||||
}
|
}
|
||||||
if exists {
|
|
||||||
// Update existing
|
|
||||||
_, err = r.db.Exec(`
|
|
||||||
UPDATE user SET pwd = ?, role_id = ?, exp_time = ?, flow = ?, in_flow = ?, out_flow = ?, flow_reset_time = ?, num = ?, updated_time = ?, status = ?
|
|
||||||
WHERE id = ?
|
|
||||||
`, u.Pwd, u.RoleID, u.ExpTime, u.Flow, u.InFlow, u.OutFlow, u.FlowResetTime, u.Num, now, u.Status, u.ID)
|
|
||||||
if err != nil {
|
|
||||||
return count, err
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
// Insert new
|
|
||||||
_, err = r.db.Exec(`
|
|
||||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
|
||||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
||||||
`, u.ID, u.User, u.Pwd, u.RoleID, u.ExpTime, u.Flow, u.InFlow, u.OutFlow, u.FlowResetTime, u.Num, u.CreatedTime, now, u.Status)
|
|
||||||
if err != nil {
|
|
||||||
return count, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
count++
|
count++
|
||||||
}
|
}
|
||||||
return count, nil
|
return count, nil
|
||||||
@@ -2660,16 +2674,10 @@ func (r *Repository) UsernameExists(username string) (bool, error) {
|
|||||||
return count > 0, nil
|
return count > 0, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Repository) importNodes(nodes []NodeBackup, now int64) (int, error) {
|
func (r *Repository) importNodes(db Execer, nodes []NodeBackup, now int64) (int, error) {
|
||||||
count := 0
|
count := 0
|
||||||
for _, n := range nodes {
|
for _, n := range nodes {
|
||||||
// Check if node exists
|
_, err := db.Exec(`
|
||||||
_, err := r.GetNodeByID(n.ID)
|
|
||||||
if err != nil {
|
|
||||||
return count, err
|
|
||||||
}
|
|
||||||
// Use upsert pattern
|
|
||||||
_, err = r.db.Exec(`
|
|
||||||
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config)
|
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config)
|
||||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||||
ON CONFLICT(id) DO UPDATE SET
|
ON CONFLICT(id) DO UPDATE SET
|
||||||
@@ -2702,10 +2710,10 @@ func (r *Repository) importNodes(nodes []NodeBackup, now int64) (int, error) {
|
|||||||
return count, nil
|
return count, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Repository) importTunnels(tunnels []TunnelBackup, now int64) (int, error) {
|
func (r *Repository) importTunnels(db Execer, tunnels []TunnelBackup, now int64) (int, error) {
|
||||||
count := 0
|
count := 0
|
||||||
for _, t := range tunnels {
|
for _, t := range tunnels {
|
||||||
_, err := r.db.Exec(`
|
_, err := db.Exec(`
|
||||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||||
ON CONFLICT(id) DO UPDATE SET
|
ON CONFLICT(id) DO UPDATE SET
|
||||||
@@ -2722,10 +2730,9 @@ func (r *Repository) importTunnels(tunnels []TunnelBackup, now int64) (int, erro
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return count, err
|
return count, err
|
||||||
}
|
}
|
||||||
// Import chain tunnels
|
|
||||||
if len(t.ChainTunnels) > 0 {
|
if len(t.ChainTunnels) > 0 {
|
||||||
for _, ct := range t.ChainTunnels {
|
for _, ct := range t.ChainTunnels {
|
||||||
_, err = r.db.Exec(`
|
_, err = db.Exec(`
|
||||||
INSERT INTO chain_tunnel(id, tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
INSERT INTO chain_tunnel(id, tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?)
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?)
|
||||||
ON CONFLICT(id) DO UPDATE SET
|
ON CONFLICT(id) DO UPDATE SET
|
||||||
@@ -2746,10 +2753,10 @@ func (r *Repository) importTunnels(tunnels []TunnelBackup, now int64) (int, erro
|
|||||||
return count, nil
|
return count, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Repository) importForwards(forwards []ForwardBackup, now int64) (int, error) {
|
func (r *Repository) importForwards(db Execer, forwards []ForwardBackup, now int64) (int, error) {
|
||||||
count := 0
|
count := 0
|
||||||
for _, f := range forwards {
|
for _, f := range forwards {
|
||||||
_, err := r.db.Exec(`
|
_, err := db.Exec(`
|
||||||
INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||||
ON CONFLICT(id) DO UPDATE SET
|
ON CONFLICT(id) DO UPDATE SET
|
||||||
@@ -2773,14 +2780,14 @@ func (r *Repository) importForwards(forwards []ForwardBackup, now int64) (int, e
|
|||||||
return count, nil
|
return count, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Repository) importUserTunnels(userTunnels []UserTunnelBackup, now int64) (int, error) {
|
func (r *Repository) importUserTunnels(db Execer, userTunnels []UserTunnelBackup, now int64) (int, error) {
|
||||||
count := 0
|
count := 0
|
||||||
for _, ut := range userTunnels {
|
for _, ut := range userTunnels {
|
||||||
var speedID interface{}
|
var speedID interface{}
|
||||||
if ut.SpeedID > 0 {
|
if ut.SpeedID > 0 {
|
||||||
speedID = ut.SpeedID
|
speedID = ut.SpeedID
|
||||||
}
|
}
|
||||||
_, err := r.db.Exec(`
|
_, err := db.Exec(`
|
||||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||||
ON CONFLICT(id) DO UPDATE SET
|
ON CONFLICT(id) DO UPDATE SET
|
||||||
@@ -2803,10 +2810,10 @@ func (r *Repository) importUserTunnels(userTunnels []UserTunnelBackup, now int64
|
|||||||
return count, nil
|
return count, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Repository) importSpeedLimits(speedLimits []SpeedLimitBackup, now int64) (int, error) {
|
func (r *Repository) importSpeedLimits(db Execer, speedLimits []SpeedLimitBackup, now int64) (int, error) {
|
||||||
count := 0
|
count := 0
|
||||||
for _, sl := range speedLimits {
|
for _, sl := range speedLimits {
|
||||||
_, err := r.db.Exec(`
|
_, err := db.Exec(`
|
||||||
INSERT INTO speed_limit(id, name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
INSERT INTO speed_limit(id, name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
||||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?)
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?)
|
||||||
ON CONFLICT(id) DO UPDATE SET
|
ON CONFLICT(id) DO UPDATE SET
|
||||||
@@ -2825,10 +2832,10 @@ func (r *Repository) importSpeedLimits(speedLimits []SpeedLimitBackup, now int64
|
|||||||
return count, nil
|
return count, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Repository) importTunnelGroups(tunnelGroups []TunnelGroupBackup, now int64) (int, error) {
|
func (r *Repository) importTunnelGroups(db Execer, tunnelGroups []TunnelGroupBackup, now int64) (int, error) {
|
||||||
count := 0
|
count := 0
|
||||||
for _, tg := range tunnelGroups {
|
for _, tg := range tunnelGroups {
|
||||||
_, err := r.db.Exec(`
|
_, err := db.Exec(`
|
||||||
INSERT INTO tunnel_group(id, name, created_time, updated_time, status)
|
INSERT INTO tunnel_group(id, name, created_time, updated_time, status)
|
||||||
VALUES(?, ?, ?, ?, ?)
|
VALUES(?, ?, ?, ?, ?)
|
||||||
ON CONFLICT(id) DO UPDATE SET
|
ON CONFLICT(id) DO UPDATE SET
|
||||||
@@ -2839,13 +2846,12 @@ func (r *Repository) importTunnelGroups(tunnelGroups []TunnelGroupBackup, now in
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return count, err
|
return count, err
|
||||||
}
|
}
|
||||||
// Update tunnel group memberships
|
_, err = db.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, tg.ID)
|
||||||
_, err = r.db.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, tg.ID)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return count, err
|
return count, err
|
||||||
}
|
}
|
||||||
for _, tunnelID := range tg.Tunnels {
|
for _, tunnelID := range tg.Tunnels {
|
||||||
_, err = r.db.Exec(`
|
_, err = db.Exec(`
|
||||||
INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time)
|
INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time)
|
||||||
VALUES(?, ?, ?)
|
VALUES(?, ?, ?)
|
||||||
`, tg.ID, tunnelID, now)
|
`, tg.ID, tunnelID, now)
|
||||||
@@ -2858,10 +2864,10 @@ func (r *Repository) importTunnelGroups(tunnelGroups []TunnelGroupBackup, now in
|
|||||||
return count, nil
|
return count, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Repository) importUserGroups(userGroups []UserGroupBackup, now int64) (int, error) {
|
func (r *Repository) importUserGroups(db Execer, userGroups []UserGroupBackup, now int64) (int, error) {
|
||||||
count := 0
|
count := 0
|
||||||
for _, ug := range userGroups {
|
for _, ug := range userGroups {
|
||||||
_, err := r.db.Exec(`
|
_, err := db.Exec(`
|
||||||
INSERT INTO user_group(id, name, created_time, updated_time, status)
|
INSERT INTO user_group(id, name, created_time, updated_time, status)
|
||||||
VALUES(?, ?, ?, ?, ?)
|
VALUES(?, ?, ?, ?, ?)
|
||||||
ON CONFLICT(id) DO UPDATE SET
|
ON CONFLICT(id) DO UPDATE SET
|
||||||
@@ -2872,13 +2878,12 @@ func (r *Repository) importUserGroups(userGroups []UserGroupBackup, now int64) (
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return count, err
|
return count, err
|
||||||
}
|
}
|
||||||
// Update user group memberships
|
_, err = db.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, ug.ID)
|
||||||
_, err = r.db.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, ug.ID)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return count, err
|
return count, err
|
||||||
}
|
}
|
||||||
for _, userID := range ug.Users {
|
for _, userID := range ug.Users {
|
||||||
_, err = r.db.Exec(`
|
_, err = db.Exec(`
|
||||||
INSERT INTO user_group_user(user_group_id, user_id, created_time)
|
INSERT INTO user_group_user(user_group_id, user_id, created_time)
|
||||||
VALUES(?, ?, ?)
|
VALUES(?, ?, ?)
|
||||||
`, ug.ID, userID, now)
|
`, ug.ID, userID, now)
|
||||||
@@ -2891,10 +2896,10 @@ func (r *Repository) importUserGroups(userGroups []UserGroupBackup, now int64) (
|
|||||||
return count, nil
|
return count, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Repository) importPermissions(permissions []PermissionBackup, now int64) (int, error) {
|
func (r *Repository) importPermissions(db Execer, permissions []PermissionBackup, now int64) (int, error) {
|
||||||
count := 0
|
count := 0
|
||||||
for _, p := range permissions {
|
for _, p := range permissions {
|
||||||
_, err := r.db.Exec(`
|
_, err := db.Exec(`
|
||||||
INSERT INTO group_permission(id, user_group_id, tunnel_group_id, created_time, created_by_group)
|
INSERT INTO group_permission(id, user_group_id, tunnel_group_id, created_time, created_by_group)
|
||||||
VALUES(?, ?, ?, ?, ?)
|
VALUES(?, ?, ?, ?, ?)
|
||||||
ON CONFLICT(id) DO UPDATE SET
|
ON CONFLICT(id) DO UPDATE SET
|
||||||
@@ -2905,9 +2910,8 @@ func (r *Repository) importPermissions(permissions []PermissionBackup, now int64
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return count, err
|
return count, err
|
||||||
}
|
}
|
||||||
// Import grants
|
|
||||||
for _, g := range p.Grants {
|
for _, g := range p.Grants {
|
||||||
_, err = r.db.Exec(`
|
_, err = db.Exec(`
|
||||||
INSERT INTO group_permission_grant(id, user_group_id, tunnel_group_id, user_tunnel_id, created_time, created_by_group)
|
INSERT INTO group_permission_grant(id, user_group_id, tunnel_group_id, user_tunnel_id, created_time, created_by_group)
|
||||||
VALUES(?, ?, ?, ?, ?, ?)
|
VALUES(?, ?, ?, ?, ?, ?)
|
||||||
ON CONFLICT(id) DO UPDATE SET
|
ON CONFLICT(id) DO UPDATE SET
|
||||||
@@ -2923,7 +2927,7 @@ func (r *Repository) importPermissions(permissions []PermissionBackup, now int64
|
|||||||
return count, nil
|
return count, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Repository) importConfigs(configs map[string]string, now int64) (int, error) {
|
func (r *Repository) importConfigs(db Execer, configs map[string]string, now int64) (int, error) {
|
||||||
count := 0
|
count := 0
|
||||||
for name, value := range configs {
|
for name, value := range configs {
|
||||||
err := r.UpsertConfig(name, value, now)
|
err := r.UpsertConfig(name, value, now)
|
||||||
|
|||||||
Reference in New Issue
Block a user