Files
flvx/go-backend/internal/store/repo/repository.go
T
sagitchu 3e11549370 fix(backend): sync user tunnel status and relax forward speedId permission check
- Return actual user_tunnel.status in admin permission list instead of hardcoded 1
- Allow non-admin users to update forwards when keeping the same speedId selection
- Add contract tests for user tunnel status mapping and forward permission edge case

Entire-Checkpoint: deb90fb942ee
2026-03-08 00:56:13 +08:00

2978 lines
86 KiB
Go

package repo
import (
"database/sql"
"errors"
"fmt"
"log"
"os"
"path/filepath"
"sort"
"strings"
"time"
gsqlite "github.com/glebarez/sqlite"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/clause"
"gorm.io/gorm/logger"
"go-backend/internal/store/model"
)
// ─── Type aliases for backward compatibility ─────────────────────────
// Handlers still reference repo.User, repo.BackupData, etc.
type User = model.User
type ViteConfig = model.ViteConfig
type Announcement = model.Announcement
type UserTunnelDetail = model.UserTunnelDetail
type UserForwardDetail = model.UserForwardDetail
type StatisticsFlow = model.StatisticsFlow
type Node = model.Node
type PeerShare = model.PeerShare
type PeerShareRuntime = model.PeerShareRuntime
type FederationTunnelBinding = model.FederationTunnelBinding
type BackupData = model.BackupData
type UserBackup = model.UserBackup
type NodeBackup = model.NodeBackup
type TunnelBackup = model.TunnelBackup
type ChainTunnelBackup = model.ChainTunnelBackup
type ForwardBackup = model.ForwardBackup
type ForwardPortBackup = model.ForwardPortBackup
type UserTunnelBackup = model.UserTunnelBackup
type SpeedLimitBackup = model.SpeedLimitBackup
type TunnelGroupBackup = model.TunnelGroupBackup
type UserGroupBackup = model.UserGroupBackup
type PermissionBackup = model.PermissionBackup
type PermissionGrantBackup = model.PermissionGrantBackup
type ImportResult = model.ImportResult
// ─── Repository ──────────────────────────────────────────────────────
type Repository struct {
db *gorm.DB
}
func (r *Repository) DB() *gorm.DB {
if r == nil {
return nil
}
return r.db
}
// ─── Open / Close ────────────────────────────────────────────────────
func Open(path string) (*Repository, error) {
if err := ensureParentDir(path); err != nil {
return nil, err
}
dsn := "file:" + path +
"?_pragma=busy_timeout(5000)" +
"&_pragma=journal_mode(WAL)" +
"&_pragma=synchronous(NORMAL)"
db, err := gorm.Open(gsqlite.Open(dsn), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
if err != nil {
return nil, err
}
sqlDB, err := db.DB()
if err != nil {
return nil, err
}
sqlDB.SetMaxOpenConns(1)
if err := prepareSQLiteLegacyColumns(db); err != nil {
_ = sqlDB.Close()
return nil, fmt.Errorf("prepare sqlite legacy schema: %w", err)
}
if err := autoMigrateAll(db); err != nil {
_ = sqlDB.Close()
return nil, fmt.Errorf("auto migrate: %w", err)
}
seedData(db)
if err := migrateSchema(db); err != nil {
_ = sqlDB.Close()
return nil, err
}
return &Repository{db: db}, nil
}
func OpenPostgres(dsn string) (*Repository, error) {
if strings.TrimSpace(dsn) == "" {
return nil, fmt.Errorf("empty postgres dsn")
}
db, err := gorm.Open(postgres.Open(dsn), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
if err != nil {
return nil, err
}
sqlDB, err := db.DB()
if err != nil {
return nil, err
}
if err := sqlDB.Ping(); err != nil {
_ = sqlDB.Close()
return nil, err
}
if err := preparePostgresLegacySchema(db); err != nil {
_ = sqlDB.Close()
return nil, fmt.Errorf("prepare postgres legacy schema: %w", err)
}
if err := autoMigrateAll(db); err != nil {
_ = sqlDB.Close()
return nil, fmt.Errorf("auto migrate: %w", err)
}
seedData(db)
if err := migrateSchema(db); err != nil {
_ = sqlDB.Close()
return nil, err
}
return &Repository{db: db}, nil
}
func (r *Repository) Close() error {
if r == nil || r.db == nil {
return nil
}
sqlDB, err := r.db.DB()
if err != nil {
return err
}
return sqlDB.Close()
}
func autoMigrateAll(db *gorm.DB) error {
models := []interface{}{
&model.User{},
&model.Forward{},
&model.ForwardPort{},
&model.Node{},
&model.SpeedLimit{},
&model.StatisticsFlow{},
&model.Tunnel{},
&model.ChainTunnel{},
&model.UserTunnel{},
&model.TunnelGroup{},
&model.UserGroup{},
&model.TunnelGroupTunnel{},
&model.UserGroupUser{},
&model.GroupPermission{},
&model.GroupPermissionGrant{},
&model.ViteConfig{},
&model.PeerShare{},
&model.PeerShareRuntime{},
&model.FederationTunnelBinding{},
&model.Announcement{},
&model.SchemaVersion{},
}
if db.Dialector.Name() != "sqlite" {
return db.AutoMigrate(models...)
}
m := db.Migrator()
hasNode := m.HasTable(&model.Node{})
hasTunnel := m.HasTable(&model.Tunnel{})
for _, item := range models {
if hasNode {
if _, ok := item.(*model.Node); ok {
continue
}
}
if hasTunnel {
if _, ok := item.(*model.Tunnel); ok {
continue
}
}
if err := db.AutoMigrate(item); err != nil {
return err
}
}
return nil
}
// preparePostgresLegacySchema renames unique constraints that were created by
// the old schema.sql (which used inline UNIQUE column syntax) to the names
// expected by GORM's NamingStrategy ("uni_<table>_<column>"). Without this,
// GORM's AutoMigrate emits "DROP CONSTRAINT uni_..." against constraints that
// don't exist under that name, crashing startup on upgraded PostgreSQL installs.
func preparePostgresLegacySchema(db *gorm.DB) error {
if db == nil || db.Dialector.Name() != "postgres" {
return nil
}
type rename struct{ table, oldName, newName string }
renames := []rename{
{"vite_config", "vite_config_name_key", "uni_vite_config_name"},
{"peer_share", "peer_share_token_key", "uni_peer_share_token"},
{"peer_share_runtime", "peer_share_runtime_reservation_id_key", "uni_peer_share_runtime_reservation_id"},
{"peer_share_runtime", "peer_share_runtime_resource_key_key", "uni_peer_share_runtime_resource_key"},
{"federation_tunnel_binding", "federation_tunnel_binding_resource_key_key", "uni_federation_tunnel_binding_resource_key"},
}
for _, r := range renames {
var count int64
if err := db.Raw(
`SELECT COUNT(*) FROM information_schema.table_constraints
WHERE constraint_schema = current_schema()
AND table_name = ?
AND constraint_name = ?
AND constraint_type = 'UNIQUE'`,
r.table, r.oldName,
).Scan(&count).Error; err != nil {
return fmt.Errorf("check constraint %s.%s: %w", r.table, r.oldName, err)
}
if count == 0 {
continue
}
if err := db.Exec(
fmt.Sprintf(`ALTER TABLE %q RENAME CONSTRAINT %q TO %q`, r.table, r.oldName, r.newName),
).Error; err != nil {
return fmt.Errorf("rename constraint %s.%s→%s: %w", r.table, r.oldName, r.newName, err)
}
}
return nil
}
func prepareSQLiteLegacyColumns(db *gorm.DB) error {
if db == nil || db.Dialector.Name() != "sqlite" {
return nil
}
m := db.Migrator()
if m.HasTable(&model.Node{}) {
for _, field := range []string{"ServerIPV4", "ServerIPV6", "ExtraIPs", "TCPListenAddr", "UDPListenAddr", "Inx", "IsRemote", "RemoteURL", "RemoteToken", "RemoteConfig"} {
if m.HasColumn(&model.Node{}, field) {
continue
}
if err := m.AddColumn(&model.Node{}, field); err != nil {
return fmt.Errorf("add node.%s: %w", field, err)
}
}
}
if m.HasTable(&model.Tunnel{}) {
for _, field := range []string{"Inx", "IPPreference"} {
if m.HasColumn(&model.Tunnel{}, field) {
continue
}
if err := m.AddColumn(&model.Tunnel{}, field); err != nil {
return fmt.Errorf("add tunnel.%s: %w", field, err)
}
}
}
return nil
}
func seedData(db *gorm.DB) {
adminUser := model.User{
ID: 1, User: "admin_user", Pwd: "3c85cdebade1c51cf64ca9f3c09d182d",
RoleID: 0, ExpTime: 2727251700000, Flow: 99999, InFlow: 0, OutFlow: 0,
FlowResetTime: 1, Num: 99999, CreatedTime: 1748914865000,
UpdatedTime: sql.NullInt64{Int64: 1754011744252, Valid: true}, Status: 1,
}
db.Where("id = ?", 1).FirstOrCreate(&adminUser)
appNameConfig := model.ViteConfig{ID: 1, Name: "app_name", Value: "flux", Time: 1755147963000}
db.Where("id = ?", 1).FirstOrCreate(&appNameConfig)
}
// ─── User Queries ────────────────────────────────────────────────────
func (r *Repository) GetUserByUsername(username string) (*model.User, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var user model.User
err := r.db.Where(`"user" = ?`, username).First(&user).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &user, nil
}
func (r *Repository) GetUserByID(id int64) (*model.User, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var user model.User
err := r.db.Where("id = ?", id).First(&user).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &user, nil
}
func (r *Repository) UsernameExists(username string) (bool, error) {
var count int64
err := r.db.Model(&model.User{}).Where(`"user" = ?`, username).Count(&count).Error
if err != nil {
return false, err
}
return count > 0, nil
}
func (r *Repository) UsernameExistsExceptID(username string, exceptID int64) (bool, error) {
if r == nil || r.db == nil {
return false, errors.New("repository not initialized")
}
var count int64
err := r.db.Model(&model.User{}).Where(`"user" = ? AND id != ?`, username, exceptID).Count(&count).Error
if err != nil {
return false, err
}
return count > 0, nil
}
func (r *Repository) UpdateUserNameAndPassword(userID int64, username, passwordMD5 string, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.User{}).Where("id = ?", userID).Updates(map[string]interface{}{
"user": username,
"pwd": passwordMD5,
"updated_time": now,
}).Error
}
// ─── Config Queries ──────────────────────────────────────────────────
func (r *Repository) GetConfigByName(name string) (*model.ViteConfig, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var cfg model.ViteConfig
err := r.db.Where("name = ?", name).First(&cfg).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &cfg, nil
}
func (r *Repository) ListConfigs() (map[string]string, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var configs []model.ViteConfig
if err := r.db.Find(&configs).Error; err != nil {
return nil, err
}
result := make(map[string]string)
for _, c := range configs {
result[c.Name] = c.Value
}
return result, nil
}
func (r *Repository) UpsertConfig(name, value string, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "name"}},
DoUpdates: clause.AssignmentColumns([]string{"value", "time"}),
}).Create(&model.ViteConfig{Name: name, Value: value, Time: now}).Error
}
// ─── Announcement Queries ────────────────────────────────────────────
func (r *Repository) GetAnnouncement() (*model.Announcement, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var ann model.Announcement
err := r.db.Order("id DESC").First(&ann).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &ann, nil
}
func (r *Repository) UpsertAnnouncement(content string, enabled int, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
var count int64
if err := r.db.Model(&model.Announcement{}).Count(&count).Error; err != nil {
return err
}
if count == 0 {
return r.db.Create(&model.Announcement{
Content: content, Enabled: enabled,
CreatedTime: now, UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
}).Error
}
return r.db.Model(&model.Announcement{}).Where("1=1").Updates(map[string]interface{}{
"content": content, "enabled": enabled, "updated_time": now,
}).Error
}
// ─── User Package Queries ────────────────────────────────────────────
func (r *Repository) GetUserPackageTunnels(userID int64) ([]model.UserTunnelDetail, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var items []model.UserTunnelDetail
err := r.db.Model(&model.UserTunnel{}).
Select("user_tunnel.id, user_tunnel.user_id, user_tunnel.tunnel_id, tunnel.name AS tunnel_name, user_tunnel.status, tunnel.flow AS tunnel_flow, user_tunnel.flow, user_tunnel.in_flow, user_tunnel.out_flow, user_tunnel.num, user_tunnel.flow_reset_time, user_tunnel.exp_time, user_tunnel.speed_id, speed_limit.name AS speed_limit, speed_limit.speed").
Joins("LEFT JOIN tunnel ON tunnel.id = user_tunnel.tunnel_id").
Joins("LEFT JOIN speed_limit ON speed_limit.id = user_tunnel.speed_id").
Where("user_tunnel.user_id = ?", userID).
Order("user_tunnel.id ASC").
Find(&items).Error
if err != nil {
return nil, err
}
if items == nil {
items = make([]model.UserTunnelDetail, 0)
}
return items, nil
}
func (r *Repository) GetUserPackageForwards(userID int64) ([]model.UserForwardDetail, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
type fwdRow struct {
ID int64
Name string
TunnelID int64
TunnelName string
RemoteAddr string
InFlow int64
OutFlow int64
Status int
CreatedAt int64
}
var rows []fwdRow
err := r.db.Model(&model.Forward{}).
Select("forward.id, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, forward.remote_addr, forward.in_flow, forward.out_flow, forward.status, forward.created_time AS created_at").
Joins("LEFT JOIN tunnel ON tunnel.id = forward.tunnel_id").
Where("forward.user_id = ?", userID).
Order("forward.id ASC").
Find(&rows).Error
if err != nil {
return nil, err
}
items := make([]model.UserForwardDetail, 0, len(rows))
for _, row := range rows {
inIP, inPort, err := resolveForwardIngress(r.db, row.ID, row.TunnelID)
if err != nil {
return nil, err
}
items = append(items, model.UserForwardDetail{
ID: row.ID, Name: row.Name, TunnelID: row.TunnelID,
TunnelName: row.TunnelName, InIP: inIP, InPort: inPort,
RemoteAddr: row.RemoteAddr, InFlow: row.InFlow, OutFlow: row.OutFlow,
Status: row.Status, CreatedAt: row.CreatedAt,
})
}
return items, nil
}
// ─── Statistics Queries ──────────────────────────────────────────────
func (r *Repository) GetStatisticsFlows(userID int64, limit int) ([]model.StatisticsFlow, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var items []model.StatisticsFlow
err := r.db.Where("user_id = ?", userID).Order("id DESC").Limit(limit).Find(&items).Error
if err != nil {
return nil, err
}
if items == nil {
items = make([]model.StatisticsFlow, 0)
}
return items, nil
}
// ─── Node Queries ────────────────────────────────────────────────────
func (r *Repository) NodeExistsBySecret(secret string) (bool, error) {
if r == nil || r.db == nil {
return false, errors.New("repository not initialized")
}
var count int64
err := r.db.Model(&model.Node{}).Where("secret = ?", secret).Count(&count).Error
if err != nil {
return false, err
}
return count > 0, nil
}
func (r *Repository) GetNodeBySecret(secret string) (*model.Node, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var n model.Node
err := r.db.Where("secret = ?", secret).First(&n).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &n, nil
}
func (r *Repository) GetNodeByID(id int64) (*model.Node, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var n model.Node
err := r.db.Where("id = ?", id).First(&n).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &n, nil
}
func (r *Repository) UpdateNodeOnline(nodeID int64, status int, version string, httpVal, tlsVal, socksVal int) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.Node{}).Where("id = ?", nodeID).Updates(map[string]interface{}{
"status": status, "version": version, "http": httpVal, "tls": tlsVal,
"socks": socksVal, "updated_time": unixMilliNow(),
}).Error
}
func (r *Repository) UpdateNodeStatus(nodeID int64, status int) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.Node{}).Where("id = ?", nodeID).Updates(map[string]interface{}{
"status": status, "updated_time": unixMilliNow(),
}).Error
}
// ─── Flow ────────────────────────────────────────────────────────────
func (r *Repository) AddFlow(forwardID, userID int64, userTunnelID int64, inFlow, outFlow 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.Model(&model.Forward{}).Where("id = ?", forwardID).
UpdateColumns(map[string]interface{}{
"in_flow": gorm.Expr("in_flow + ?", inFlow),
"out_flow": gorm.Expr("out_flow + ?", outFlow),
}).Error; err != nil {
return err
}
if err := tx.Model(&model.User{}).Where("id = ?", userID).
UpdateColumns(map[string]interface{}{
"in_flow": gorm.Expr("in_flow + ?", inFlow),
"out_flow": gorm.Expr("out_flow + ?", outFlow),
}).Error; err != nil {
return err
}
if userTunnelID > 0 {
if err := tx.Model(&model.UserTunnel{}).Where("id = ?", userTunnelID).
UpdateColumns(map[string]interface{}{
"in_flow": gorm.Expr("in_flow + ?", inFlow),
"out_flow": gorm.Expr("out_flow + ?", outFlow),
}).Error; err != nil {
return err
}
}
return nil
})
}
// ─── List Methods (return map[string]interface{}) ────────────────────
func (r *Repository) ListNodes() ([]map[string]interface{}, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var nodes []model.Node
if err := r.db.Order("inx ASC, id ASC").Find(&nodes).Error; err != nil {
return nil, err
}
items := make([]map[string]interface{}, 0, len(nodes))
for _, n := range nodes {
items = append(items, map[string]interface{}{
"id": n.ID, "inx": n.Inx, "name": n.Name,
"ip": n.ServerIP, "serverIp": n.ServerIP,
"serverIpV4": nullableString(n.ServerIPV4),
"serverIpV6": nullableString(n.ServerIPV6),
"extraIPs": nullableString(n.ExtraIPs),
"port": n.Port,
"tcpListenAddr": n.TCPListenAddr,
"udpListenAddr": n.UDPListenAddr,
"version": nullableString(n.Version),
"http": n.HTTP, "tls": n.TLS, "socks": n.Socks,
"status": n.Status, "isRemote": n.IsRemote,
"remoteUrl": nullableString(n.RemoteURL),
"remoteToken": nullableString(n.RemoteToken),
"remoteConfig": nullableString(n.RemoteConfig),
})
}
return items, nil
}
func (r *Repository) ListUsers() ([]map[string]interface{}, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var users []model.User
if err := r.db.Where("role_id != ?", 0).Order("id DESC").Find(&users).Error; err != nil {
return nil, err
}
items := make([]map[string]interface{}, 0, len(users))
for _, u := range users {
items = append(items, map[string]interface{}{
"id": u.ID, "user": u.User, "name": u.User,
"roleId": u.RoleID, "status": u.Status,
"flow": u.Flow, "num": u.Num, "expTime": u.ExpTime,
"flowResetTime": u.FlowResetTime, "createdTime": u.CreatedTime,
"updatedTime": nullableInt64(u.UpdatedTime),
"inFlow": u.InFlow, "outFlow": u.OutFlow,
})
}
return items, nil
}
func (r *Repository) ListSpeedLimits() ([]map[string]interface{}, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var limits []model.SpeedLimit
if err := r.db.Order("id DESC").Find(&limits).Error; err != nil {
return nil, err
}
items := make([]map[string]interface{}, 0, len(limits))
for _, sl := range limits {
item := map[string]interface{}{
"id": sl.ID, "name": sl.Name, "speed": sl.Speed,
"status": sl.Status, "createdTime": sl.CreatedTime,
"updatedTime": nullableInt64(sl.UpdatedTime),
}
items = append(items, item)
}
return items, nil
}
func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
type fwdRow struct {
ID int64
UserID int64
UserName string
Name string
TunnelID int64
TunnelName string
RemoteAddr string
Strategy string
InFlow int64
OutFlow int64
CreatedTime int64
Status int
Inx int
SpeedID sql.NullInt64
}
var rows []fwdRow
err := r.db.Model(&model.Forward{}).
Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx, forward.speed_id").
Joins("LEFT JOIN tunnel ON tunnel.id = forward.tunnel_id").
Order("forward.inx ASC, forward.id ASC").
Find(&rows).Error
if err != nil {
return nil, err
}
items := make([]map[string]interface{}, 0, len(rows))
for _, row := range rows {
inIP, inPort, err := resolveForwardIngress(r.db, row.ID, row.TunnelID)
if err != nil {
return nil, err
}
item := map[string]interface{}{
"id": row.ID, "userId": row.UserID, "userName": row.UserName,
"name": row.Name, "tunnelId": row.TunnelID, "tunnelName": row.TunnelName,
"inIp": nullableForwardIngress(inIP), "inPort": nullableInt64(inPort),
"remoteAddr": row.RemoteAddr, "strategy": row.Strategy,
"inFlow": row.InFlow, "outFlow": row.OutFlow,
"createdTime": row.CreatedTime, "status": row.Status, "inx": int64(row.Inx),
}
if row.SpeedID.Valid {
item["speedId"] = row.SpeedID.Int64
}
items = append(items, item)
}
return items, nil
}
func (r *Repository) ListUserAccessibleTunnels(userID int64) ([]map[string]interface{}, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
type row struct {
ID int64
Name string
}
var rows []row
err := r.db.Model(&model.UserTunnel{}).
Select("tunnel.id, tunnel.name").
Joins("JOIN tunnel ON tunnel.id = user_tunnel.tunnel_id").
Where("user_tunnel.user_id = ? AND tunnel.status = 1", userID).
Order("tunnel.inx ASC, tunnel.id ASC").
Find(&rows).Error
if err != nil {
return nil, err
}
tunnelIDs := make([]int64, 0, len(rows))
for _, rw := range rows {
tunnelIDs = append(tunnelIDs, rw.ID)
}
portRangeMap := r.getTunnelEntryPortRanges(tunnelIDs)
items := make([]map[string]interface{}, 0, len(rows))
for _, rw := range rows {
item := map[string]interface{}{"id": rw.ID, "name": rw.Name}
if pr, ok := portRangeMap[rw.ID]; ok {
item["portRangeMin"] = pr.min
item["portRangeMax"] = pr.max
}
items = append(items, item)
}
return items, nil
}
func (r *Repository) ListEnabledTunnelSummaries() ([]map[string]interface{}, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
type row struct {
ID int64
Name string
}
var rows []row
err := r.db.Model(&model.Tunnel{}).Select("id, name").Where("status = 1").Order("inx ASC, id ASC").Find(&rows).Error
if err != nil {
return nil, err
}
tunnelIDs := make([]int64, 0, len(rows))
for _, rw := range rows {
tunnelIDs = append(tunnelIDs, rw.ID)
}
portRangeMap := r.getTunnelEntryPortRanges(tunnelIDs)
items := make([]map[string]interface{}, 0, len(rows))
for _, rw := range rows {
item := map[string]interface{}{"id": rw.ID, "name": rw.Name}
if pr, ok := portRangeMap[rw.ID]; ok {
item["portRangeMin"] = pr.min
item["portRangeMax"] = pr.max
}
items = append(items, item)
}
return items, nil
}
type tunnelPortRange struct {
min int
max int
}
func (r *Repository) getTunnelEntryPortRanges(tunnelIDs []int64) map[int64]tunnelPortRange {
result := make(map[int64]tunnelPortRange)
if len(tunnelIDs) == 0 {
return result
}
type entryNode struct {
TunnelID int64
NodeID int64
}
var entries []entryNode
r.db.Model(&model.ChainTunnel{}).
Select("tunnel_id, node_id").
Where("tunnel_id IN (?) AND chain_type = ?", tunnelIDs, "1").
Find(&entries)
nodeIDs := make([]int64, 0, len(entries))
nodeSet := make(map[int64]struct{})
for _, e := range entries {
if _, exists := nodeSet[e.NodeID]; !exists {
nodeSet[e.NodeID] = struct{}{}
nodeIDs = append(nodeIDs, e.NodeID)
}
}
type nodePort struct {
ID int64
Port string
}
var nodePorts []nodePort
if len(nodeIDs) > 0 {
r.db.Model(&model.Node{}).Select("id, port").Where("id IN (?)", nodeIDs).Find(&nodePorts)
}
nodePortMap := make(map[int64]string)
for _, np := range nodePorts {
nodePortMap[np.ID] = np.Port
}
for _, e := range entries {
portSpec := nodePortMap[e.NodeID]
if portSpec == "" {
continue
}
minP, maxP := parsePortRangeMinMax(portSpec)
if minP <= 0 || maxP <= 0 {
continue
}
pr, exists := result[e.TunnelID]
if !exists {
result[e.TunnelID] = tunnelPortRange{min: minP, max: maxP}
} else {
if minP < pr.min {
pr.min = minP
}
if maxP > pr.max {
pr.max = maxP
}
result[e.TunnelID] = pr
}
}
return result
}
func parsePortRangeMinMax(input string) (int, int) {
input = strings.TrimSpace(input)
if input == "" {
return 0, 0
}
minPort, maxPort := 0, 0
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, end := parseIntPort(r[0]), parseIntPort(r[1])
if start <= 0 || end <= 0 {
continue
}
if end < start {
start, end = end, start
}
if minPort == 0 || start < minPort {
minPort = start
}
if maxPort == 0 || end > maxPort {
maxPort = end
}
continue
}
p := parseIntPort(part)
if p <= 0 {
continue
}
if minPort == 0 || p < minPort {
minPort = p
}
if maxPort == 0 || p > maxPort {
maxPort = p
}
}
return minPort, maxPort
}
func parseIntPort(s string) int {
var p int
fmt.Sscanf(strings.TrimSpace(s), "%d", &p)
return p
}
func (r *Repository) ListTunnels() ([]map[string]interface{}, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var tunnels []model.Tunnel
if err := r.db.Order("inx ASC, id ASC").Find(&tunnels).Error; err != nil {
return nil, err
}
tunnelMap := make(map[int64]map[string]interface{})
orderedIDs := make([]int64, 0, len(tunnels))
for _, t := range tunnels {
tunnelMap[t.ID] = map[string]interface{}{
"id": t.ID, "inx": t.Inx, "name": t.Name,
"type": t.Type, "flow": t.Flow, "trafficRatio": t.TrafficRatio,
"status": t.Status, "createdTime": t.CreatedTime,
"inIp": nullableString(t.InIP),
"ipPreference": t.IPPreference,
"inNodeId": make([]map[string]interface{}, 0),
"outNodeId": make([]map[string]interface{}, 0),
"chainNodes": make([][]map[string]interface{}, 0),
}
orderedIDs = append(orderedIDs, t.ID)
}
// Build node IP map
nodeIPMap := map[int64]string{}
var nodeList []model.Node
if err := r.db.Select("id, server_ip").Find(&nodeList).Error; err == nil {
for _, n := range nodeList {
nodeIPMap[n.ID] = n.ServerIP
}
}
// Load chain tunnels
var chains []model.ChainTunnel
if err := r.db.Order("tunnel_id ASC, chain_type ASC, inx ASC, id ASC").Find(&chains).Error; err != nil {
return nil, err
}
chainBucket := map[int64]map[int][]map[string]interface{}{}
inNodeIPs := map[int64][]string{}
for _, c := range chains {
t, ok := tunnelMap[c.TunnelID]
if !ok {
continue
}
chainTypeInt := 0
fmt.Sscanf(c.ChainType, "%d", &chainTypeInt)
inx := int64(0)
if c.Inx.Valid {
inx = c.Inx.Int64
}
nodeObj := map[string]interface{}{
"nodeId": c.NodeID,
"chainType": chainTypeInt,
"inx": inx,
}
if c.Protocol.Valid {
nodeObj["protocol"] = c.Protocol.String
}
if c.Strategy.Valid {
nodeObj["strategy"] = c.Strategy.String
}
if c.ConnectIP.Valid {
nodeObj["connectIp"] = c.ConnectIP.String
}
switch chainTypeInt {
case 1:
t["inNodeId"] = append(t["inNodeId"].([]map[string]interface{}), nodeObj)
if ip, ok := nodeIPMap[c.NodeID]; ok && ip != "" {
inNodeIPs[c.TunnelID] = append(inNodeIPs[c.TunnelID], ip)
}
case 2:
if _, ok := chainBucket[c.TunnelID]; !ok {
chainBucket[c.TunnelID] = map[int][]map[string]interface{}{}
}
chainBucket[c.TunnelID][int(inx)] = append(chainBucket[c.TunnelID][int(inx)], nodeObj)
case 3:
t["outNodeId"] = append(t["outNodeId"].([]map[string]interface{}), nodeObj)
}
}
for tunnelID, groups := range chainBucket {
t := tunnelMap[tunnelID]
if t == nil {
continue
}
keys := make([]int, 0, len(groups))
for k := range groups {
keys = append(keys, k)
}
sort.Ints(keys)
ordered := make([][]map[string]interface{}, 0, len(keys))
for _, k := range keys {
ordered = append(ordered, groups[k])
}
t["chainNodes"] = ordered
if s, ok := t["inIp"].(string); !ok || strings.TrimSpace(s) == "" {
if ips := inNodeIPs[tunnelID]; len(ips) > 0 {
t["inIp"] = strings.Join(ips, ",")
}
}
}
result := make([]map[string]interface{}, 0, len(orderedIDs))
for _, id := range orderedIDs {
if t, ok := tunnelMap[id]; ok {
result = append(result, t)
}
}
return result, nil
}
// ─── Group Queries ───────────────────────────────────────────────────
func (r *Repository) ListTunnelGroups() ([]map[string]interface{}, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var groups []model.TunnelGroup
if err := r.db.Order("id ASC").Find(&groups).Error; err != nil {
return nil, err
}
result := make([]map[string]interface{}, 0, len(groups))
for _, g := range groups {
ids, names, err := r.listTunnelGroupMembers(g.ID)
if err != nil {
return nil, err
}
result = append(result, map[string]interface{}{
"id": g.ID, "name": g.Name, "status": g.Status,
"tunnelIds": ids, "tunnelNames": names,
"createdTime": g.CreatedTime,
})
}
return result, nil
}
func (r *Repository) ListUserGroups() ([]map[string]interface{}, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var groups []model.UserGroup
if err := r.db.Order("id ASC").Find(&groups).Error; err != nil {
return nil, err
}
result := make([]map[string]interface{}, 0, len(groups))
for _, g := range groups {
ids, names, err := r.listUserGroupMembers(g.ID)
if err != nil {
return nil, err
}
result = append(result, map[string]interface{}{
"id": g.ID, "name": g.Name, "status": g.Status,
"userIds": ids, "userNames": names,
"createdTime": g.CreatedTime,
})
}
return result, nil
}
func (r *Repository) ListGroupPermissions() ([]map[string]interface{}, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
type permRow struct {
ID int64
UserGroupID int64
UserGroupName sql.NullString
TunnelGroupID int64
TunnelGroupName sql.NullString
CreatedTime int64
}
var rows []permRow
err := r.db.Model(&model.GroupPermission{}).
Select("group_permission.id, group_permission.user_group_id, user_group.name AS user_group_name, group_permission.tunnel_group_id, tunnel_group.name AS tunnel_group_name, group_permission.created_time").
Joins("LEFT JOIN user_group ON user_group.id = group_permission.user_group_id").
Joins("LEFT JOIN tunnel_group ON tunnel_group.id = group_permission.tunnel_group_id").
Order("group_permission.id ASC").
Find(&rows).Error
if err != nil {
return nil, err
}
result := make([]map[string]interface{}, 0, len(rows))
for _, r := range rows {
result = append(result, map[string]interface{}{
"id": r.ID, "userGroupId": r.UserGroupID,
"userGroupName": nullableString(r.UserGroupName),
"tunnelGroupId": r.TunnelGroupID,
"tunnelGroupName": nullableString(r.TunnelGroupName),
"createdTime": r.CreatedTime,
})
}
return result, nil
}
func (r *Repository) listTunnelGroupMembers(groupID int64) ([]int64, []string, error) {
type row struct {
ID int64
Name string
}
var rows []row
err := r.db.Model(&model.TunnelGroupTunnel{}).
Select("tunnel.id, tunnel.name").
Joins("JOIN tunnel ON tunnel.id = tunnel_group_tunnel.tunnel_id").
Where("tunnel_group_tunnel.tunnel_group_id = ?", groupID).
Order("tunnel.id ASC").
Find(&rows).Error
if err != nil {
return nil, nil, err
}
ids := make([]int64, 0, len(rows))
names := make([]string, 0, len(rows))
for _, r := range rows {
ids = append(ids, r.ID)
names = append(names, r.Name)
}
return ids, names, nil
}
func (r *Repository) listUserGroupMembers(groupID int64) ([]int64, []string, error) {
type row struct {
ID int64
Name string
}
var rows []row
err := r.db.Model(&model.UserGroupUser{}).
Select(`"user".id, "user"."user" AS name`).
Joins(`JOIN "user" ON "user".id = user_group_user.user_id`).
Where("user_group_user.user_group_id = ?", groupID).
Order(`"user".id ASC`).
Find(&rows).Error
if err != nil {
return nil, nil, err
}
ids := make([]int64, 0, len(rows))
names := make([]string, 0, len(rows))
for _, r := range rows {
ids = append(ids, r.ID)
names = append(names, r.Name)
}
return ids, names, nil
}
// ─── PeerShare CRUD ──────────────────────────────────────────────────
func (r *Repository) CreatePeerShare(share *model.PeerShare) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Create(share).Error
}
func (r *Repository) UpdatePeerShare(share *model.PeerShare) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.PeerShare{}).Where("id = ?", share.ID).Updates(map[string]interface{}{
"name": share.Name, "max_bandwidth": share.MaxBandwidth,
"expiry_time": share.ExpiryTime, "port_range_start": share.PortRangeStart,
"port_range_end": share.PortRangeEnd, "is_active": share.IsActive,
"updated_time": share.UpdatedTime, "allowed_domains": share.AllowedDomains,
"allowed_ips": share.AllowedIPs,
}).Error
}
func (r *Repository) DeletePeerShare(id int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Transaction(func(tx *gorm.DB) error {
tx.Where("share_id = ?", id).Delete(&model.PeerShareRuntime{})
return tx.Where("id = ?", id).Delete(&model.PeerShare{}).Error
})
}
func (r *Repository) GetPeerShare(id int64) (*model.PeerShare, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var s model.PeerShare
err := r.db.Where("id = ?", id).First(&s).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &s, nil
}
func (r *Repository) GetPeerShareByToken(token string) (*model.PeerShare, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var s model.PeerShare
err := r.db.Where("token = ?", token).First(&s).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &s, nil
}
func (r *Repository) ListPeerShares() ([]model.PeerShare, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var shares []model.PeerShare
err := r.db.Order("id DESC").Find(&shares).Error
return shares, err
}
func (r *Repository) AddPeerShareCurrentFlow(shareID int64, delta int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
if shareID <= 0 || delta <= 0 {
return nil
}
return r.db.Model(&model.PeerShare{}).Where("id = ?", shareID).
UpdateColumns(map[string]interface{}{
"current_flow": gorm.Expr("current_flow + ?", delta),
"updated_time": unixMilliNow(),
}).Error
}
func (r *Repository) ResetPeerShareCurrentFlow(shareID int64, updatedTime int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
if shareID <= 0 {
return nil
}
if updatedTime <= 0 {
updatedTime = unixMilliNow()
}
return r.db.Model(&model.PeerShare{}).Where("id = ?", shareID).Updates(map[string]interface{}{
"current_flow": 0, "updated_time": updatedTime,
}).Error
}
// ─── PeerShareRuntime CRUD ───────────────────────────────────────────
func (r *Repository) CreatePeerShareRuntime(item *model.PeerShareRuntime) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
if item == nil {
return errors.New("runtime item is nil")
}
return r.db.Create(item).Error
}
func (r *Repository) UpdatePeerShareRuntime(item *model.PeerShareRuntime) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
if item == nil {
return errors.New("runtime item is nil")
}
return r.db.Model(&model.PeerShareRuntime{}).Where("id = ?", item.ID).Updates(map[string]interface{}{
"binding_id": item.BindingID, "role": item.Role,
"chain_name": item.ChainName, "service_name": item.ServiceName,
"protocol": item.Protocol, "strategy": item.Strategy,
"port": item.Port, "target": item.Target,
"applied": item.Applied, "status": item.Status,
"updated_time": item.UpdatedTime,
}).Error
}
func (r *Repository) MarkPeerShareRuntimeReleased(id int64, updatedTime int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.PeerShareRuntime{}).Where("id = ?", id).Updates(map[string]interface{}{
"status": 0, "updated_time": updatedTime,
}).Error
}
func (r *Repository) GetPeerShareRuntimeByResourceKey(shareID int64, resourceKey string) (*model.PeerShareRuntime, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var item model.PeerShareRuntime
err := r.db.Where("share_id = ? AND resource_key = ?", shareID, resourceKey).First(&item).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &item, nil
}
func (r *Repository) GetPeerShareRuntimeByReservationID(shareID int64, reservationID string) (*model.PeerShareRuntime, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var item model.PeerShareRuntime
err := r.db.Where("share_id = ? AND reservation_id = ?", shareID, reservationID).First(&item).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &item, nil
}
func (r *Repository) GetPeerShareRuntimeByBindingID(shareID int64, bindingID string) (*model.PeerShareRuntime, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var item model.PeerShareRuntime
err := r.db.Where("share_id = ? AND binding_id = ?", shareID, bindingID).First(&item).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &item, nil
}
func (r *Repository) GetPeerShareRuntimeByID(id int64) (*model.PeerShareRuntime, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var item model.PeerShareRuntime
err := r.db.Where("id = ?", id).First(&item).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &item, nil
}
func (r *Repository) ListActivePeerShareRuntimesByShareID(shareID int64) ([]model.PeerShareRuntime, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var out []model.PeerShareRuntime
err := r.db.Where("share_id = ? AND status = 1", shareID).Order("port ASC, id ASC").Find(&out).Error
if err != nil {
return nil, err
}
if out == nil {
out = make([]model.PeerShareRuntime, 0)
}
return out, nil
}
func (r *Repository) ListActivePeerShareRuntimePorts(shareID int64, nodeID int64) ([]int, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var ports []int
err := r.db.Model(&model.PeerShareRuntime{}).
Where("share_id = ? AND node_id = ? AND status = 1 AND port > 0", shareID, nodeID).
Pluck("port", &ports).Error
if err != nil {
return nil, err
}
if ports == nil {
ports = make([]int, 0)
}
return ports, nil
}
func (r *Repository) ListActiveForwardPeerShareRuntimesByServiceName(serviceName string) ([]model.PeerShareRuntime, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var items []model.PeerShareRuntime
err := r.db.Where("service_name = ? AND status = 1 AND role = ?", serviceName, "forward").
Order("id ASC").
Find(&items).Error
if err != nil {
return nil, err
}
if items == nil {
items = make([]model.PeerShareRuntime, 0)
}
return items, nil
}
func (r *Repository) ListActiveForwardPeerShareRuntimesByNodeAndServiceName(nodeID int64, serviceName string) ([]model.PeerShareRuntime, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
serviceName = strings.TrimSpace(serviceName)
if serviceName == "" {
return []model.PeerShareRuntime{}, nil
}
var items []model.PeerShareRuntime
err := r.db.Where("node_id = ? AND service_name = ? AND status = 1 AND role = ?", nodeID, serviceName, "forward").
Order("id ASC").
Find(&items).Error
if err != nil {
return nil, err
}
if items == nil {
items = make([]model.PeerShareRuntime, 0)
}
return items, nil
}
func (r *Repository) ListActiveForwardPeerShareRuntimeServiceNamesByNode(nodeID int64) ([]string, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var names []string
err := r.db.Model(&model.PeerShareRuntime{}).
Where("node_id = ? AND status = 1 AND role = ? AND service_name <> ''", nodeID, "forward").
Pluck("service_name", &names).Error
if err != nil {
return nil, err
}
if names == nil {
names = make([]string, 0)
}
return names, nil
}
func (r *Repository) HasRecentUnboundForwardPeerShareRuntimeOnNode(nodeID int64, minUpdatedTime int64) (bool, error) {
if r == nil || r.db == nil {
return false, errors.New("repository not initialized")
}
var count int64
err := r.db.Model(&model.PeerShareRuntime{}).
Where("node_id = ? AND status = 1 AND role = ? AND applied = 0 AND updated_time >= ? AND (service_name = '' OR service_name IS NULL)", nodeID, "forward", minUpdatedTime).
Count(&count).Error
if err != nil {
return false, err
}
return count > 0, nil
}
func (r *Repository) GetActiveForwardPeerShareRuntimeByPort(shareID int64, port int) (*model.PeerShareRuntime, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var item model.PeerShareRuntime
err := r.db.Where("share_id = ? AND port = ? AND status = 1 AND role = ?", shareID, port, "forward").First(&item).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &item, nil
}
func (r *Repository) GetActiveForwardPeerShareRuntimeByServiceName(shareID int64, serviceName string) (*model.PeerShareRuntime, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
serviceName = strings.TrimSpace(serviceName)
if shareID <= 0 || serviceName == "" {
return nil, nil
}
var item model.PeerShareRuntime
err := r.db.Where("share_id = ? AND service_name = ? AND status = 1 AND role = ?", shareID, serviceName, "forward").
Order("id ASC").
First(&item).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &item, nil
}
func (r *Repository) ExistsActivePeerShareRuntimeOnNodePort(nodeID int64, port int) (bool, error) {
if r == nil || r.db == nil {
return false, errors.New("repository not initialized")
}
var count int64
err := r.db.Model(&model.PeerShareRuntime{}).
Where("node_id = ? AND port = ? AND status = 1", nodeID, port).
Count(&count).Error
if err != nil {
return false, err
}
return count > 0, nil
}
func (r *Repository) UpdatePeerShareRuntimeServiceName(id int64, serviceName string, updatedTime int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.PeerShareRuntime{}).Where("id = ?", id).Updates(map[string]interface{}{
"service_name": serviceName,
"applied": 1,
"updated_time": updatedTime,
}).Error
}
func (r *Repository) MarkPeerShareRuntimeReleasedByPort(shareID int64, port int, updatedTime int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
if shareID <= 0 || port <= 0 {
return nil
}
if updatedTime <= 0 {
updatedTime = unixMilliNow()
}
return r.db.Model(&model.PeerShareRuntime{}).Where("share_id = ? AND port = ? AND status = 1", shareID, port).Updates(map[string]interface{}{
"status": 0,
"applied": 0,
"service_name": "",
"updated_time": updatedTime,
}).Error
}
func (r *Repository) MarkForwardPeerShareRuntimeReleasedByServiceName(shareID int64, serviceName string, updatedTime int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
serviceName = strings.TrimSpace(serviceName)
if shareID <= 0 || serviceName == "" {
return nil
}
if updatedTime <= 0 {
updatedTime = unixMilliNow()
}
return r.db.Model(&model.PeerShareRuntime{}).
Where("share_id = ? AND status = 1 AND role = ? AND service_name = ?", shareID, "forward", serviceName).
Updates(map[string]interface{}{
"status": 0,
"applied": 0,
"service_name": "",
"updated_time": updatedTime,
}).Error
}
// ─── FederationTunnelBinding ─────────────────────────────────────────
func (r *Repository) UpsertFederationTunnelBinding(item *model.FederationTunnelBinding) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
if item == nil {
return errors.New("binding item is nil")
}
return r.db.Clauses(clause.OnConflict{
Columns: []clause.Column{
{Name: "tunnel_id"}, {Name: "node_id"}, {Name: "chain_type"}, {Name: "hop_inx"},
},
DoUpdates: clause.AssignmentColumns([]string{
"remote_url", "resource_key", "remote_binding_id",
"allocated_port", "status", "updated_time",
}),
}).Create(item).Error
}
func (r *Repository) ListActiveFederationTunnelBindingsByTunnel(tunnelID int64) ([]model.FederationTunnelBinding, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var out []model.FederationTunnelBinding
err := r.db.Where("tunnel_id = ? AND status = 1", tunnelID).
Order("chain_type ASC, hop_inx ASC, id ASC").Find(&out).Error
if err != nil {
return nil, err
}
if out == nil {
out = make([]model.FederationTunnelBinding, 0)
}
return out, nil
}
func (r *Repository) DeleteFederationTunnelBindingsByTunnel(tunnelID int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Where("tunnel_id = ?", tunnelID).Delete(&model.FederationTunnelBinding{}).Error
}
// ─── Export Methods ──────────────────────────────────────────────────
func (r *Repository) ExportAll() (*model.BackupData, error) {
backup := &model.BackupData{Version: "1.0", ExportedAt: unixMilliNow()}
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
}
func (r *Repository) ExportPartial(types []string) (*model.BackupData, error) {
backup := &model.BackupData{Version: "1.0", ExportedAt: unixMilliNow()}
typeSet := make(map[string]bool)
for _, t := range types {
typeSet[t] = true
}
if typeSet["users"] {
v, err := r.exportUsers()
if err != nil {
return nil, fmt.Errorf("export users failed: %w", err)
}
backup.Users = v
}
if typeSet["nodes"] {
v, err := r.exportNodes()
if err != nil {
return nil, fmt.Errorf("export nodes failed: %w", err)
}
backup.Nodes = v
}
if typeSet["tunnels"] {
v, err := r.exportTunnels()
if err != nil {
return nil, fmt.Errorf("export tunnels failed: %w", err)
}
backup.Tunnels = v
}
if typeSet["forwards"] {
v, err := r.exportForwards()
if err != nil {
return nil, fmt.Errorf("export forwards failed: %w", err)
}
backup.Forwards = v
}
if typeSet["userTunnels"] {
v, err := r.exportUserTunnels()
if err != nil {
return nil, fmt.Errorf("export user tunnels failed: %w", err)
}
backup.UserTunnels = v
}
if typeSet["speedLimits"] {
v, err := r.exportSpeedLimits()
if err != nil {
return nil, fmt.Errorf("export speed limits failed: %w", err)
}
backup.SpeedLimits = v
}
if typeSet["tunnelGroups"] {
v, err := r.exportTunnelGroups()
if err != nil {
return nil, fmt.Errorf("export tunnel groups failed: %w", err)
}
backup.TunnelGroups = v
}
if typeSet["userGroups"] {
v, err := r.exportUserGroups()
if err != nil {
return nil, fmt.Errorf("export user groups failed: %w", err)
}
backup.UserGroups = v
}
if typeSet["permissions"] {
v, err := r.exportPermissions()
if err != nil {
return nil, fmt.Errorf("export permissions failed: %w", err)
}
backup.Permissions = v
}
if typeSet["configs"] {
v, err := r.ListConfigs()
if err != nil {
return nil, fmt.Errorf("export configs failed: %w", err)
}
backup.Configs = v
}
return backup, nil
}
func (r *Repository) exportUsers() ([]model.UserBackup, error) {
var users []model.User
if err := r.db.Order("id ASC").Find(&users).Error; err != nil {
return nil, err
}
out := make([]model.UserBackup, 0, len(users))
for _, u := range users {
b := model.UserBackup{
ID: u.ID, User: u.User, Pwd: u.Pwd, RoleID: u.RoleID,
ExpTime: u.ExpTime, Flow: u.Flow, InFlow: u.InFlow, OutFlow: u.OutFlow,
FlowResetTime: u.FlowResetTime, Num: u.Num,
CreatedTime: u.CreatedTime, Status: u.Status,
}
if u.UpdatedTime.Valid {
b.UpdatedTime = u.UpdatedTime.Int64
}
out = append(out, b)
}
return out, nil
}
func (r *Repository) exportNodes() ([]model.NodeBackup, error) {
var nodes []model.Node
if err := r.db.Order("inx ASC, id ASC").Find(&nodes).Error; err != nil {
return nil, err
}
out := make([]model.NodeBackup, 0, len(nodes))
for _, n := range nodes {
b := model.NodeBackup{
ID: n.ID, Name: n.Name, Secret: n.Secret, ServerIP: n.ServerIP,
Port: n.Port, HTTP: n.HTTP, TLS: n.TLS, Socks: n.Socks,
CreatedTime: n.CreatedTime, Status: n.Status,
TCPListenAddr: n.TCPListenAddr, UDPListenAddr: n.UDPListenAddr,
Inx: n.Inx, IsRemote: n.IsRemote,
}
if n.UpdatedTime.Valid {
b.UpdatedTime = n.UpdatedTime.Int64
}
if n.ServerIPV4.Valid {
b.ServerIPv4 = n.ServerIPV4.String
}
if n.ServerIPV6.Valid {
b.ServerIPv6 = n.ServerIPV6.String
}
if n.InterfaceName.Valid {
b.InterfaceName = n.InterfaceName.String
}
if n.Version.Valid {
b.Version = n.Version.String
}
if n.RemoteURL.Valid {
b.RemoteURL = n.RemoteURL.String
}
if n.RemoteToken.Valid {
b.RemoteToken = n.RemoteToken.String
}
if n.RemoteConfig.Valid {
b.RemoteConfig = n.RemoteConfig.String
}
out = append(out, b)
}
return out, nil
}
func (r *Repository) exportTunnels() ([]model.TunnelBackup, error) {
var tunnels []model.Tunnel
if err := r.db.Order("inx ASC, id ASC").Find(&tunnels).Error; err != nil {
return nil, err
}
out := make([]model.TunnelBackup, 0, len(tunnels))
for _, t := range tunnels {
b := model.TunnelBackup{
ID: t.ID, Name: t.Name, TrafficRatio: t.TrafficRatio,
Type: t.Type, Protocol: t.Protocol, Flow: t.Flow,
CreatedTime: t.CreatedTime, UpdatedTime: t.UpdatedTime,
Status: t.Status, Inx: t.Inx, IPPreference: t.IPPreference,
}
if t.InIP.Valid {
b.InIP = t.InIP.String
}
chains, err := r.exportChainTunnels(t.ID)
if err != nil {
return nil, err
}
b.ChainTunnels = chains
out = append(out, b)
}
return out, nil
}
func (r *Repository) exportChainTunnels(tunnelID int64) ([]model.ChainTunnelBackup, error) {
var chains []model.ChainTunnel
if err := r.db.Where("tunnel_id = ?", tunnelID).Order("inx ASC, id ASC").Find(&chains).Error; err != nil {
return nil, err
}
out := make([]model.ChainTunnelBackup, 0, len(chains))
for _, c := range chains {
b := model.ChainTunnelBackup{
ID: c.ID, TunnelID: c.TunnelID, ChainType: c.ChainType, NodeID: c.NodeID,
}
if c.Port.Valid {
b.Port = int(c.Port.Int64)
}
if c.Strategy.Valid {
b.Strategy = c.Strategy.String
}
if c.Inx.Valid {
b.Inx = int(c.Inx.Int64)
}
if c.Protocol.Valid {
b.Protocol = c.Protocol.String
}
out = append(out, b)
}
return out, nil
}
func (r *Repository) exportForwards() ([]model.ForwardBackup, error) {
var forwards []model.Forward
if err := r.db.Order("id ASC").Find(&forwards).Error; err != nil {
return nil, err
}
out := make([]model.ForwardBackup, 0, len(forwards))
for _, f := range forwards {
b := model.ForwardBackup{
ID: f.ID, UserID: f.UserID, UserName: f.UserName, Name: f.Name,
TunnelID: f.TunnelID, RemoteAddr: f.RemoteAddr, Strategy: f.Strategy,
InFlow: f.InFlow, OutFlow: f.OutFlow, CreatedTime: f.CreatedTime,
UpdatedTime: f.UpdatedTime, Status: f.Status, Inx: f.Inx,
}
ports, err := r.exportForwardPorts(f.ID)
if err != nil {
return nil, err
}
portsCopy := append([]model.ForwardPortBackup(nil), ports...)
b.ForwardPorts = &portsCopy
out = append(out, b)
}
return out, nil
}
func (r *Repository) exportForwardPorts(forwardID int64) ([]model.ForwardPortBackup, error) {
var fps []model.ForwardPort
if err := r.db.Where("forward_id = ?", forwardID).Order("id ASC").Find(&fps).Error; err != nil {
return nil, err
}
out := make([]model.ForwardPortBackup, 0, len(fps))
for _, fp := range fps {
out = append(out, model.ForwardPortBackup{NodeID: fp.NodeID, Port: fp.Port})
}
return out, nil
}
func (r *Repository) exportUserTunnels() ([]model.UserTunnelBackup, error) {
var uts []model.UserTunnel
if err := r.db.Order("id ASC").Find(&uts).Error; err != nil {
return nil, err
}
out := make([]model.UserTunnelBackup, 0, len(uts))
for _, ut := range uts {
b := model.UserTunnelBackup{
ID: ut.ID, UserID: ut.UserID, TunnelID: ut.TunnelID,
Num: ut.Num, Flow: ut.Flow, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
FlowResetTime: ut.FlowResetTime, ExpTime: ut.ExpTime, Status: ut.Status,
}
if ut.SpeedID.Valid {
b.SpeedID = ut.SpeedID.Int64
}
out = append(out, b)
}
return out, nil
}
func (r *Repository) exportSpeedLimits() ([]model.SpeedLimitBackup, error) {
var sls []model.SpeedLimit
if err := r.db.Order("id ASC").Find(&sls).Error; err != nil {
return nil, err
}
out := make([]model.SpeedLimitBackup, 0, len(sls))
for _, sl := range sls {
b := model.SpeedLimitBackup{
ID: sl.ID, Name: sl.Name, Speed: int64(sl.Speed),
CreatedTime: sl.CreatedTime, Status: sl.Status,
}
if sl.UpdatedTime.Valid {
b.UpdatedTime = sl.UpdatedTime.Int64
}
out = append(out, b)
}
return out, nil
}
func (r *Repository) exportTunnelGroups() ([]model.TunnelGroupBackup, error) {
var groups []model.TunnelGroup
if err := r.db.Order("id ASC").Find(&groups).Error; err != nil {
return nil, err
}
out := make([]model.TunnelGroupBackup, 0, len(groups))
for _, tg := range groups {
b := model.TunnelGroupBackup{
ID: tg.ID, Name: tg.Name, CreatedTime: tg.CreatedTime,
UpdatedTime: tg.UpdatedTime, Status: tg.Status,
}
var tunnelIDs []int64
r.db.Model(&model.TunnelGroupTunnel{}).Where("tunnel_group_id = ?", tg.ID).Pluck("tunnel_id", &tunnelIDs)
b.Tunnels = tunnelIDs
out = append(out, b)
}
return out, nil
}
func (r *Repository) exportUserGroups() ([]model.UserGroupBackup, error) {
var groups []model.UserGroup
if err := r.db.Order("id ASC").Find(&groups).Error; err != nil {
return nil, err
}
out := make([]model.UserGroupBackup, 0, len(groups))
for _, ug := range groups {
b := model.UserGroupBackup{
ID: ug.ID, Name: ug.Name, CreatedTime: ug.CreatedTime,
UpdatedTime: ug.UpdatedTime, Status: ug.Status,
}
var userIDs []int64
r.db.Model(&model.UserGroupUser{}).Where("user_group_id = ?", ug.ID).Pluck("user_id", &userIDs)
b.Users = userIDs
out = append(out, b)
}
return out, nil
}
func (r *Repository) exportPermissions() ([]model.PermissionBackup, error) {
var perms []model.GroupPermission
if err := r.db.Order("id ASC").Find(&perms).Error; err != nil {
return nil, err
}
out := make([]model.PermissionBackup, 0, len(perms))
for _, p := range perms {
b := model.PermissionBackup{
ID: p.ID, UserGroupID: p.UserGroupID, TunnelGroupID: p.TunnelGroupID,
CreatedTime: p.CreatedTime,
}
var grants []model.GroupPermissionGrant
r.db.Where("user_group_id = ? AND tunnel_group_id = ?", p.UserGroupID, p.TunnelGroupID).Find(&grants)
for _, g := range grants {
b.Grants = append(b.Grants, model.PermissionGrantBackup{
ID: g.ID, UserGroupID: g.UserGroupID, TunnelGroupID: g.TunnelGroupID,
UserTunnelID: g.UserTunnelID, CreatedTime: g.CreatedTime,
CreatedByGroup: g.CreatedByGroup,
})
}
out = append(out, b)
}
return out, nil
}
// ─── Import Methods ──────────────────────────────────────────────────
func (r *Repository) Import(backup *model.BackupData, types []string) (*model.ImportResult, error) {
result := &model.ImportResult{}
typeSet := make(map[string]bool)
for _, t := range types {
typeSet[t] = true
}
err := r.db.Transaction(func(tx *gorm.DB) error {
now := unixMilliNow()
if typeSet["users"] && len(backup.Users) > 0 {
count, err := importUsers(tx, backup.Users, now)
if err != nil {
return fmt.Errorf("import users failed: %w", err)
}
result.UsersImported = count
}
if typeSet["nodes"] && len(backup.Nodes) > 0 {
count, err := importNodes(tx, backup.Nodes, now)
if err != nil {
return fmt.Errorf("import nodes failed: %w", err)
}
result.NodesImported = count
}
if typeSet["tunnels"] && len(backup.Tunnels) > 0 {
count, err := importTunnels(tx, backup.Tunnels, now)
if err != nil {
return fmt.Errorf("import tunnels failed: %w", err)
}
result.TunnelsImported = count
}
if typeSet["forwards"] && len(backup.Forwards) > 0 {
count, err := importForwards(tx, backup.Forwards, now)
if err != nil {
return fmt.Errorf("import forwards failed: %w", err)
}
result.ForwardsImported = count
}
if typeSet["userTunnels"] && len(backup.UserTunnels) > 0 {
count, err := importUserTunnels(tx, backup.UserTunnels, now)
if err != nil {
return fmt.Errorf("import user tunnels failed: %w", err)
}
result.UserTunnelsImported = count
}
if typeSet["speedLimits"] && len(backup.SpeedLimits) > 0 {
count, err := importSpeedLimits(tx, backup.SpeedLimits, now)
if err != nil {
return fmt.Errorf("import speed limits failed: %w", err)
}
result.SpeedLimitsImported = count
}
if typeSet["tunnelGroups"] && len(backup.TunnelGroups) > 0 {
count, err := importTunnelGroups(tx, backup.TunnelGroups, now)
if err != nil {
return fmt.Errorf("import tunnel groups failed: %w", err)
}
result.TunnelGroupsImported = count
}
if typeSet["userGroups"] && len(backup.UserGroups) > 0 {
count, err := importUserGroups(tx, backup.UserGroups, now)
if err != nil {
return fmt.Errorf("import user groups failed: %w", err)
}
result.UserGroupsImported = count
}
if typeSet["permissions"] && len(backup.Permissions) > 0 {
count, err := importPermissions(tx, backup.Permissions, now)
if err != nil {
return fmt.Errorf("import permissions failed: %w", err)
}
result.PermissionsImported = count
}
if typeSet["configs"] && len(backup.Configs) > 0 {
count, err := importConfigs(tx, backup.Configs, now)
if err != nil {
return fmt.Errorf("import configs failed: %w", err)
}
result.ConfigsImported = count
}
return nil
})
if err != nil {
return nil, err
}
return result, nil
}
func importUsers(tx *gorm.DB, users []model.UserBackup, now int64) (int, error) {
count := 0
for _, u := range users {
item := model.User{
ID: u.ID,
User: u.User,
Pwd: u.Pwd,
RoleID: u.RoleID,
ExpTime: u.ExpTime,
Flow: u.Flow,
InFlow: u.InFlow,
OutFlow: u.OutFlow,
FlowResetTime: u.FlowResetTime,
Num: u.Num,
CreatedTime: u.CreatedTime,
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
Status: u.Status,
}
err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{
"user", "pwd", "role_id", "exp_time", "flow", "in_flow", "out_flow",
"flow_reset_time", "num", "updated_time", "status",
}),
}).Create(&item).Error
if err != nil {
return count, err
}
count++
}
return count, nil
}
func importNodes(tx *gorm.DB, nodes []model.NodeBackup, now int64) (int, error) {
count := 0
for _, n := range nodes {
item := model.Node{
ID: n.ID,
Name: n.Name,
Secret: n.Secret,
ServerIP: n.ServerIP,
ServerIPV4: sql.NullString{String: n.ServerIPv4, Valid: true},
ServerIPV6: sql.NullString{String: n.ServerIPv6, Valid: true},
Port: n.Port,
InterfaceName: sql.NullString{String: n.InterfaceName, Valid: true},
Version: sql.NullString{String: n.Version, Valid: true},
HTTP: n.HTTP,
TLS: n.TLS,
Socks: n.Socks,
CreatedTime: n.CreatedTime,
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
Status: n.Status,
TCPListenAddr: n.TCPListenAddr,
UDPListenAddr: n.UDPListenAddr,
Inx: n.Inx,
IsRemote: n.IsRemote,
RemoteURL: sql.NullString{String: n.RemoteURL, Valid: true},
RemoteToken: sql.NullString{String: n.RemoteToken, Valid: true},
RemoteConfig: sql.NullString{String: n.RemoteConfig, Valid: true},
}
err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{
"name", "secret", "server_ip", "server_ip_v4", "server_ip_v6", "port", "interface_name", "version",
"http", "tls", "socks", "updated_time", "status", "tcp_listen_addr", "udp_listen_addr",
"inx", "is_remote", "remote_url", "remote_token", "remote_config",
}),
}).Create(&item).Error
if err != nil {
return count, err
}
count++
}
return count, nil
}
func importTunnels(tx *gorm.DB, tunnels []model.TunnelBackup, now int64) (int, error) {
count := 0
for _, t := range tunnels {
item := model.Tunnel{
ID: t.ID,
Name: t.Name,
TrafficRatio: t.TrafficRatio,
Type: t.Type,
Protocol: t.Protocol,
Flow: t.Flow,
CreatedTime: t.CreatedTime,
UpdatedTime: now,
Status: t.Status,
InIP: sql.NullString{String: t.InIP, Valid: true},
Inx: t.Inx,
IPPreference: t.IPPreference,
}
err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{
"name", "traffic_ratio", "type", "protocol", "flow", "updated_time", "status", "in_ip", "inx", "ip_preference",
}),
}).Create(&item).Error
if err != nil {
return count, err
}
for _, ct := range t.ChainTunnels {
chainItem := model.ChainTunnel{
ID: ct.ID,
TunnelID: ct.TunnelID,
ChainType: ct.ChainType,
NodeID: ct.NodeID,
Port: sql.NullInt64{Int64: int64(ct.Port), Valid: true},
Strategy: sql.NullString{String: ct.Strategy, Valid: true},
Inx: sql.NullInt64{Int64: int64(ct.Inx), Valid: true},
Protocol: sql.NullString{String: ct.Protocol, Valid: true},
}
err = tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{
"chain_type", "node_id", "port", "strategy", "inx", "protocol",
}),
}).Create(&chainItem).Error
if err != nil {
return count, err
}
}
count++
}
return count, nil
}
func importForwards(tx *gorm.DB, forwards []model.ForwardBackup, now int64) (int, error) {
count := 0
for _, f := range forwards {
item := model.Forward{
ID: f.ID,
UserID: f.UserID,
UserName: f.UserName,
Name: f.Name,
TunnelID: f.TunnelID,
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
InFlow: f.InFlow,
OutFlow: f.OutFlow,
CreatedTime: f.CreatedTime,
UpdatedTime: now,
Status: f.Status,
Inx: f.Inx,
}
err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{
"user_id", "user_name", "name", "tunnel_id", "remote_addr", "strategy",
"in_flow", "out_flow", "updated_time", "status", "inx",
}),
}).Create(&item).Error
if err != nil {
return count, err
}
if f.ForwardPorts != nil {
if err := tx.Where("forward_id = ?", f.ID).Delete(&model.ForwardPort{}).Error; err != nil {
return count, err
}
for _, fp := range *f.ForwardPorts {
if err := tx.Create(&model.ForwardPort{ForwardID: f.ID, NodeID: fp.NodeID, Port: fp.Port}).Error; err != nil {
return count, err
}
}
}
count++
}
return count, nil
}
func importUserTunnels(tx *gorm.DB, userTunnels []model.UserTunnelBackup, _ int64) (int, error) {
count := 0
for _, ut := range userTunnels {
item := model.UserTunnel{
ID: ut.ID,
UserID: ut.UserID,
TunnelID: ut.TunnelID,
SpeedID: sql.NullInt64{Int64: ut.SpeedID, Valid: ut.SpeedID > 0},
Num: ut.Num,
Flow: ut.Flow,
InFlow: ut.InFlow,
OutFlow: ut.OutFlow,
FlowResetTime: ut.FlowResetTime,
ExpTime: ut.ExpTime,
Status: ut.Status,
}
err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{
"user_id", "tunnel_id", "speed_id", "num", "flow", "in_flow", "out_flow",
"flow_reset_time", "exp_time", "status",
}),
}).Create(&item).Error
if err != nil {
return count, err
}
count++
}
return count, nil
}
func importSpeedLimits(tx *gorm.DB, speedLimits []model.SpeedLimitBackup, now int64) (int, error) {
count := 0
for _, sl := range speedLimits {
item := model.SpeedLimit{
ID: sl.ID,
Name: sl.Name,
Speed: int(sl.Speed),
TunnelID: sql.NullInt64{Int64: 0, Valid: false},
TunnelName: sql.NullString{String: "", Valid: false},
CreatedTime: sl.CreatedTime,
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
Status: sl.Status,
}
err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{
"name", "speed", "tunnel_id", "tunnel_name", "updated_time", "status",
}),
}).Create(&item).Error
if err != nil {
return count, err
}
count++
}
return count, nil
}
func importTunnelGroups(tx *gorm.DB, tunnelGroups []model.TunnelGroupBackup, now int64) (int, error) {
count := 0
for _, tg := range tunnelGroups {
item := model.TunnelGroup{
ID: tg.ID,
Name: tg.Name,
CreatedTime: tg.CreatedTime,
UpdatedTime: now,
Status: tg.Status,
}
err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{"name", "updated_time", "status"}),
}).Create(&item).Error
if err != nil {
return count, err
}
if err := tx.Where("tunnel_group_id = ?", tg.ID).Delete(&model.TunnelGroupTunnel{}).Error; err != nil {
return count, err
}
for _, tunnelID := range tg.Tunnels {
if err := tx.Create(&model.TunnelGroupTunnel{TunnelGroupID: tg.ID, TunnelID: tunnelID, CreatedTime: now}).Error; err != nil {
return count, err
}
}
count++
}
return count, nil
}
func importUserGroups(tx *gorm.DB, userGroups []model.UserGroupBackup, now int64) (int, error) {
count := 0
for _, ug := range userGroups {
item := model.UserGroup{
ID: ug.ID,
Name: ug.Name,
CreatedTime: ug.CreatedTime,
UpdatedTime: now,
Status: ug.Status,
}
err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{"name", "updated_time", "status"}),
}).Create(&item).Error
if err != nil {
return count, err
}
if err := tx.Where("user_group_id = ?", ug.ID).Delete(&model.UserGroupUser{}).Error; err != nil {
return count, err
}
for _, userID := range ug.Users {
if err := tx.Create(&model.UserGroupUser{UserGroupID: ug.ID, UserID: userID, CreatedTime: now}).Error; err != nil {
return count, err
}
}
count++
}
return count, nil
}
func importPermissions(tx *gorm.DB, permissions []model.PermissionBackup, _ int64) (int, error) {
count := 0
for _, p := range permissions {
item := model.GroupPermission{
ID: p.ID,
UserGroupID: p.UserGroupID,
TunnelGroupID: p.TunnelGroupID,
CreatedTime: p.CreatedTime,
}
err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{"user_group_id", "tunnel_group_id"}),
}).Create(&item).Error
if err != nil {
return count, err
}
for _, g := range p.Grants {
grantItem := model.GroupPermissionGrant{
ID: g.ID,
UserGroupID: g.UserGroupID,
TunnelGroupID: g.TunnelGroupID,
UserTunnelID: g.UserTunnelID,
CreatedTime: g.CreatedTime,
CreatedByGroup: g.CreatedByGroup,
}
err = tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{"user_tunnel_id", "created_by_group"}),
}).Create(&grantItem).Error
if err != nil {
return count, err
}
}
count++
}
return count, nil
}
func importConfigs(tx *gorm.DB, configs map[string]string, now int64) (int, error) {
count := 0
for name, value := range configs {
err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "name"}},
DoUpdates: clause.AssignmentColumns([]string{"value", "time"}),
}).Create(&model.ViteConfig{Name: name, Value: value, Time: now}).Error
if err != nil {
return count, err
}
count++
}
return count, nil
}
// ─── Jobs Queries (background stats / expiry) ───────────────────────
func (r *Repository) PurgeOldStatisticsFlows(cutoffMs int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Where("created_time < ?", cutoffMs).Delete(&model.StatisticsFlow{}).Error
}
func (r *Repository) ListAllUserFlowSnapshots() ([]model.UserFlowSnapshot, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var users []model.User
err := r.db.Order("id ASC").Find(&users).Error
if err != nil {
return nil, err
}
out := make([]model.UserFlowSnapshot, len(users))
for i, u := range users {
out[i] = model.UserFlowSnapshot{UserID: u.ID, InFlow: u.InFlow, OutFlow: u.OutFlow}
}
return out, nil
}
func (r *Repository) GetLastStatisticsFlowTotal(userID int64) (sql.NullInt64, error) {
if r == nil || r.db == nil {
return sql.NullInt64{}, errors.New("repository not initialized")
}
var sf model.StatisticsFlow
err := r.db.Where("user_id = ?", userID).Order("id DESC").First(&sf).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return sql.NullInt64{}, nil
}
if err != nil {
return sql.NullInt64{}, err
}
return sql.NullInt64{Int64: sf.TotalFlow, Valid: true}, nil
}
func (r *Repository) CreateStatisticsFlow(userID, flow, totalFlow int64, timeText string, createdTime int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Create(&model.StatisticsFlow{
UserID: userID, Flow: flow, TotalFlow: totalFlow,
Time: timeText, CreatedTime: createdTime,
}).Error
}
func (r *Repository) ResetUserMonthlyFlow(day int, lastDay int) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
updates := map[string]interface{}{"in_flow": 0, "out_flow": 0}
if day == lastDay {
return r.db.Model(&model.User{}).
Where("flow_reset_time != 0 AND (flow_reset_time = ? OR flow_reset_time > ?)", day, lastDay).
Updates(updates).Error
}
return r.db.Model(&model.User{}).
Where("flow_reset_time != 0 AND flow_reset_time = ?", day).
Updates(updates).Error
}
func (r *Repository) ResetUserTunnelMonthlyFlow(day int, lastDay int) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
updates := map[string]interface{}{"in_flow": 0, "out_flow": 0}
if day == lastDay {
return r.db.Model(&model.UserTunnel{}).
Where("flow_reset_time != 0 AND (flow_reset_time = ? OR flow_reset_time > ?)", day, lastDay).
Updates(updates).Error
}
return r.db.Model(&model.UserTunnel{}).
Where("flow_reset_time != 0 AND flow_reset_time = ?", day).
Updates(updates).Error
}
func (r *Repository) ListExpiredActiveUserIDs(nowMs int64) ([]int64, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var ids []int64
err := r.db.Model(&model.User{}).
Where("role_id != 0 AND status = 1 AND exp_time > 0 AND exp_time < ?", nowMs).
Pluck("id", &ids).Error
if err != nil {
return nil, err
}
return ids, nil
}
func (r *Repository) DisableUser(userID int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.User{}).Where("id = ?", userID).Update("status", 0).Error
}
func (r *Repository) ListExpiredActiveUserTunnels(nowMs int64) ([]model.ExpiredUserTunnel, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var uts []model.UserTunnel
err := r.db.Where("status = 1 AND exp_time > 0 AND exp_time < ?", nowMs).Find(&uts).Error
if err != nil {
return nil, err
}
out := make([]model.ExpiredUserTunnel, len(uts))
for i, ut := range uts {
out[i] = model.ExpiredUserTunnel{ID: ut.ID, UserID: ut.UserID, TunnelID: ut.TunnelID}
}
return out, nil
}
func (r *Repository) DisableUserTunnel(id int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.UserTunnel{}).Where("id = ?", id).Update("status", 0).Error
}
func (r *Repository) GetUserTunnelByID(id int64) (*model.UserTunnel, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var ut model.UserTunnel
err := r.db.Where("id = ?", id).First(&ut).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &ut, nil
}
// ─── Migration ───────────────────────────────────────────────────────
const currentSchemaVersion = 4
var ensurePostgresIDDefaultsFn = ensurePostgresIDDefaults
var migrateViteConfigValueColumnTypeFn = migrateViteConfigValueColumnType
var migrateSpeedLimitTunnelBindingFn = migrateSpeedLimitTunnelBinding
func getSchemaVersion(db *gorm.DB) int {
var v model.SchemaVersion
if err := db.First(&v).Error; err != nil {
db.Create(&model.SchemaVersion{Version: 0})
return 0
}
return v.Version
}
func setSchemaVersion(db *gorm.DB, ver int) {
db.Model(&model.SchemaVersion{}).Where("1=1").Update("version", ver)
}
func migrateSchema(db *gorm.DB) error {
if db == nil {
return errors.New("nil db")
}
if err := ensurePostgresIDDefaultsFn(db); err != nil {
return err
}
ver := getSchemaVersion(db)
if ver >= currentSchemaVersion {
return nil
}
// Normalize strategy columns
normalizeStrategy := func(modelRef interface{}, table, defaultValue string) error {
result := db.Model(modelRef).Where("strategy IS NULL").Update("strategy", defaultValue)
if result.Error != nil {
msg := strings.ToLower(result.Error.Error())
if strings.Contains(msg, "no such table") || (strings.Contains(msg, "relation") && strings.Contains(msg, "does not exist")) {
return nil
}
return fmt.Errorf("normalize %s.strategy: %w", table, result.Error)
}
return nil
}
if err := normalizeStrategy(&model.Forward{}, "forward", "fifo"); err != nil {
return err
}
if err := normalizeStrategy(&model.ChainTunnel{}, "chain_tunnel", "round"); err != nil {
return err
}
if err := normalizeStrategy(&model.PeerShareRuntime{}, "peer_share_runtime", "round"); err != nil {
return err
}
if ver < 3 {
if err := migrateViteConfigValueColumnTypeFn(db); err != nil {
return err
}
}
if ver < 4 {
if err := migrateSpeedLimitTunnelBindingFn(db); err != nil {
return err
}
}
setSchemaVersion(db, currentSchemaVersion)
return nil
}
func migrateViteConfigValueColumnType(db *gorm.DB) error {
if db == nil {
return errors.New("nil db")
}
if !db.Migrator().HasTable(&model.ViteConfig{}) {
return nil
}
if db.Dialector.Name() != "postgres" {
return nil
}
type columnRow struct {
DataType string `gorm:"column:data_type"`
}
var row columnRow
if err := db.Raw(
`SELECT data_type FROM information_schema.columns
WHERE table_schema = current_schema()
AND table_name = ?
AND column_name = ?`,
"vite_config", "value",
).Scan(&row).Error; err != nil {
return fmt.Errorf("inspect vite_config.value type: %w", err)
}
if strings.EqualFold(row.DataType, "text") {
return nil
}
if err := db.Exec(`ALTER TABLE "vite_config" ALTER COLUMN "value" TYPE TEXT`).Error; err != nil {
return fmt.Errorf("alter vite_config.value to text: %w", err)
}
return nil
}
func migrateSpeedLimitTunnelBinding(db *gorm.DB) error {
if db == nil {
return errors.New("nil db")
}
if !db.Migrator().HasTable(&model.SpeedLimit{}) {
return nil
}
if err := db.Model(&model.SpeedLimit{}).
Where("tunnel_id IS NOT NULL OR tunnel_name IS NOT NULL").
UpdateColumns(map[string]interface{}{
"tunnel_id": nil,
"tunnel_name": nil,
}).Error; err != nil {
return fmt.Errorf("clear speed_limit tunnel binding: %w", err)
}
return nil
}
func ensurePostgresIDDefaults(db *gorm.DB) error {
if db.Dialector.Name() != "postgres" {
return nil
}
type idRow struct {
TableSchema string
TableName string
}
var rows []idRow
err := db.Table("information_schema.table_constraints AS tc").
Select("c.table_schema, c.table_name").
Joins("JOIN information_schema.key_column_usage AS kcu ON tc.constraint_name = kcu.constraint_name AND tc.table_schema = kcu.table_schema").
Joins("JOIN information_schema.columns AS c ON c.table_schema = kcu.table_schema AND c.table_name = kcu.table_name AND c.column_name = kcu.column_name").
Where("tc.constraint_type = ?", "PRIMARY KEY").
Where("kcu.column_name = ?", "id").
Where("c.data_type IN ?", []string{"integer", "bigint"}).
Where("c.is_identity = ?", "NO").
Where("c.table_schema = current_schema()").
Order("c.table_name ASC").
Scan(&rows).Error
if err != nil {
return fmt.Errorf("discover postgres id columns: %w", err)
}
for _, r := range rows {
if err := ensurePostgresTableIDDefault(db, r.TableSchema, r.TableName); err != nil {
return fmt.Errorf("repair %s.%s id default: %w", r.TableSchema, r.TableName, err)
}
}
return nil
}
func ensurePostgresTableIDDefault(db *gorm.DB, schemaName, tableName string) error {
type defaultRow struct {
ColumnDefault sql.NullString `gorm:"column:column_default"`
}
var row defaultRow
err := db.Table("information_schema.columns").
Select("column_default").
Where("table_schema = ? AND table_name = ? AND column_name = 'id'", schemaName, tableName).
Limit(1).
Scan(&row).Error
if err != nil {
return err
}
defaultExpr := row.ColumnDefault
hasNextvalDefault := defaultExpr.Valid && strings.Contains(strings.ToLower(defaultExpr.String), "nextval(")
seqRef := ""
if hasNextvalDefault {
seqRef = extractNextvalRegclass(defaultExpr.String)
}
if !hasNextvalDefault || seqRef == "" {
seqName := tableName + "_id_seq"
if err := db.Exec(fmt.Sprintf("CREATE SEQUENCE IF NOT EXISTS %s.%s", quoteSQLIdentifier(schemaName), quoteSQLIdentifier(seqName))).Error; err != nil {
return err
}
seqRef = schemaName + "." + seqName
if err := db.Exec(fmt.Sprintf(
"ALTER TABLE %s.%s ALTER COLUMN id SET DEFAULT nextval(%s::regclass)",
quoteSQLIdentifier(schemaName), quoteSQLIdentifier(tableName), quoteSQLLiteral(seqRef),
)).Error; err != nil {
return err
}
if err := db.Exec(fmt.Sprintf(
"ALTER SEQUENCE %s.%s OWNED BY %s.%s.id",
quoteSQLIdentifier(schemaName), quoteSQLIdentifier(seqName),
quoteSQLIdentifier(schemaName), quoteSQLIdentifier(tableName),
)).Error; err != nil {
return err
}
}
return syncPostgresTableIDSequence(db, schemaName, tableName, seqRef)
}
func syncPostgresTableIDSequence(db *gorm.DB, schemaName, tableName, seqRef string) error {
type maxRow struct {
MaxID int64 `gorm:"column:max_id"`
}
var row maxRow
qualifiedTable := fmt.Sprintf("%s.%s", quoteSQLIdentifier(schemaName), quoteSQLIdentifier(tableName))
err := db.Table(qualifiedTable).
Select("COALESCE(MAX(id), 0) AS max_id").
Scan(&row).Error
if err != nil {
return err
}
maxID := row.MaxID
setVal := maxID
isCalled := true
if maxID <= 0 {
setVal = 1
isCalled = false
}
return db.Exec(`SELECT setval(?::regclass, ?, ?)`, seqRef, setVal, isCalled).Error
}
func extractNextvalRegclass(defaultExpr string) string {
nextvalIdx := strings.Index(strings.ToLower(defaultExpr), "nextval(")
if nextvalIdx < 0 {
return ""
}
expr := defaultExpr[nextvalIdx:]
firstQuote := strings.Index(expr, "'")
if firstQuote < 0 {
return ""
}
expr = expr[firstQuote+1:]
secondQuote := strings.Index(expr, "'")
if secondQuote < 0 {
return ""
}
return strings.TrimSpace(expr[:secondQuote])
}
func quoteSQLIdentifier(ident string) string {
return `"` + strings.ReplaceAll(ident, `"`, `""`) + `"`
}
func quoteSQLLiteral(value string) string {
return "'" + strings.ReplaceAll(value, "'", "''") + "'"
}
// ─── Helper Functions ────────────────────────────────────────────────
func resolveForwardIngress(db *gorm.DB, forwardID int64, tunnelID int64) (string, sql.NullInt64, error) {
var tunnelInIP sql.NullString
db.Model(&model.Tunnel{}).Select("in_ip").Where("id = ?", tunnelID).Limit(1).Scan(&tunnelInIP)
type fpRow struct {
Port sql.NullInt64
ServerIP sql.NullString
InIP sql.NullString
}
var fpRows []fpRow
err := db.Model(&model.ForwardPort{}).
Select("forward_port.port, node.server_ip, forward_port.in_ip").
Joins("LEFT JOIN node ON node.id = forward_port.node_id").
Where("forward_port.forward_id = ?", forwardID).
Order("forward_port.id ASC").
Find(&fpRows).Error
if err != nil {
return "", sql.NullInt64{}, err
}
ports := make([]int64, 0)
entries := make([]string, 0)
seenPorts := make(map[int64]struct{})
seenPairs := make(map[string]struct{})
for _, row := range fpRows {
if !row.Port.Valid {
continue
}
if _, ok := seenPorts[row.Port.Int64]; !ok {
seenPorts[row.Port.Int64] = struct{}{}
ports = append(ports, row.Port.Int64)
}
var ip string
if row.InIP.Valid && strings.TrimSpace(row.InIP.String) != "" {
ip = strings.TrimSpace(row.InIP.String)
} else if row.ServerIP.Valid && strings.TrimSpace(row.ServerIP.String) != "" {
ip = strings.TrimSpace(row.ServerIP.String)
}
if ip != "" {
pair := fmt.Sprintf("%s:%d", ip, row.Port.Int64)
if _, ok := seenPairs[pair]; !ok {
seenPairs[pair] = struct{}{}
entries = append(entries, pair)
}
}
}
if len(ports) == 0 {
return "", sql.NullInt64{}, nil
}
inPort := sql.NullInt64{Int64: ports[0], Valid: true}
return strings.Join(entries, ","), inPort, nil
}
func nullableString(v sql.NullString) interface{} {
if v.Valid {
return v.String
}
return nil
}
func nullableForwardIngress(v string) interface{} {
v = strings.TrimSpace(v)
if v == "" {
return nil
}
return v
}
func nullableInt64(v sql.NullInt64) interface{} {
if v.Valid {
return v.Int64
}
return nil
}
func unixMilliNow() int64 {
return time.Now().UnixMilli()
}
func ensureParentDir(dbPath string) error {
if dbPath == "" {
return fmt.Errorf("empty db path")
}
dir := filepath.Dir(dbPath)
if dir == "" || dir == "." {
return nil
}
return osMkdirAll(dir)
}
var osMkdirAll = func(path string) error {
return os.MkdirAll(path, 0o755)
}
// Suppress unused import warning for log
var _ = log.Printf