From b11283d48848f8824cd71ec6113cd757033848fe Mon Sep 17 00:00:00 2001 From: sagit Date: Fri, 13 Feb 2026 03:06:23 +0000 Subject: [PATCH 01/12] feat(backend): add backup and restore functionality Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-opencode) Co-authored-by: Sisyphus --- go-backend/internal/http/handler/handler.go | 77 ++ .../internal/store/sqlite/repository.go | 1051 +++++++++++++++++ 2 files changed, 1128 insertions(+) diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index 4bfec49..4fcd49f 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -166,6 +166,9 @@ func (h *Handler) Register(mux *http.ServeMux) { mux.HandleFunc("/api/v1/federation/runtime/diagnose", h.authPeer(h.federationRuntimeDiagnose)) 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) @@ -1108,3 +1111,77 @@ func (h *Handler) verifyCloudflareTurnstile(token, secretKey string) bool { } return body.Success } + +type backupExportRequest struct { + Types []string `json:"types"` +} + +func (h *Handler) backupExport(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + + var req backupExportRequest + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.Err(500, "请求参数错误")) + return + } + + var backup interface{} + var err error + + if len(req.Types) == 0 { + backup, err = h.repo.ExportAll() + } else { + backup, err = h.repo.ExportPartial(req.Types) + } + + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + + w.Header().Set("Content-Disposition", "attachment; filename=backup.json") + w.Header().Set("Content-Type", "application/json") + if err := json.NewEncoder(w).Encode(backup); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } +} + +type backupImportRequest struct { + Types []string `json:"types"` +} + +func (h *Handler) backupImport(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + + var req backupImportRequest + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.Err(500, "请求参数错误")) + return + } + + if len(req.Types) == 0 { + response.WriteJSON(w, response.Err(500, "请选择要导入的数据类型")) + return + } + + var backup sqlite.BackupData + if err := decodeJSON(r.Body, &backup); err != nil { + response.WriteJSON(w, response.Err(500, "备份数据格式错误")) + return + } + + result, err := h.repo.Import(&backup, req.Types) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + + response.WriteJSON(w, response.OK(result)) +} diff --git a/go-backend/internal/store/sqlite/repository.go b/go-backend/internal/store/sqlite/repository.go index 7f6bc00..7251077 100644 --- a/go-backend/internal/store/sqlite/repository.go +++ b/go-backend/internal/store/sqlite/repository.go @@ -1655,3 +1655,1054 @@ func (r *Repository) DeleteFederationTunnelBindingsByTunnel(tunnelID int64) erro var osMkdirAll = func(path string) error { return os.MkdirAll(path, 0o755) } + +// ============ Backup/Export Data Structures ============ + +// BackupData represents the full backup structure +type BackupData struct { + Version string `json:"version"` + ExportedAt int64 `json:"exportedAt"` + Users []UserBackup `json:"users,omitempty"` + Nodes []NodeBackup `json:"nodes,omitempty"` + Tunnels []TunnelBackup `json:"tunnels,omitempty"` + Forwards []ForwardBackup `json:"forwards,omitempty"` + UserTunnels []UserTunnelBackup `json:"userTunnels,omitempty"` + SpeedLimits []SpeedLimitBackup `json:"speedLimits,omitempty"` + TunnelGroups []TunnelGroupBackup `json:"tunnelGroups,omitempty"` + UserGroups []UserGroupBackup `json:"userGroups,omitempty"` + Permissions []PermissionBackup `json:"permissions,omitempty"` + Configs map[string]string `json:"configs,omitempty"` +} + +type UserBackup struct { + ID int64 `json:"id"` + User string `json:"user"` + Pwd string `json:"pwd"` + RoleID int `json:"roleId"` + ExpTime int64 `json:"expTime"` + Flow int64 `json:"flow"` + InFlow int64 `json:"inFlow"` + OutFlow int64 `json:"outFlow"` + FlowResetTime int64 `json:"flowResetTime"` + Num int `json:"num"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime,omitempty"` + Status int `json:"status"` +} + +type NodeBackup struct { + ID int64 `json:"id"` + Name string `json:"name"` + Secret string `json:"secret"` + ServerIP string `json:"serverIp"` + ServerIPv4 string `json:"serverIpV4,omitempty"` + ServerIPv6 string `json:"serverIpV6,omitempty"` + Port string `json:"port"` + InterfaceName string `json:"interfaceName,omitempty"` + Version string `json:"version,omitempty"` + HTTP int `json:"http"` + TLS int `json:"tls"` + Socks int `json:"socks"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime,omitempty"` + Status int `json:"status"` + TCPListenAddr string `json:"tcpListenAddr"` + UDPListenAddr string `json:"udpListenAddr"` + Inx int `json:"inx"` + IsRemote int `json:"isRemote"` + RemoteURL string `json:"remoteUrl,omitempty"` + RemoteToken string `json:"remoteToken,omitempty"` + RemoteConfig string `json:"remoteConfig,omitempty"` +} + +type TunnelBackup struct { + ID int64 `json:"id"` + Name string `json:"name"` + TrafficRatio float64 `json:"trafficRatio"` + Type int `json:"type"` + Protocol string `json:"protocol"` + Flow int64 `json:"flow"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime"` + Status int `json:"status"` + InIP string `json:"inIp,omitempty"` + Inx int `json:"inx"` + ChainTunnels []ChainTunnelBackup `json:"chainTunnels,omitempty"` +} + +type ChainTunnelBackup struct { + ID int64 `json:"id"` + TunnelID int64 `json:"tunnelId"` + ChainType string `json:"chainType"` + NodeID int64 `json:"nodeId"` + Port int `json:"port,omitempty"` + Strategy string `json:"strategy,omitempty"` + Inx int `json:"inx,omitempty"` + Protocol string `json:"protocol,omitempty"` +} + +type ForwardBackup struct { + ID int64 `json:"id"` + UserID int64 `json:"userId"` + UserName string `json:"userName"` + Name string `json:"name"` + TunnelID int64 `json:"tunnelId"` + RemoteAddr string `json:"remoteAddr"` + Strategy string `json:"strategy"` + InFlow int64 `json:"inFlow"` + OutFlow int64 `json:"outFlow"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime"` + Status int `json:"status"` + Inx int `json:"inx"` +} + +type UserTunnelBackup struct { + ID int64 `json:"id"` + UserID int64 `json:"userId"` + TunnelID int64 `json:"tunnelId"` + SpeedID int64 `json:"speedId,omitempty"` + Num int `json:"num"` + Flow int64 `json:"flow"` + InFlow int64 `json:"inFlow"` + OutFlow int64 `json:"outFlow"` + FlowResetTime int64 `json:"flowResetTime"` + ExpTime int64 `json:"expTime"` + Status int `json:"status"` +} + +type SpeedLimitBackup struct { + ID int64 `json:"id"` + Name string `json:"name"` + Speed int64 `json:"speed"` + TunnelID int64 `json:"tunnelId"` + TunnelName string `json:"tunnelName"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime,omitempty"` + Status int `json:"status"` +} + +type TunnelGroupBackup struct { + ID int64 `json:"id"` + Name string `json:"name"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime"` + Status int `json:"status"` + Tunnels []int64 `json:"tunnels,omitempty"` +} + +type UserGroupBackup struct { + ID int64 `json:"id"` + Name string `json:"name"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime"` + Status int `json:"status"` + Users []int64 `json:"users,omitempty"` +} + +type PermissionBackup struct { + ID int64 `json:"id"` + UserGroupID int64 `json:"userGroupId"` + TunnelGroupID int64 `json:"tunnelGroupId"` + CreatedTime int64 `json:"createdTime"` + CreatedByGroup int `json:"createdByGroup"` + Grants []PermissionGrantBackup `json:"grants,omitempty"` +} + +type PermissionGrantBackup struct { + ID int64 `json:"id"` + UserGroupID int64 `json:"userGroupId"` + TunnelGroupID int64 `json:"tunnelGroupId"` + UserTunnelID int64 `json:"userTunnelId"` + CreatedTime int64 `json:"createdTime"` + CreatedByGroup int `json:"createdByGroup"` +} + +// ============ Export Methods ============ + +// ExportAll exports all data as BackupData +func (r *Repository) ExportAll() (*BackupData, error) { + backup := &BackupData{ + Version: "1.0", + ExportedAt: unixMilliNow(), + } + + // Export all data types + users, err := r.exportUsers() + if err != nil { + return nil, fmt.Errorf("export users failed: %w", err) + } + backup.Users = users + + nodes, err := r.exportNodes() + if err != nil { + return nil, fmt.Errorf("export nodes failed: %w", err) + } + backup.Nodes = nodes + + tunnels, err := r.exportTunnels() + if err != nil { + return nil, fmt.Errorf("export tunnels failed: %w", err) + } + backup.Tunnels = tunnels + + forwards, err := r.exportForwards() + if err != nil { + return nil, fmt.Errorf("export forwards failed: %w", err) + } + backup.Forwards = forwards + + userTunnels, err := r.exportUserTunnels() + if err != nil { + return nil, fmt.Errorf("export user tunnels failed: %w", err) + } + backup.UserTunnels = userTunnels + + speedLimits, err := r.exportSpeedLimits() + if err != nil { + return nil, fmt.Errorf("export speed limits failed: %w", err) + } + backup.SpeedLimits = speedLimits + + tunnelGroups, err := r.exportTunnelGroups() + if err != nil { + return nil, fmt.Errorf("export tunnel groups failed: %w", err) + } + backup.TunnelGroups = tunnelGroups + + userGroups, err := r.exportUserGroups() + if err != nil { + return nil, fmt.Errorf("export user groups failed: %w", err) + } + backup.UserGroups = userGroups + + permissions, err := r.exportPermissions() + if err != nil { + return nil, fmt.Errorf("export permissions failed: %w", err) + } + backup.Permissions = permissions + + configs, err := r.ListConfigs() + if err != nil { + return nil, fmt.Errorf("export configs failed: %w", err) + } + backup.Configs = configs + + return backup, nil +} + +// ExportPartial exports selected data types +func (r *Repository) ExportPartial(types []string) (*BackupData, error) { + backup := &BackupData{ + Version: "1.0", + ExportedAt: unixMilliNow(), + } + + typeSet := make(map[string]bool) + for _, t := range types { + typeSet[t] = true + } + + if typeSet["users"] { + users, err := r.exportUsers() + if err != nil { + return nil, fmt.Errorf("export users failed: %w", err) + } + backup.Users = users + } + if typeSet["nodes"] { + nodes, err := r.exportNodes() + if err != nil { + return nil, fmt.Errorf("export nodes failed: %w", err) + } + backup.Nodes = nodes + } + if typeSet["tunnels"] { + tunnels, err := r.exportTunnels() + if err != nil { + return nil, fmt.Errorf("export tunnels failed: %w", err) + } + backup.Tunnels = tunnels + } + if typeSet["forwards"] { + forwards, err := r.exportForwards() + if err != nil { + return nil, fmt.Errorf("export forwards failed: %w", err) + } + backup.Forwards = forwards + } + if typeSet["userTunnels"] { + userTunnels, err := r.exportUserTunnels() + if err != nil { + return nil, fmt.Errorf("export user tunnels failed: %w", err) + } + backup.UserTunnels = userTunnels + } + if typeSet["speedLimits"] { + speedLimits, err := r.exportSpeedLimits() + if err != nil { + return nil, fmt.Errorf("export speed limits failed: %w", err) + } + backup.SpeedLimits = speedLimits + } + if typeSet["tunnelGroups"] { + tunnelGroups, err := r.exportTunnelGroups() + if err != nil { + return nil, fmt.Errorf("export tunnel groups failed: %w", err) + } + backup.TunnelGroups = tunnelGroups + } + if typeSet["userGroups"] { + userGroups, err := r.exportUserGroups() + if err != nil { + return nil, fmt.Errorf("export user groups failed: %w", err) + } + backup.UserGroups = userGroups + } + if typeSet["permissions"] { + permissions, err := r.exportPermissions() + if err != nil { + return nil, fmt.Errorf("export permissions failed: %w", err) + } + backup.Permissions = permissions + } + if typeSet["configs"] { + configs, err := r.ListConfigs() + if err != nil { + return nil, fmt.Errorf("export configs failed: %w", err) + } + backup.Configs = configs + } + + return backup, nil +} + +func (r *Repository) exportUsers() ([]UserBackup, error) { + rows, err := r.db.Query(` + SELECT id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status + FROM user ORDER BY id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + var users []UserBackup + for rows.Next() { + var u UserBackup + var updatedTime sql.NullInt64 + if err := rows.Scan(&u.ID, &u.User, &u.Pwd, &u.RoleID, &u.ExpTime, &u.Flow, &u.InFlow, &u.OutFlow, &u.FlowResetTime, &u.Num, &u.CreatedTime, &updatedTime, &u.Status); err != nil { + return nil, err + } + if updatedTime.Valid { + u.UpdatedTime = updatedTime.Int64 + } + users = append(users, u) + } + return users, rows.Err() +} + +func (r *Repository) exportNodes() ([]NodeBackup, error) { + rows, err := r.db.Query(` + SELECT 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 + FROM node ORDER BY inx ASC, id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + var nodes []NodeBackup + for rows.Next() { + var n NodeBackup + var updatedTime sql.NullInt64 + var serverIPv4, serverIPv6, interfaceName, version, remoteURL, remoteToken, remoteConfig sql.NullString + if err := rows.Scan(&n.ID, &n.Name, &n.Secret, &n.ServerIP, &serverIPv4, &serverIPv6, &n.Port, &interfaceName, &version, &n.HTTP, &n.TLS, &n.Socks, &n.CreatedTime, &updatedTime, &n.Status, &n.TCPListenAddr, &n.UDPListenAddr, &n.Inx, &n.IsRemote, &remoteURL, &remoteToken, &remoteConfig); err != nil { + return nil, err + } + if updatedTime.Valid { + n.UpdatedTime = updatedTime.Int64 + } + if serverIPv4.Valid { + n.ServerIPv4 = serverIPv4.String + } + if serverIPv6.Valid { + n.ServerIPv6 = serverIPv6.String + } + if interfaceName.Valid { + n.InterfaceName = interfaceName.String + } + if version.Valid { + n.Version = version.String + } + if remoteURL.Valid { + n.RemoteURL = remoteURL.String + } + if remoteToken.Valid { + n.RemoteToken = remoteToken.String + } + if remoteConfig.Valid { + n.RemoteConfig = remoteConfig.String + } + nodes = append(nodes, n) + } + return nodes, rows.Err() +} + +func (r *Repository) exportTunnels() ([]TunnelBackup, error) { + rows, err := r.db.Query(` + SELECT id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx + FROM tunnel ORDER BY inx ASC, id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + var tunnels []TunnelBackup + for rows.Next() { + var t TunnelBackup + var inIP sql.NullString + if err := rows.Scan(&t.ID, &t.Name, &t.TrafficRatio, &t.Type, &t.Protocol, &t.Flow, &t.CreatedTime, &t.UpdatedTime, &t.Status, &inIP, &t.Inx); err != nil { + return nil, err + } + if inIP.Valid { + t.InIP = inIP.String + } + // Export chain tunnels + chainTunnels, err := r.exportChainTunnels(t.ID) + if err != nil { + return nil, err + } + t.ChainTunnels = chainTunnels + tunnels = append(tunnels, t) + } + return tunnels, rows.Err() +} + +func (r *Repository) exportChainTunnels(tunnelID int64) ([]ChainTunnelBackup, error) { + rows, err := r.db.Query(` + SELECT id, tunnel_id, chain_type, node_id, port, strategy, inx, protocol + FROM chain_tunnel WHERE tunnel_id = ? ORDER BY inx ASC, id ASC + `, tunnelID) + if err != nil { + return nil, err + } + defer rows.Close() + + var chainTunnels []ChainTunnelBackup + for rows.Next() { + var ct ChainTunnelBackup + var port sql.NullInt64 + if err := rows.Scan(&ct.ID, &ct.TunnelID, &ct.ChainType, &ct.NodeID, &port, &ct.Strategy, &ct.Inx, &ct.Protocol); err != nil { + return nil, err + } + if port.Valid { + ct.Port = int(port.Int64) + } + chainTunnels = append(chainTunnels, ct) + } + return chainTunnels, rows.Err() +} + +func (r *Repository) exportForwards() ([]ForwardBackup, error) { + rows, err := r.db.Query(` + SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx + FROM forward ORDER BY id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + var forwards []ForwardBackup + for rows.Next() { + var f ForwardBackup + if err := rows.Scan(&f.ID, &f.UserID, &f.UserName, &f.Name, &f.TunnelID, &f.RemoteAddr, &f.Strategy, &f.InFlow, &f.OutFlow, &f.CreatedTime, &f.UpdatedTime, &f.Status, &f.Inx); err != nil { + return nil, err + } + forwards = append(forwards, f) + } + return forwards, rows.Err() +} + +func (r *Repository) exportUserTunnels() ([]UserTunnelBackup, error) { + 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 + FROM user_tunnel ORDER BY id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + var userTunnels []UserTunnelBackup + for rows.Next() { + var ut UserTunnelBackup + var speedID sql.NullInt64 + if err := rows.Scan(&ut.ID, &ut.UserID, &ut.TunnelID, &speedID, &ut.Num, &ut.Flow, &ut.InFlow, &ut.OutFlow, &ut.FlowResetTime, &ut.ExpTime, &ut.Status); err != nil { + return nil, err + } + if speedID.Valid { + ut.SpeedID = speedID.Int64 + } + userTunnels = append(userTunnels, ut) + } + return userTunnels, rows.Err() +} + +func (r *Repository) exportSpeedLimits() ([]SpeedLimitBackup, error) { + rows, err := r.db.Query(` + SELECT id, name, speed, tunnel_id, tunnel_name, created_time, updated_time, status + FROM speed_limit ORDER BY id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + var speedLimits []SpeedLimitBackup + for rows.Next() { + var sl SpeedLimitBackup + var updatedTime sql.NullInt64 + if err := rows.Scan(&sl.ID, &sl.Name, &sl.Speed, &sl.TunnelID, &sl.TunnelName, &sl.CreatedTime, &updatedTime, &sl.Status); err != nil { + return nil, err + } + if updatedTime.Valid { + sl.UpdatedTime = updatedTime.Int64 + } + speedLimits = append(speedLimits, sl) + } + return speedLimits, rows.Err() +} + +func (r *Repository) exportTunnelGroups() ([]TunnelGroupBackup, error) { + rows, err := r.db.Query(` + SELECT id, name, created_time, updated_time, status + FROM tunnel_group ORDER BY id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + var groups []TunnelGroupBackup + for rows.Next() { + var tg TunnelGroupBackup + if err := rows.Scan(&tg.ID, &tg.Name, &tg.CreatedTime, &tg.UpdatedTime, &tg.Status); err != nil { + return nil, err + } + // Get tunnel IDs for this group + tunnelRows, err := r.db.Query(`SELECT tunnel_id FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, tg.ID) + if err != nil { + return nil, err + } + for tunnelRows.Next() { + var tunnelID int64 + if err := tunnelRows.Scan(&tunnelID); err != nil { + tunnelRows.Close() + return nil, err + } + tg.Tunnels = append(tg.Tunnels, tunnelID) + } + tunnelRows.Close() + groups = append(groups, tg) + } + return groups, rows.Err() +} + +func (r *Repository) exportUserGroups() ([]UserGroupBackup, error) { + rows, err := r.db.Query(` + SELECT id, name, created_time, updated_time, status + FROM user_group ORDER BY id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + var groups []UserGroupBackup + for rows.Next() { + var ug UserGroupBackup + if err := rows.Scan(&ug.ID, &ug.Name, &ug.CreatedTime, &ug.UpdatedTime, &ug.Status); err != nil { + return nil, err + } + // Get user IDs for this group + userRows, err := r.db.Query(`SELECT user_id FROM user_group_user WHERE user_group_id = ?`, ug.ID) + if err != nil { + return nil, err + } + for userRows.Next() { + var userID int64 + if err := userRows.Scan(&userID); err != nil { + userRows.Close() + return nil, err + } + ug.Users = append(ug.Users, userID) + } + userRows.Close() + groups = append(groups, ug) + } + return groups, rows.Err() +} + +func (r *Repository) exportPermissions() ([]PermissionBackup, error) { + rows, err := r.db.Query(` + SELECT id, user_group_id, tunnel_group_id, created_time, created_by_group + FROM group_permission ORDER BY id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + 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 { + return nil, err + } + // 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 { + return nil, err + } + for grantRows.Next() { + var g PermissionGrantBackup + if err := grantRows.Scan(&g.ID, &g.UserGroupID, &g.TunnelGroupID, &g.UserTunnelID, &g.CreatedTime, &g.CreatedByGroup); err != nil { + grantRows.Close() + return nil, err + } + p.Grants = append(p.Grants, g) + } + grantRows.Close() + permissions = append(permissions, p) + } + return permissions, rows.Err() +} + +// ============ Import Methods ============ + +// ImportResult contains the result of an import operation +type ImportResult struct { + UsersImported int `json:"usersImported"` + NodesImported int `json:"nodesImported"` + TunnelsImported int `json:"tunnelsImported"` + ForwardsImported int `json:"forwardsImported"` + UserTunnelsImported int `json:"userTunnelsImported"` + SpeedLimitsImported int `json:"speedLimitsImported"` + TunnelGroupsImported int `json:"tunnelGroupsImported"` + UserGroupsImported int `json:"userGroupsImported"` + PermissionsImported int `json:"permissionsImported"` + ConfigsImported int `json:"configsImported"` +} + +// Import imports data from BackupData +func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult, error) { + result := &ImportResult{} + + typeSet := make(map[string]bool) + for _, t := range types { + typeSet[t] = true + } + + now := unixMilliNow() + + if typeSet["users"] && len(backup.Users) > 0 { + count, err := r.importUsers(backup.Users, now) + if err != nil { + return nil, fmt.Errorf("import users failed: %w", err) + } + result.UsersImported = count + } + + if typeSet["nodes"] && len(backup.Nodes) > 0 { + count, err := r.importNodes(backup.Nodes, now) + if err != nil { + return nil, fmt.Errorf("import nodes failed: %w", err) + } + result.NodesImported = count + } + + if typeSet["tunnels"] && len(backup.Tunnels) > 0 { + count, err := r.importTunnels(backup.Tunnels, now) + if err != nil { + return nil, fmt.Errorf("import tunnels failed: %w", err) + } + result.TunnelsImported = count + } + + if typeSet["forwards"] && len(backup.Forwards) > 0 { + count, err := r.importForwards(backup.Forwards, now) + if err != nil { + return nil, fmt.Errorf("import forwards failed: %w", err) + } + result.ForwardsImported = count + } + + if typeSet["userTunnels"] && len(backup.UserTunnels) > 0 { + count, err := r.importUserTunnels(backup.UserTunnels, now) + if err != nil { + return nil, fmt.Errorf("import user tunnels failed: %w", err) + } + result.UserTunnelsImported = count + } + + if typeSet["speedLimits"] && len(backup.SpeedLimits) > 0 { + count, err := r.importSpeedLimits(backup.SpeedLimits, now) + if err != nil { + return nil, fmt.Errorf("import speed limits failed: %w", err) + } + result.SpeedLimitsImported = count + } + + if typeSet["tunnelGroups"] && len(backup.TunnelGroups) > 0 { + count, err := r.importTunnelGroups(backup.TunnelGroups, now) + if err != nil { + return nil, fmt.Errorf("import tunnel groups failed: %w", err) + } + result.TunnelGroupsImported = count + } + + if typeSet["userGroups"] && len(backup.UserGroups) > 0 { + count, err := r.importUserGroups(backup.UserGroups, now) + if err != nil { + return nil, fmt.Errorf("import user groups failed: %w", err) + } + result.UserGroupsImported = count + } + + if typeSet["permissions"] && len(backup.Permissions) > 0 { + count, err := r.importPermissions(backup.Permissions, now) + if err != nil { + return nil, fmt.Errorf("import permissions failed: %w", err) + } + result.PermissionsImported = count + } + + if typeSet["configs"] && len(backup.Configs) > 0 { + count, err := r.importConfigs(backup.Configs, now) + if err != nil { + return nil, fmt.Errorf("import configs failed: %w", err) + } + result.ConfigsImported = count + } + + return result, nil +} + +func (r *Repository) importUsers(users []UserBackup, now int64) (int, error) { + count := 0 + for _, u := range users { + // Check if user exists + exists, err := r.UsernameExists(u.User) + if err != nil { + 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++ + } + return count, nil +} + +func (r *Repository) UsernameExists(username string) (bool, error) { + var count int + err := r.db.QueryRow(`SELECT COUNT(1) FROM user WHERE user = ?`, username).Scan(&count) + if err != nil { + return false, err + } + return count > 0, nil +} + +func (r *Repository) importNodes(nodes []NodeBackup, now int64) (int, error) { + count := 0 + for _, n := range nodes { + // Check if node exists + _, 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) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + name = excluded.name, + secret = excluded.secret, + server_ip = excluded.server_ip, + server_ip_v4 = excluded.server_ip_v4, + server_ip_v6 = excluded.server_ip_v6, + port = excluded.port, + interface_name = excluded.interface_name, + version = excluded.version, + http = excluded.http, + tls = excluded.tls, + socks = excluded.socks, + updated_time = excluded.updated_time, + status = excluded.status, + tcp_listen_addr = excluded.tcp_listen_addr, + udp_listen_addr = excluded.udp_listen_addr, + inx = excluded.inx, + is_remote = excluded.is_remote, + remote_url = excluded.remote_url, + remote_token = excluded.remote_token, + remote_config = excluded.remote_config + `, n.ID, n.Name, n.Secret, n.ServerIP, n.ServerIPv4, n.ServerIPv6, n.Port, n.InterfaceName, n.Version, n.HTTP, n.TLS, n.Socks, n.CreatedTime, now, n.Status, n.TCPListenAddr, n.UDPListenAddr, n.Inx, n.IsRemote, n.RemoteURL, n.RemoteToken, n.RemoteConfig) + if err != nil { + return count, err + } + count++ + } + return count, nil +} + +func (r *Repository) importTunnels(tunnels []TunnelBackup, now int64) (int, error) { + count := 0 + for _, t := range tunnels { + _, err := r.db.Exec(` + INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + name = excluded.name, + traffic_ratio = excluded.traffic_ratio, + type = excluded.type, + protocol = excluded.protocol, + flow = excluded.flow, + updated_time = excluded.updated_time, + status = excluded.status, + in_ip = excluded.in_ip, + inx = excluded.inx + `, t.ID, t.Name, t.TrafficRatio, t.Type, t.Protocol, t.Flow, t.CreatedTime, now, t.Status, t.InIP, t.Inx) + if err != nil { + return count, err + } + // Import chain tunnels + if len(t.ChainTunnels) > 0 { + for _, ct := range t.ChainTunnels { + _, err = r.db.Exec(` + INSERT INTO chain_tunnel(id, tunnel_id, chain_type, node_id, port, strategy, inx, protocol) + VALUES(?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + chain_type = excluded.chain_type, + node_id = excluded.node_id, + port = excluded.port, + strategy = excluded.strategy, + inx = excluded.inx, + protocol = excluded.protocol + `, ct.ID, ct.TunnelID, ct.ChainType, ct.NodeID, ct.Port, ct.Strategy, ct.Inx, ct.Protocol) + if err != nil { + return count, err + } + } + } + count++ + } + return count, nil +} + +func (r *Repository) importForwards(forwards []ForwardBackup, now int64) (int, error) { + count := 0 + for _, f := range forwards { + _, err := r.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) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + user_id = excluded.user_id, + user_name = excluded.user_name, + name = excluded.name, + tunnel_id = excluded.tunnel_id, + remote_addr = excluded.remote_addr, + strategy = excluded.strategy, + in_flow = excluded.in_flow, + out_flow = excluded.out_flow, + updated_time = excluded.updated_time, + status = excluded.status, + inx = excluded.inx + `, f.ID, f.UserID, f.UserName, f.Name, f.TunnelID, f.RemoteAddr, f.Strategy, f.InFlow, f.OutFlow, f.CreatedTime, now, f.Status, f.Inx) + if err != nil { + return count, err + } + count++ + } + return count, nil +} + +func (r *Repository) importUserTunnels(userTunnels []UserTunnelBackup, now int64) (int, error) { + count := 0 + for _, ut := range userTunnels { + var speedID interface{} + if ut.SpeedID > 0 { + speedID = ut.SpeedID + } + _, err := r.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) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + user_id = excluded.user_id, + tunnel_id = excluded.tunnel_id, + speed_id = excluded.speed_id, + num = excluded.num, + flow = excluded.flow, + in_flow = excluded.in_flow, + out_flow = excluded.out_flow, + flow_reset_time = excluded.flow_reset_time, + exp_time = excluded.exp_time, + status = excluded.status + `, ut.ID, ut.UserID, ut.TunnelID, speedID, ut.Num, ut.Flow, ut.InFlow, ut.OutFlow, ut.FlowResetTime, ut.ExpTime, ut.Status) + if err != nil { + return count, err + } + count++ + } + return count, nil +} + +func (r *Repository) importSpeedLimits(speedLimits []SpeedLimitBackup, now int64) (int, error) { + count := 0 + for _, sl := range speedLimits { + _, err := r.db.Exec(` + INSERT INTO speed_limit(id, name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) + VALUES(?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + name = excluded.name, + speed = excluded.speed, + tunnel_id = excluded.tunnel_id, + tunnel_name = excluded.tunnel_name, + updated_time = excluded.updated_time, + status = excluded.status + `, sl.ID, sl.Name, sl.Speed, sl.TunnelID, sl.TunnelName, sl.CreatedTime, now, sl.Status) + if err != nil { + return count, err + } + count++ + } + return count, nil +} + +func (r *Repository) importTunnelGroups(tunnelGroups []TunnelGroupBackup, now int64) (int, error) { + count := 0 + for _, tg := range tunnelGroups { + _, err := r.db.Exec(` + INSERT INTO tunnel_group(id, name, created_time, updated_time, status) + VALUES(?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + name = excluded.name, + updated_time = excluded.updated_time, + status = excluded.status + `, tg.ID, tg.Name, tg.CreatedTime, now, tg.Status) + if err != nil { + return count, err + } + // Update tunnel group memberships + _, err = r.db.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, tg.ID) + if err != nil { + return count, err + } + for _, tunnelID := range tg.Tunnels { + _, err = r.db.Exec(` + INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) + VALUES(?, ?, ?) + `, tg.ID, tunnelID, now) + if err != nil { + return count, err + } + } + count++ + } + return count, nil +} + +func (r *Repository) importUserGroups(userGroups []UserGroupBackup, now int64) (int, error) { + count := 0 + for _, ug := range userGroups { + _, err := r.db.Exec(` + INSERT INTO user_group(id, name, created_time, updated_time, status) + VALUES(?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + name = excluded.name, + updated_time = excluded.updated_time, + status = excluded.status + `, ug.ID, ug.Name, ug.CreatedTime, now, ug.Status) + if err != nil { + return count, err + } + // Update user group memberships + _, err = r.db.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, ug.ID) + if err != nil { + return count, err + } + for _, userID := range ug.Users { + _, err = r.db.Exec(` + INSERT INTO user_group_user(user_group_id, user_id, created_time) + VALUES(?, ?, ?) + `, ug.ID, userID, now) + if err != nil { + return count, err + } + } + count++ + } + return count, nil +} + +func (r *Repository) importPermissions(permissions []PermissionBackup, now int64) (int, error) { + count := 0 + for _, p := range permissions { + _, err := r.db.Exec(` + INSERT INTO group_permission(id, user_group_id, tunnel_group_id, created_time, created_by_group) + VALUES(?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + user_group_id = excluded.user_group_id, + tunnel_group_id = excluded.tunnel_group_id, + created_by_group = excluded.created_by_group + `, p.ID, p.UserGroupID, p.TunnelGroupID, p.CreatedTime, p.CreatedByGroup) + if err != nil { + return count, err + } + // Import grants + for _, g := range p.Grants { + _, err = r.db.Exec(` + INSERT INTO group_permission_grant(id, user_group_id, tunnel_group_id, user_tunnel_id, created_time, created_by_group) + VALUES(?, ?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + user_tunnel_id = excluded.user_tunnel_id, + created_by_group = excluded.created_by_group + `, g.ID, g.UserGroupID, g.TunnelGroupID, g.UserTunnelID, g.CreatedTime, g.CreatedByGroup) + if err != nil { + return count, err + } + } + count++ + } + return count, nil +} + +func (r *Repository) importConfigs(configs map[string]string, now int64) (int, error) { + count := 0 + for name, value := range configs { + err := r.UpsertConfig(name, value, now) + if err != nil { + return count, err + } + count++ + } + return count, nil +} From f879a58bb465d47f75264fd36caadf7680dc05d7 Mon Sep 17 00:00:00 2001 From: sagit Date: Fri, 13 Feb 2026 03:06:32 +0000 Subject: [PATCH 02/12] feat(frontend): add backup export and import UI Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-opencode) Co-authored-by: Sisyphus --- vite-frontend/src/api/index.ts | 43 ++++++ vite-frontend/src/pages/config.tsx | 205 ++++++++++++++++++++++++++++- 2 files changed, 245 insertions(+), 3 deletions(-) diff --git a/vite-frontend/src/api/index.ts b/vite-frontend/src/api/index.ts index f0b56d2..0863be0 100644 --- a/vite-frontend/src/api/index.ts +++ b/vite-frontend/src/api/index.ts @@ -210,3 +210,46 @@ export const importRemoteNode = (data: { remoteUrl: string; token: string; }) => Network.post("/federation/node/import", data); + +import axios from "axios"; + +export interface BackupTypes { + users?: boolean; + nodes?: boolean; + tunnels?: boolean; + forwards?: boolean; + userTunnels?: boolean; + speedLimits?: boolean; + tunnelGroups?: boolean; + userGroups?: boolean; + permissions?: boolean; + configs?: boolean; +} + +export const exportBackup = async (types: string[] = []) => { + const token = window.localStorage.getItem("token"); + const baseURL = axios.defaults.baseURL || "/api/v1/"; + + const response = await axios.post(`${baseURL}/backup/export`, { types }, { + headers: { + Authorization: token, + "Content-Type": "application/json", + }, + responseType: "blob", + }); + + const url = window.URL.createObjectURL(new Blob([response.data])); + const link = document.createElement("a"); + link.href = url; + const timestamp = new Date().toISOString().slice(0, 19).replace(/[:-]/g, ""); + link.setAttribute("download", `backup_${timestamp}.json`); + document.body.appendChild(link); + link.click(); + document.body.removeChild(link); + window.URL.revokeObjectURL(url); +}; + +export const importBackup = (data: { + types: string[]; + [key: string]: any; +}) => Network.post("/backup/import", data); diff --git a/vite-frontend/src/pages/config.tsx b/vite-frontend/src/pages/config.tsx index 4ee5498..71bc198 100644 --- a/vite-frontend/src/pages/config.tsx +++ b/vite-frontend/src/pages/config.tsx @@ -1,4 +1,4 @@ -import { useState, useEffect } from "react"; +import { useState, useEffect, useRef } from "react"; import { useNavigate } from "react-router-dom"; import { Button } from "@heroui/button"; import { Card, CardBody, CardHeader } from "@heroui/card"; @@ -7,9 +7,10 @@ import { Spinner } from "@heroui/spinner"; import { Divider } from "@heroui/divider"; import { Switch } from "@heroui/switch"; import { Select, SelectItem } from "@heroui/select"; +import { Checkbox, CheckboxGroup } from "@heroui/checkbox"; import toast from "react-hot-toast"; -import { updateConfigs } from "@/api"; +import { updateConfigs, exportBackup, importBackup } from "@/api"; import { SettingsIcon } from "@/components/icons"; import { isAdmin } from "@/utils/auth"; import { @@ -130,12 +131,19 @@ export default function ConfigPage() { useState>(initialConfigs); const [loading, setLoading] = useState( Object.keys(initialConfigs).length === 0, - ); // 如果有缓存数据,不显示loading + ); const [saving, setSaving] = useState(false); const [hasChanges, setHasChanges] = useState(false); const [originalConfigs, setOriginalConfigs] = useState>(initialConfigs); + const [exportTypes, setExportTypes] = useState([]); + const [importTypes, setImportTypes] = useState([]); + const [exporting, setExporting] = useState(false); + const [importing, setImporting] = useState(false); + const [importFileName, setImportFileName] = useState(""); + const fileInputRef = useRef(null); + // 权限检查 useEffect(() => { if (!isAdmin()) { @@ -331,6 +339,60 @@ export default function ConfigPage() { } }; + const handleExport = async () => { + if (exportTypes.length === 0) { + toast.error("请至少选择一种数据类型"); + return; + } + setExporting(true); + try { + await exportBackup(exportTypes); + toast.success("导出成功"); + } catch { + toast.error("导出失败,请重试"); + } finally { + setExporting(false); + } + }; + + const handleFileChange = async (e: React.ChangeEvent) => { + const file = e.target.files?.[0]; + if (!file) return; + + if (importTypes.length === 0) { + toast.error("请先选择要导入的数据类型"); + return; + } + + setImportFileName(file.name); + setImporting(true); + + try { + const text = await file.text(); + const data = JSON.parse(text); + + const response = await importBackup({ + types: importTypes, + ...data, + }); + + if (response.code === 0) { + toast.success(`导入成功: ${JSON.stringify(response.data)}`); + setImportTypes([]); + setImportFileName(""); + } else { + toast.error("导入失败: " + response.msg); + } + } catch { + toast.error("导入失败,请检查文件格式"); + } finally { + setImporting(false); + if (fileInputRef.current) { + fileInputRef.current.value = ""; + } + } + }; + if (loading) { return (
@@ -427,6 +489,143 @@ export default function ConfigPage() { )} + + {/* 备份与恢复 */} + + +
+
+

数据备份与恢复

+

+ 导出或导入系统数据,支持选择特定数据类型 +

+
+
+
+ + + + + {/* 导出部分 */} +
+

导出数据

+

+ 选择要导出的数据类型,导出为 JSON 格式文件 +

+ + setExportTypes(values as string[])} + > + 用户 + 节点 + 隧道 + 转发 + 用户隧道权限 + 限速规则 + 隧道分组 + 用户分组 + 分组权限 + 系统配置 + + +
+ + + +
+
+ + + + {/* 导入部分 */} +
+

导入数据

+

+ 选择要导入的数据类型,支持从备份文件恢复数据 +

+ + setImportTypes(values as string[])} + > + 用户 + 节点 + 隧道 + 转发 + 用户隧道权限 + 限速规则 + 隧道分组 + 用户分组 + 分组权限 + 系统配置 + + + + +
+ + {importFileName && ( + + 已选择: {importFileName} + + )} +
+
+
+
); } From dd206ced14c29caf1470e616b155888c7f742bff Mon Sep 17 00:00:00 2001 From: sagit Date: Fri, 13 Feb 2026 03:26:21 +0000 Subject: [PATCH 03/12] fix(permission): revoke inherited tunnel access after group unbind/removal --- go-backend/internal/http/handler/mutations.go | 178 +++++++++++++- .../group_permission_contract_test.go | 219 ++++++++++++++++++ 2 files changed, 394 insertions(+), 3 deletions(-) create mode 100644 go-backend/tests/contract/group_permission_contract_test.go diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index 49cb324..e4719c4 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -1713,10 +1713,19 @@ func (h *Handler) groupUserAssign(w http.ResponseWriter, r *http.Request) { return } defer func() { _ = tx.Rollback() }() + previousUserIDs, err := queryInt64ListTx(tx, `SELECT user_id FROM user_group_user WHERE user_group_id = ?`, req.GroupID) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } _, _ = tx.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, req.GroupID) for _, uid := range req.UserIDs { _, _ = tx.Exec(`INSERT INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?) ON CONFLICT DO NOTHING`, req.GroupID, uid, time.Now().UnixMilli()) } + if err := revokeGroupGrantsForRemovedUsersTx(tx, req.GroupID, previousUserIDs, req.UserIDs); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } if err := tx.Commit(); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return @@ -1748,10 +1757,35 @@ func (h *Handler) groupPermissionRemove(w http.ResponseWriter, r *http.Request) if id <= 0 { return } + tx, err := h.repo.DB().Begin() + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + defer func() { _ = tx.Rollback() }() + var ug, tg int64 - _ = h.repo.DB().QueryRow(`SELECT user_group_id, tunnel_group_id FROM group_permission WHERE id = ?`, id).Scan(&ug, &tg) - _, _ = h.repo.DB().Exec(`DELETE FROM group_permission WHERE id = ?`, id) - _, _ = h.repo.DB().Exec(`DELETE FROM group_permission_grant WHERE user_group_id = ? AND tunnel_group_id = ?`, ug, tg) + err = tx.QueryRow(`SELECT user_group_id, tunnel_group_id FROM group_permission WHERE id = ?`, id).Scan(&ug, &tg) + if err != nil && err != sql.ErrNoRows { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + + if _, err := tx.Exec(`DELETE FROM group_permission WHERE id = ?`, id); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if err == nil { + if err := revokeGroupPermissionPairTx(tx, ug, tg); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + } + + if err := tx.Commit(); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } response.WriteJSON(w, response.OKEmpty()) } @@ -1908,6 +1942,144 @@ func queryInt64List(db *store.DB, q string, args ...interface{}) ([]int64, error return out, rows.Err() } +func queryInt64ListTx(tx *store.Tx, q string, args ...interface{}) ([]int64, error) { + rows, err := tx.Query(q, args...) + if err != nil { + return nil, err + } + defer rows.Close() + out := make([]int64, 0) + for rows.Next() { + var v int64 + if err := rows.Scan(&v); err != nil { + return nil, err + } + out = append(out, v) + } + return out, rows.Err() +} + +func revokeGroupGrantsForRemovedUsersTx(tx *store.Tx, userGroupID int64, previousUserIDs, currentUserIDs []int64) error { + currentSet := make(map[int64]struct{}, len(currentUserIDs)) + for _, uid := range currentUserIDs { + if uid > 0 { + currentSet[uid] = struct{}{} + } + } + + removedUserIDs := make([]int64, 0) + for _, uid := range previousUserIDs { + if uid <= 0 { + continue + } + if _, ok := currentSet[uid]; !ok { + removedUserIDs = append(removedUserIDs, uid) + } + } + if len(removedUserIDs) == 0 { + return nil + } + + for _, userID := range removedUserIDs { + rows, err := tx.Query(` + SELECT g.user_tunnel_id, g.created_by_group + FROM group_permission_grant g + JOIN user_tunnel ut ON ut.id = g.user_tunnel_id + WHERE g.user_group_id = ? AND ut.user_id = ? + `, userGroupID, userID) + if err != nil { + return err + } + + groupCreatedTunnelIDs := make(map[int64]struct{}) + for rows.Next() { + var userTunnelID int64 + var createdByGroup int + if err := rows.Scan(&userTunnelID, &createdByGroup); err != nil { + rows.Close() + return err + } + if createdByGroup == 1 && userTunnelID > 0 { + groupCreatedTunnelIDs[userTunnelID] = struct{}{} + } + } + if err := rows.Err(); err != nil { + rows.Close() + return err + } + rows.Close() + + if _, err := tx.Exec(` + DELETE FROM group_permission_grant + WHERE user_group_id = ? + AND user_tunnel_id IN (SELECT id FROM user_tunnel WHERE user_id = ?) + `, userGroupID, userID); err != nil { + return err + } + + for userTunnelID := range groupCreatedTunnelIDs { + var remaining int + if err := tx.QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&remaining); err != nil { + return err + } + if remaining == 0 { + if _, err := tx.Exec(`DELETE FROM user_tunnel WHERE id = ?`, userTunnelID); err != nil { + return err + } + } + } + } + + return nil +} + +func revokeGroupPermissionPairTx(tx *store.Tx, userGroupID, tunnelGroupID int64) error { + rows, err := tx.Query(` + SELECT user_tunnel_id, created_by_group + FROM group_permission_grant + WHERE user_group_id = ? AND tunnel_group_id = ? + `, userGroupID, tunnelGroupID) + if err != nil { + return err + } + + groupCreatedTunnelIDs := make(map[int64]struct{}) + for rows.Next() { + var userTunnelID int64 + var createdByGroup int + if err := rows.Scan(&userTunnelID, &createdByGroup); err != nil { + rows.Close() + return err + } + if createdByGroup == 1 && userTunnelID > 0 { + groupCreatedTunnelIDs[userTunnelID] = struct{}{} + } + } + if err := rows.Err(); err != nil { + rows.Close() + return err + } + rows.Close() + + if _, err := tx.Exec(`DELETE FROM group_permission_grant WHERE user_group_id = ? AND tunnel_group_id = ?`, userGroupID, tunnelGroupID); err != nil { + return err + } + + for userTunnelID := range groupCreatedTunnelIDs { + var remaining int + if err := tx.QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&remaining); err != nil { + return err + } + if remaining == 0 { + if _, err := tx.Exec(`DELETE FROM user_tunnel WHERE id = ?`, userTunnelID); err != nil { + return err + } + } + } + + return nil +} + func queryPairs(db *store.DB, q string, args ...interface{}) ([][2]int64, error) { rows, err := db.Query(q, args...) if err != nil { diff --git a/go-backend/tests/contract/group_permission_contract_test.go b/go-backend/tests/contract/group_permission_contract_test.go new file mode 100644 index 0000000..530acdd --- /dev/null +++ b/go-backend/tests/contract/group_permission_contract_test.go @@ -0,0 +1,219 @@ +package contract_test + +import ( + "bytes" + "net/http" + "net/http/httptest" + "testing" + "time" + + "go-backend/internal/auth" +) + +func TestGroupUserUnbindRevokesInheritedTunnelPermission(t *testing.T) { + secret := "contract-jwt-secret" + router, repo := setupContractRouter(t, secret) + now := time.Now().UnixMilli() + + if _, err := repo.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(200, 'group_user_contract', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) + `, now, now); err != nil { + t.Fatalf("insert test user: %v", err) + } + + tunnelRes, err := repo.DB().Exec(` + INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES('group-contract-tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0) + `, now, now) + if err != nil { + t.Fatalf("insert tunnel: %v", err) + } + tunnelID, err := tunnelRes.LastInsertId() + if err != nil { + t.Fatalf("read tunnel id: %v", err) + } + + ugRes, err := repo.DB().Exec(`INSERT INTO user_group(name, created_time, updated_time, status) VALUES('ug-contract', ?, ?, 1)`, now, now) + if err != nil { + t.Fatalf("insert user_group: %v", err) + } + userGroupID, err := ugRes.LastInsertId() + if err != nil { + t.Fatalf("read user_group id: %v", err) + } + + tgRes, err := repo.DB().Exec(`INSERT INTO tunnel_group(name, created_time, updated_time, status) VALUES('tg-contract', ?, ?, 1)`, now, now) + if err != nil { + t.Fatalf("insert tunnel_group: %v", err) + } + tunnelGroupID, err := tgRes.LastInsertId() + if err != nil { + t.Fatalf("read tunnel_group id: %v", err) + } + + if _, err := repo.DB().Exec(`INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) VALUES(?, ?, ?)`, tunnelGroupID, tunnelID, now); err != nil { + t.Fatalf("insert tunnel_group_tunnel: %v", err) + } + if _, err := repo.DB().Exec(`INSERT INTO group_permission(user_group_id, tunnel_group_id, created_time) VALUES(?, ?, ?)`, userGroupID, tunnelGroupID, now); err != nil { + t.Fatalf("insert group_permission: %v", err) + } + + adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate admin token: %v", err) + } + + bindReq := httptest.NewRequest(http.MethodPost, "/api/v1/group/user/assign", bytes.NewBufferString(`{"groupId":`+jsonNumber(userGroupID)+`,"userIds":[200]}`)) + bindReq.Header.Set("Authorization", adminToken) + bindRes := httptest.NewRecorder() + router.ServeHTTP(bindRes, bindReq) + assertCode(t, bindRes, 0) + + var userTunnelID int64 + if err := repo.DB().QueryRow(`SELECT id FROM user_tunnel WHERE user_id = 200 AND tunnel_id = ?`, tunnelID).Scan(&userTunnelID); err != nil { + t.Fatalf("query user_tunnel after bind: %v", err) + } + + var grantCount int + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&grantCount); err != nil { + t.Fatalf("query group_permission_grant after bind: %v", err) + } + if grantCount == 0 { + t.Fatalf("expected non-zero grants after bind") + } + + unbindReq := httptest.NewRequest(http.MethodPost, "/api/v1/group/user/assign", bytes.NewBufferString(`{"groupId":`+jsonNumber(userGroupID)+`,"userIds":[]}`)) + unbindReq.Header.Set("Authorization", adminToken) + unbindRes := httptest.NewRecorder() + router.ServeHTTP(unbindRes, unbindReq) + assertCode(t, unbindRes, 0) + + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&grantCount); err != nil { + t.Fatalf("query group_permission_grant after unbind: %v", err) + } + if grantCount != 0 { + t.Fatalf("expected grants revoked after unbind, got %d", grantCount) + } + + var userTunnelCount int + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM user_tunnel WHERE id = ?`, userTunnelID).Scan(&userTunnelCount); err != nil { + t.Fatalf("query user_tunnel after unbind: %v", err) + } + if userTunnelCount != 0 { + t.Fatalf("expected user_tunnel revoked after unbind, got %d", userTunnelCount) + } +} + +func TestGroupPermissionRemoveRevokesInheritedTunnelPermission(t *testing.T) { + secret := "contract-jwt-secret" + router, repo := setupContractRouter(t, secret) + now := time.Now().UnixMilli() + + if _, err := repo.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(201, 'group_user_permission_remove', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) + `, now, now); err != nil { + t.Fatalf("insert test user: %v", err) + } + + tunnelRes, err := repo.DB().Exec(` + INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES('group-remove-tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0) + `, now, now) + if err != nil { + t.Fatalf("insert tunnel: %v", err) + } + tunnelID, err := tunnelRes.LastInsertId() + if err != nil { + t.Fatalf("read tunnel id: %v", err) + } + + ugRes, err := repo.DB().Exec(`INSERT INTO user_group(name, created_time, updated_time, status) VALUES('ug-remove-contract', ?, ?, 1)`, now, now) + if err != nil { + t.Fatalf("insert user_group: %v", err) + } + userGroupID, err := ugRes.LastInsertId() + if err != nil { + t.Fatalf("read user_group id: %v", err) + } + + tgRes, err := repo.DB().Exec(`INSERT INTO tunnel_group(name, created_time, updated_time, status) VALUES('tg-remove-contract', ?, ?, 1)`, now, now) + if err != nil { + t.Fatalf("insert tunnel_group: %v", err) + } + tunnelGroupID, err := tgRes.LastInsertId() + if err != nil { + t.Fatalf("read tunnel_group id: %v", err) + } + + adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate admin token: %v", err) + } + + assignTunnelReq := httptest.NewRequest(http.MethodPost, "/api/v1/group/tunnel/assign", bytes.NewBufferString(`{"groupId":`+jsonNumber(tunnelGroupID)+`,"tunnelIds":[`+jsonNumber(tunnelID)+`]}`)) + assignTunnelReq.Header.Set("Authorization", adminToken) + assignTunnelRes := httptest.NewRecorder() + router.ServeHTTP(assignTunnelRes, assignTunnelReq) + assertCode(t, assignTunnelRes, 0) + + assignUserReq := httptest.NewRequest(http.MethodPost, "/api/v1/group/user/assign", bytes.NewBufferString(`{"groupId":`+jsonNumber(userGroupID)+`,"userIds":[201]}`)) + assignUserReq.Header.Set("Authorization", adminToken) + assignUserRes := httptest.NewRecorder() + router.ServeHTTP(assignUserRes, assignUserReq) + assertCode(t, assignUserRes, 0) + + assignPermissionReq := httptest.NewRequest(http.MethodPost, "/api/v1/group/permission/assign", bytes.NewBufferString(`{"userGroupId":`+jsonNumber(userGroupID)+`,"tunnelGroupId":`+jsonNumber(tunnelGroupID)+`}`)) + assignPermissionReq.Header.Set("Authorization", adminToken) + assignPermissionRes := httptest.NewRecorder() + router.ServeHTTP(assignPermissionRes, assignPermissionReq) + assertCode(t, assignPermissionRes, 0) + + var permissionID int64 + if err := repo.DB().QueryRow(`SELECT id FROM group_permission WHERE user_group_id = ? AND tunnel_group_id = ?`, userGroupID, tunnelGroupID).Scan(&permissionID); err != nil { + t.Fatalf("query group_permission id: %v", err) + } + + var userTunnelID int64 + if err := repo.DB().QueryRow(`SELECT id FROM user_tunnel WHERE user_id = 201 AND tunnel_id = ?`, tunnelID).Scan(&userTunnelID); err != nil { + t.Fatalf("query user_tunnel after assign: %v", err) + } + + var grantCount int + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&grantCount); err != nil { + t.Fatalf("query group_permission_grant after assign: %v", err) + } + if grantCount == 0 { + t.Fatalf("expected non-zero grants after permission assign") + } + + removeReq := httptest.NewRequest(http.MethodPost, "/api/v1/group/permission/remove", bytes.NewBufferString(`{"id":`+jsonNumber(permissionID)+`}`)) + removeReq.Header.Set("Authorization", adminToken) + removeRes := httptest.NewRecorder() + router.ServeHTTP(removeRes, removeReq) + assertCode(t, removeRes, 0) + + var permissionCount int + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM group_permission WHERE id = ?`, permissionID).Scan(&permissionCount); err != nil { + t.Fatalf("query group_permission after remove: %v", err) + } + if permissionCount != 0 { + t.Fatalf("expected group_permission removed, got %d", permissionCount) + } + + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&grantCount); err != nil { + t.Fatalf("query group_permission_grant after remove: %v", err) + } + if grantCount != 0 { + t.Fatalf("expected grants removed after permission remove, got %d", grantCount) + } + + var userTunnelCount int + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM user_tunnel WHERE id = ?`, userTunnelID).Scan(&userTunnelCount); err != nil { + t.Fatalf("query user_tunnel after permission remove: %v", err) + } + if userTunnelCount != 0 { + t.Fatalf("expected user_tunnel revoked after permission remove, got %d", userTunnelCount) + } +} From 2d2ca389e33722fe021438a69bc8c4860655e1c2 Mon Sep 17 00:00:00 2001 From: sagit Date: Fri, 13 Feb 2026 05:44:50 +0000 Subject: [PATCH 04/12] 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 --- go-backend/internal/http/handler/handler.go | 9 +- .../internal/store/sqlite/repository.go | 156 +++++++++--------- 2 files changed, 88 insertions(+), 77 deletions(-) diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index 7e24ada..8a26877 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -1203,6 +1203,12 @@ func (h *Handler) backupImport(w http.ResponseWriter, r *http.Request) { return } + autoBackup, err := h.repo.ExportAll() + if err != nil { + response.WriteJSON(w, response.Err(-2, fmt.Sprintf("导入前自动备份失败: %v", err))) + return + } + var backup sqlite.BackupData if err := decodeJSON(r.Body, &backup); err != nil { 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) if err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) + response.WriteJSON(w, response.Err(-2, fmt.Sprintf("导入失败: %v", err))) return } + result.AutoBackup = autoBackup response.WriteJSON(w, response.OK(result)) } diff --git a/go-backend/internal/store/sqlite/repository.go b/go-backend/internal/store/sqlite/repository.go index afc7809..d72a6ec 100644 --- a/go-backend/internal/store/sqlite/repository.go +++ b/go-backend/internal/store/sqlite/repository.go @@ -24,6 +24,14 @@ var embeddedSchema string //go:embed sql/data.sql 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 { db *store.DB } @@ -2513,19 +2521,20 @@ func (r *Repository) exportPermissions() ([]PermissionBackup, error) { // ImportResult contains the result of an import operation type ImportResult struct { - UsersImported int `json:"usersImported"` - NodesImported int `json:"nodesImported"` - TunnelsImported int `json:"tunnelsImported"` - ForwardsImported int `json:"forwardsImported"` - UserTunnelsImported int `json:"userTunnelsImported"` - SpeedLimitsImported int `json:"speedLimitsImported"` - TunnelGroupsImported int `json:"tunnelGroupsImported"` - UserGroupsImported int `json:"userGroupsImported"` - PermissionsImported int `json:"permissionsImported"` - ConfigsImported int `json:"configsImported"` + UsersImported int `json:"usersImported"` + NodesImported int `json:"nodesImported"` + TunnelsImported int `json:"tunnelsImported"` + ForwardsImported int `json:"forwardsImported"` + UserTunnelsImported int `json:"userTunnelsImported"` + SpeedLimitsImported int `json:"speedLimitsImported"` + TunnelGroupsImported int `json:"tunnelGroupsImported"` + UserGroupsImported int `json:"userGroupsImported"` + PermissionsImported int `json:"permissionsImported"` + 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) { result := &ImportResult{} @@ -2534,10 +2543,16 @@ func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult, 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() 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 { 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 { - count, err := r.importNodes(backup.Nodes, now) + count, err := r.importNodes(tx, backup.Nodes, now) if err != nil { 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 { - count, err := r.importTunnels(backup.Tunnels, now) + count, err := r.importTunnels(tx, backup.Tunnels, now) if err != nil { 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 { - count, err := r.importForwards(backup.Forwards, now) + count, err := r.importForwards(tx, backup.Forwards, now) if err != nil { 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 { - count, err := r.importUserTunnels(backup.UserTunnels, now) + count, err := r.importUserTunnels(tx, backup.UserTunnels, now) if err != nil { 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 { - count, err := r.importSpeedLimits(backup.SpeedLimits, now) + count, err := r.importSpeedLimits(tx, backup.SpeedLimits, now) if err != nil { 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 { - count, err := r.importTunnelGroups(backup.TunnelGroups, now) + count, err := r.importTunnelGroups(tx, backup.TunnelGroups, now) if err != nil { 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 { - count, err := r.importUserGroups(backup.UserGroups, now) + count, err := r.importUserGroups(tx, backup.UserGroups, now) if err != nil { 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 { - count, err := r.importPermissions(backup.Permissions, now) + count, err := r.importPermissions(tx, backup.Permissions, now) if err != nil { 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 { - count, err := r.importConfigs(backup.Configs, now) + count, err := r.importConfigs(tx, backup.Configs, now) if err != nil { return nil, fmt.Errorf("import configs failed: %w", err) } result.ConfigsImported = count } + if err := tx.Commit(); err != nil { + return nil, fmt.Errorf("failed to commit transaction: %w", err) + } + 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 for _, u := range users { - // Check if user exists - exists, err := r.UsernameExists(u.User) + _, err := 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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + 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 { 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++ } return count, nil @@ -2660,16 +2674,10 @@ func (r *Repository) UsernameExists(username string) (bool, error) { 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 for _, n := range nodes { - // Check if node exists - _, err := r.GetNodeByID(n.ID) - if err != nil { - return count, err - } - // Use upsert pattern - _, err = r.db.Exec(` + _, err := 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) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT(id) DO UPDATE SET @@ -2702,10 +2710,10 @@ func (r *Repository) importNodes(nodes []NodeBackup, now int64) (int, error) { 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 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) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT(id) DO UPDATE SET @@ -2722,10 +2730,9 @@ func (r *Repository) importTunnels(tunnels []TunnelBackup, now int64) (int, erro if err != nil { return count, err } - // Import chain tunnels if len(t.ChainTunnels) > 0 { 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) VALUES(?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT(id) DO UPDATE SET @@ -2746,10 +2753,10 @@ func (r *Repository) importTunnels(tunnels []TunnelBackup, now int64) (int, erro 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 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) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT(id) DO UPDATE SET @@ -2773,14 +2780,14 @@ func (r *Repository) importForwards(forwards []ForwardBackup, now int64) (int, e 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 for _, ut := range userTunnels { var speedID interface{} if ut.SpeedID > 0 { 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) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT(id) DO UPDATE SET @@ -2803,10 +2810,10 @@ func (r *Repository) importUserTunnels(userTunnels []UserTunnelBackup, now int64 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 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) VALUES(?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT(id) DO UPDATE SET @@ -2825,10 +2832,10 @@ func (r *Repository) importSpeedLimits(speedLimits []SpeedLimitBackup, now int64 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 for _, tg := range tunnelGroups { - _, err := r.db.Exec(` + _, err := db.Exec(` INSERT INTO tunnel_group(id, name, created_time, updated_time, status) VALUES(?, ?, ?, ?, ?) ON CONFLICT(id) DO UPDATE SET @@ -2839,13 +2846,12 @@ func (r *Repository) importTunnelGroups(tunnelGroups []TunnelGroupBackup, now in if err != nil { return count, err } - // Update tunnel group memberships - _, err = r.db.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, tg.ID) + _, err = db.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, tg.ID) if err != nil { return count, err } for _, tunnelID := range tg.Tunnels { - _, err = r.db.Exec(` + _, err = db.Exec(` INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) VALUES(?, ?, ?) `, tg.ID, tunnelID, now) @@ -2858,10 +2864,10 @@ func (r *Repository) importTunnelGroups(tunnelGroups []TunnelGroupBackup, now in 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 for _, ug := range userGroups { - _, err := r.db.Exec(` + _, err := db.Exec(` INSERT INTO user_group(id, name, created_time, updated_time, status) VALUES(?, ?, ?, ?, ?) ON CONFLICT(id) DO UPDATE SET @@ -2872,13 +2878,12 @@ func (r *Repository) importUserGroups(userGroups []UserGroupBackup, now int64) ( if err != nil { return count, err } - // Update user group memberships - _, err = r.db.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, ug.ID) + _, err = db.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, ug.ID) if err != nil { return count, err } for _, userID := range ug.Users { - _, err = r.db.Exec(` + _, err = db.Exec(` INSERT INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?) `, ug.ID, userID, now) @@ -2891,10 +2896,10 @@ func (r *Repository) importUserGroups(userGroups []UserGroupBackup, now int64) ( 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 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) VALUES(?, ?, ?, ?, ?) ON CONFLICT(id) DO UPDATE SET @@ -2905,9 +2910,8 @@ func (r *Repository) importPermissions(permissions []PermissionBackup, now int64 if err != nil { return count, err } - // Import 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) VALUES(?, ?, ?, ?, ?, ?) ON CONFLICT(id) DO UPDATE SET @@ -2923,7 +2927,7 @@ func (r *Repository) importPermissions(permissions []PermissionBackup, now int64 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 for name, value := range configs { err := r.UpsertConfig(name, value, now) From ae8a3db3df1e17b70c651dca4a196805f74cf1c9 Mon Sep 17 00:00:00 2001 From: sagit Date: Fri, 13 Feb 2026 06:01:42 +0000 Subject: [PATCH 05/12] feat(federation): add remote node command support Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-opencode) Co-authored-by: Sisyphus --- go-backend/internal/http/client/federation.go | 51 +++++++++++++++++++ .../internal/http/handler/federation.go | 50 ++++++++++++++++++ go-backend/internal/http/handler/handler.go | 1 + go-backend/internal/http/middleware/auth.go | 2 + 4 files changed, 104 insertions(+) diff --git a/go-backend/internal/http/client/federation.go b/go-backend/internal/http/client/federation.go index 128055a..09d6a64 100644 --- a/go-backend/internal/http/client/federation.go +++ b/go-backend/internal/http/client/federation.go @@ -77,6 +77,18 @@ type RuntimeDiagnoseRequest struct { Timeout int `json:"timeout"` } +type RuntimeNodeCommandRequest struct { + CommandType string `json:"commandType"` + Data interface{} `json:"data"` +} + +type RuntimeNodeCommandResponse struct { + Type string `json:"type"` + Success bool `json:"success"` + Message string `json:"message"` + Data map[string]interface{} `json:"data,omitempty"` +} + func NewFederationClient() *FederationClient { return &FederationClient{ client: &http.Client{ @@ -333,3 +345,42 @@ func (c *FederationClient) Diagnose(url, token, localDomain string, reqData Runt return res.Data, nil } + +func (c *FederationClient) Command(url, token, localDomain string, reqData RuntimeNodeCommandRequest) (*RuntimeNodeCommandResponse, error) { + url = strings.TrimSuffix(url, "/") + bodyBytes, _ := json.Marshal(reqData) + req, err := http.NewRequest("POST", url+"/api/v1/federation/runtime/command", strings.NewReader(string(bodyBytes))) + if err != nil { + return nil, err + } + req.Header.Set("Authorization", "Bearer "+token) + if localDomain != "" { + req.Header.Set("X-Panel-Domain", localDomain) + } + req.Header.Set("Content-Type", "application/json") + + resp, err := c.client.Do(req) + if err != nil { + return nil, err + } + defer resp.Body.Close() + + if resp.StatusCode != 200 { + body, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("remote error %d: %s", resp.StatusCode, string(body)) + } + + var res struct { + Code int `json:"code"` + Msg string `json:"msg"` + Data RuntimeNodeCommandResponse `json:"data"` + } + if err := json.NewDecoder(resp.Body).Decode(&res); err != nil { + return nil, err + } + if res.Code != 0 { + return nil, fmt.Errorf("remote api error: %s", res.Msg) + } + + return &res.Data, nil +} diff --git a/go-backend/internal/http/handler/federation.go b/go-backend/internal/http/handler/federation.go index 09abef2..831f64f 100644 --- a/go-backend/internal/http/handler/federation.go +++ b/go-backend/internal/http/handler/federation.go @@ -91,6 +91,11 @@ type federationRuntimeDiagnoseRequest struct { Timeout int `json:"timeout"` } +type federationRuntimeCommandRequest struct { + CommandType string `json:"commandType"` + Data interface{} `json:"data"` +} + type peerShareUsedPort struct { RuntimeID int64 `json:"runtimeId"` Port int `json:"port"` @@ -1199,6 +1204,51 @@ func (h *Handler) federationRuntimeDiagnose(w http.ResponseWriter, r *http.Reque response.WriteJSON(w, response.OK(res.Data)) } +func (h *Handler) federationRuntimeCommand(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("Invalid method")) + return + } + + token := extractBearerToken(r) + share, err := h.repo.GetPeerShareByToken(token) + if err != nil || share == nil { + response.WriteJSON(w, response.Err(401, "Unauthorized")) + return + } + + var req federationRuntimeCommandRequest + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("Invalid JSON")) + return + } + cmd := strings.TrimSpace(req.CommandType) + if cmd == "" { + response.WriteJSON(w, response.ErrDefault("commandType is required")) + return + } + if !isFederationRuntimeCommandAllowed(cmd) { + response.WriteJSON(w, response.ErrDefault("command not allowed")) + return + } + + res, err := h.sendNodeCommand(share.NodeID, cmd, req.Data, false, false) + if err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + response.WriteJSON(w, response.OK(res)) +} + +func isFederationRuntimeCommandAllowed(commandType string) bool { + switch strings.ToLower(strings.TrimSpace(commandType)) { + case "addservice", "updateservice", "deleteservice", "pauseservice", "resumeservice", "addchains", "deletechains", "addlimiters", "deletelimiters", "tcpping", "reload": + return true + default: + return false + } +} + func (h *Handler) pickPeerSharePort(share *sqlite.PeerShare, requestedPort int) (int, error) { if share == nil { return 0, fmt.Errorf("share not found") diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index 204613f..2cca255 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -169,6 +169,7 @@ func (h *Handler) Register(mux *http.ServeMux) { mux.HandleFunc("/api/v1/federation/runtime/apply-role", h.authPeer(h.federationRuntimeApplyRole)) mux.HandleFunc("/api/v1/federation/runtime/release-role", h.authPeer(h.federationRuntimeReleaseRole)) mux.HandleFunc("/api/v1/federation/runtime/diagnose", h.authPeer(h.federationRuntimeDiagnose)) + mux.HandleFunc("/api/v1/federation/runtime/command", h.authPeer(h.federationRuntimeCommand)) mux.HandleFunc("/api/v1/federation/node/import", h.nodeImport) mux.HandleFunc("/flow/test", h.flowTest) diff --git a/go-backend/internal/http/middleware/auth.go b/go-backend/internal/http/middleware/auth.go index bf3aebf..7e60afc 100644 --- a/go-backend/internal/http/middleware/auth.go +++ b/go-backend/internal/http/middleware/auth.go @@ -93,6 +93,8 @@ func shouldSkip(path string) bool { return true case path == "/api/v1/federation/runtime/diagnose": return true + case path == "/api/v1/federation/runtime/command": + return true default: return false } From 229ae9e454df7472d7740c67f66c4f393a29465e Mon Sep 17 00:00:00 2001 From: sagit Date: Fri, 13 Feb 2026 06:01:52 +0000 Subject: [PATCH 06/12] feat(backend): route node commands to remote panels Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-opencode) Co-authored-by: Sisyphus --- .../internal/http/handler/control_plane.go | 50 ++++++++++++++++++- go-backend/internal/http/handler/mutations.go | 2 +- 2 files changed, 50 insertions(+), 2 deletions(-) diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index 0e3d084..d61a86a 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -457,7 +457,17 @@ func (h *Handler) applyNodeProtocolChange(nodeID int64, httpVal, tlsVal, socksVa } func (h *Handler) sendNodeCommand(nodeID int64, commandType string, data interface{}, tolerateExists bool, tolerateNotFound bool) (ws.CommandResult, error) { - result, err := h.wsServer.SendCommand(nodeID, commandType, data, 12*time.Second) + var ( + result ws.CommandResult + err error + ) + + node, nodeErr := h.getNodeRecord(nodeID) + if nodeErr == nil && node != nil && node.IsRemote == 1 { + result, err = h.sendRemoteNodeCommand(node, commandType, data) + } else { + result, err = h.wsServer.SendCommand(nodeID, commandType, data, 12*time.Second) + } if err == nil { return result, nil } @@ -475,6 +485,44 @@ func (h *Handler) sendNodeCommand(nodeID int64, commandType string, data interfa return result, err } +func (h *Handler) sendRemoteNodeCommand(node *nodeRecord, commandType string, data interface{}) (ws.CommandResult, error) { + if node == nil { + return ws.CommandResult{}, errors.New("节点不存在") + } + remoteURL := strings.TrimSpace(node.RemoteURL) + remoteToken := strings.TrimSpace(node.RemoteToken) + if remoteURL == "" || remoteToken == "" { + return ws.CommandResult{}, errors.New("远程节点缺少共享配置") + } + + fc := client.NewFederationClient() + res, err := fc.Command(remoteURL, remoteToken, h.federationLocalDomain(), client.RuntimeNodeCommandRequest{ + CommandType: commandType, + Data: data, + }) + if err != nil { + return ws.CommandResult{}, err + } + if res == nil { + return ws.CommandResult{}, errors.New("远程节点未返回命令结果") + } + + result := ws.CommandResult{ + Type: res.Type, + Success: res.Success, + Message: res.Message, + Data: res.Data, + } + if !result.Success { + msg := strings.TrimSpace(result.Message) + if msg == "" { + msg = "命令执行失败" + } + return result, errors.New(msg) + } + return result, nil +} + func (h *Handler) diagnoseForwardRuntime(forward *forwardRecord) (map[string]interface{}, error) { if forward == nil { return nil, errForwardNotFound diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index cff0744..dfa9c3b 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -2061,7 +2061,7 @@ func (h *Handler) prepareTunnelCreateState(tx *store.Tx, req map[string]interfac } return nil, err } - if node.Status != 1 { + if node.IsRemote != 1 && node.Status != 1 { return nil, errors.New("部分节点不在线") } state.Nodes[nodeID] = node From 149a841a49e2f70fcc39749d8da997045b41009a Mon Sep 17 00:00:00 2001 From: sagit Date: Fri, 13 Feb 2026 06:02:01 +0000 Subject: [PATCH 07/12] test(federation): add tests for remote node command and offline status Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-opencode) Co-authored-by: Sisyphus --- .../http/handler/federation_runtime_test.go | 65 +++++++++++++++++++ .../federation_dual_panel_contract_test.go | 19 ++++++ 2 files changed, 84 insertions(+) diff --git a/go-backend/internal/http/handler/federation_runtime_test.go b/go-backend/internal/http/handler/federation_runtime_test.go index b29f758..17c1520 100644 --- a/go-backend/internal/http/handler/federation_runtime_test.go +++ b/go-backend/internal/http/handler/federation_runtime_test.go @@ -152,6 +152,71 @@ func TestPrepareTunnelCreateStateRemoteAutoPortDefersToFederation(t *testing.T) } } +func TestPrepareTunnelCreateStateAllowsOfflineRemoteMiddleNode(t *testing.T) { + repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + if err != nil { + t.Fatalf("open repo: %v", err) + } + defer repo.Close() + + h := &Handler{repo: repo} + now := time.Now().UnixMilli() + + insertNode := func(name string, status int, portRange string, isRemote int) int64 { + res, execErr := repo.DB().Exec(` + INSERT INTO node(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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, name, name+"-secret", "10.0.0.1", "10.0.0.1", "", portRange, "", "v1", 1, 1, 1, now, now, status, "[::]", "[::]", 0, isRemote, "http://peer", "peer-token", `{"shareId":2}`) + if execErr != nil { + t.Fatalf("insert node %s: %v", name, execErr) + } + id, idErr := res.LastInsertId() + if idErr != nil { + t.Fatalf("node id %s: %v", name, idErr) + } + return id + } + + entryID := insertNode("entry-local", 1, "32000-32010", 0) + remoteMiddleID := insertNode("middle-remote", 0, "33000-33010", 1) + outID := insertNode("out-local", 1, "34000-34010", 0) + + tx, err := repo.DB().Begin() + if err != nil { + t.Fatalf("begin tx: %v", err) + } + defer tx.Rollback() + + req := map[string]interface{}{ + "name": "remote-middle-offline-status", + "inNodeId": []interface{}{ + map[string]interface{}{"nodeId": float64(entryID), "protocol": "tls", "strategy": "round"}, + }, + "chainNodes": []interface{}{ + []interface{}{ + map[string]interface{}{"nodeId": float64(remoteMiddleID), "protocol": "tls", "strategy": "round", "port": float64(0)}, + }, + }, + "outNodeId": []interface{}{ + map[string]interface{}{"nodeId": float64(outID), "protocol": "tls", "strategy": "round", "port": float64(0)}, + }, + } + + state, err := h.prepareTunnelCreateState(tx, req, 2, 0) + if err != nil { + t.Fatalf("prepare state should allow offline remote middle node: %v", err) + } + if len(state.ChainHops) != 1 || len(state.ChainHops[0]) != 1 { + t.Fatalf("expected one middle hop node, got %+v", state.ChainHops) + } + if state.ChainHops[0][0].NodeID != remoteMiddleID { + t.Fatalf("expected remote middle node id %d, got %d", remoteMiddleID, state.ChainHops[0][0].NodeID) + } + if state.Nodes[remoteMiddleID] == nil || state.Nodes[remoteMiddleID].IsRemote != 1 { + t.Fatalf("expected remote middle node metadata in state") + } +} + func TestFederationRuntimeReservePortRejectsWhenShareFlowExceeded(t *testing.T) { repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) if err != nil { diff --git a/go-backend/tests/contract/federation_dual_panel_contract_test.go b/go-backend/tests/contract/federation_dual_panel_contract_test.go index b3bece9..0959f65 100644 --- a/go-backend/tests/contract/federation_dual_panel_contract_test.go +++ b/go-backend/tests/contract/federation_dual_panel_contract_test.go @@ -78,6 +78,8 @@ func TestFederationDualPanelMiddleExitAutoPortContract(t *testing.T) { middleRemoteNodeID := queryRemoteNodeIDByToken(t, consumerRepo, "share-middle-token") exitRemoteNodeID := queryRemoteNodeIDByToken(t, consumerRepo, "share-exit-token") + stopEntry := startMockNodeSession(t, providerServer.URL, "provider-entry-secret") + defer stopEntry() stopMiddle := startMockNodeSession(t, providerServer.URL, "provider-middle-secret") defer stopMiddle() stopExit := startMockNodeSession(t, providerServer.URL, "provider-exit-secret") @@ -148,6 +150,23 @@ func TestFederationDualPanelMiddleExitAutoPortContract(t *testing.T) { assertTunnelPortInRange(t, consumerRepo, secondTunnelID, 2, middleRemoteNodeID, 44000, 44010) assertTunnelPortInRange(t, consumerRepo, secondTunnelID, 3, exitRemoteNodeID, 45000, 45010) + forwardPayload := map[string]interface{}{ + "name": "dual-panel-remote-entry-forward", + "tunnelId": secondTunnelID, + "remoteAddr": "1.1.1.1:443", + "strategy": "fifo", + } + forwardBody, err := json.Marshal(forwardPayload) + if err != nil { + t.Fatalf("marshal forward payload: %v", err) + } + forwardReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(forwardBody)) + forwardReq.Header.Set("Authorization", consumerAdminToken) + forwardReq.Header.Set("Content-Type", "application/json") + forwardRes := httptest.NewRecorder() + consumerRouter.ServeHTTP(forwardRes, forwardReq) + assertCode(t, forwardRes, 0) + assertCount(t, providerRepo, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND status = 1 AND applied = 1`, middleShareID, 1) assertCount(t, providerRepo, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND status = 1 AND applied = 1`, exitShareID, 1) assertCount(t, providerRepo, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ?`, entryShareID, 0) From 0191f29cf145dd5f14cf3309d0f4ad5f3efb409c Mon Sep 17 00:00:00 2001 From: sagit Date: Fri, 13 Feb 2026 06:02:09 +0000 Subject: [PATCH 08/12] chore(docker): update network subnet in compose files Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-opencode) Co-authored-by: Sisyphus --- docker-compose-v4.yml | 2 +- docker-compose-v6.yml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/docker-compose-v4.yml b/docker-compose-v4.yml index ea85469..3e4bbe0 100644 --- a/docker-compose-v4.yml +++ b/docker-compose-v4.yml @@ -87,4 +87,4 @@ networks: driver: bridge ipam: config: - - subnet: 172.20.0.0/16 + - subnet: 172.80.0.0/16 diff --git a/docker-compose-v6.yml b/docker-compose-v6.yml index cb7ae47..b77fc2f 100644 --- a/docker-compose-v6.yml +++ b/docker-compose-v6.yml @@ -88,5 +88,5 @@ networks: enable_ipv6: true ipam: config: - - subnet: 172.20.0.0/16 + - subnet: 172.80.0.0/16 - subnet: fd00:dead:beef::/48 From 641aa66afc66635ec14c3c1051d05141587d9b8f Mon Sep 17 00:00:00 2001 From: sagit Date: Fri, 13 Feb 2026 07:25:13 +0000 Subject: [PATCH 09/12] feat(backup): restore backup export/import flow --- go-backend/internal/http/handler/handler.go | 84 ++ .../internal/store/sqlite/repository.go | 1055 +++++++++++++++++ vite-frontend/src/api/index.ts | 43 + vite-frontend/src/pages/config.tsx | 205 +++- 4 files changed, 1384 insertions(+), 3 deletions(-) diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index 204613f..51c4093 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -171,6 +171,9 @@ func (h *Handler) Register(mux *http.ServeMux) { mux.HandleFunc("/api/v1/federation/runtime/diagnose", h.authPeer(h.federationRuntimeDiagnose)) 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) @@ -1140,3 +1143,84 @@ func (h *Handler) verifyCloudflareTurnstile(token, secretKey string) bool { } return body.Success } + +type backupExportRequest struct { + Types []string `json:"types"` +} + +func (h *Handler) backupExport(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + + var req backupExportRequest + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.Err(500, "请求参数错误")) + return + } + + var backup interface{} + var err error + + if len(req.Types) == 0 { + backup, err = h.repo.ExportAll() + } else { + backup, err = h.repo.ExportPartial(req.Types) + } + + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + + w.Header().Set("Content-Disposition", "attachment; filename=backup.json") + w.Header().Set("Content-Type", "application/json") + if err := json.NewEncoder(w).Encode(backup); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } +} + +type backupImportRequest struct { + Types []string `json:"types"` + sqlite.BackupData +} + +func (h *Handler) backupImport(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + + var req backupImportRequest + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.Err(500, "请求参数错误")) + return + } + + if len(req.Types) == 0 { + response.WriteJSON(w, response.Err(500, "请选择要导入的数据类型")) + return + } + + autoBackup, err := h.repo.ExportAll() + if err != nil { + response.WriteJSON(w, response.Err(-2, fmt.Sprintf("导入前自动备份失败: %v", err))) + return + } + + if req.BackupData.Version == "" { + response.WriteJSON(w, response.Err(500, "备份数据格式错误")) + return + } + + result, err := h.repo.Import(&req.BackupData, req.Types) + if err != nil { + response.WriteJSON(w, response.Err(-2, fmt.Sprintf("导入失败: %v", err))) + return + } + + result.AutoBackup = autoBackup + response.WriteJSON(w, response.OK(result)) +} diff --git a/go-backend/internal/store/sqlite/repository.go b/go-backend/internal/store/sqlite/repository.go index 0ff75fe..d72a6ec 100644 --- a/go-backend/internal/store/sqlite/repository.go +++ b/go-backend/internal/store/sqlite/repository.go @@ -24,6 +24,14 @@ var embeddedSchema string //go:embed sql/data.sql 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 { db *store.DB } @@ -1883,3 +1891,1050 @@ func (r *Repository) DeleteFederationTunnelBindingsByTunnel(tunnelID int64) erro var osMkdirAll = func(path string) error { return os.MkdirAll(path, 0o755) } + +// ============ Backup/Export Data Structures ============ + +// BackupData represents the full backup structure +type BackupData struct { + Version string `json:"version"` + ExportedAt int64 `json:"exportedAt"` + Users []UserBackup `json:"users,omitempty"` + Nodes []NodeBackup `json:"nodes,omitempty"` + Tunnels []TunnelBackup `json:"tunnels,omitempty"` + Forwards []ForwardBackup `json:"forwards,omitempty"` + UserTunnels []UserTunnelBackup `json:"userTunnels,omitempty"` + SpeedLimits []SpeedLimitBackup `json:"speedLimits,omitempty"` + TunnelGroups []TunnelGroupBackup `json:"tunnelGroups,omitempty"` + UserGroups []UserGroupBackup `json:"userGroups,omitempty"` + Permissions []PermissionBackup `json:"permissions,omitempty"` + Configs map[string]string `json:"configs,omitempty"` +} + +type UserBackup struct { + ID int64 `json:"id"` + User string `json:"user"` + Pwd string `json:"pwd"` + RoleID int `json:"roleId"` + ExpTime int64 `json:"expTime"` + Flow int64 `json:"flow"` + InFlow int64 `json:"inFlow"` + OutFlow int64 `json:"outFlow"` + FlowResetTime int64 `json:"flowResetTime"` + Num int `json:"num"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime,omitempty"` + Status int `json:"status"` +} + +type NodeBackup struct { + ID int64 `json:"id"` + Name string `json:"name"` + Secret string `json:"secret"` + ServerIP string `json:"serverIp"` + ServerIPv4 string `json:"serverIpV4,omitempty"` + ServerIPv6 string `json:"serverIpV6,omitempty"` + Port string `json:"port"` + InterfaceName string `json:"interfaceName,omitempty"` + Version string `json:"version,omitempty"` + HTTP int `json:"http"` + TLS int `json:"tls"` + Socks int `json:"socks"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime,omitempty"` + Status int `json:"status"` + TCPListenAddr string `json:"tcpListenAddr"` + UDPListenAddr string `json:"udpListenAddr"` + Inx int `json:"inx"` + IsRemote int `json:"isRemote"` + RemoteURL string `json:"remoteUrl,omitempty"` + RemoteToken string `json:"remoteToken,omitempty"` + RemoteConfig string `json:"remoteConfig,omitempty"` +} + +type TunnelBackup struct { + ID int64 `json:"id"` + Name string `json:"name"` + TrafficRatio float64 `json:"trafficRatio"` + Type int `json:"type"` + Protocol string `json:"protocol"` + Flow int64 `json:"flow"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime"` + Status int `json:"status"` + InIP string `json:"inIp,omitempty"` + Inx int `json:"inx"` + ChainTunnels []ChainTunnelBackup `json:"chainTunnels,omitempty"` +} + +type ChainTunnelBackup struct { + ID int64 `json:"id"` + TunnelID int64 `json:"tunnelId"` + ChainType string `json:"chainType"` + NodeID int64 `json:"nodeId"` + Port int `json:"port,omitempty"` + Strategy string `json:"strategy,omitempty"` + Inx int `json:"inx,omitempty"` + Protocol string `json:"protocol,omitempty"` +} + +type ForwardBackup struct { + ID int64 `json:"id"` + UserID int64 `json:"userId"` + UserName string `json:"userName"` + Name string `json:"name"` + TunnelID int64 `json:"tunnelId"` + RemoteAddr string `json:"remoteAddr"` + Strategy string `json:"strategy"` + InFlow int64 `json:"inFlow"` + OutFlow int64 `json:"outFlow"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime"` + Status int `json:"status"` + Inx int `json:"inx"` +} + +type UserTunnelBackup struct { + ID int64 `json:"id"` + UserID int64 `json:"userId"` + TunnelID int64 `json:"tunnelId"` + SpeedID int64 `json:"speedId,omitempty"` + Num int `json:"num"` + Flow int64 `json:"flow"` + InFlow int64 `json:"inFlow"` + OutFlow int64 `json:"outFlow"` + FlowResetTime int64 `json:"flowResetTime"` + ExpTime int64 `json:"expTime"` + Status int `json:"status"` +} + +type SpeedLimitBackup struct { + ID int64 `json:"id"` + Name string `json:"name"` + Speed int64 `json:"speed"` + TunnelID int64 `json:"tunnelId"` + TunnelName string `json:"tunnelName"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime,omitempty"` + Status int `json:"status"` +} + +type TunnelGroupBackup struct { + ID int64 `json:"id"` + Name string `json:"name"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime"` + Status int `json:"status"` + Tunnels []int64 `json:"tunnels,omitempty"` +} + +type UserGroupBackup struct { + ID int64 `json:"id"` + Name string `json:"name"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime"` + Status int `json:"status"` + Users []int64 `json:"users,omitempty"` +} + +type PermissionBackup struct { + ID int64 `json:"id"` + UserGroupID int64 `json:"userGroupId"` + TunnelGroupID int64 `json:"tunnelGroupId"` + CreatedTime int64 `json:"createdTime"` + CreatedByGroup int `json:"createdByGroup"` + Grants []PermissionGrantBackup `json:"grants,omitempty"` +} + +type PermissionGrantBackup struct { + ID int64 `json:"id"` + UserGroupID int64 `json:"userGroupId"` + TunnelGroupID int64 `json:"tunnelGroupId"` + UserTunnelID int64 `json:"userTunnelId"` + CreatedTime int64 `json:"createdTime"` + CreatedByGroup int `json:"createdByGroup"` +} + +// ============ Export Methods ============ + +// ExportAll exports all data as BackupData +func (r *Repository) ExportAll() (*BackupData, error) { + backup := &BackupData{ + Version: "1.0", + ExportedAt: unixMilliNow(), + } + + // Export all data types + users, err := r.exportUsers() + if err != nil { + return nil, fmt.Errorf("export users failed: %w", err) + } + backup.Users = users + + nodes, err := r.exportNodes() + if err != nil { + return nil, fmt.Errorf("export nodes failed: %w", err) + } + backup.Nodes = nodes + + tunnels, err := r.exportTunnels() + if err != nil { + return nil, fmt.Errorf("export tunnels failed: %w", err) + } + backup.Tunnels = tunnels + + forwards, err := r.exportForwards() + if err != nil { + return nil, fmt.Errorf("export forwards failed: %w", err) + } + backup.Forwards = forwards + + userTunnels, err := r.exportUserTunnels() + if err != nil { + return nil, fmt.Errorf("export user tunnels failed: %w", err) + } + backup.UserTunnels = userTunnels + + speedLimits, err := r.exportSpeedLimits() + if err != nil { + return nil, fmt.Errorf("export speed limits failed: %w", err) + } + backup.SpeedLimits = speedLimits + + tunnelGroups, err := r.exportTunnelGroups() + if err != nil { + return nil, fmt.Errorf("export tunnel groups failed: %w", err) + } + backup.TunnelGroups = tunnelGroups + + userGroups, err := r.exportUserGroups() + if err != nil { + return nil, fmt.Errorf("export user groups failed: %w", err) + } + backup.UserGroups = userGroups + + permissions, err := r.exportPermissions() + if err != nil { + return nil, fmt.Errorf("export permissions failed: %w", err) + } + backup.Permissions = permissions + + configs, err := r.ListConfigs() + if err != nil { + return nil, fmt.Errorf("export configs failed: %w", err) + } + backup.Configs = configs + + return backup, nil +} + +// ExportPartial exports selected data types +func (r *Repository) ExportPartial(types []string) (*BackupData, error) { + backup := &BackupData{ + Version: "1.0", + ExportedAt: unixMilliNow(), + } + + typeSet := make(map[string]bool) + for _, t := range types { + typeSet[t] = true + } + + if typeSet["users"] { + users, err := r.exportUsers() + if err != nil { + return nil, fmt.Errorf("export users failed: %w", err) + } + backup.Users = users + } + if typeSet["nodes"] { + nodes, err := r.exportNodes() + if err != nil { + return nil, fmt.Errorf("export nodes failed: %w", err) + } + backup.Nodes = nodes + } + if typeSet["tunnels"] { + tunnels, err := r.exportTunnels() + if err != nil { + return nil, fmt.Errorf("export tunnels failed: %w", err) + } + backup.Tunnels = tunnels + } + if typeSet["forwards"] { + forwards, err := r.exportForwards() + if err != nil { + return nil, fmt.Errorf("export forwards failed: %w", err) + } + backup.Forwards = forwards + } + if typeSet["userTunnels"] { + userTunnels, err := r.exportUserTunnels() + if err != nil { + return nil, fmt.Errorf("export user tunnels failed: %w", err) + } + backup.UserTunnels = userTunnels + } + if typeSet["speedLimits"] { + speedLimits, err := r.exportSpeedLimits() + if err != nil { + return nil, fmt.Errorf("export speed limits failed: %w", err) + } + backup.SpeedLimits = speedLimits + } + if typeSet["tunnelGroups"] { + tunnelGroups, err := r.exportTunnelGroups() + if err != nil { + return nil, fmt.Errorf("export tunnel groups failed: %w", err) + } + backup.TunnelGroups = tunnelGroups + } + if typeSet["userGroups"] { + userGroups, err := r.exportUserGroups() + if err != nil { + return nil, fmt.Errorf("export user groups failed: %w", err) + } + backup.UserGroups = userGroups + } + if typeSet["permissions"] { + permissions, err := r.exportPermissions() + if err != nil { + return nil, fmt.Errorf("export permissions failed: %w", err) + } + backup.Permissions = permissions + } + if typeSet["configs"] { + configs, err := r.ListConfigs() + if err != nil { + return nil, fmt.Errorf("export configs failed: %w", err) + } + backup.Configs = configs + } + + return backup, nil +} + +func (r *Repository) exportUsers() ([]UserBackup, error) { + rows, err := r.db.Query(` + SELECT id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status + FROM user ORDER BY id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + var users []UserBackup + for rows.Next() { + var u UserBackup + var updatedTime sql.NullInt64 + if err := rows.Scan(&u.ID, &u.User, &u.Pwd, &u.RoleID, &u.ExpTime, &u.Flow, &u.InFlow, &u.OutFlow, &u.FlowResetTime, &u.Num, &u.CreatedTime, &updatedTime, &u.Status); err != nil { + return nil, err + } + if updatedTime.Valid { + u.UpdatedTime = updatedTime.Int64 + } + users = append(users, u) + } + return users, rows.Err() +} + +func (r *Repository) exportNodes() ([]NodeBackup, error) { + rows, err := r.db.Query(` + SELECT 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 + FROM node ORDER BY inx ASC, id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + var nodes []NodeBackup + for rows.Next() { + var n NodeBackup + var updatedTime sql.NullInt64 + var serverIPv4, serverIPv6, interfaceName, version, remoteURL, remoteToken, remoteConfig sql.NullString + if err := rows.Scan(&n.ID, &n.Name, &n.Secret, &n.ServerIP, &serverIPv4, &serverIPv6, &n.Port, &interfaceName, &version, &n.HTTP, &n.TLS, &n.Socks, &n.CreatedTime, &updatedTime, &n.Status, &n.TCPListenAddr, &n.UDPListenAddr, &n.Inx, &n.IsRemote, &remoteURL, &remoteToken, &remoteConfig); err != nil { + return nil, err + } + if updatedTime.Valid { + n.UpdatedTime = updatedTime.Int64 + } + if serverIPv4.Valid { + n.ServerIPv4 = serverIPv4.String + } + if serverIPv6.Valid { + n.ServerIPv6 = serverIPv6.String + } + if interfaceName.Valid { + n.InterfaceName = interfaceName.String + } + if version.Valid { + n.Version = version.String + } + if remoteURL.Valid { + n.RemoteURL = remoteURL.String + } + if remoteToken.Valid { + n.RemoteToken = remoteToken.String + } + if remoteConfig.Valid { + n.RemoteConfig = remoteConfig.String + } + nodes = append(nodes, n) + } + return nodes, rows.Err() +} + +func (r *Repository) exportTunnels() ([]TunnelBackup, error) { + rows, err := r.db.Query(` + SELECT id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx + FROM tunnel ORDER BY inx ASC, id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + var tunnels []TunnelBackup + for rows.Next() { + var t TunnelBackup + var inIP sql.NullString + if err := rows.Scan(&t.ID, &t.Name, &t.TrafficRatio, &t.Type, &t.Protocol, &t.Flow, &t.CreatedTime, &t.UpdatedTime, &t.Status, &inIP, &t.Inx); err != nil { + return nil, err + } + if inIP.Valid { + t.InIP = inIP.String + } + // Export chain tunnels + chainTunnels, err := r.exportChainTunnels(t.ID) + if err != nil { + return nil, err + } + t.ChainTunnels = chainTunnels + tunnels = append(tunnels, t) + } + return tunnels, rows.Err() +} + +func (r *Repository) exportChainTunnels(tunnelID int64) ([]ChainTunnelBackup, error) { + rows, err := r.db.Query(` + SELECT id, tunnel_id, chain_type, node_id, port, strategy, inx, protocol + FROM chain_tunnel WHERE tunnel_id = ? ORDER BY inx ASC, id ASC + `, tunnelID) + if err != nil { + return nil, err + } + defer rows.Close() + + var chainTunnels []ChainTunnelBackup + for rows.Next() { + var ct ChainTunnelBackup + var port sql.NullInt64 + if err := rows.Scan(&ct.ID, &ct.TunnelID, &ct.ChainType, &ct.NodeID, &port, &ct.Strategy, &ct.Inx, &ct.Protocol); err != nil { + return nil, err + } + if port.Valid { + ct.Port = int(port.Int64) + } + chainTunnels = append(chainTunnels, ct) + } + return chainTunnels, rows.Err() +} + +func (r *Repository) exportForwards() ([]ForwardBackup, error) { + rows, err := r.db.Query(` + SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx + FROM forward ORDER BY id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + var forwards []ForwardBackup + for rows.Next() { + var f ForwardBackup + if err := rows.Scan(&f.ID, &f.UserID, &f.UserName, &f.Name, &f.TunnelID, &f.RemoteAddr, &f.Strategy, &f.InFlow, &f.OutFlow, &f.CreatedTime, &f.UpdatedTime, &f.Status, &f.Inx); err != nil { + return nil, err + } + forwards = append(forwards, f) + } + return forwards, rows.Err() +} + +func (r *Repository) exportUserTunnels() ([]UserTunnelBackup, error) { + 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 + FROM user_tunnel ORDER BY id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + var userTunnels []UserTunnelBackup + for rows.Next() { + var ut UserTunnelBackup + var speedID sql.NullInt64 + if err := rows.Scan(&ut.ID, &ut.UserID, &ut.TunnelID, &speedID, &ut.Num, &ut.Flow, &ut.InFlow, &ut.OutFlow, &ut.FlowResetTime, &ut.ExpTime, &ut.Status); err != nil { + return nil, err + } + if speedID.Valid { + ut.SpeedID = speedID.Int64 + } + userTunnels = append(userTunnels, ut) + } + return userTunnels, rows.Err() +} + +func (r *Repository) exportSpeedLimits() ([]SpeedLimitBackup, error) { + rows, err := r.db.Query(` + SELECT id, name, speed, tunnel_id, tunnel_name, created_time, updated_time, status + FROM speed_limit ORDER BY id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + var speedLimits []SpeedLimitBackup + for rows.Next() { + var sl SpeedLimitBackup + var updatedTime sql.NullInt64 + if err := rows.Scan(&sl.ID, &sl.Name, &sl.Speed, &sl.TunnelID, &sl.TunnelName, &sl.CreatedTime, &updatedTime, &sl.Status); err != nil { + return nil, err + } + if updatedTime.Valid { + sl.UpdatedTime = updatedTime.Int64 + } + speedLimits = append(speedLimits, sl) + } + return speedLimits, rows.Err() +} + +func (r *Repository) exportTunnelGroups() ([]TunnelGroupBackup, error) { + rows, err := r.db.Query(` + SELECT id, name, created_time, updated_time, status + FROM tunnel_group ORDER BY id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + var groups []TunnelGroupBackup + for rows.Next() { + var tg TunnelGroupBackup + if err := rows.Scan(&tg.ID, &tg.Name, &tg.CreatedTime, &tg.UpdatedTime, &tg.Status); err != nil { + return nil, err + } + // Get tunnel IDs for this group + tunnelRows, err := r.db.Query(`SELECT tunnel_id FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, tg.ID) + if err != nil { + return nil, err + } + for tunnelRows.Next() { + var tunnelID int64 + if err := tunnelRows.Scan(&tunnelID); err != nil { + tunnelRows.Close() + return nil, err + } + tg.Tunnels = append(tg.Tunnels, tunnelID) + } + tunnelRows.Close() + groups = append(groups, tg) + } + return groups, rows.Err() +} + +func (r *Repository) exportUserGroups() ([]UserGroupBackup, error) { + rows, err := r.db.Query(` + SELECT id, name, created_time, updated_time, status + FROM user_group ORDER BY id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + var groups []UserGroupBackup + for rows.Next() { + var ug UserGroupBackup + if err := rows.Scan(&ug.ID, &ug.Name, &ug.CreatedTime, &ug.UpdatedTime, &ug.Status); err != nil { + return nil, err + } + // Get user IDs for this group + userRows, err := r.db.Query(`SELECT user_id FROM user_group_user WHERE user_group_id = ?`, ug.ID) + if err != nil { + return nil, err + } + for userRows.Next() { + var userID int64 + if err := userRows.Scan(&userID); err != nil { + userRows.Close() + return nil, err + } + ug.Users = append(ug.Users, userID) + } + userRows.Close() + groups = append(groups, ug) + } + return groups, rows.Err() +} + +func (r *Repository) exportPermissions() ([]PermissionBackup, error) { + rows, err := r.db.Query(` + SELECT id, user_group_id, tunnel_group_id, created_time, created_by_group + FROM group_permission ORDER BY id ASC + `) + if err != nil { + return nil, err + } + defer rows.Close() + + 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 { + return nil, err + } + // 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 { + return nil, err + } + for grantRows.Next() { + var g PermissionGrantBackup + if err := grantRows.Scan(&g.ID, &g.UserGroupID, &g.TunnelGroupID, &g.UserTunnelID, &g.CreatedTime, &g.CreatedByGroup); err != nil { + grantRows.Close() + return nil, err + } + p.Grants = append(p.Grants, g) + } + grantRows.Close() + permissions = append(permissions, p) + } + return permissions, rows.Err() +} + +// ============ Import Methods ============ + +// ImportResult contains the result of an import operation +type ImportResult struct { + UsersImported int `json:"usersImported"` + NodesImported int `json:"nodesImported"` + TunnelsImported int `json:"tunnelsImported"` + ForwardsImported int `json:"forwardsImported"` + UserTunnelsImported int `json:"userTunnelsImported"` + SpeedLimitsImported int `json:"speedLimitsImported"` + TunnelGroupsImported int `json:"tunnelGroupsImported"` + UserGroupsImported int `json:"userGroupsImported"` + PermissionsImported int `json:"permissionsImported"` + ConfigsImported int `json:"configsImported"` + AutoBackup *BackupData `json:"autoBackup,omitempty"` +} + +// Import imports data from BackupData with transaction support +func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult, error) { + result := &ImportResult{} + + typeSet := make(map[string]bool) + for _, t := range types { + 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() + + if typeSet["users"] && len(backup.Users) > 0 { + count, err := r.importUsers(tx, backup.Users, now) + if err != nil { + return nil, fmt.Errorf("import users failed: %w", err) + } + result.UsersImported = count + } + + if typeSet["nodes"] && len(backup.Nodes) > 0 { + count, err := r.importNodes(tx, backup.Nodes, now) + if err != nil { + return nil, fmt.Errorf("import nodes failed: %w", err) + } + result.NodesImported = count + } + + if typeSet["tunnels"] && len(backup.Tunnels) > 0 { + count, err := r.importTunnels(tx, backup.Tunnels, now) + if err != nil { + return nil, fmt.Errorf("import tunnels failed: %w", err) + } + result.TunnelsImported = count + } + + if typeSet["forwards"] && len(backup.Forwards) > 0 { + count, err := r.importForwards(tx, backup.Forwards, now) + if err != nil { + return nil, fmt.Errorf("import forwards failed: %w", err) + } + result.ForwardsImported = count + } + + if typeSet["userTunnels"] && len(backup.UserTunnels) > 0 { + count, err := r.importUserTunnels(tx, backup.UserTunnels, now) + if err != nil { + return nil, fmt.Errorf("import user tunnels failed: %w", err) + } + result.UserTunnelsImported = count + } + + if typeSet["speedLimits"] && len(backup.SpeedLimits) > 0 { + count, err := r.importSpeedLimits(tx, backup.SpeedLimits, now) + if err != nil { + return nil, fmt.Errorf("import speed limits failed: %w", err) + } + result.SpeedLimitsImported = count + } + + if typeSet["tunnelGroups"] && len(backup.TunnelGroups) > 0 { + count, err := r.importTunnelGroups(tx, backup.TunnelGroups, now) + if err != nil { + return nil, fmt.Errorf("import tunnel groups failed: %w", err) + } + result.TunnelGroupsImported = count + } + + if typeSet["userGroups"] && len(backup.UserGroups) > 0 { + count, err := r.importUserGroups(tx, backup.UserGroups, now) + if err != nil { + return nil, fmt.Errorf("import user groups failed: %w", err) + } + result.UserGroupsImported = count + } + + if typeSet["permissions"] && len(backup.Permissions) > 0 { + count, err := r.importPermissions(tx, backup.Permissions, now) + if err != nil { + return nil, fmt.Errorf("import permissions failed: %w", err) + } + result.PermissionsImported = count + } + + if typeSet["configs"] && len(backup.Configs) > 0 { + count, err := r.importConfigs(tx, backup.Configs, now) + if err != nil { + return nil, fmt.Errorf("import configs failed: %w", err) + } + result.ConfigsImported = count + } + + if err := tx.Commit(); err != nil { + return nil, fmt.Errorf("failed to commit transaction: %w", err) + } + + return result, nil +} + +func (r *Repository) importUsers(db Execer, users []UserBackup, now int64) (int, error) { + count := 0 + for _, u := range users { + _, err := 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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + 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 { + return count, err + } + count++ + } + return count, nil +} + +func (r *Repository) UsernameExists(username string) (bool, error) { + var count int + err := r.db.QueryRow(`SELECT COUNT(1) FROM user WHERE user = ?`, username).Scan(&count) + if err != nil { + return false, err + } + return count > 0, nil +} + +func (r *Repository) importNodes(db Execer, nodes []NodeBackup, now int64) (int, error) { + count := 0 + for _, n := range nodes { + _, err := 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) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + name = excluded.name, + secret = excluded.secret, + server_ip = excluded.server_ip, + server_ip_v4 = excluded.server_ip_v4, + server_ip_v6 = excluded.server_ip_v6, + port = excluded.port, + interface_name = excluded.interface_name, + version = excluded.version, + http = excluded.http, + tls = excluded.tls, + socks = excluded.socks, + updated_time = excluded.updated_time, + status = excluded.status, + tcp_listen_addr = excluded.tcp_listen_addr, + udp_listen_addr = excluded.udp_listen_addr, + inx = excluded.inx, + is_remote = excluded.is_remote, + remote_url = excluded.remote_url, + remote_token = excluded.remote_token, + remote_config = excluded.remote_config + `, n.ID, n.Name, n.Secret, n.ServerIP, n.ServerIPv4, n.ServerIPv6, n.Port, n.InterfaceName, n.Version, n.HTTP, n.TLS, n.Socks, n.CreatedTime, now, n.Status, n.TCPListenAddr, n.UDPListenAddr, n.Inx, n.IsRemote, n.RemoteURL, n.RemoteToken, n.RemoteConfig) + if err != nil { + return count, err + } + count++ + } + return count, nil +} + +func (r *Repository) importTunnels(db Execer, tunnels []TunnelBackup, now int64) (int, error) { + count := 0 + for _, t := range tunnels { + _, err := db.Exec(` + INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + name = excluded.name, + traffic_ratio = excluded.traffic_ratio, + type = excluded.type, + protocol = excluded.protocol, + flow = excluded.flow, + updated_time = excluded.updated_time, + status = excluded.status, + in_ip = excluded.in_ip, + inx = excluded.inx + `, t.ID, t.Name, t.TrafficRatio, t.Type, t.Protocol, t.Flow, t.CreatedTime, now, t.Status, t.InIP, t.Inx) + if err != nil { + return count, err + } + if len(t.ChainTunnels) > 0 { + for _, ct := range t.ChainTunnels { + _, err = db.Exec(` + INSERT INTO chain_tunnel(id, tunnel_id, chain_type, node_id, port, strategy, inx, protocol) + VALUES(?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + chain_type = excluded.chain_type, + node_id = excluded.node_id, + port = excluded.port, + strategy = excluded.strategy, + inx = excluded.inx, + protocol = excluded.protocol + `, ct.ID, ct.TunnelID, ct.ChainType, ct.NodeID, ct.Port, ct.Strategy, ct.Inx, ct.Protocol) + if err != nil { + return count, err + } + } + } + count++ + } + return count, nil +} + +func (r *Repository) importForwards(db Execer, forwards []ForwardBackup, now int64) (int, error) { + count := 0 + for _, f := range forwards { + _, 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) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + user_id = excluded.user_id, + user_name = excluded.user_name, + name = excluded.name, + tunnel_id = excluded.tunnel_id, + remote_addr = excluded.remote_addr, + strategy = excluded.strategy, + in_flow = excluded.in_flow, + out_flow = excluded.out_flow, + updated_time = excluded.updated_time, + status = excluded.status, + inx = excluded.inx + `, f.ID, f.UserID, f.UserName, f.Name, f.TunnelID, f.RemoteAddr, f.Strategy, f.InFlow, f.OutFlow, f.CreatedTime, now, f.Status, f.Inx) + if err != nil { + return count, err + } + count++ + } + return count, nil +} + +func (r *Repository) importUserTunnels(db Execer, userTunnels []UserTunnelBackup, now int64) (int, error) { + count := 0 + for _, ut := range userTunnels { + var speedID interface{} + if ut.SpeedID > 0 { + speedID = ut.SpeedID + } + _, 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) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + user_id = excluded.user_id, + tunnel_id = excluded.tunnel_id, + speed_id = excluded.speed_id, + num = excluded.num, + flow = excluded.flow, + in_flow = excluded.in_flow, + out_flow = excluded.out_flow, + flow_reset_time = excluded.flow_reset_time, + exp_time = excluded.exp_time, + status = excluded.status + `, ut.ID, ut.UserID, ut.TunnelID, speedID, ut.Num, ut.Flow, ut.InFlow, ut.OutFlow, ut.FlowResetTime, ut.ExpTime, ut.Status) + if err != nil { + return count, err + } + count++ + } + return count, nil +} + +func (r *Repository) importSpeedLimits(db Execer, speedLimits []SpeedLimitBackup, now int64) (int, error) { + count := 0 + for _, sl := range speedLimits { + _, err := db.Exec(` + INSERT INTO speed_limit(id, name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) + VALUES(?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + name = excluded.name, + speed = excluded.speed, + tunnel_id = excluded.tunnel_id, + tunnel_name = excluded.tunnel_name, + updated_time = excluded.updated_time, + status = excluded.status + `, sl.ID, sl.Name, sl.Speed, sl.TunnelID, sl.TunnelName, sl.CreatedTime, now, sl.Status) + if err != nil { + return count, err + } + count++ + } + return count, nil +} + +func (r *Repository) importTunnelGroups(db Execer, tunnelGroups []TunnelGroupBackup, now int64) (int, error) { + count := 0 + for _, tg := range tunnelGroups { + _, err := db.Exec(` + INSERT INTO tunnel_group(id, name, created_time, updated_time, status) + VALUES(?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + name = excluded.name, + updated_time = excluded.updated_time, + status = excluded.status + `, tg.ID, tg.Name, tg.CreatedTime, now, tg.Status) + if err != nil { + return count, err + } + _, err = db.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, tg.ID) + if err != nil { + return count, err + } + for _, tunnelID := range tg.Tunnels { + _, err = db.Exec(` + INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) + VALUES(?, ?, ?) + `, tg.ID, tunnelID, now) + if err != nil { + return count, err + } + } + count++ + } + return count, nil +} + +func (r *Repository) importUserGroups(db Execer, userGroups []UserGroupBackup, now int64) (int, error) { + count := 0 + for _, ug := range userGroups { + _, err := db.Exec(` + INSERT INTO user_group(id, name, created_time, updated_time, status) + VALUES(?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + name = excluded.name, + updated_time = excluded.updated_time, + status = excluded.status + `, ug.ID, ug.Name, ug.CreatedTime, now, ug.Status) + if err != nil { + return count, err + } + _, err = db.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, ug.ID) + if err != nil { + return count, err + } + for _, userID := range ug.Users { + _, err = db.Exec(` + INSERT INTO user_group_user(user_group_id, user_id, created_time) + VALUES(?, ?, ?) + `, ug.ID, userID, now) + if err != nil { + return count, err + } + } + count++ + } + return count, nil +} + +func (r *Repository) importPermissions(db Execer, permissions []PermissionBackup, now int64) (int, error) { + count := 0 + for _, p := range permissions { + _, err := db.Exec(` + INSERT INTO group_permission(id, user_group_id, tunnel_group_id, created_time, created_by_group) + VALUES(?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + user_group_id = excluded.user_group_id, + tunnel_group_id = excluded.tunnel_group_id, + created_by_group = excluded.created_by_group + `, p.ID, p.UserGroupID, p.TunnelGroupID, p.CreatedTime, p.CreatedByGroup) + if err != nil { + return count, err + } + for _, g := range p.Grants { + _, err = db.Exec(` + INSERT INTO group_permission_grant(id, user_group_id, tunnel_group_id, user_tunnel_id, created_time, created_by_group) + VALUES(?, ?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET + user_tunnel_id = excluded.user_tunnel_id, + created_by_group = excluded.created_by_group + `, g.ID, g.UserGroupID, g.TunnelGroupID, g.UserTunnelID, g.CreatedTime, g.CreatedByGroup) + if err != nil { + return count, err + } + } + count++ + } + return count, nil +} + +func (r *Repository) importConfigs(db Execer, configs map[string]string, now int64) (int, error) { + count := 0 + for name, value := range configs { + err := r.UpsertConfig(name, value, now) + if err != nil { + return count, err + } + count++ + } + return count, nil +} diff --git a/vite-frontend/src/api/index.ts b/vite-frontend/src/api/index.ts index 846283d..265cc00 100644 --- a/vite-frontend/src/api/index.ts +++ b/vite-frontend/src/api/index.ts @@ -228,3 +228,46 @@ export const importRemoteNode = (data: { remoteUrl: string; token: string; }) => Network.post("/federation/node/import", data); + +import axios from "axios"; + +export interface BackupTypes { + users?: boolean; + nodes?: boolean; + tunnels?: boolean; + forwards?: boolean; + userTunnels?: boolean; + speedLimits?: boolean; + tunnelGroups?: boolean; + userGroups?: boolean; + permissions?: boolean; + configs?: boolean; +} + +export const exportBackup = async (types: string[] = []) => { + const token = window.localStorage.getItem("token"); + const baseURL = axios.defaults.baseURL || "/api/v1/"; + + const response = await axios.post(`${baseURL}/backup/export`, { types }, { + headers: { + Authorization: token, + "Content-Type": "application/json", + }, + responseType: "blob", + }); + + const url = window.URL.createObjectURL(new Blob([response.data])); + const link = document.createElement("a"); + link.href = url; + const timestamp = new Date().toISOString().slice(0, 19).replace(/[:-]/g, ""); + link.setAttribute("download", `backup_${timestamp}.json`); + document.body.appendChild(link); + link.click(); + document.body.removeChild(link); + window.URL.revokeObjectURL(url); +}; + +export const importBackup = (data: { + types: string[]; + [key: string]: any; +}) => Network.post("/backup/import", data); diff --git a/vite-frontend/src/pages/config.tsx b/vite-frontend/src/pages/config.tsx index 4ee5498..71bc198 100644 --- a/vite-frontend/src/pages/config.tsx +++ b/vite-frontend/src/pages/config.tsx @@ -1,4 +1,4 @@ -import { useState, useEffect } from "react"; +import { useState, useEffect, useRef } from "react"; import { useNavigate } from "react-router-dom"; import { Button } from "@heroui/button"; import { Card, CardBody, CardHeader } from "@heroui/card"; @@ -7,9 +7,10 @@ import { Spinner } from "@heroui/spinner"; import { Divider } from "@heroui/divider"; import { Switch } from "@heroui/switch"; import { Select, SelectItem } from "@heroui/select"; +import { Checkbox, CheckboxGroup } from "@heroui/checkbox"; import toast from "react-hot-toast"; -import { updateConfigs } from "@/api"; +import { updateConfigs, exportBackup, importBackup } from "@/api"; import { SettingsIcon } from "@/components/icons"; import { isAdmin } from "@/utils/auth"; import { @@ -130,12 +131,19 @@ export default function ConfigPage() { useState>(initialConfigs); const [loading, setLoading] = useState( Object.keys(initialConfigs).length === 0, - ); // 如果有缓存数据,不显示loading + ); const [saving, setSaving] = useState(false); const [hasChanges, setHasChanges] = useState(false); const [originalConfigs, setOriginalConfigs] = useState>(initialConfigs); + const [exportTypes, setExportTypes] = useState([]); + const [importTypes, setImportTypes] = useState([]); + const [exporting, setExporting] = useState(false); + const [importing, setImporting] = useState(false); + const [importFileName, setImportFileName] = useState(""); + const fileInputRef = useRef(null); + // 权限检查 useEffect(() => { if (!isAdmin()) { @@ -331,6 +339,60 @@ export default function ConfigPage() { } }; + const handleExport = async () => { + if (exportTypes.length === 0) { + toast.error("请至少选择一种数据类型"); + return; + } + setExporting(true); + try { + await exportBackup(exportTypes); + toast.success("导出成功"); + } catch { + toast.error("导出失败,请重试"); + } finally { + setExporting(false); + } + }; + + const handleFileChange = async (e: React.ChangeEvent) => { + const file = e.target.files?.[0]; + if (!file) return; + + if (importTypes.length === 0) { + toast.error("请先选择要导入的数据类型"); + return; + } + + setImportFileName(file.name); + setImporting(true); + + try { + const text = await file.text(); + const data = JSON.parse(text); + + const response = await importBackup({ + types: importTypes, + ...data, + }); + + if (response.code === 0) { + toast.success(`导入成功: ${JSON.stringify(response.data)}`); + setImportTypes([]); + setImportFileName(""); + } else { + toast.error("导入失败: " + response.msg); + } + } catch { + toast.error("导入失败,请检查文件格式"); + } finally { + setImporting(false); + if (fileInputRef.current) { + fileInputRef.current.value = ""; + } + } + }; + if (loading) { return (
@@ -427,6 +489,143 @@ export default function ConfigPage() { )} + + {/* 备份与恢复 */} + + +
+
+

数据备份与恢复

+

+ 导出或导入系统数据,支持选择特定数据类型 +

+
+
+
+ + + + + {/* 导出部分 */} +
+

导出数据

+

+ 选择要导出的数据类型,导出为 JSON 格式文件 +

+ + setExportTypes(values as string[])} + > + 用户 + 节点 + 隧道 + 转发 + 用户隧道权限 + 限速规则 + 隧道分组 + 用户分组 + 分组权限 + 系统配置 + + +
+ + + +
+
+ + + + {/* 导入部分 */} +
+

导入数据

+

+ 选择要导入的数据类型,支持从备份文件恢复数据 +

+ + setImportTypes(values as string[])} + > + 用户 + 节点 + 隧道 + 转发 + 用户隧道权限 + 限速规则 + 隧道分组 + 用户分组 + 分组权限 + 系统配置 + + + + +
+ + {importFileName && ( + + 已选择: {importFileName} + + )} +
+
+
+
); } From 3b294c6b9e41a71c984d5a39536d6a4025785107 Mon Sep 17 00:00:00 2001 From: sagit Date: Fri, 13 Feb 2026 07:25:27 +0000 Subject: [PATCH 10/12] chore(frontend): fix HeroUI deps and apply lint cleanup --- vite-frontend/package.json | 2 +- vite-frontend/src/api/index.ts | 48 +- vite-frontend/src/pages/config.tsx | 22 +- vite-frontend/src/pages/index.tsx | 44 +- vite-frontend/src/pages/node.tsx | 694 ++++++++++++---------- vite-frontend/src/pages/panel-sharing.tsx | 291 +++++++-- 6 files changed, 673 insertions(+), 428 deletions(-) diff --git a/vite-frontend/package.json b/vite-frontend/package.json index b0b2eaf..aea3441 100644 --- a/vite-frontend/package.json +++ b/vite-frontend/package.json @@ -40,7 +40,7 @@ "@heroui/system": "2.4.19", "@heroui/table": "^2.2.24", "@heroui/tabs": "^2.2.27", - "@heroui/theme": "2.4.19", + "@heroui/theme": "2.4.24", "@heroui/use-theme": "2.1.10", "@marsidev/react-turnstile": "^1.1.0", "@nextui-org/system": "^2.4.6", diff --git a/vite-frontend/src/api/index.ts b/vite-frontend/src/api/index.ts index 265cc00..7abb7f7 100644 --- a/vite-frontend/src/api/index.ts +++ b/vite-frontend/src/api/index.ts @@ -1,3 +1,5 @@ +import axios from "axios"; + import Network from "./network"; // 登陆相关接口 @@ -42,9 +44,17 @@ export const checkNodeStatus = (nodeId?: number) => { }; export const upgradeNode = (id: number, version?: string) => - Network.post("/node/upgrade", { id, version: version || "" }, { timeout: 5 * 60 * 1000 }); + Network.post( + "/node/upgrade", + { id, version: version || "" }, + { timeout: 5 * 60 * 1000 }, + ); export const batchUpgradeNodes = (ids: number[], version?: string) => - Network.post("/node/batch-upgrade", { ids, version: version || "" }, { timeout: 15 * 60 * 1000 }); + Network.post( + "/node/batch-upgrade", + { ids, version: version || "" }, + { timeout: 15 * 60 * 1000 }, + ); export const getNodeReleases = () => Network.post("/node/releases"); export const rollbackNode = (id: number) => Network.post("/node/rollback", { id }); @@ -224,12 +234,8 @@ export const resetPeerShareFlow = (id: number) => Network.post("/federation/share/reset-flow", { id }); export const getPeerRemoteUsageList = () => Network.post("/federation/share/remote-usage/list"); -export const importRemoteNode = (data: { - remoteUrl: string; - token: string; -}) => Network.post("/federation/node/import", data); - -import axios from "axios"; +export const importRemoteNode = (data: { remoteUrl: string; token: string }) => + Network.post("/federation/node/import", data); export interface BackupTypes { users?: boolean; @@ -247,19 +253,25 @@ export interface BackupTypes { export const exportBackup = async (types: string[] = []) => { const token = window.localStorage.getItem("token"); const baseURL = axios.defaults.baseURL || "/api/v1/"; - - const response = await axios.post(`${baseURL}/backup/export`, { types }, { - headers: { - Authorization: token, - "Content-Type": "application/json", + + const response = await axios.post( + `${baseURL}/backup/export`, + { types }, + { + headers: { + Authorization: token, + "Content-Type": "application/json", + }, + responseType: "blob", }, - responseType: "blob", - }); + ); const url = window.URL.createObjectURL(new Blob([response.data])); const link = document.createElement("a"); + link.href = url; const timestamp = new Date().toISOString().slice(0, 19).replace(/[:-]/g, ""); + link.setAttribute("download", `backup_${timestamp}.json`); document.body.appendChild(link); link.click(); @@ -267,7 +279,5 @@ export const exportBackup = async (types: string[] = []) => { window.URL.revokeObjectURL(url); }; -export const importBackup = (data: { - types: string[]; - [key: string]: any; -}) => Network.post("/backup/import", data); +export const importBackup = (data: { types: string[]; [key: string]: any }) => + Network.post("/backup/import", data); diff --git a/vite-frontend/src/pages/config.tsx b/vite-frontend/src/pages/config.tsx index 71bc198..163931a 100644 --- a/vite-frontend/src/pages/config.tsx +++ b/vite-frontend/src/pages/config.tsx @@ -342,6 +342,7 @@ export default function ConfigPage() { const handleExport = async () => { if (exportTypes.length === 0) { toast.error("请至少选择一种数据类型"); + return; } setExporting(true); @@ -357,10 +358,12 @@ export default function ConfigPage() { const handleFileChange = async (e: React.ChangeEvent) => { const file = e.target.files?.[0]; + if (!file) return; if (importTypes.length === 0) { toast.error("请先选择要导入的数据类型"); + return; } @@ -512,13 +515,13 @@ export default function ConfigPage() {

选择要导出的数据类型,导出为 JSON 格式文件

- + setExportTypes(values as string[])} > @@ -561,10 +564,7 @@ export default function ConfigPage() { > 全选 - @@ -580,11 +580,11 @@ export default function ConfigPage() {

setImportTypes(values as string[])} > @@ -601,18 +601,18 @@ export default function ConfigPage() {
- - -
- )} -
+ {/* 操作按钮 */} +
{!isRemoteNode && ( +
+ + + +
+ )} +
+ {!isRemoteNode && ( + + )} - )} - +
-
- - - )} - + + + )} + ); })} @@ -2041,6 +2097,7 @@ export default function NodePage() { selectedKeys={selectedVersion ? [selectedVersion] : []} onSelectionChange={(keys) => { const selected = Array.from(keys)[0] as string; + setSelectedVersion(selected || ""); }} > @@ -2053,7 +2110,12 @@ export default function NodePage() { ? new Date(r.publishedAt).toLocaleDateString() : ""} {r.prerelease && ( - + 预览 )} diff --git a/vite-frontend/src/pages/panel-sharing.tsx b/vite-frontend/src/pages/panel-sharing.tsx index df05e73..b216042 100644 --- a/vite-frontend/src/pages/panel-sharing.tsx +++ b/vite-frontend/src/pages/panel-sharing.tsx @@ -12,6 +12,7 @@ import { } from "@heroui/modal"; import { Select, SelectItem } from "@heroui/select"; import { toast } from "react-hot-toast"; + import { getNodeList, createPeerShare, @@ -128,6 +129,7 @@ export default function PanelSharingPage() { setLoading(true); try { const res = await getPeerShareList(); + if (res.code === 0) { setShares(res.data || []); } else { @@ -141,10 +143,12 @@ export default function PanelSharingPage() { const loadNodes = useCallback(async () => { try { const res = await getNodeList(); + if (res.code === 0) { const localNodes: Node[] = (res.data || []).filter( (node: Node) => (node?.isRemote ?? 0) !== 1, ); + setNodes(localNodes); setShareForm((prev) => { if (!prev.nodeId) { @@ -153,6 +157,7 @@ export default function PanelSharingPage() { const hasSelectedNode = localNodes.some( (node: Node) => String(node.id) === prev.nodeId, ); + return hasSelectedNode ? prev : { ...prev, nodeId: "" }; }); } @@ -165,6 +170,7 @@ export default function PanelSharingPage() { setRemoteUsageLoading(true); try { const res = await getPeerRemoteUsageList(); + if (res.code === 0) { setRemoteUsageNodes(res.data || []); } else { @@ -179,6 +185,7 @@ export default function PanelSharingPage() { if (selectedTab === "my-shares") { loadShares(); loadNodes(); + return; } if (selectedTab === "remote-nodes") { @@ -189,15 +196,19 @@ export default function PanelSharingPage() { const handleCreateShare = async () => { if (!shareForm.name || !shareForm.nodeId) { toast.error("请填写必要信息"); + return; } const nodeId = parseInt(shareForm.nodeId, 10); + if (Number.isNaN(nodeId) || !nodes.some((node) => node.id === nodeId)) { toast.error("仅可选择本地节点"); + return; } if (shareForm.maxBandwidth < 0) { toast.error("流量上限不能为负数"); + return; } try { @@ -213,6 +224,7 @@ export default function PanelSharingPage() { allowedDomains: shareForm.allowedDomains, allowedIps: shareForm.allowedIps, }); + if (res.code === 0) { toast.success("创建成功"); setCreateShareOpen(false); @@ -228,6 +240,7 @@ export default function PanelSharingPage() { const handleDeleteShare = async (id: number) => { try { const res = await deletePeerShare(id); + if (res.code === 0) { toast.success("删除成功"); loadShares(); @@ -242,6 +255,7 @@ export default function PanelSharingPage() { const handleResetShareFlow = async (id: number) => { try { const res = await resetPeerShareFlow(id); + if (res.code === 0) { toast.success("共享流量已重置"); loadShares(); @@ -257,7 +271,10 @@ export default function PanelSharingPage() { setEditForm({ id: share.id, name: share.name, - maxBandwidth: share.maxBandwidth > 0 ? Math.round(share.maxBandwidth / (1024 * 1024 * 1024)) : 0, + maxBandwidth: + share.maxBandwidth > 0 + ? Math.round(share.maxBandwidth / (1024 * 1024 * 1024)) + : 0, expiryTime: share.expiryTime, portRangeStart: share.portRangeStart, portRangeEnd: share.portRangeEnd, @@ -270,10 +287,12 @@ export default function PanelSharingPage() { const handleEditShare = async () => { if (!editForm.name) { toast.error("名称不能为空"); + return; } if (editForm.maxBandwidth < 0) { toast.error("流量上限不能为负数"); + return; } try { @@ -287,6 +306,7 @@ export default function PanelSharingPage() { allowedDomains: editForm.allowedDomains, allowedIps: editForm.allowedIps, }); + if (res.code === 0) { toast.success("编辑成功"); setEditShareOpen(false); @@ -302,19 +322,22 @@ export default function PanelSharingPage() { const handleImportNode = async () => { if (!importForm.remoteUrl || !importForm.token) { toast.error("请填写完整信息"); + return; } try { // Automatically add http/https if missing let url = importForm.remoteUrl.trim(); + if (!url.startsWith("http")) { url = "http://" + url; } - + const res = await importRemoteNode({ remoteUrl: url, token: importForm.token.trim(), }); + if (res.code === 0) { toast.success("导入成功,请前往节点列表查看"); setImportNodeOpen(false); @@ -341,6 +364,7 @@ export default function PanelSharingPage() { if (bytes < 1024 * 1024) return (bytes / 1024).toFixed(2) + " KB"; if (bytes < 1024 * 1024 * 1024) return (bytes / (1024 * 1024)).toFixed(2) + " MB"; + return (bytes / (1024 * 1024 * 1024)).toFixed(2) + " GB"; }; @@ -351,6 +375,7 @@ export default function PanelSharingPage() { if (chainType === 3) { return "出口节点"; } + return "未知链路"; }; @@ -370,11 +395,14 @@ export default function PanelSharingPage() {
-
- + {loading ? (
加载中...
) : shares.length === 0 ? ( @@ -382,7 +410,10 @@ export default function PanelSharingPage() { ) : (
{shares.map((share) => ( - +

{share.name}

@@ -400,29 +431,67 @@ export default function PanelSharingPage() { > 重置流量 - +
-

端口范围: {share.portRangeStart} - {share.portRangeEnd}

-

流量上限: {share.maxBandwidth > 0 ? formatFlowGB(share.maxBandwidth) : "不限制"}

+

+ 端口范围: {share.portRangeStart} -{" "} + {share.portRangeEnd} +

+

+ 流量上限:{" "} + {share.maxBandwidth > 0 + ? formatFlowGB(share.maxBandwidth) + : "不限制"} +

当前流量: {formatFlowGB(share.currentFlow || 0)}

-

远程占用端口: {share.usedPorts && share.usedPorts.length > 0 ? share.usedPorts.join(", ") : "暂无"}

- {share.usedPortDetails && share.usedPortDetails.length > 0 && ( -
- {share.usedPortDetails.map((item) => ( - - {item.port} / {item.role || "reserved"} - - ))} -
+

+ 远程占用端口:{" "} + {share.usedPorts && share.usedPorts.length > 0 + ? share.usedPorts.join(", ") + : "暂无"} +

+ {share.usedPortDetails && + share.usedPortDetails.length > 0 && ( +
+ {share.usedPortDetails.map((item) => ( + + {item.port} / {item.role || "reserved"} + + ))} +
+ )} + {share.allowedDomains && ( +

允许域名: {share.allowedDomains}

)} - {share.allowedDomains &&

允许域名: {share.allowedDomains}

} - {share.allowedIps &&

允许API IP: {share.allowedIps}

} -

过期时间: {share.expiryTime === 0 ? "永久" : new Date(share.expiryTime).toLocaleDateString()}

+ {share.allowedIps && ( +

允许API IP: {share.allowedIps}

+ )} +

+ 过期时间:{" "} + {share.expiryTime === 0 + ? "永久" + : new Date(share.expiryTime).toLocaleDateString()} +

- +
@@ -436,7 +505,10 @@ export default function PanelSharingPage() {
-
@@ -446,15 +518,22 @@ export default function PanelSharingPage() { ) : remoteUsageNodes.length === 0 ? (

暂无远程节点占用记录。

-

导入远程节点并创建隧道后,这里会显示远端端口占用情况。

+

+ 导入远程节点并创建隧道后,这里会显示远端端口占用情况。 +

) : (
{remoteUsageNodes.map((node) => ( - +

{node.nodeName}

- 绑定 {node.activeBindingNum || 0} + + 绑定 {node.activeBindingNum || 0} +
{node.syncError && ( @@ -470,18 +549,40 @@ export default function PanelSharingPage() { )} {node.remoteUrl &&

远程地址: {node.remoteUrl}

}

共享ID: {node.shareId || "-"}

-

端口范围: {node.portRangeStart > 0 && node.portRangeEnd > 0 ? `${node.portRangeStart} - ${node.portRangeEnd}` : "-"}

-

共享流量: {node.maxBandwidth > 0 ? `${formatFlowGB(node.currentFlow || 0)} / ${formatFlowGB(node.maxBandwidth)}` : `${formatFlowGB(node.currentFlow || 0)} / 不限制`}

-

远端占用端口: {node.usedPorts && node.usedPorts.length > 0 ? node.usedPorts.join(", ") : "暂无"}

+

+ 端口范围:{" "} + {node.portRangeStart > 0 && node.portRangeEnd > 0 + ? `${node.portRangeStart} - ${node.portRangeEnd}` + : "-"} +

+

+ 共享流量:{" "} + {node.maxBandwidth > 0 + ? `${formatFlowGB(node.currentFlow || 0)} / ${formatFlowGB(node.maxBandwidth)}` + : `${formatFlowGB(node.currentFlow || 0)} / 不限制`} +

+

+ 远端占用端口:{" "} + {node.usedPorts && node.usedPorts.length > 0 + ? node.usedPorts.join(", ") + : "暂无"} +

{node.bindings && node.bindings.length > 0 && (
{node.bindings.map((binding) => ( -

- 隧道 {binding.tunnelName || `#${binding.tunnelId}`} +

+ 隧道{" "} + {binding.tunnelName || `#${binding.tunnelId}`} {" · "} 端口 {binding.allocatedPort} {" · "} - {formatChainType(binding.chainType, binding.hopInx)} + {formatChainType( + binding.chainType, + binding.hopInx, + )}

))}
@@ -505,13 +606,17 @@ export default function PanelSharingPage() { label="名称" placeholder="备注名称" value={shareForm.name} - onChange={(e) => setShareForm({ ...shareForm, name: e.target.value })} + onChange={(e) => + setShareForm({ ...shareForm, name: e.target.value }) + } /> setShareForm({ ...shareForm, portRangeEnd: parseInt(e.target.value) })} + onChange={(e) => + setShareForm({ + ...shareForm, + portRangeEnd: parseInt(e.target.value), + }) + } />
setShareForm({ ...shareForm, expiryDays: parseInt(e.target.value) })} + onChange={(e) => + setShareForm({ + ...shareForm, + expiryDays: parseInt(e.target.value), + }) + } /> setShareForm({ ...shareForm, maxBandwidth: parseInt(e.target.value, 10) || 0 })} + onChange={(e) => + setShareForm({ + ...shareForm, + maxBandwidth: parseInt(e.target.value, 10) || 0, + }) + } /> setShareForm({ ...shareForm, allowedDomains: e.target.value })} + onChange={(e) => + setShareForm({ ...shareForm, allowedDomains: e.target.value }) + } /> setShareForm({ ...shareForm, allowedIps: e.target.value })} + onChange={(e) => + setShareForm({ ...shareForm, allowedIps: e.target.value }) + } /> - + @@ -578,54 +709,88 @@ export default function PanelSharingPage() { label="名称" placeholder="备注名称" value={editForm.name} - onChange={(e) => setEditForm({ ...editForm, name: e.target.value })} + onChange={(e) => + setEditForm({ ...editForm, name: e.target.value }) + } />
setEditForm({ ...editForm, portRangeStart: parseInt(e.target.value) || 0 })} + onChange={(e) => + setEditForm({ + ...editForm, + portRangeStart: parseInt(e.target.value) || 0, + }) + } /> setEditForm({ ...editForm, portRangeEnd: parseInt(e.target.value) || 0 })} + onChange={(e) => + setEditForm({ + ...editForm, + portRangeEnd: parseInt(e.target.value) || 0, + }) + } />
setEditForm({ ...editForm, maxBandwidth: parseInt(e.target.value, 10) || 0 })} + onChange={(e) => + setEditForm({ + ...editForm, + maxBandwidth: parseInt(e.target.value, 10) || 0, + }) + } /> 0 ? new Date(editForm.expiryTime).toISOString().slice(0, 16) : ""} - onChange={(e) => setEditForm({ ...editForm, expiryTime: e.target.value ? new Date(e.target.value).getTime() : 0 })} + value={ + editForm.expiryTime > 0 + ? new Date(editForm.expiryTime).toISOString().slice(0, 16) + : "" + } + onChange={(e) => + setEditForm({ + ...editForm, + expiryTime: e.target.value + ? new Date(e.target.value).getTime() + : 0, + }) + } /> setEditForm({ ...editForm, allowedDomains: e.target.value })} + onChange={(e) => + setEditForm({ ...editForm, allowedDomains: e.target.value }) + } /> setEditForm({ ...editForm, allowedIps: e.target.value })} + onChange={(e) => + setEditForm({ ...editForm, allowedIps: e.target.value }) + } /> - + @@ -639,18 +804,24 @@ export default function PanelSharingPage() { label="远程面板地址" placeholder="http://panel.example.com:8088" value={importForm.remoteUrl} - onChange={(e) => setImportForm({ ...importForm, remoteUrl: e.target.value })} + onChange={(e) => + setImportForm({ ...importForm, remoteUrl: e.target.value }) + } /> setImportForm({ ...importForm, token: e.target.value })} + onChange={(e) => + setImportForm({ ...importForm, token: e.target.value }) + } /> - + From 5a9715eb2639a8e219189148975de7d480c1f155 Mon Sep 17 00:00:00 2001 From: sagit Date: Fri, 13 Feb 2026 08:17:34 +0000 Subject: [PATCH 11/12] 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 12/12] 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")