Files
flvx/go-backend/internal/store/repo/repository_mutations.go
T

1615 lines
47 KiB
Go

package repo
import (
"crypto/rand"
"database/sql"
"errors"
"fmt"
"math/big"
"sort"
"strconv"
"strings"
"time"
"go-backend/internal/store/model"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
func (r *Repository) UserExists(username string) (bool, error) {
if r == nil || r.db == nil {
return false, errors.New("repository not initialized")
}
var cnt int64
err := r.db.Model(&model.User{}).Where(`"user" = ?`, username).Count(&cnt).Error
return cnt > 0, err
}
func (r *Repository) UserExistsExcluding(username string, excludeID int64) (bool, error) {
if r == nil || r.db == nil {
return false, errors.New("repository not initialized")
}
var cnt int64
err := r.db.Model(&model.User{}).
Where(`"user" = ? AND id != ?`, username, excludeID).
Count(&cnt).Error
return cnt > 0, err
}
func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, flow, flowResetTime int64, num, status, maxConn int, now int64) (int64, error) {
if r == nil || r.db == nil {
return 0, errors.New("repository not initialized")
}
user := model.User{
User: username,
Pwd: pwdHash,
RoleID: roleID,
ExpTime: expTime,
Flow: flow,
InFlow: 0,
OutFlow: 0,
FlowResetTime: flowResetTime,
Num: num,
MaxConn: maxConn,
CreatedTime: now,
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
Status: status,
}
if err := r.db.Create(&user).Error; err != nil {
return 0, err
}
return user.ID, nil
}
func (r *Repository) GetUserRoleID(userID int64) (int, error) {
if r == nil || r.db == nil {
return 0, errors.New("repository not initialized")
}
var user model.User
err := r.db.Select("role_id").Where("id = ?", userID).First(&user).Error
if err != nil {
return 0, normalizeNotFoundErr(err)
}
return user.RoleID, nil
}
func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string, flow int64, num int, expTime, flowResetTime int64, status, maxConn int, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.User{}).
Where("id = ?", id).
Updates(map[string]interface{}{
"user": username,
"pwd": pwdHash,
"flow": flow,
"num": num,
"exp_time": expTime,
"flow_reset_time": flowResetTime,
"status": status,
"max_conn": maxConn,
"updated_time": sql.NullInt64{Int64: now, Valid: true},
}).Error
}
func (r *Repository) UpdateUserWithoutPassword(id int64, username string, flow int64, num int, expTime, flowResetTime int64, status, maxConn int, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.User{}).
Where("id = ?", id).
Updates(map[string]interface{}{
"user": username,
"flow": flow,
"num": num,
"exp_time": expTime,
"flow_reset_time": flowResetTime,
"status": status,
"max_conn": maxConn,
"updated_time": sql.NullInt64{Int64: now, Valid: true},
}).Error
}
func (r *Repository) PropagateUserFlowToTunnels(userID int64, flow int64, num int, expTime, flowResetTime int64) {
if r == nil || r.db == nil {
return
}
_ = r.db.Model(&model.UserTunnel{}).
Where("user_id = ?", userID).
Updates(map[string]interface{}{
"flow": flow,
"num": num,
"exp_time": expTime,
"flow_reset_time": flowResetTime,
}).Error
}
func (r *Repository) DeleteUserCascade(userID int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Transaction(func(tx *gorm.DB) error {
forwardIDs := tx.Model(&model.Forward{}).Select("id").Where("user_id = ?", userID)
if err := tx.Where("forward_id IN (?)", forwardIDs).Delete(&model.ForwardPort{}).Error; err != nil {
return err
}
if err := tx.Where("user_id = ?", userID).Delete(&model.Forward{}).Error; err != nil {
return err
}
userTunnelIDs := tx.Model(&model.UserTunnel{}).Select("id").Where("user_id = ?", userID)
if err := tx.Where("user_tunnel_id IN (?)", userTunnelIDs).Delete(&model.GroupPermissionGrant{}).Error; err != nil {
return err
}
if err := tx.Where("user_id = ?", userID).Delete(&model.UserTunnel{}).Error; err != nil {
return err
}
if err := tx.Where("user_id = ?", userID).Delete(&model.UserGroupUser{}).Error; err != nil {
return err
}
if err := tx.Where("user_id = ?", userID).Delete(&model.StatisticsFlow{}).Error; err != nil {
return err
}
if err := tx.Where("user_id = ?", userID).Delete(&model.UserQuota{}).Error; err != nil {
return err
}
return tx.Where("id = ?", userID).Delete(&model.User{}).Error
})
}
func (r *Repository) ResetUserFlowByUser(userID int64, now int64) {
if r == nil || r.db == nil {
return
}
_ = r.db.Model(&model.User{}).
Where("id = ?", userID).
Updates(map[string]interface{}{
"in_flow": 0,
"out_flow": 0,
"updated_time": sql.NullInt64{Int64: now, Valid: true},
}).Error
_ = r.db.Model(&model.UserTunnel{}).
Where("user_id = ?", userID).
Updates(map[string]interface{}{"in_flow": 0, "out_flow": 0}).Error
}
func (r *Repository) ResetUserFlowByUserTunnel(userTunnelID int64) {
if r == nil || r.db == nil {
return
}
_ = r.db.Model(&model.UserTunnel{}).
Where("id = ?", userTunnelID).
Updates(map[string]interface{}{"in_flow": 0, "out_flow": 0}).Error
}
func (r *Repository) GetUsernameByID(userID int64) string {
if r == nil || r.db == nil {
return ""
}
var user model.User
if err := r.db.Select("user").Where("id = ?", userID).First(&user).Error; err != nil {
return ""
}
return user.User
}
func (r *Repository) GetUserDefaultsForTunnel(userID int64) (flow int64, num int, expTime int64, flowReset int64, err error) {
if r == nil || r.db == nil {
return 0, 0, 0, 0, errors.New("repository not initialized")
}
var user model.User
err = r.db.Select("flow", "num", "exp_time", "flow_reset_time").Where("id = ?", userID).First(&user).Error
if err != nil {
return 0, 0, 0, 0, normalizeNotFoundErr(err)
}
return user.Flow, user.Num, user.ExpTime, user.FlowResetTime, nil
}
func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serverIPV6, port, interfaceName, version, remark, expiryTime, renewalCycle interface{}, httpFlag, tlsFlag, socksFlag int, now int64, status int, tcpAddr, udpAddr string, inx, isRemote int, remoteURL, remoteToken, remoteConfig, extraIPs interface{}) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
node := model.Node{
Name: name,
Remark: nullStringFromInterface(remark),
ExpiryTime: nullInt64FromInterface(expiryTime),
RenewalCycle: nullStringFromInterface(renewalCycle),
Secret: secret,
ServerIP: serverIP,
ServerIPV4: nullStringFromInterface(serverIPV4),
ServerIPV6: nullStringFromInterface(serverIPV6),
ExtraIPs: nullStringFromInterface(extraIPs),
Port: stringFromInterface(port),
InterfaceName: nullStringFromInterface(interfaceName),
Version: nullStringFromInterface(version),
HTTP: httpFlag,
TLS: tlsFlag,
Socks: socksFlag,
CreatedTime: now,
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
Status: status,
TCPListenAddr: tcpAddr,
UDPListenAddr: udpAddr,
Inx: inx,
IsRemote: isRemote,
RemoteURL: nullStringFromInterface(remoteURL),
RemoteToken: nullStringFromInterface(remoteToken),
RemoteConfig: nullStringFromInterface(remoteConfig),
}
return r.db.Create(&node).Error
}
func (r *Repository) GetNodeStatusFields(nodeID int64) (status, httpFlag, tlsFlag, socksFlag int, err error) {
if r == nil || r.db == nil {
return 0, 0, 0, 0, errors.New("repository not initialized")
}
var node model.Node
err = r.db.Select("status", "http", "tls", "socks").Where("id = ?", nodeID).First(&node).Error
if err != nil {
return 0, 0, 0, 0, normalizeNotFoundErr(err)
}
return node.Status, node.HTTP, node.TLS, node.Socks, nil
}
func (r *Repository) UpdateNode(id int64, name, serverIP string, serverIPV4, serverIPV6, port, interfaceName, extraIPs, remark, expiryTime, renewalCycle interface{}, httpFlag, tlsFlag, socksFlag int, tcpAddr, udpAddr string, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.Node{}).
Where("id = ?", id).
Updates(map[string]interface{}{
"name": name,
"remark": nullStringFromInterface(remark),
"expiry_time": nullInt64FromInterface(expiryTime),
"renewal_cycle": nullStringFromInterface(renewalCycle),
"server_ip": serverIP,
"server_ip_v4": nullStringFromInterface(serverIPV4),
"server_ip_v6": nullStringFromInterface(serverIPV6),
"extra_ips": nullStringFromInterface(extraIPs),
"port": stringFromInterface(port),
"interface_name": nullStringFromInterface(interfaceName),
"http": httpFlag,
"tls": tlsFlag,
"socks": socksFlag,
"tcp_listen_addr": tcpAddr,
"udp_listen_addr": udpAddr,
"updated_time": sql.NullInt64{Int64: now, Valid: true},
"expiry_reminder_dismissed": 0,
}).Error
}
func (r *Repository) GetNodeSecret(nodeID int64) (string, error) {
if r == nil || r.db == nil {
return "", errors.New("repository not initialized")
}
var node model.Node
err := r.db.Select("secret").Where("id = ?", nodeID).First(&node).Error
if err != nil {
return "", normalizeNotFoundErr(err)
}
return node.Secret, nil
}
func (r *Repository) GetViteConfigValue(name string) (string, error) {
if r == nil || r.db == nil {
return "", errors.New("repository not initialized")
}
var cfg model.ViteConfig
err := r.db.Select("value").Where("name = ?", name).First(&cfg).Error
if err != nil {
return "", normalizeNotFoundErr(err)
}
return cfg.Value, nil
}
func (r *Repository) UpdateNodeOrder(nodeID int64, inx int, now int64) {
if r == nil || r.db == nil {
return
}
_ = r.db.Model(&model.Node{}).
Where("id = ?", nodeID).
Updates(map[string]interface{}{
"inx": inx,
"updated_time": sql.NullInt64{Int64: now, Valid: true},
}).Error
}
func (r *Repository) UpdateNodeExpiryReminderDismissed(nodeID int64, dismissed int) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.Node{}).
Where("id = ?", nodeID).
Update("expiry_reminder_dismissed", dismissed).Error
}
func (r *Repository) DeleteNodeCascade(nodeID int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Transaction(func(tx *gorm.DB) error {
if err := tx.Where("node_id = ?", nodeID).Delete(&model.ForwardPort{}).Error; err != nil {
return err
}
if err := tx.Where("node_id = ?", nodeID).Delete(&model.ChainTunnel{}).Error; err != nil {
return err
}
if err := tx.Where("node_id = ?", nodeID).Delete(&model.FederationTunnelBinding{}).Error; err != nil {
return err
}
return tx.Where("id = ?", nodeID).Delete(&model.Node{}).Error
})
}
func (r *Repository) GetNodeRemoteFields(nodeID int64) (isRemote int, remoteURL, remoteToken sql.NullString, err error) {
if r == nil || r.db == nil {
return 0, sql.NullString{}, sql.NullString{}, errors.New("repository not initialized")
}
return r.GetNodeRemoteFieldsTx(r.db, nodeID)
}
func (r *Repository) GetNodeRemoteFieldsTx(tx *gorm.DB, nodeID int64) (isRemote int, remoteURL, remoteToken sql.NullString, err error) {
if tx == nil {
return 0, sql.NullString{}, sql.NullString{}, errors.New("database unavailable")
}
var node model.Node
err = tx.Select("is_remote", "remote_url", "remote_token").Where("id = ?", nodeID).First(&node).Error
if err != nil {
return 0, sql.NullString{}, sql.NullString{}, normalizeNotFoundErr(err)
}
return node.IsRemote, node.RemoteURL, node.RemoteToken, nil
}
func (r *Repository) GetNodePortRange(nodeID int64) (string, error) {
if r == nil || r.db == nil {
return "", errors.New("repository not initialized")
}
var node model.Node
err := r.db.Select("port").Where("id = ?", nodeID).First(&node).Error
if err != nil {
return "", normalizeNotFoundErr(err)
}
return node.Port, nil
}
func (r *Repository) TunnelNameExists(name string) (bool, error) {
if r == nil || r.db == nil {
return false, errors.New("repository not initialized")
}
var cnt int64
err := r.db.Model(&model.Tunnel{}).Where("name = ?", name).Count(&cnt).Error
return cnt > 0, err
}
func (r *Repository) BeginTx() *gorm.DB {
if r == nil || r.db == nil {
return nil
}
return r.db.Begin()
}
func (r *Repository) UpdateTunnelOrder(tunnelID int64, inx int, now int64) {
if r == nil || r.db == nil {
return
}
_ = r.db.Model(&model.Tunnel{}).
Where("id = ?", tunnelID).
Updates(map[string]interface{}{"inx": inx, "updated_time": now}).Error
}
func (r *Repository) UpdateTunnelTx(tx *gorm.DB, tunnelID int64, name string, typeVal int, flow int64, trafficRatio float64, status int, inIP, ipPreference string, protocol string, now int64) error {
if tx == nil {
return errors.New("database unavailable")
}
return tx.Model(&model.Tunnel{}).
Where("id = ?", tunnelID).
Updates(map[string]interface{}{
"name": name,
"type": typeVal,
"flow": flow,
"traffic_ratio": trafficRatio,
"status": status,
"in_ip": nullStringFromInterface(inIP),
"ip_preference": ipPreference,
"protocol": protocol,
"updated_time": now,
}).Error
}
func (r *Repository) DeleteChainTunnelsByTunnelTx(tx *gorm.DB, tunnelID int64) error {
if tx == nil {
return errors.New("database unavailable")
}
return tx.Where("tunnel_id = ?", tunnelID).Delete(&model.ChainTunnel{}).Error
}
func (r *Repository) CreateChainTunnelTx(tx *gorm.DB, tunnelID int64, chainType string, nodeID int64, port sql.NullInt64, strategy string, inx int, protocol string, connectIp string) error {
if tx == nil {
return errors.New("database unavailable")
}
ct := model.ChainTunnel{
TunnelID: tunnelID,
ChainType: chainType,
NodeID: nodeID,
Port: port,
Strategy: nullStringFromInterface(strategy),
Inx: nullInt64FromInterface(inx),
Protocol: nullStringFromInterface(protocol),
ConnectIP: sql.NullString{String: connectIp, Valid: connectIp != ""},
}
return tx.Create(&ct).Error
}
func (r *Repository) IsRemoteNodeTx(tx *gorm.DB, nodeID int64) (bool, error) {
if tx == nil {
return false, errors.New("database unavailable")
}
if nodeID <= 0 {
return false, errors.New("节点不存在")
}
var node model.Node
err := tx.Select("is_remote").Where("id = ?", nodeID).First(&node).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, errors.New("节点不存在")
}
if err != nil {
return false, err
}
return node.IsRemote == 1, nil
}
func (r *Repository) PickNodePortTx(tx *gorm.DB, nodeID int64, allocated map[int64]int, excludeTunnelID int64) (int, error) {
return r.pickNodePortTx(tx, nodeID, allocated, excludeTunnelID, false)
}
func (r *Repository) PickRandomNodePortTx(tx *gorm.DB, nodeID int64, allocated map[int64]int, excludeTunnelID int64) (int, error) {
return r.pickNodePortTx(tx, nodeID, allocated, excludeTunnelID, true)
}
func (r *Repository) pickNodePortTx(tx *gorm.DB, nodeID int64, allocated map[int64]int, excludeTunnelID int64, randomPick bool) (int, error) {
if tx == nil {
return 0, errors.New("database unavailable")
}
if nodeID <= 0 {
return 0, errors.New("节点不存在")
}
if port, ok := allocated[nodeID]; ok && port > 0 {
return port, nil
}
var node model.Node
err := tx.Select("port").Where("id = ?", nodeID).First(&node).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return 0, errors.New("节点不存在")
}
if err != nil {
return 0, err
}
candidates := parsePortRangeSpec(node.Port)
if len(candidates) == 0 {
return 0, errors.New("节点端口已满,无可用端口")
}
used := make(map[int]struct{})
chainQuery := tx.Model(&model.ChainTunnel{}).Where("node_id = ? AND port > 0", nodeID)
if excludeTunnelID > 0 {
chainQuery = chainQuery.Where("tunnel_id != ?", excludeTunnelID)
}
var chainPorts []int
if err := chainQuery.Pluck("port", &chainPorts).Error; err != nil {
return 0, err
}
for _, p := range chainPorts {
if p > 0 {
used[p] = struct{}{}
}
}
var forwardPorts []int
if err := tx.Model(&model.ForwardPort{}).
Where("node_id = ? AND port > 0", nodeID).
Pluck("port", &forwardPorts).Error; err != nil {
return 0, err
}
for _, p := range forwardPorts {
if p > 0 {
used[p] = struct{}{}
}
}
var available []int
for _, candidate := range candidates {
if candidate <= 0 {
continue
}
if _, ok := used[candidate]; ok {
continue
}
available = append(available, candidate)
}
if len(available) == 0 {
return 0, errors.New("节点端口已满,无可用端口")
}
if !randomPick {
allocated[nodeID] = available[0]
return available[0], nil
}
idx, err := rand.Int(rand.Reader, big.NewInt(int64(len(available))))
if err != nil {
allocated[nodeID] = available[0]
return available[0], nil
}
port := available[idx.Int64()]
allocated[nodeID] = port
return port, nil
}
func (r *Repository) GetTunnelIPPreference(tunnelID int64) string {
if r == nil || r.db == nil {
return ""
}
var tunnel model.Tunnel
if err := r.db.Select("ip_preference").Where("id = ?", tunnelID).First(&tunnel).Error; err != nil {
return ""
}
return tunnel.IPPreference
}
func (r *Repository) DeleteTunnelCascade(tunnelID int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Transaction(func(tx *gorm.DB) error {
forwardIDs := tx.Model(&model.Forward{}).Select("id").Where("tunnel_id = ?", tunnelID)
if err := tx.Where("forward_id IN (?)", forwardIDs).Delete(&model.ForwardPort{}).Error; err != nil {
return err
}
if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.Forward{}).Error; err != nil {
return err
}
if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.UserTunnel{}).Error; err != nil {
return err
}
if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.ChainTunnel{}).Error; err != nil {
return err
}
if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.FederationTunnelBinding{}).Error; err != nil {
return err
}
return tx.Where("id = ?", tunnelID).Delete(&model.Tunnel{}).Error
})
}
func (r *Repository) TunnelEntryNodeIDs(tunnelID int64) ([]int64, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var ids []int64
err := r.db.Model(&model.ChainTunnel{}).
Where("tunnel_id = ? AND chain_type = ?", tunnelID, "1").
Order("inx ASC, id ASC").
Pluck("node_id", &ids).Error
if err != nil {
return nil, err
}
return ids, nil
}
func (r *Repository) DeleteUserTunnel(id int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Where("id = ?", id).Delete(&model.UserTunnel{}).Error
}
func (r *Repository) UpdateUserTunnel(id int64, flow int64, num int, expTime, flowResetTime int64, speedID interface{}, status int) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.UserTunnel{}).
Where("id = ?", id).
Updates(map[string]interface{}{
"flow": flow,
"num": num,
"exp_time": expTime,
"flow_reset_time": flowResetTime,
"speed_id": nullInt64FromInterface(speedID),
"status": status,
}).Error
}
func (r *Repository) GetUserTunnelUserAndTunnel(id int64) (userID, tunnelID int64, err error) {
if r == nil || r.db == nil {
return 0, 0, errors.New("repository not initialized")
}
var ut model.UserTunnel
err = r.db.Select("user_id", "tunnel_id").Where("id = ?", id).First(&ut).Error
if err != nil {
return 0, 0, normalizeNotFoundErr(err)
}
return ut.UserID, ut.TunnelID, nil
}
func (r *Repository) GetExistingUserTunnel(userID, tunnelID int64) (id int64, flow, num, expTime, flowReset int64, speedID sql.NullInt64, status int, err error) {
if r == nil || r.db == nil {
return 0, 0, 0, 0, 0, sql.NullInt64{}, 0, errors.New("repository not initialized")
}
var ut model.UserTunnel
err = r.db.Select("id", "flow", "num", "exp_time", "flow_reset_time", "speed_id", "status").
Where("user_id = ? AND tunnel_id = ?", userID, tunnelID).
First(&ut).Error
if err != nil {
return 0, 0, 0, 0, 0, sql.NullInt64{}, 0, normalizeNotFoundErr(err)
}
return ut.ID, ut.Flow, int64(ut.Num), ut.ExpTime, ut.FlowResetTime, ut.SpeedID, ut.Status, nil
}
func (r *Repository) InsertUserTunnel(userID, tunnelID int64, speedID interface{}, num int, flow, flowResetTime, expTime int64, status int) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
ut := model.UserTunnel{
UserID: userID,
TunnelID: tunnelID,
SpeedID: nullInt64FromInterface(speedID),
Num: num,
Flow: flow,
InFlow: 0,
OutFlow: 0,
FlowResetTime: flowResetTime,
ExpTime: expTime,
Status: status,
}
return r.db.Create(&ut).Error
}
func (r *Repository) UpdateUserTunnelFields(id int64, speedID interface{}, flow int64, num int, expTime, flowResetTime int64, status int) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.UserTunnel{}).
Where("id = ?", id).
Updates(map[string]interface{}{
"speed_id": nullInt64FromInterface(speedID),
"flow": flow,
"num": num,
"exp_time": expTime,
"flow_reset_time": flowResetTime,
"status": status,
}).Error
}
func (r *Repository) GetMinForwardPort(forwardID int64) sql.NullInt64 {
if r == nil || r.db == nil {
return sql.NullInt64{}
}
var p sql.NullInt64
_ = r.db.Model(&model.ForwardPort{}).
Select("MIN(port)").
Where("forward_id = ?", forwardID).
Scan(&p).Error
return p
}
func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.Forward{}).
Where("id = ?", id).
Updates(map[string]interface{}{
"name": name,
"tunnel_id": tunnelID,
"remote_addr": remoteAddr,
"strategy": strategy,
"speed_id": nullInt64FromInterface(speedID),
"max_conn": maxConn,
"ip_max_conn": ipMaxConn,
"ip_speed_id": nullInt64FromInterface(ipSpeedID),
"proxy_protocol": proxyProtocol,
"updated_time": now,
}).Error
}
func (r *Repository) UpdateForwardOrder(forwardID int64, inx int, now int64) {
if r == nil || r.db == nil {
return
}
_ = r.db.Model(&model.Forward{}).
Where("id = ?", forwardID).
Updates(map[string]interface{}{"inx": inx, "updated_time": now}).Error
}
func (r *Repository) UpdateForwardTunnel(forwardID, tunnelID int64, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.Forward{}).
Where("id = ?", forwardID).
Updates(map[string]interface{}{"tunnel_id": tunnelID, "updated_time": now}).Error
}
func (r *Repository) DeleteForwardCascade(forwardID int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Transaction(func(tx *gorm.DB) error {
if err := tx.Where("forward_id = ?", forwardID).Delete(&model.ForwardPort{}).Error; err != nil {
return err
}
return tx.Where("id = ?", forwardID).Delete(&model.Forward{}).Error
})
}
func (r *Repository) ReplaceForwardPorts(forwardID int64, entries []struct {
NodeID int64
Port int
InIP string
}) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Transaction(func(tx *gorm.DB) error {
if err := tx.Where("forward_id = ?", forwardID).Delete(&model.ForwardPort{}).Error; err != nil {
return err
}
if len(entries) == 0 {
return nil
}
rows := make([]model.ForwardPort, 0, len(entries))
for _, e := range entries {
rows = append(rows, model.ForwardPort{
ForwardID: forwardID,
NodeID: e.NodeID,
Port: e.Port,
InIP: sql.NullString{String: e.InIP, Valid: e.InIP != ""},
})
}
return tx.Create(&rows).Error
})
}
func (r *Repository) UpdateForwardPortBindIP(forwardID, nodeID int64, port int, inIP string) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
if forwardID <= 0 || nodeID <= 0 || port <= 0 {
return nil
}
return r.db.Model(&model.ForwardPort{}).
Where("forward_id = ? AND node_id = ? AND port = ?", forwardID, nodeID, port).
Update("in_ip", sql.NullString{String: inIP, Valid: strings.TrimSpace(inIP) != ""}).Error
}
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int, now int64) {
if r == nil || r.db == nil {
return
}
_ = r.db.Model(&model.Forward{}).
Where("id = ?", id).
Updates(map[string]interface{}{
"user_id": userID,
"user_name": userName,
"name": name,
"tunnel_id": tunnelID,
"remote_addr": remoteAddr,
"strategy": strategy,
"status": status,
"speed_id": nullInt64FromInterface(speedID),
"max_conn": maxConn,
"ip_max_conn": ipMaxConn,
"ip_speed_id": nullInt64FromInterface(ipSpeedID),
"proxy_protocol": proxyProtocol,
"updated_time": now,
}).Error
}
func (r *Repository) GetUsedPortsOnNodeAsMap(nodeID int64) (map[int]bool, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
used := make(map[int]bool)
var forwardPorts []int
if err := r.db.Model(&model.ForwardPort{}).Where("node_id = ?", nodeID).Pluck("port", &forwardPorts).Error; err != nil {
return nil, err
}
for _, p := range forwardPorts {
used[p] = true
}
var chainPorts []int
if err := r.db.Model(&model.ChainTunnel{}).Where("node_id = ? AND port > 0", nodeID).Pluck("port", &chainPorts).Error; err != nil {
return nil, err
}
for _, p := range chainPorts {
used[p] = true
}
return used, nil
}
func (r *Repository) CreateSpeedLimit(name string, speed int, now int64, status int) (int64, error) {
if r == nil || r.db == nil {
return 0, errors.New("repository not initialized")
}
sl := model.SpeedLimit{
Name: name,
Speed: speed,
TunnelID: sql.NullInt64{Int64: 0, Valid: false},
TunnelName: sql.NullString{String: "", Valid: false},
CreatedTime: now,
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
Status: status,
}
if err := r.db.Create(&sl).Error; err != nil {
return 0, err
}
return sl.ID, nil
}
func (r *Repository) UpdateSpeedLimit(id int64, name string, speed int, status int, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
updates := map[string]interface{}{
"name": name,
"speed": speed,
"status": status,
"tunnel_id": nil,
"tunnel_name": nil,
"updated_time": sql.NullInt64{
Int64: now,
Valid: true,
},
}
return r.db.Model(&model.SpeedLimit{}).
Where("id = ?", id).
Updates(updates).Error
}
func (r *Repository) DeleteSpeedLimit(id int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Where("id = ?", id).Delete(&model.SpeedLimit{}).Error
}
func (r *Repository) GroupCreate(table, name string, status int, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
switch table {
case "tunnel_group":
return r.db.Create(&model.TunnelGroup{
Name: name,
CreatedTime: now,
UpdatedTime: now,
Status: status,
}).Error
case "user_group":
return r.db.Create(&model.UserGroup{
Name: name,
CreatedTime: now,
UpdatedTime: now,
Status: status,
}).Error
default:
return errors.New("invalid group table")
}
}
func (r *Repository) GroupUpdate(table string, id int64, name string, status int, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
updates := map[string]interface{}{"name": name, "status": status, "updated_time": now}
switch table {
case "tunnel_group":
return r.db.Model(&model.TunnelGroup{}).Where("id = ?", id).Updates(updates).Error
case "user_group":
return r.db.Model(&model.UserGroup{}).Where("id = ?", id).Updates(updates).Error
default:
return errors.New("invalid group table")
}
}
func (r *Repository) GroupDeleteCascade(table string, id int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Transaction(func(tx *gorm.DB) error {
switch table {
case "tunnel_group":
if err := tx.Where("tunnel_group_id = ?", id).Delete(&model.TunnelGroupTunnel{}).Error; err != nil {
return err
}
if err := tx.Where("tunnel_group_id = ?", id).Delete(&model.GroupPermission{}).Error; err != nil {
return err
}
if err := tx.Where("tunnel_group_id = ?", id).Delete(&model.GroupPermissionGrant{}).Error; err != nil {
return err
}
return tx.Where("id = ?", id).Delete(&model.TunnelGroup{}).Error
case "user_group":
if err := tx.Where("user_group_id = ?", id).Delete(&model.UserGroupUser{}).Error; err != nil {
return err
}
if err := tx.Where("user_group_id = ?", id).Delete(&model.GroupPermission{}).Error; err != nil {
return err
}
if err := tx.Where("user_group_id = ?", id).Delete(&model.GroupPermissionGrant{}).Error; err != nil {
return err
}
return tx.Where("id = ?", id).Delete(&model.UserGroup{}).Error
default:
return errors.New("invalid group table")
}
})
}
func (r *Repository) ListUserIDsByUserGroupTx(tx *gorm.DB, userGroupID int64) ([]int64, error) {
if tx == nil {
return nil, errors.New("database unavailable")
}
var ids []int64
err := tx.Model(&model.UserGroupUser{}).
Where("user_group_id = ?", userGroupID).
Pluck("user_id", &ids).Error
if err != nil {
return nil, err
}
if ids == nil {
ids = make([]int64, 0)
}
return ids, nil
}
func (r *Repository) ReplaceTunnelGroupMembersTx(tx *gorm.DB, groupID int64, tunnelIDs []int64, now int64) error {
if tx == nil {
return errors.New("database unavailable")
}
if err := tx.Where("tunnel_group_id = ?", groupID).Delete(&model.TunnelGroupTunnel{}).Error; err != nil {
return err
}
if len(tunnelIDs) == 0 {
return nil
}
rows := make([]model.TunnelGroupTunnel, 0, len(tunnelIDs))
for _, tunnelID := range tunnelIDs {
if tunnelID <= 0 {
continue
}
rows = append(rows, model.TunnelGroupTunnel{TunnelGroupID: groupID, TunnelID: tunnelID, CreatedTime: now})
}
if len(rows) == 0 {
return nil
}
return tx.Clauses(clause.OnConflict{DoNothing: true}).Create(&rows).Error
}
func (r *Repository) ReplaceUserGroupMembersTx(tx *gorm.DB, groupID int64, userIDs []int64, now int64) error {
if tx == nil {
return errors.New("database unavailable")
}
if err := tx.Where("user_group_id = ?", groupID).Delete(&model.UserGroupUser{}).Error; err != nil {
return err
}
if len(userIDs) == 0 {
return nil
}
rows := make([]model.UserGroupUser, 0, len(userIDs))
for _, userID := range userIDs {
if userID <= 0 {
continue
}
rows = append(rows, model.UserGroupUser{UserGroupID: groupID, UserID: userID, CreatedTime: now})
}
if len(rows) == 0 {
return nil
}
return tx.Clauses(clause.OnConflict{DoNothing: true}).Create(&rows).Error
}
func (r *Repository) GetGroupPermissionPairByIDTx(tx *gorm.DB, id int64) (userGroupID int64, tunnelGroupID int64, exists bool, err error) {
if tx == nil {
return 0, 0, false, errors.New("database unavailable")
}
var gp model.GroupPermission
err = tx.Select("user_group_id", "tunnel_group_id").Where("id = ?", id).First(&gp).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return 0, 0, false, nil
}
if err != nil {
return 0, 0, false, err
}
return gp.UserGroupID, gp.TunnelGroupID, true, nil
}
func (r *Repository) DeleteGroupPermissionByIDTx(tx *gorm.DB, id int64) error {
if tx == nil {
return errors.New("database unavailable")
}
return tx.Where("id = ?", id).Delete(&model.GroupPermission{}).Error
}
// RevokedUserTunnelPair holds the (userID, tunnelID) of a deleted user_tunnel row,
// so the handler layer can clean up associated forwarding rules.
type RevokedUserTunnelPair struct {
UserID int64
TunnelID int64
}
func (r *Repository) RevokeGroupGrantsForRemovedUsersTx(tx *gorm.DB, userGroupID int64, previousUserIDs, currentUserIDs []int64) ([]RevokedUserTunnelPair, error) {
if tx == nil {
return nil, errors.New("database unavailable")
}
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, nil
}
type grantRow struct {
UserTunnelID int64
CreatedByGroup int
}
var revoked []RevokedUserTunnelPair
for _, userID := range removedUserIDs {
var rows []grantRow
if err := tx.Model(&model.GroupPermissionGrant{}).
Select("group_permission_grant.user_tunnel_id, group_permission_grant.created_by_group").
Joins("JOIN user_tunnel ON user_tunnel.id = group_permission_grant.user_tunnel_id").
Where("group_permission_grant.user_group_id = ? AND user_tunnel.user_id = ?", userGroupID, userID).
Find(&rows).Error; err != nil {
return revoked, err
}
groupCreatedTunnelIDs := make(map[int64]struct{})
for _, row := range rows {
if row.CreatedByGroup == 1 && row.UserTunnelID > 0 {
groupCreatedTunnelIDs[row.UserTunnelID] = struct{}{}
}
}
userTunnelIDs := tx.Model(&model.UserTunnel{}).Select("id").Where("user_id = ?", userID)
if err := tx.Where("user_group_id = ? AND user_tunnel_id IN (?)", userGroupID, userTunnelIDs).
Delete(&model.GroupPermissionGrant{}).Error; err != nil {
return revoked, err
}
for userTunnelID := range groupCreatedTunnelIDs {
var remaining int64
if err := tx.Model(&model.GroupPermissionGrant{}).Where("user_tunnel_id = ?", userTunnelID).Count(&remaining).Error; err != nil {
return revoked, err
}
if remaining == 0 {
var ut model.UserTunnel
if lookupErr := tx.Select("user_id", "tunnel_id").Where("id = ?", userTunnelID).First(&ut).Error; lookupErr == nil {
revoked = append(revoked, RevokedUserTunnelPair{UserID: ut.UserID, TunnelID: ut.TunnelID})
}
if err := tx.Where("id = ?", userTunnelID).Delete(&model.UserTunnel{}).Error; err != nil {
return revoked, err
}
}
}
}
return revoked, nil
}
func (r *Repository) RevokeGroupPermissionPairTx(tx *gorm.DB, userGroupID, tunnelGroupID int64) ([]RevokedUserTunnelPair, error) {
if tx == nil {
return nil, errors.New("database unavailable")
}
type grantRow struct {
UserTunnelID int64
CreatedByGroup int
}
var rows []grantRow
if err := tx.Model(&model.GroupPermissionGrant{}).
Select("user_tunnel_id, created_by_group").
Where("user_group_id = ? AND tunnel_group_id = ?", userGroupID, tunnelGroupID).
Find(&rows).Error; err != nil {
return nil, err
}
groupCreatedTunnelIDs := make(map[int64]struct{})
for _, row := range rows {
if row.CreatedByGroup == 1 && row.UserTunnelID > 0 {
groupCreatedTunnelIDs[row.UserTunnelID] = struct{}{}
}
}
if err := tx.Where("user_group_id = ? AND tunnel_group_id = ?", userGroupID, tunnelGroupID).
Delete(&model.GroupPermissionGrant{}).Error; err != nil {
return nil, err
}
var revoked []RevokedUserTunnelPair
for userTunnelID := range groupCreatedTunnelIDs {
var remaining int64
if err := tx.Model(&model.GroupPermissionGrant{}).Where("user_tunnel_id = ?", userTunnelID).Count(&remaining).Error; err != nil {
return revoked, err
}
if remaining == 0 {
var ut model.UserTunnel
if lookupErr := tx.Select("user_id", "tunnel_id").Where("id = ?", userTunnelID).First(&ut).Error; lookupErr == nil {
revoked = append(revoked, RevokedUserTunnelPair{UserID: ut.UserID, TunnelID: ut.TunnelID})
}
if err := tx.Where("id = ?", userTunnelID).Delete(&model.UserTunnel{}).Error; err != nil {
return revoked, err
}
}
}
return revoked, nil
}
func (r *Repository) ReplaceFederationTunnelBindingsTx(tx *gorm.DB, tunnelID int64, bindings []FederationTunnelBinding) error {
if tx == nil {
return errors.New("database unavailable")
}
if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.FederationTunnelBinding{}).Error; err != nil {
return err
}
if len(bindings) == 0 {
return nil
}
rows := make([]model.FederationTunnelBinding, 0, len(bindings))
now := time.Now().UnixMilli()
for _, b := range bindings {
created := b.CreatedTime
if created <= 0 {
created = now
}
updated := b.UpdatedTime
if updated <= 0 {
updated = created
}
rows = append(rows, model.FederationTunnelBinding{
TunnelID: tunnelID,
NodeID: b.NodeID,
ChainType: b.ChainType,
HopInx: b.HopInx,
RemoteURL: b.RemoteURL,
ResourceKey: b.ResourceKey,
RemoteBindingID: b.RemoteBindingID,
AllocatedPort: b.AllocatedPort,
Status: b.Status,
CreatedTime: created,
UpdatedTime: updated,
})
}
return tx.Create(&rows).Error
}
func (r *Repository) InsertGroupPermission(userGroupID, tunnelGroupID int64, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
gp := model.GroupPermission{UserGroupID: userGroupID, TunnelGroupID: tunnelGroupID, CreatedTime: now}
return r.db.Clauses(clause.OnConflict{DoNothing: true}).Create(&gp).Error
}
func (r *Repository) InsertGroupPermissionGrant(userGroupID, tunnelGroupID, userTunnelID int64, createdByGroup int, now int64) {
if r == nil || r.db == nil {
return
}
g := model.GroupPermissionGrant{
UserGroupID: userGroupID,
TunnelGroupID: tunnelGroupID,
UserTunnelID: userTunnelID,
CreatedByGroup: createdByGroup,
CreatedTime: now,
}
_ = r.db.Clauses(clause.OnConflict{DoNothing: true}).Create(&g).Error
}
func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool, error) {
if r == nil || r.db == nil {
return 0, false, errors.New("repository not initialized")
}
var existing model.UserTunnel
err := r.db.Select("id").Where("user_id = ? AND tunnel_id = ?", userID, tunnelID).First(&existing).Error
if err == nil {
return existing.ID, false, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return 0, false, err
}
var user model.User
if err := r.db.Select("flow, num, exp_time, flow_reset_time").Where("id = ?", userID).First(&user).Error; err != nil {
return 0, false, err
}
flow := user.Flow
num := user.Num
expTime := user.ExpTime
flowReset := user.FlowResetTime
ut := model.UserTunnel{
UserID: userID,
TunnelID: tunnelID,
Num: num,
Flow: flow,
InFlow: 0,
OutFlow: 0,
FlowResetTime: flowReset,
ExpTime: expTime,
Status: 1,
}
if err := r.db.Create(&ut).Error; err != nil {
return 0, false, err
}
return ut.ID, true, nil
}
func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, inIp string, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int) (int64, error) {
if r == nil || r.db == nil {
return 0, errors.New("repository not initialized")
}
var forwardID int64
err := r.db.Transaction(func(tx *gorm.DB) error {
fwd := model.Forward{
UserID: userID,
UserName: userName,
Name: name,
TunnelID: tunnelID,
RemoteAddr: remoteAddr,
Strategy: strategy,
InFlow: 0,
OutFlow: 0,
CreatedTime: now,
UpdatedTime: now,
Status: 1,
Inx: inx,
MaxConn: maxConn,
SpeedID: nullInt64FromInterface(speedID),
IPMaxConn: ipMaxConn,
IPSpeedID: nullInt64FromInterface(ipSpeedID),
ProxyProtocol: proxyProtocol,
}
if err := tx.Create(&fwd).Error; err != nil {
return err
}
forwardID = fwd.ID
for _, nodeID := range entryNodeIDs {
fp := model.ForwardPort{
ForwardID: forwardID,
NodeID: nodeID,
Port: port,
InIP: sql.NullString{String: inIp, Valid: inIp != ""},
}
if err := tx.Create(&fp).Error; err != nil {
return err
}
}
return nil
})
return forwardID, err
}
func (r *Repository) BatchUpdateForwardStatus(ids []int64, status int) (int, int) {
if r == nil || r.db == nil {
return 0, len(ids)
}
s := 0
f := 0
now := time.Now().UnixMilli()
for _, id := range ids {
if err := r.db.Model(&model.Forward{}).Where("id = ?", id).Updates(map[string]interface{}{"status": status, "updated_time": now}).Error; err != nil {
f++
} else {
s++
}
}
return s, f
}
func (r *Repository) CreateTunnelTx(tx *gorm.DB, name string, trafficRatio float64, typeVal int, flow int64, now int64, status int, inIP interface{}, inx int, ipPreference string) (int64, error) {
inIPVal := nullStringFromInterface(inIP)
tunnel := model.Tunnel{
Name: name,
TrafficRatio: trafficRatio,
Type: typeVal,
Protocol: "tls",
Flow: flow,
CreatedTime: now,
UpdatedTime: now,
Status: status,
InIP: inIPVal,
Inx: inx,
IPPreference: ipPreference,
}
if err := tx.Create(&tunnel).Error; err != nil {
return 0, err
}
return tunnel.ID, nil
}
func normalizeNotFoundErr(err error) error {
if errors.Is(err, gorm.ErrRecordNotFound) {
return sql.ErrNoRows
}
return err
}
func nullStringFromInterface(v interface{}) sql.NullString {
switch t := v.(type) {
case nil:
return sql.NullString{}
case sql.NullString:
return t
case *sql.NullString:
if t == nil {
return sql.NullString{}
}
return *t
case string:
if t == "" {
return sql.NullString{}
}
return sql.NullString{String: t, Valid: true}
case *string:
if t == nil || *t == "" {
return sql.NullString{}
}
return sql.NullString{String: *t, Valid: true}
default:
return sql.NullString{}
}
}
func nullInt64FromInterface(v interface{}) sql.NullInt64 {
switch t := v.(type) {
case nil:
return sql.NullInt64{}
case sql.NullInt64:
return t
case *sql.NullInt64:
if t == nil {
return sql.NullInt64{}
}
return *t
case int64:
return sql.NullInt64{Int64: t, Valid: true}
case *int64:
if t == nil {
return sql.NullInt64{}
}
return sql.NullInt64{Int64: *t, Valid: true}
case int:
return sql.NullInt64{Int64: int64(t), Valid: true}
case int32:
return sql.NullInt64{Int64: int64(t), Valid: true}
case int16:
return sql.NullInt64{Int64: int64(t), Valid: true}
case int8:
return sql.NullInt64{Int64: int64(t), Valid: true}
case uint64:
return sql.NullInt64{Int64: int64(t), Valid: true}
case uint:
return sql.NullInt64{Int64: int64(t), Valid: true}
case uint32:
return sql.NullInt64{Int64: int64(t), Valid: true}
case uint16:
return sql.NullInt64{Int64: int64(t), Valid: true}
case uint8:
return sql.NullInt64{Int64: int64(t), Valid: true}
case float64:
return sql.NullInt64{Int64: int64(t), Valid: true}
default:
return sql.NullInt64{}
}
}
func stringFromInterface(v interface{}) string {
switch t := v.(type) {
case nil:
return ""
case string:
return t
case sql.NullString:
if t.Valid {
return t.String
}
return ""
case *sql.NullString:
if t != nil && t.Valid {
return t.String
}
return ""
case []byte:
return string(t)
default:
return ""
}
}
func parsePortRangeSpec(input string) []int {
input = strings.TrimSpace(input)
if input == "" {
return nil
}
set := make(map[int]struct{})
parts := strings.Split(input, ",")
for _, part := range parts {
part = strings.TrimSpace(part)
if part == "" {
continue
}
if strings.Contains(part, "-") {
r := strings.SplitN(part, "-", 2)
if len(r) != 2 {
continue
}
start, err1 := strconv.Atoi(strings.TrimSpace(r[0]))
end, err2 := strconv.Atoi(strings.TrimSpace(r[1]))
if err1 != nil || err2 != nil || start <= 0 || end <= 0 {
continue
}
if end < start {
start, end = end, start
}
for p := start; p <= end; p++ {
set[p] = struct{}{}
}
continue
}
p, err := strconv.Atoi(part)
if err != nil || p <= 0 {
continue
}
set[p] = struct{}{}
}
out := make([]int, 0, len(set))
for p := range set {
out = append(out, p)
}
sort.Ints(out)
return out
}
func (r *Repository) AddUserToGroups(userID int64, groupIDs []int64, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
if len(groupIDs) == 0 {
return nil
}
rows := make([]model.UserGroupUser, 0, len(groupIDs))
for _, gid := range groupIDs {
if gid <= 0 {
continue
}
rows = append(rows, model.UserGroupUser{UserGroupID: gid, UserID: userID, CreatedTime: now})
}
if len(rows) == 0 {
return nil
}
return r.db.Clauses(clause.OnConflict{DoNothing: true}).Create(&rows).Error
}
func (r *Repository) ReplaceUserGroupsByUserID(userID int64, newGroupIDs []int64, now int64) (affectedGroupIDs []int64, err error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var oldGroupIDs []int64
if err = r.db.Model(&model.UserGroupUser{}).
Where("user_id = ?", userID).
Pluck("user_group_id", &oldGroupIDs).Error; err != nil {
return nil, err
}
seen := make(map[int64]struct{})
for _, id := range oldGroupIDs {
seen[id] = struct{}{}
}
for _, id := range newGroupIDs {
if id > 0 {
seen[id] = struct{}{}
}
}
for id := range seen {
affectedGroupIDs = append(affectedGroupIDs, id)
}
if err = r.db.Where("user_id = ?", userID).Delete(&model.UserGroupUser{}).Error; err != nil {
return nil, err
}
if len(newGroupIDs) == 0 {
return affectedGroupIDs, nil
}
rows := make([]model.UserGroupUser, 0, len(newGroupIDs))
for _, gid := range newGroupIDs {
if gid <= 0 {
continue
}
rows = append(rows, model.UserGroupUser{UserGroupID: gid, UserID: userID, CreatedTime: now})
}
if len(rows) > 0 {
if err = r.db.Clauses(clause.OnConflict{DoNothing: true}).Create(&rows).Error; err != nil {
return nil, err
}
}
return affectedGroupIDs, nil
}
func (r *Repository) AdvanceNodeRenewalCycles(now int64) (int, error) {
if r == nil || r.db == nil {
return 0, nil
}
var nodes []model.Node
if err := r.db.Where("renewal_cycle IS NOT NULL AND renewal_cycle != '' AND expiry_time IS NOT NULL").Find(&nodes).Error; err != nil {
return 0, fmt.Errorf("list nodes with renewal cycle: %w", err)
}
advanced := 0
for _, node := range nodes {
if !node.ExpiryTime.Valid || node.ExpiryTime.Int64 <= 0 {
continue
}
cycleMonths := 0
switch node.RenewalCycle.String {
case "month":
cycleMonths = 1
case "quarter":
cycleMonths = 3
case "year":
cycleMonths = 12
default:
continue
}
anchorTime := node.ExpiryTime.Int64
for anchorTime <= now {
nextAnchor := advanceByMonths(anchorTime, cycleMonths)
if nextAnchor <= anchorTime {
break
}
anchorTime = nextAnchor
}
if anchorTime == node.ExpiryTime.Int64 {
continue
}
if err := r.db.Model(&model.Node{}).Where("id = ?", node.ID).Update("expiry_time", anchorTime).Error; err != nil {
continue
}
advanced++
}
return advanced, nil
}
func advanceByMonths(timestamp int64, months int) int64 {
t := time.Unix(timestamp/1000, 0)
next := t.AddDate(0, months, 0)
return next.UnixMilli()
}