mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-30 16:26:36 +08:00
3065 lines
87 KiB
Go
3065 lines
87 KiB
Go
package sqlite
|
|
|
|
import (
|
|
"database/sql"
|
|
_ "embed"
|
|
"errors"
|
|
"fmt"
|
|
"log"
|
|
"os"
|
|
"path/filepath"
|
|
"sort"
|
|
"strings"
|
|
"time"
|
|
|
|
_ "github.com/jackc/pgx/v5/stdlib"
|
|
"go-backend/internal/store"
|
|
pgstore "go-backend/internal/store/postgres"
|
|
_ "modernc.org/sqlite"
|
|
)
|
|
|
|
//go:embed sql/schema.sql
|
|
var embeddedSchema string
|
|
|
|
//go:embed sql/data.sql
|
|
var embeddedSeedData string
|
|
|
|
// Execer is an interface that both *store.DB and *store.Tx satisfy.
|
|
// Used to allow import functions to work with both regular DB and transactions.
|
|
type Execer interface {
|
|
Exec(query string, args ...any) (sql.Result, error)
|
|
Query(query string, args ...any) (*sql.Rows, error)
|
|
QueryRow(query string, args ...any) *sql.Row
|
|
}
|
|
|
|
type Repository struct {
|
|
db *store.DB
|
|
}
|
|
|
|
func (r *Repository) DB() *store.DB {
|
|
if r == nil {
|
|
return nil
|
|
}
|
|
return r.db
|
|
}
|
|
|
|
type User struct {
|
|
ID int64
|
|
User string
|
|
Pwd string
|
|
RoleID int
|
|
ExpTime int64
|
|
Flow int64
|
|
InFlow int64
|
|
OutFlow int64
|
|
FlowResetTime int64
|
|
Num int
|
|
CreatedTime int64
|
|
UpdatedTime sql.NullInt64
|
|
Status int
|
|
}
|
|
|
|
type ViteConfig struct {
|
|
ID int64 `json:"id"`
|
|
Name string `json:"name"`
|
|
Value string `json:"value"`
|
|
Time int64 `json:"time"`
|
|
}
|
|
|
|
type UserTunnelDetail struct {
|
|
ID int64
|
|
UserID int64
|
|
TunnelID int64
|
|
TunnelName string
|
|
TunnelFlow int
|
|
Flow int64
|
|
InFlow int64
|
|
OutFlow int64
|
|
Num int
|
|
FlowResetTime int64
|
|
ExpTime int64
|
|
SpeedID sql.NullInt64
|
|
SpeedLimit sql.NullString
|
|
Speed sql.NullInt64
|
|
}
|
|
|
|
type UserForwardDetail struct {
|
|
ID int64
|
|
Name string
|
|
TunnelID int64
|
|
TunnelName string
|
|
InIP string
|
|
InPort sql.NullInt64
|
|
RemoteAddr string
|
|
InFlow int64
|
|
OutFlow int64
|
|
Status int
|
|
CreatedAt int64
|
|
}
|
|
|
|
type StatisticsFlow struct {
|
|
ID int64 `json:"id"`
|
|
UserID int64 `json:"userId"`
|
|
Flow int64 `json:"flow"`
|
|
TotalFlow int64 `json:"totalFlow"`
|
|
Time string `json:"time"`
|
|
}
|
|
|
|
type Node struct {
|
|
ID int64
|
|
Secret string
|
|
Version sql.NullString
|
|
HTTP int
|
|
TLS int
|
|
Socks int
|
|
Status int
|
|
IsRemote int
|
|
RemoteURL sql.NullString
|
|
RemoteToken sql.NullString
|
|
RemoteConfig sql.NullString
|
|
}
|
|
|
|
type PeerShare struct {
|
|
ID int64 `json:"id"`
|
|
Name string `json:"name"`
|
|
NodeID int64 `json:"nodeId"`
|
|
Token string `json:"token"`
|
|
MaxBandwidth int64 `json:"maxBandwidth"`
|
|
ExpiryTime int64 `json:"expiryTime"`
|
|
PortRangeStart int `json:"portRangeStart"`
|
|
PortRangeEnd int `json:"portRangeEnd"`
|
|
CurrentFlow int64 `json:"currentFlow"`
|
|
IsActive int `json:"isActive"`
|
|
CreatedTime int64 `json:"createdTime"`
|
|
UpdatedTime int64 `json:"updatedTime"`
|
|
AllowedDomains string `json:"allowedDomains"`
|
|
AllowedIPs string `json:"allowedIps"`
|
|
}
|
|
|
|
type PeerShareRuntime struct {
|
|
ID int64
|
|
ShareID int64
|
|
NodeID int64
|
|
ReservationID string
|
|
ResourceKey string
|
|
BindingID string
|
|
Role string
|
|
ChainName string
|
|
ServiceName string
|
|
Protocol string
|
|
Strategy string
|
|
Port int
|
|
Target string
|
|
Applied int
|
|
Status int
|
|
CreatedTime int64
|
|
UpdatedTime int64
|
|
}
|
|
|
|
type FederationTunnelBinding struct {
|
|
ID int64
|
|
TunnelID int64
|
|
NodeID int64
|
|
ChainType int
|
|
HopInx int
|
|
RemoteURL string
|
|
ResourceKey string
|
|
RemoteBindingID string
|
|
AllocatedPort int
|
|
Status int
|
|
CreatedTime int64
|
|
UpdatedTime int64
|
|
}
|
|
|
|
func Open(path string) (*Repository, error) {
|
|
if err := ensureParentDir(path); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// Use _pragma DSN parameters so every connection from the pool gets
|
|
// the same settings (busy_timeout and synchronous are per-connection).
|
|
dsn := "file:" + path +
|
|
"?_pragma=busy_timeout(5000)" +
|
|
"&_pragma=journal_mode(WAL)" +
|
|
"&_pragma=synchronous(NORMAL)"
|
|
raw, err := sql.Open("sqlite", dsn)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
db := store.Wrap(raw, store.DialectSQLite)
|
|
|
|
if err := db.Ping(); err != nil {
|
|
_ = db.Close()
|
|
return nil, err
|
|
}
|
|
|
|
if err := bootstrapSchema(db, embeddedSchema, embeddedSeedData); err != nil {
|
|
_ = db.Close()
|
|
return nil, err
|
|
}
|
|
|
|
if err := migrateSchema(db); err != nil {
|
|
_ = db.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")
|
|
}
|
|
|
|
raw, err := sql.Open("pgx", dsn)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
db := store.Wrap(raw, store.DialectPostgres)
|
|
|
|
if err := db.Ping(); err != nil {
|
|
_ = db.Close()
|
|
return nil, err
|
|
}
|
|
|
|
if err := bootstrapSchema(db, pgstore.EmbeddedSchema, pgstore.EmbeddedSeedData); err != nil {
|
|
_ = db.Close()
|
|
return nil, err
|
|
}
|
|
|
|
if err := migrateSchema(db); err != nil {
|
|
_ = db.Close()
|
|
return nil, err
|
|
}
|
|
|
|
return &Repository{db: db}, nil
|
|
}
|
|
|
|
func (r *Repository) Close() error {
|
|
if r == nil || r.db == nil {
|
|
return nil
|
|
}
|
|
return r.db.Close()
|
|
}
|
|
|
|
func (r *Repository) GetUserByUsername(username string) (*User, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
row := r.db.QueryRow(`
|
|
SELECT id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status
|
|
FROM user WHERE user = ? LIMIT 1
|
|
`, username)
|
|
user := &User{}
|
|
if err := row.Scan(
|
|
&user.ID, &user.User, &user.Pwd, &user.RoleID, &user.ExpTime,
|
|
&user.Flow, &user.InFlow, &user.OutFlow, &user.FlowResetTime,
|
|
&user.Num, &user.CreatedTime, &user.UpdatedTime, &user.Status,
|
|
); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, nil
|
|
}
|
|
return nil, err
|
|
}
|
|
return user, nil
|
|
}
|
|
|
|
func (r *Repository) GetConfigByName(name string) (*ViteConfig, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
row := r.db.QueryRow(`SELECT id, name, value, time FROM vite_config WHERE name = ? LIMIT 1`, name)
|
|
cfg := &ViteConfig{}
|
|
if err := row.Scan(&cfg.ID, &cfg.Name, &cfg.Value, &cfg.Time); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, 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")
|
|
}
|
|
|
|
rows, err := r.db.Query(`SELECT name, value FROM vite_config`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
result := make(map[string]string)
|
|
for rows.Next() {
|
|
var name, value string
|
|
if err := rows.Scan(&name, &value); err != nil {
|
|
return nil, err
|
|
}
|
|
result[name] = value
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
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")
|
|
}
|
|
|
|
_, err := r.db.Exec(`
|
|
INSERT INTO vite_config(name, value, time)
|
|
VALUES(?, ?, ?)
|
|
ON CONFLICT(name) DO UPDATE SET value=excluded.value, time=excluded.time
|
|
`, name, value, now)
|
|
return err
|
|
}
|
|
|
|
func (r *Repository) GetUserByID(id int64) (*User, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
row := r.db.QueryRow(`
|
|
SELECT id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status
|
|
FROM user WHERE id = ? LIMIT 1
|
|
`, id)
|
|
user := &User{}
|
|
if err := row.Scan(
|
|
&user.ID, &user.User, &user.Pwd, &user.RoleID, &user.ExpTime,
|
|
&user.Flow, &user.InFlow, &user.OutFlow, &user.FlowResetTime,
|
|
&user.Num, &user.CreatedTime, &user.UpdatedTime, &user.Status,
|
|
); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, nil
|
|
}
|
|
return nil, err
|
|
}
|
|
return user, nil
|
|
}
|
|
|
|
func (r *Repository) UsernameExistsExceptID(username string, exceptID int64) (bool, error) {
|
|
if r == nil || r.db == nil {
|
|
return false, errors.New("repository not initialized")
|
|
}
|
|
|
|
row := r.db.QueryRow(`SELECT COUNT(1) FROM user WHERE user = ? AND id != ?`, username, exceptID)
|
|
var count int
|
|
if err := row.Scan(&count); 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")
|
|
}
|
|
_, err := r.db.Exec(`UPDATE user SET user = ?, pwd = ?, updated_time = ? WHERE id = ?`, username, passwordMD5, now, userID)
|
|
return err
|
|
}
|
|
|
|
func (r *Repository) GetUserPackageTunnels(userID int64) ([]UserTunnelDetail, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
rows, err := r.db.Query(`
|
|
SELECT ut.id, ut.user_id, ut.tunnel_id, t.name, t.flow, ut.flow, ut.in_flow, ut.out_flow,
|
|
ut.num, ut.flow_reset_time, ut.exp_time, ut.speed_id, sl.name, sl.speed
|
|
FROM user_tunnel ut
|
|
LEFT JOIN tunnel t ON t.id = ut.tunnel_id
|
|
LEFT JOIN speed_limit sl ON sl.id = ut.speed_id
|
|
WHERE ut.user_id = ?
|
|
ORDER BY ut.id ASC
|
|
`, userID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
items := make([]UserTunnelDetail, 0)
|
|
for rows.Next() {
|
|
var item UserTunnelDetail
|
|
if err := rows.Scan(
|
|
&item.ID, &item.UserID, &item.TunnelID, &item.TunnelName, &item.TunnelFlow,
|
|
&item.Flow, &item.InFlow, &item.OutFlow, &item.Num, &item.FlowResetTime,
|
|
&item.ExpTime, &item.SpeedID, &item.SpeedLimit, &item.Speed,
|
|
); err != nil {
|
|
return nil, err
|
|
}
|
|
items = append(items, item)
|
|
}
|
|
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return items, nil
|
|
}
|
|
|
|
func (r *Repository) GetUserPackageForwards(userID int64) ([]UserForwardDetail, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
rows, err := r.db.Query(`
|
|
SELECT f.id, f.name, f.tunnel_id, COALESCE(t.name, ''), f.remote_addr, f.in_flow, f.out_flow, f.status, f.created_time
|
|
FROM forward f
|
|
LEFT JOIN tunnel t ON t.id = f.tunnel_id
|
|
WHERE f.user_id = ?
|
|
ORDER BY f.id ASC
|
|
`, userID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
items := make([]UserForwardDetail, 0)
|
|
for rows.Next() {
|
|
var item UserForwardDetail
|
|
if err := rows.Scan(
|
|
&item.ID, &item.Name, &item.TunnelID, &item.TunnelName, &item.RemoteAddr,
|
|
&item.InFlow, &item.OutFlow, &item.Status, &item.CreatedAt,
|
|
); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
inIP, inPort, err := resolveForwardIngress(r.db, item.ID, item.TunnelID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
item.InIP = inIP
|
|
item.InPort = inPort
|
|
|
|
items = append(items, item)
|
|
}
|
|
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return items, nil
|
|
}
|
|
|
|
func (r *Repository) GetStatisticsFlows(userID int64, limit int) ([]StatisticsFlow, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
rows, err := r.db.Query(`
|
|
SELECT id, user_id, flow, total_flow, time
|
|
FROM statistics_flow
|
|
WHERE user_id = ?
|
|
ORDER BY id DESC
|
|
LIMIT ?
|
|
`, userID, limit)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
items := make([]StatisticsFlow, 0)
|
|
for rows.Next() {
|
|
var item StatisticsFlow
|
|
if err := rows.Scan(&item.ID, &item.UserID, &item.Flow, &item.TotalFlow, &item.Time); err != nil {
|
|
return nil, err
|
|
}
|
|
items = append(items, item)
|
|
}
|
|
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return items, nil
|
|
}
|
|
|
|
func (r *Repository) NodeExistsBySecret(secret string) (bool, error) {
|
|
if r == nil || r.db == nil {
|
|
return false, errors.New("repository not initialized")
|
|
}
|
|
|
|
row := r.db.QueryRow(`SELECT COUNT(1) FROM node WHERE secret = ?`, secret)
|
|
var count int
|
|
if err := row.Scan(&count); err != nil {
|
|
return false, err
|
|
}
|
|
return count > 0, nil
|
|
}
|
|
|
|
func (r *Repository) GetNodeBySecret(secret string) (*Node, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
row := r.db.QueryRow(`SELECT id, secret, version, http, tls, socks, status, is_remote, remote_url, remote_token, remote_config FROM node WHERE secret = ? LIMIT 1`, secret)
|
|
var n Node
|
|
if err := row.Scan(&n.ID, &n.Secret, &n.Version, &n.HTTP, &n.TLS, &n.Socks, &n.Status, &n.IsRemote, &n.RemoteURL, &n.RemoteToken, &n.RemoteConfig); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, nil
|
|
}
|
|
return nil, err
|
|
}
|
|
return &n, nil
|
|
}
|
|
|
|
func (r *Repository) GetNodeByID(id int64) (*Node, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
row := r.db.QueryRow(`SELECT id, secret, version, http, tls, socks, status, is_remote, remote_url, remote_token, remote_config FROM node WHERE id = ? LIMIT 1`, id)
|
|
var n Node
|
|
if err := row.Scan(&n.ID, &n.Secret, &n.Version, &n.HTTP, &n.TLS, &n.Socks, &n.Status, &n.IsRemote, &n.RemoteURL, &n.RemoteToken, &n.RemoteConfig); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, 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")
|
|
}
|
|
_, err := r.db.Exec(`UPDATE node SET status = ?, version = ?, http = ?, tls = ?, socks = ?, updated_time = ? WHERE id = ?`,
|
|
status, version, httpVal, tlsVal, socksVal, unixMilliNow(), nodeID)
|
|
return err
|
|
}
|
|
|
|
func (r *Repository) UpdateNodeStatus(nodeID int64, status int) error {
|
|
if r == nil || r.db == nil {
|
|
return errors.New("repository not initialized")
|
|
}
|
|
_, err := r.db.Exec(`UPDATE node SET status = ?, updated_time = ? WHERE id = ?`, status, unixMilliNow(), nodeID)
|
|
return err
|
|
}
|
|
|
|
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")
|
|
}
|
|
|
|
tx, err := r.db.Begin()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer func() {
|
|
if err != nil {
|
|
_ = tx.Rollback()
|
|
}
|
|
}()
|
|
|
|
if _, err = tx.Exec(`UPDATE forward SET in_flow = in_flow + ?, out_flow = out_flow + ? WHERE id = ?`, inFlow, outFlow, forwardID); err != nil {
|
|
return err
|
|
}
|
|
if _, err = tx.Exec(`UPDATE user SET in_flow = in_flow + ?, out_flow = out_flow + ? WHERE id = ?`, inFlow, outFlow, userID); err != nil {
|
|
return err
|
|
}
|
|
if userTunnelID > 0 {
|
|
if _, err = tx.Exec(`UPDATE user_tunnel SET in_flow = in_flow + ?, out_flow = out_flow + ? WHERE id = ?`, inFlow, outFlow, userTunnelID); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
|
|
err = tx.Commit()
|
|
return err
|
|
}
|
|
|
|
func (r *Repository) ListNodes() ([]map[string]interface{}, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
rows, err := r.db.Query(`
|
|
SELECT id, inx, name, server_ip, server_ip_v4, server_ip_v6, port, tcp_listen_addr, udp_listen_addr, version, http, tls, socks, status, is_remote, remote_url, remote_token, remote_config
|
|
FROM node
|
|
ORDER BY inx ASC, id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
items := make([]map[string]interface{}, 0)
|
|
for rows.Next() {
|
|
var id, inx int64
|
|
var name, serverIP, port string
|
|
var serverIPV4, serverIPV6, tcpListen, udpListen, version, remoteURL, remoteToken, remoteConfig sql.NullString
|
|
var httpVal, tlsVal, socksVal, status, isRemote int
|
|
|
|
if err := rows.Scan(&id, &inx, &name, &serverIP, &serverIPV4, &serverIPV6, &port, &tcpListen, &udpListen, &version, &httpVal, &tlsVal, &socksVal, &status, &isRemote, &remoteURL, &remoteToken, &remoteConfig); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
items = append(items, map[string]interface{}{
|
|
"id": id,
|
|
"inx": inx,
|
|
"name": name,
|
|
"ip": serverIP,
|
|
"serverIp": serverIP,
|
|
"serverIpV4": nullableString(serverIPV4),
|
|
"serverIpV6": nullableString(serverIPV6),
|
|
"port": port,
|
|
"tcpListenAddr": nullableString(tcpListen),
|
|
"udpListenAddr": nullableString(udpListen),
|
|
"version": nullableString(version),
|
|
"http": httpVal,
|
|
"tls": tlsVal,
|
|
"socks": socksVal,
|
|
"status": status,
|
|
"isRemote": isRemote,
|
|
"remoteUrl": nullableString(remoteURL),
|
|
"remoteToken": nullableString(remoteToken),
|
|
"remoteConfig": nullableString(remoteConfig),
|
|
})
|
|
}
|
|
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
return items, nil
|
|
}
|
|
|
|
func (r *Repository) ListUsers() ([]map[string]interface{}, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
rows, err := r.db.Query(`
|
|
SELECT id, user, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status
|
|
FROM user
|
|
WHERE role_id != 0
|
|
ORDER BY id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
items := make([]map[string]interface{}, 0)
|
|
for rows.Next() {
|
|
var id int64
|
|
var user string
|
|
var roleID int
|
|
var expTime, flow, inFlow, outFlow, flowResetTime, createdTime int64
|
|
var num, status int
|
|
var updatedTime sql.NullInt64
|
|
|
|
if err := rows.Scan(&id, &user, &roleID, &expTime, &flow, &inFlow, &outFlow, &flowResetTime, &num, &createdTime, &updatedTime, &status); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
items = append(items, map[string]interface{}{
|
|
"id": id,
|
|
"user": user,
|
|
"name": user,
|
|
"roleId": roleID,
|
|
"status": status,
|
|
"flow": flow,
|
|
"num": num,
|
|
"expTime": expTime,
|
|
"flowResetTime": flowResetTime,
|
|
"createdTime": createdTime,
|
|
"updatedTime": nullableInt64(updatedTime),
|
|
"inFlow": inFlow,
|
|
"outFlow": outFlow,
|
|
})
|
|
}
|
|
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
return items, nil
|
|
}
|
|
|
|
func (r *Repository) ListSpeedLimits() ([]map[string]interface{}, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
rows, err := r.db.Query(`
|
|
SELECT id, name, speed, tunnel_id, tunnel_name, status, created_time, updated_time
|
|
FROM speed_limit
|
|
ORDER BY id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
items := make([]map[string]interface{}, 0)
|
|
for rows.Next() {
|
|
var id, tunnelID, createdTime int64
|
|
var name, tunnelName string
|
|
var speed, status int
|
|
var updatedTime sql.NullInt64
|
|
if err := rows.Scan(&id, &name, &speed, &tunnelID, &tunnelName, &status, &createdTime, &updatedTime); err != nil {
|
|
return nil, err
|
|
}
|
|
items = append(items, map[string]interface{}{
|
|
"id": id,
|
|
"name": name,
|
|
"speed": speed,
|
|
"tunnelId": tunnelID,
|
|
"tunnelName": tunnelName,
|
|
"status": status,
|
|
"createdTime": createdTime,
|
|
"updatedTime": nullableInt64(updatedTime),
|
|
})
|
|
}
|
|
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
return items, nil
|
|
}
|
|
|
|
func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
rows, err := r.db.Query(`
|
|
SELECT f.id, f.user_id, f.user_name, f.name, f.tunnel_id, COALESCE(t.name, ''), f.remote_addr, COALESCE(f.strategy, 'fifo'),
|
|
f.in_flow, f.out_flow, f.created_time, f.status, f.inx
|
|
FROM forward f
|
|
LEFT JOIN tunnel t ON t.id = f.tunnel_id
|
|
ORDER BY f.inx ASC, f.id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
items := make([]map[string]interface{}, 0)
|
|
for rows.Next() {
|
|
var id, userID, tunnelID, inFlow, outFlow, createdTime, inx int64
|
|
var userName, name, tunnelName, remoteAddr, strategy string
|
|
var status int
|
|
|
|
if err := rows.Scan(&id, &userID, &userName, &name, &tunnelID, &tunnelName, &remoteAddr, &strategy, &inFlow, &outFlow, &createdTime, &status, &inx); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
inIP, inPort, err := resolveForwardIngress(r.db, id, tunnelID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
items = append(items, map[string]interface{}{
|
|
"id": id,
|
|
"userId": userID,
|
|
"userName": userName,
|
|
"name": name,
|
|
"tunnelId": tunnelID,
|
|
"tunnelName": tunnelName,
|
|
"inIp": nullableForwardIngress(inIP),
|
|
"inPort": nullableInt64(inPort),
|
|
"remoteAddr": remoteAddr,
|
|
"strategy": strategy,
|
|
"inFlow": inFlow,
|
|
"outFlow": outFlow,
|
|
"createdTime": createdTime,
|
|
"status": status,
|
|
"inx": inx,
|
|
})
|
|
}
|
|
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
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")
|
|
}
|
|
|
|
rows, err := r.db.Query(`
|
|
SELECT t.id, t.name
|
|
FROM user_tunnel ut
|
|
JOIN tunnel t ON t.id = ut.tunnel_id
|
|
WHERE ut.user_id = ? AND t.status = 1
|
|
ORDER BY t.inx ASC, t.id ASC
|
|
`, userID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
items := make([]map[string]interface{}, 0)
|
|
for rows.Next() {
|
|
var id int64
|
|
var name string
|
|
if err := rows.Scan(&id, &name); err != nil {
|
|
return nil, err
|
|
}
|
|
items = append(items, map[string]interface{}{"id": id, "name": name})
|
|
}
|
|
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
return items, nil
|
|
}
|
|
|
|
func (r *Repository) ListEnabledTunnelSummaries() ([]map[string]interface{}, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
rows, err := r.db.Query(`
|
|
SELECT id, name
|
|
FROM tunnel
|
|
WHERE status = 1
|
|
ORDER BY inx ASC, id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
items := make([]map[string]interface{}, 0)
|
|
for rows.Next() {
|
|
var id int64
|
|
var name string
|
|
if err := rows.Scan(&id, &name); err != nil {
|
|
return nil, err
|
|
}
|
|
items = append(items, map[string]interface{}{"id": id, "name": name})
|
|
}
|
|
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
return items, nil
|
|
}
|
|
|
|
func (r *Repository) ListTunnels() ([]map[string]interface{}, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
rows, err := r.db.Query(`
|
|
SELECT id, inx, name, type, flow, traffic_ratio, status, created_time, in_ip
|
|
FROM tunnel
|
|
ORDER BY inx ASC, id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
tunnelMap := make(map[int64]map[string]interface{})
|
|
orderedIDs := make([]int64, 0)
|
|
|
|
for rows.Next() {
|
|
var id, inx, flow, createdTime int64
|
|
var name string
|
|
var typ, status int
|
|
var trafficRatio float64
|
|
var inIP sql.NullString
|
|
if err := rows.Scan(&id, &inx, &name, &typ, &flow, &trafficRatio, &status, &createdTime, &inIP); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
tunnelMap[id] = map[string]interface{}{
|
|
"id": id,
|
|
"inx": inx,
|
|
"name": name,
|
|
"type": typ,
|
|
"flow": flow,
|
|
"trafficRatio": trafficRatio,
|
|
"status": status,
|
|
"createdTime": createdTime,
|
|
"inIp": nullableString(inIP),
|
|
"inNodeId": make([]map[string]interface{}, 0),
|
|
"outNodeId": make([]map[string]interface{}, 0),
|
|
"chainNodes": make([][]map[string]interface{}, 0),
|
|
}
|
|
orderedIDs = append(orderedIDs, id)
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
nodeIPMap := map[int64]string{}
|
|
nRows, err := r.db.Query(`SELECT id, server_ip FROM node`)
|
|
if err == nil {
|
|
for nRows.Next() {
|
|
var id int64
|
|
var ip string
|
|
if scanErr := nRows.Scan(&id, &ip); scanErr == nil {
|
|
nodeIPMap[id] = ip
|
|
}
|
|
}
|
|
_ = nRows.Close()
|
|
}
|
|
|
|
chainRows, err := r.db.Query(`
|
|
SELECT tunnel_id, CAST(chain_type AS INTEGER), node_id, protocol, strategy, COALESCE(inx, 0)
|
|
FROM chain_tunnel
|
|
ORDER BY tunnel_id ASC, CAST(chain_type AS INTEGER) ASC, inx ASC, id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer chainRows.Close()
|
|
|
|
chainBucket := map[int64]map[int][]map[string]interface{}{}
|
|
inNodeIPs := map[int64][]string{}
|
|
|
|
for chainRows.Next() {
|
|
var tunnelID, nodeID, inx int64
|
|
var chainType int
|
|
var protocol, strategy sql.NullString
|
|
if err := chainRows.Scan(&tunnelID, &chainType, &nodeID, &protocol, &strategy, &inx); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
t, ok := tunnelMap[tunnelID]
|
|
if !ok {
|
|
continue
|
|
}
|
|
|
|
nodeObj := map[string]interface{}{
|
|
"nodeId": nodeID,
|
|
"chainType": chainType,
|
|
"inx": inx,
|
|
}
|
|
if protocol.Valid {
|
|
nodeObj["protocol"] = protocol.String
|
|
}
|
|
if strategy.Valid {
|
|
nodeObj["strategy"] = strategy.String
|
|
}
|
|
|
|
switch chainType {
|
|
case 1:
|
|
t["inNodeId"] = append(t["inNodeId"].([]map[string]interface{}), nodeObj)
|
|
if ip, ok := nodeIPMap[nodeID]; ok && ip != "" {
|
|
inNodeIPs[tunnelID] = append(inNodeIPs[tunnelID], ip)
|
|
}
|
|
case 2:
|
|
if _, ok := chainBucket[tunnelID]; !ok {
|
|
chainBucket[tunnelID] = map[int][]map[string]interface{}{}
|
|
}
|
|
chainBucket[tunnelID][int(inx)] = append(chainBucket[tunnelID][int(inx)], nodeObj)
|
|
case 3:
|
|
t["outNodeId"] = append(t["outNodeId"].([]map[string]interface{}), nodeObj)
|
|
}
|
|
}
|
|
if err := chainRows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
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
|
|
}
|
|
|
|
func (r *Repository) ListTunnelGroups() ([]map[string]interface{}, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
rows, err := r.db.Query(`SELECT id, name, status, created_time FROM tunnel_group ORDER BY id ASC`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
result := make([]map[string]interface{}, 0)
|
|
for rows.Next() {
|
|
var id, createdTime int64
|
|
var name string
|
|
var status int
|
|
if err := rows.Scan(&id, &name, &status, &createdTime); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
ids, names, err := r.listTunnelGroupMembers(id)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
result = append(result, map[string]interface{}{
|
|
"id": id,
|
|
"name": name,
|
|
"status": status,
|
|
"tunnelIds": ids,
|
|
"tunnelNames": names,
|
|
"createdTime": createdTime,
|
|
})
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
return result, nil
|
|
}
|
|
|
|
func (r *Repository) ListUserGroups() ([]map[string]interface{}, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
rows, err := r.db.Query(`SELECT id, name, status, created_time FROM user_group ORDER BY id ASC`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
result := make([]map[string]interface{}, 0)
|
|
for rows.Next() {
|
|
var id, createdTime int64
|
|
var name string
|
|
var status int
|
|
if err := rows.Scan(&id, &name, &status, &createdTime); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
ids, names, err := r.listUserGroupMembers(id)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
result = append(result, map[string]interface{}{
|
|
"id": id,
|
|
"name": name,
|
|
"status": status,
|
|
"userIds": ids,
|
|
"userNames": names,
|
|
"createdTime": createdTime,
|
|
})
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
return result, nil
|
|
}
|
|
|
|
func (r *Repository) ListGroupPermissions() ([]map[string]interface{}, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
|
|
rows, err := r.db.Query(`
|
|
SELECT gp.id, gp.user_group_id, ug.name, gp.tunnel_group_id, tg.name, gp.created_time
|
|
FROM group_permission gp
|
|
LEFT JOIN user_group ug ON ug.id = gp.user_group_id
|
|
LEFT JOIN tunnel_group tg ON tg.id = gp.tunnel_group_id
|
|
ORDER BY gp.id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
result := make([]map[string]interface{}, 0)
|
|
for rows.Next() {
|
|
var id, userGroupID, tunnelGroupID, createdTime int64
|
|
var userGroupName, tunnelGroupName sql.NullString
|
|
if err := rows.Scan(&id, &userGroupID, &userGroupName, &tunnelGroupID, &tunnelGroupName, &createdTime); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
result = append(result, map[string]interface{}{
|
|
"id": id,
|
|
"userGroupId": userGroupID,
|
|
"userGroupName": nullableString(userGroupName),
|
|
"tunnelGroupId": tunnelGroupID,
|
|
"tunnelGroupName": nullableString(tunnelGroupName),
|
|
"createdTime": createdTime,
|
|
})
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
return result, nil
|
|
}
|
|
|
|
func (r *Repository) listTunnelGroupMembers(groupID int64) ([]int64, []string, error) {
|
|
rows, err := r.db.Query(`
|
|
SELECT t.id, t.name
|
|
FROM tunnel_group_tunnel tgt
|
|
JOIN tunnel t ON t.id = tgt.tunnel_id
|
|
WHERE tgt.tunnel_group_id = ?
|
|
ORDER BY t.id ASC
|
|
`, groupID)
|
|
if err != nil {
|
|
return nil, nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
ids := make([]int64, 0)
|
|
names := make([]string, 0)
|
|
for rows.Next() {
|
|
var id int64
|
|
var name string
|
|
if err := rows.Scan(&id, &name); err != nil {
|
|
return nil, nil, err
|
|
}
|
|
ids = append(ids, id)
|
|
names = append(names, name)
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return nil, nil, err
|
|
}
|
|
return ids, names, nil
|
|
}
|
|
|
|
func (r *Repository) listUserGroupMembers(groupID int64) ([]int64, []string, error) {
|
|
rows, err := r.db.Query(`
|
|
SELECT u.id, u.user
|
|
FROM user_group_user ugu
|
|
JOIN user u ON u.id = ugu.user_id
|
|
WHERE ugu.user_group_id = ?
|
|
ORDER BY u.id ASC
|
|
`, groupID)
|
|
if err != nil {
|
|
return nil, nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
ids := make([]int64, 0)
|
|
names := make([]string, 0)
|
|
for rows.Next() {
|
|
var id int64
|
|
var name string
|
|
if err := rows.Scan(&id, &name); err != nil {
|
|
return nil, nil, err
|
|
}
|
|
ids = append(ids, id)
|
|
names = append(names, name)
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return nil, nil, err
|
|
}
|
|
return ids, names, 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 resolveForwardIngress(db *store.DB, forwardID int64, tunnelID int64) (string, sql.NullInt64, error) {
|
|
var tunnelInIP sql.NullString
|
|
if err := db.QueryRow(`SELECT in_ip FROM tunnel WHERE id = ? LIMIT 1`, tunnelID).Scan(&tunnelInIP); err != nil {
|
|
if !errors.Is(err, sql.ErrNoRows) {
|
|
return "", sql.NullInt64{}, err
|
|
}
|
|
}
|
|
|
|
rows, err := db.Query(`
|
|
SELECT fp.port, n.server_ip
|
|
FROM forward_port fp
|
|
LEFT JOIN node n ON n.id = fp.node_id
|
|
WHERE fp.forward_id = ?
|
|
ORDER BY fp.id ASC
|
|
`, forwardID)
|
|
if err != nil {
|
|
return "", sql.NullInt64{}, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
ports := make([]int64, 0)
|
|
nodePairs := make([]string, 0)
|
|
seenPorts := make(map[int64]struct{})
|
|
seenPairs := make(map[string]struct{})
|
|
|
|
for rows.Next() {
|
|
var port sql.NullInt64
|
|
var nodeIP sql.NullString
|
|
if err := rows.Scan(&port, &nodeIP); err != nil {
|
|
return "", sql.NullInt64{}, err
|
|
}
|
|
if !port.Valid {
|
|
continue
|
|
}
|
|
if _, ok := seenPorts[port.Int64]; !ok {
|
|
seenPorts[port.Int64] = struct{}{}
|
|
ports = append(ports, port.Int64)
|
|
}
|
|
if nodeIP.Valid && strings.TrimSpace(nodeIP.String) != "" {
|
|
pair := fmt.Sprintf("%s:%d", strings.TrimSpace(nodeIP.String), port.Int64)
|
|
if _, ok := seenPairs[pair]; !ok {
|
|
seenPairs[pair] = struct{}{}
|
|
nodePairs = append(nodePairs, pair)
|
|
}
|
|
}
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return "", sql.NullInt64{}, err
|
|
}
|
|
|
|
if len(ports) == 0 {
|
|
return "", sql.NullInt64{}, nil
|
|
}
|
|
|
|
inPort := sql.NullInt64{Int64: ports[0], Valid: true}
|
|
|
|
entries := make([]string, 0)
|
|
if tunnelInIP.Valid && strings.TrimSpace(tunnelInIP.String) != "" {
|
|
tunnelIPs := strings.Split(tunnelInIP.String, ",")
|
|
seen := make(map[string]struct{})
|
|
for _, ip := range tunnelIPs {
|
|
ip = strings.TrimSpace(ip)
|
|
if ip == "" {
|
|
continue
|
|
}
|
|
if _, ok := seen[ip]; ok {
|
|
continue
|
|
}
|
|
seen[ip] = struct{}{}
|
|
for _, port := range ports {
|
|
entries = append(entries, fmt.Sprintf("%s:%d", ip, port))
|
|
}
|
|
}
|
|
} else {
|
|
entries = append(entries, nodePairs...)
|
|
}
|
|
|
|
return strings.Join(entries, ","), inPort, nil
|
|
}
|
|
|
|
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)
|
|
}
|
|
|
|
func bootstrapSchema(db *store.DB, schemaSQL, seedSQL string) error {
|
|
if db == nil {
|
|
return errors.New("nil db")
|
|
}
|
|
|
|
if _, err := db.Exec(schemaSQL); err != nil {
|
|
return fmt.Errorf("apply schema.sql: %w", err)
|
|
}
|
|
|
|
if _, err := db.Exec(seedSQL); err != nil {
|
|
return fmt.Errorf("apply data.sql: %w", err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
const currentSchemaVersion = 2
|
|
|
|
var ensurePostgresIDDefaultsFn = ensurePostgresIDDefaults
|
|
|
|
func getSchemaVersion(db *store.DB) int {
|
|
_, _ = db.Exec(`CREATE TABLE IF NOT EXISTS schema_version (version INTEGER NOT NULL DEFAULT 0)`)
|
|
var v int
|
|
if err := db.QueryRow(`SELECT version FROM schema_version LIMIT 1`).Scan(&v); err != nil {
|
|
_, _ = db.Exec(`INSERT INTO schema_version(version) VALUES(0)`)
|
|
return 0
|
|
}
|
|
return v
|
|
}
|
|
|
|
func setSchemaVersion(db *store.DB, v int) {
|
|
_, _ = db.Exec(`UPDATE schema_version SET version = ?`, v)
|
|
}
|
|
|
|
func migrateSchema(db *store.DB) error {
|
|
if db == nil {
|
|
return errors.New("nil db")
|
|
}
|
|
|
|
ver := getSchemaVersion(db)
|
|
if db.Dialect() == store.DialectPostgres {
|
|
if err := ensurePostgresIDDefaultsFn(db); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
if ver >= currentSchemaVersion {
|
|
return nil
|
|
}
|
|
|
|
ensureColumn := func(table, col, typ string) {
|
|
var dummy interface{}
|
|
err := db.QueryRow(fmt.Sprintf("SELECT %s FROM %s LIMIT 1", col, table)).Scan(&dummy)
|
|
if err == nil || errors.Is(err, sql.ErrNoRows) {
|
|
return
|
|
}
|
|
if isMissingColumnError(db.Dialect(), err) {
|
|
if _, alterErr := db.Exec(fmt.Sprintf("ALTER TABLE %s ADD COLUMN %s %s", table, col, typ)); alterErr != nil {
|
|
log.Printf("failed to add column %s to %s: %v", col, table, alterErr)
|
|
}
|
|
}
|
|
}
|
|
|
|
columnsByTable := map[string]map[string]string{
|
|
"peer_share": {
|
|
"allowed_domains": "TEXT DEFAULT ''",
|
|
"allowed_ips": "TEXT DEFAULT ''",
|
|
},
|
|
"node": {
|
|
"server_ip_v4": "VARCHAR(100)",
|
|
"server_ip_v6": "VARCHAR(100)",
|
|
"inx": "INTEGER NOT NULL DEFAULT 0",
|
|
"is_remote": "INTEGER DEFAULT 0",
|
|
"remote_url": "TEXT",
|
|
"remote_token": "TEXT",
|
|
"remote_config": "TEXT",
|
|
},
|
|
"tunnel": {
|
|
"inx": "INTEGER NOT NULL DEFAULT 0",
|
|
},
|
|
"forward": {
|
|
"inx": "INTEGER NOT NULL DEFAULT 0",
|
|
},
|
|
"chain_tunnel": {
|
|
"inx": "INTEGER",
|
|
},
|
|
}
|
|
|
|
for table, columns := range columnsByTable {
|
|
for col, typ := range columns {
|
|
ensureColumn(table, col, typ)
|
|
}
|
|
}
|
|
|
|
normalizeStrategy := func(table, defaultValue string) error {
|
|
_, err := db.Exec(fmt.Sprintf("UPDATE %s SET strategy = ? WHERE strategy IS NULL", table), defaultValue)
|
|
if err != nil {
|
|
if isMissingTableError(db.Dialect(), err) {
|
|
return nil
|
|
}
|
|
return fmt.Errorf("normalize %s.strategy: %w", table, err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
if err := normalizeStrategy("forward", "fifo"); err != nil {
|
|
return err
|
|
}
|
|
if err := normalizeStrategy("chain_tunnel", "round"); err != nil {
|
|
return err
|
|
}
|
|
if err := normalizeStrategy("peer_share_runtime", "round"); err != nil {
|
|
return err
|
|
}
|
|
|
|
setSchemaVersion(db, currentSchemaVersion)
|
|
return nil
|
|
}
|
|
|
|
func ensurePostgresIDDefaults(db *store.DB) error {
|
|
rows, err := db.Query(`
|
|
SELECT c.table_schema, c.table_name
|
|
FROM information_schema.table_constraints tc
|
|
JOIN information_schema.key_column_usage kcu
|
|
ON tc.constraint_name = kcu.constraint_name
|
|
AND tc.table_schema = kcu.table_schema
|
|
JOIN information_schema.columns 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'
|
|
AND kcu.column_name = 'id'
|
|
AND c.data_type IN ('integer', 'bigint')
|
|
AND c.is_identity = 'NO'
|
|
AND c.table_schema = current_schema()
|
|
ORDER BY c.table_name ASC
|
|
`)
|
|
if err != nil {
|
|
return fmt.Errorf("discover postgres id columns: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
|
|
for rows.Next() {
|
|
var schemaName string
|
|
var tableName string
|
|
if err := rows.Scan(&schemaName, &tableName); err != nil {
|
|
return fmt.Errorf("scan postgres id table row: %w", err)
|
|
}
|
|
if err := ensurePostgresTableIDDefault(db, schemaName, tableName); err != nil {
|
|
return fmt.Errorf("repair %s.%s id default: %w", schemaName, tableName, err)
|
|
}
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return fmt.Errorf("iterate postgres id tables: %w", err)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func ensurePostgresTableIDDefault(db *store.DB, schemaName, tableName string) error {
|
|
var defaultExpr sql.NullString
|
|
if err := db.QueryRow(`
|
|
SELECT column_default
|
|
FROM information_schema.columns
|
|
WHERE table_schema = ?
|
|
AND table_name = ?
|
|
AND column_name = 'id'
|
|
LIMIT 1
|
|
`, schemaName, tableName).Scan(&defaultExpr); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil
|
|
}
|
|
return err
|
|
}
|
|
|
|
hasNextvalDefault := defaultExpr.Valid && strings.Contains(strings.ToLower(defaultExpr.String), "nextval(")
|
|
|
|
var serialSeq sql.NullString
|
|
if err := db.QueryRow(`
|
|
SELECT pg_get_serial_sequence(quote_ident(?) || '.' || quote_ident(?), 'id')
|
|
`, schemaName, tableName).Scan(&serialSeq); err != nil {
|
|
return err
|
|
}
|
|
|
|
seqRef := strings.TrimSpace(serialSeq.String)
|
|
if seqRef == "" && 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))); 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),
|
|
)); 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),
|
|
)); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
|
|
return syncPostgresTableIDSequence(db, schemaName, tableName, seqRef)
|
|
}
|
|
|
|
func syncPostgresTableIDSequence(db *store.DB, schemaName, tableName, seqRef string) error {
|
|
var maxID int64
|
|
if err := db.QueryRow(fmt.Sprintf(
|
|
"SELECT COALESCE(MAX(id), 0) FROM %s.%s",
|
|
quoteSQLIdentifier(schemaName),
|
|
quoteSQLIdentifier(tableName),
|
|
)).Scan(&maxID); err != nil {
|
|
return err
|
|
}
|
|
|
|
setVal := maxID
|
|
isCalled := true
|
|
if maxID <= 0 {
|
|
setVal = 1
|
|
isCalled = false
|
|
}
|
|
|
|
if _, err := db.Exec(`SELECT setval(?::regclass, ?, ?)`, seqRef, setVal, isCalled); err != nil {
|
|
return err
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
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, "'", "''") + "'"
|
|
}
|
|
|
|
func isMissingColumnError(dialect store.Dialect, err error) bool {
|
|
if err == nil {
|
|
return false
|
|
}
|
|
msg := strings.ToLower(err.Error())
|
|
if dialect == store.DialectPostgres {
|
|
return strings.Contains(msg, "column") && strings.Contains(msg, "does not exist")
|
|
}
|
|
return strings.Contains(msg, "no such column")
|
|
}
|
|
|
|
func isMissingTableError(dialect store.Dialect, err error) bool {
|
|
if err == nil {
|
|
return false
|
|
}
|
|
msg := strings.ToLower(err.Error())
|
|
if dialect == store.DialectPostgres {
|
|
return strings.Contains(msg, "relation") && strings.Contains(msg, "does not exist")
|
|
}
|
|
return strings.Contains(msg, "no such table")
|
|
}
|
|
|
|
func (r *Repository) CreatePeerShare(share *PeerShare) error {
|
|
if r == nil || r.db == nil {
|
|
return errors.New("repository not initialized")
|
|
}
|
|
_, err := r.db.Exec(`
|
|
INSERT INTO peer_share(name, node_id, token, max_bandwidth, expiry_time, port_range_start, port_range_end, current_flow, is_active, created_time, updated_time, allowed_domains, allowed_ips)
|
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
`, share.Name, share.NodeID, share.Token, share.MaxBandwidth, share.ExpiryTime, share.PortRangeStart, share.PortRangeEnd, share.CurrentFlow, share.IsActive, share.CreatedTime, share.UpdatedTime, share.AllowedDomains, share.AllowedIPs)
|
|
return err
|
|
}
|
|
|
|
func (r *Repository) UpdatePeerShare(share *PeerShare) error {
|
|
if r == nil || r.db == nil {
|
|
return errors.New("repository not initialized")
|
|
}
|
|
_, err := r.db.Exec(`
|
|
UPDATE peer_share SET name=?, max_bandwidth=?, expiry_time=?, port_range_start=?, port_range_end=?, is_active=?, updated_time=?, allowed_domains=?, allowed_ips=?
|
|
WHERE id=?
|
|
`, share.Name, share.MaxBandwidth, share.ExpiryTime, share.PortRangeStart, share.PortRangeEnd, share.IsActive, share.UpdatedTime, share.AllowedDomains, share.AllowedIPs, share.ID)
|
|
return err
|
|
}
|
|
|
|
func (r *Repository) DeletePeerShare(id int64) error {
|
|
if r == nil || r.db == nil {
|
|
return errors.New("repository not initialized")
|
|
}
|
|
tx, err := r.db.Begin()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer func() { _ = tx.Rollback() }()
|
|
_, _ = tx.Exec(`DELETE FROM peer_share_runtime WHERE share_id = ?`, id)
|
|
if _, err := tx.Exec(`DELETE FROM peer_share WHERE id=?`, id); err != nil {
|
|
return err
|
|
}
|
|
return tx.Commit()
|
|
}
|
|
|
|
func (r *Repository) GetPeerShare(id int64) (*PeerShare, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
row := r.db.QueryRow(`SELECT id, name, node_id, token, max_bandwidth, expiry_time, port_range_start, port_range_end, current_flow, is_active, created_time, updated_time, allowed_domains, allowed_ips FROM peer_share WHERE id = ?`, id)
|
|
var s PeerShare
|
|
if err := row.Scan(&s.ID, &s.Name, &s.NodeID, &s.Token, &s.MaxBandwidth, &s.ExpiryTime, &s.PortRangeStart, &s.PortRangeEnd, &s.CurrentFlow, &s.IsActive, &s.CreatedTime, &s.UpdatedTime, &s.AllowedDomains, &s.AllowedIPs); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, nil
|
|
}
|
|
return nil, err
|
|
}
|
|
return &s, nil
|
|
}
|
|
|
|
func (r *Repository) GetPeerShareByToken(token string) (*PeerShare, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
row := r.db.QueryRow(`SELECT id, name, node_id, token, max_bandwidth, expiry_time, port_range_start, port_range_end, current_flow, is_active, created_time, updated_time, allowed_domains, allowed_ips FROM peer_share WHERE token = ?`, token)
|
|
var s PeerShare
|
|
if err := row.Scan(&s.ID, &s.Name, &s.NodeID, &s.Token, &s.MaxBandwidth, &s.ExpiryTime, &s.PortRangeStart, &s.PortRangeEnd, &s.CurrentFlow, &s.IsActive, &s.CreatedTime, &s.UpdatedTime, &s.AllowedDomains, &s.AllowedIPs); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, nil
|
|
}
|
|
return nil, err
|
|
}
|
|
return &s, nil
|
|
}
|
|
|
|
func (r *Repository) ListPeerShares() ([]PeerShare, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
rows, err := r.db.Query(`SELECT id, name, node_id, token, max_bandwidth, expiry_time, port_range_start, port_range_end, current_flow, is_active, created_time, updated_time, allowed_domains, allowed_ips FROM peer_share ORDER BY id DESC`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var shares []PeerShare
|
|
for rows.Next() {
|
|
var s PeerShare
|
|
if err := rows.Scan(&s.ID, &s.Name, &s.NodeID, &s.Token, &s.MaxBandwidth, &s.ExpiryTime, &s.PortRangeStart, &s.PortRangeEnd, &s.CurrentFlow, &s.IsActive, &s.CreatedTime, &s.UpdatedTime, &s.AllowedDomains, &s.AllowedIPs); err != nil {
|
|
return nil, err
|
|
}
|
|
shares = append(shares, s)
|
|
}
|
|
return shares, nil
|
|
}
|
|
|
|
func (r *Repository) GetPeerShareRuntimeByResourceKey(shareID int64, resourceKey string) (*PeerShareRuntime, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
row := r.db.QueryRow(`
|
|
SELECT id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time
|
|
FROM peer_share_runtime
|
|
WHERE share_id = ? AND resource_key = ?
|
|
LIMIT 1
|
|
`, shareID, resourceKey)
|
|
var item PeerShareRuntime
|
|
if err := row.Scan(&item.ID, &item.ShareID, &item.NodeID, &item.ReservationID, &item.ResourceKey, &item.BindingID, &item.Role, &item.ChainName, &item.ServiceName, &item.Protocol, &item.Strategy, &item.Port, &item.Target, &item.Applied, &item.Status, &item.CreatedTime, &item.UpdatedTime); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, nil
|
|
}
|
|
return nil, err
|
|
}
|
|
return &item, nil
|
|
}
|
|
|
|
func (r *Repository) GetPeerShareRuntimeByReservationID(shareID int64, reservationID string) (*PeerShareRuntime, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
row := r.db.QueryRow(`
|
|
SELECT id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time
|
|
FROM peer_share_runtime
|
|
WHERE share_id = ? AND reservation_id = ?
|
|
LIMIT 1
|
|
`, shareID, reservationID)
|
|
var item PeerShareRuntime
|
|
if err := row.Scan(&item.ID, &item.ShareID, &item.NodeID, &item.ReservationID, &item.ResourceKey, &item.BindingID, &item.Role, &item.ChainName, &item.ServiceName, &item.Protocol, &item.Strategy, &item.Port, &item.Target, &item.Applied, &item.Status, &item.CreatedTime, &item.UpdatedTime); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, nil
|
|
}
|
|
return nil, err
|
|
}
|
|
return &item, nil
|
|
}
|
|
|
|
func (r *Repository) GetPeerShareRuntimeByBindingID(shareID int64, bindingID string) (*PeerShareRuntime, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
row := r.db.QueryRow(`
|
|
SELECT id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time
|
|
FROM peer_share_runtime
|
|
WHERE share_id = ? AND binding_id = ?
|
|
LIMIT 1
|
|
`, shareID, bindingID)
|
|
var item PeerShareRuntime
|
|
if err := row.Scan(&item.ID, &item.ShareID, &item.NodeID, &item.ReservationID, &item.ResourceKey, &item.BindingID, &item.Role, &item.ChainName, &item.ServiceName, &item.Protocol, &item.Strategy, &item.Port, &item.Target, &item.Applied, &item.Status, &item.CreatedTime, &item.UpdatedTime); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, nil
|
|
}
|
|
return nil, err
|
|
}
|
|
return &item, nil
|
|
}
|
|
|
|
func (r *Repository) GetPeerShareRuntimeByID(id int64) (*PeerShareRuntime, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
row := r.db.QueryRow(`
|
|
SELECT id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time
|
|
FROM peer_share_runtime
|
|
WHERE id = ?
|
|
LIMIT 1
|
|
`, id)
|
|
var item PeerShareRuntime
|
|
if err := row.Scan(&item.ID, &item.ShareID, &item.NodeID, &item.ReservationID, &item.ResourceKey, &item.BindingID, &item.Role, &item.ChainName, &item.ServiceName, &item.Protocol, &item.Strategy, &item.Port, &item.Target, &item.Applied, &item.Status, &item.CreatedTime, &item.UpdatedTime); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, nil
|
|
}
|
|
return nil, err
|
|
}
|
|
return &item, nil
|
|
}
|
|
|
|
func (r *Repository) ListActivePeerShareRuntimesByShareID(shareID int64) ([]PeerShareRuntime, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
rows, err := r.db.Query(`
|
|
SELECT id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time
|
|
FROM peer_share_runtime
|
|
WHERE share_id = ? AND status = 1
|
|
ORDER BY port ASC, id ASC
|
|
`, shareID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
out := make([]PeerShareRuntime, 0)
|
|
for rows.Next() {
|
|
var item PeerShareRuntime
|
|
if err := rows.Scan(&item.ID, &item.ShareID, &item.NodeID, &item.ReservationID, &item.ResourceKey, &item.BindingID, &item.Role, &item.ChainName, &item.ServiceName, &item.Protocol, &item.Strategy, &item.Port, &item.Target, &item.Applied, &item.Status, &item.CreatedTime, &item.UpdatedTime); err != nil {
|
|
return nil, err
|
|
}
|
|
out = append(out, item)
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
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
|
|
}
|
|
_, err := r.db.Exec(`UPDATE peer_share SET current_flow = current_flow + ?, updated_time = ? WHERE id = ?`, delta, unixMilliNow(), shareID)
|
|
return err
|
|
}
|
|
|
|
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()
|
|
}
|
|
_, err := r.db.Exec(`UPDATE peer_share SET current_flow = 0, updated_time = ? WHERE id = ?`, updatedTime, shareID)
|
|
return err
|
|
}
|
|
|
|
func (r *Repository) CreatePeerShareRuntime(item *PeerShareRuntime) error {
|
|
if r == nil || r.db == nil {
|
|
return errors.New("repository not initialized")
|
|
}
|
|
if item == nil {
|
|
return errors.New("runtime item is nil")
|
|
}
|
|
_, err := r.db.Exec(`
|
|
INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
|
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
`, item.ShareID, item.NodeID, item.ReservationID, item.ResourceKey, item.BindingID, item.Role, item.ChainName, item.ServiceName, item.Protocol, item.Strategy, item.Port, item.Target, item.Applied, item.Status, item.CreatedTime, item.UpdatedTime)
|
|
return err
|
|
}
|
|
|
|
func (r *Repository) UpdatePeerShareRuntime(item *PeerShareRuntime) error {
|
|
if r == nil || r.db == nil {
|
|
return errors.New("repository not initialized")
|
|
}
|
|
if item == nil {
|
|
return errors.New("runtime item is nil")
|
|
}
|
|
_, err := r.db.Exec(`
|
|
UPDATE peer_share_runtime
|
|
SET binding_id = ?, role = ?, chain_name = ?, service_name = ?, protocol = ?, strategy = ?, port = ?, target = ?, applied = ?, status = ?, updated_time = ?
|
|
WHERE id = ?
|
|
`, item.BindingID, item.Role, item.ChainName, item.ServiceName, item.Protocol, item.Strategy, item.Port, item.Target, item.Applied, item.Status, item.UpdatedTime, item.ID)
|
|
return err
|
|
}
|
|
|
|
func (r *Repository) MarkPeerShareRuntimeReleased(id int64, updatedTime int64) error {
|
|
if r == nil || r.db == nil {
|
|
return errors.New("repository not initialized")
|
|
}
|
|
_, err := r.db.Exec(`UPDATE peer_share_runtime SET status = 0, updated_time = ? WHERE id = ?`, updatedTime, id)
|
|
return err
|
|
}
|
|
|
|
func (r *Repository) ListActivePeerShareRuntimePorts(shareID int64, nodeID int64) ([]int, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
rows, err := r.db.Query(`SELECT port FROM peer_share_runtime WHERE share_id = ? AND node_id = ? AND status = 1 AND port > 0`, shareID, nodeID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
out := make([]int, 0)
|
|
for rows.Next() {
|
|
var port int
|
|
if err := rows.Scan(&port); err != nil {
|
|
return nil, err
|
|
}
|
|
if port > 0 {
|
|
out = append(out, port)
|
|
}
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
func (r *Repository) UpsertFederationTunnelBinding(item *FederationTunnelBinding) error {
|
|
if r == nil || r.db == nil {
|
|
return errors.New("repository not initialized")
|
|
}
|
|
if item == nil {
|
|
return errors.New("binding item is nil")
|
|
}
|
|
_, err := r.db.Exec(`
|
|
INSERT INTO federation_tunnel_binding(tunnel_id, node_id, chain_type, hop_inx, remote_url, resource_key, remote_binding_id, allocated_port, status, created_time, updated_time)
|
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
ON CONFLICT(tunnel_id, node_id, chain_type, hop_inx)
|
|
DO UPDATE SET
|
|
remote_url = excluded.remote_url,
|
|
resource_key = excluded.resource_key,
|
|
remote_binding_id = excluded.remote_binding_id,
|
|
allocated_port = excluded.allocated_port,
|
|
status = excluded.status,
|
|
updated_time = excluded.updated_time
|
|
`, item.TunnelID, item.NodeID, item.ChainType, item.HopInx, item.RemoteURL, item.ResourceKey, item.RemoteBindingID, item.AllocatedPort, item.Status, item.CreatedTime, item.UpdatedTime)
|
|
return err
|
|
}
|
|
|
|
func (r *Repository) ListActiveFederationTunnelBindingsByTunnel(tunnelID int64) ([]FederationTunnelBinding, error) {
|
|
if r == nil || r.db == nil {
|
|
return nil, errors.New("repository not initialized")
|
|
}
|
|
rows, err := r.db.Query(`
|
|
SELECT id, tunnel_id, node_id, chain_type, hop_inx, remote_url, resource_key, remote_binding_id, allocated_port, status, created_time, updated_time
|
|
FROM federation_tunnel_binding
|
|
WHERE tunnel_id = ? AND status = 1
|
|
ORDER BY chain_type ASC, hop_inx ASC, id ASC
|
|
`, tunnelID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
out := make([]FederationTunnelBinding, 0)
|
|
for rows.Next() {
|
|
var item FederationTunnelBinding
|
|
if err := rows.Scan(&item.ID, &item.TunnelID, &item.NodeID, &item.ChainType, &item.HopInx, &item.RemoteURL, &item.ResourceKey, &item.RemoteBindingID, &item.AllocatedPort, &item.Status, &item.CreatedTime, &item.UpdatedTime); err != nil {
|
|
return nil, err
|
|
}
|
|
out = append(out, item)
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
func (r *Repository) DeleteFederationTunnelBindingsByTunnel(tunnelID int64) error {
|
|
if r == nil || r.db == nil {
|
|
return errors.New("repository not initialized")
|
|
}
|
|
_, err := r.db.Exec(`DELETE FROM federation_tunnel_binding WHERE tunnel_id = ?`, tunnelID)
|
|
return err
|
|
}
|
|
|
|
var osMkdirAll = func(path string) error {
|
|
return os.MkdirAll(path, 0o755)
|
|
}
|
|
|
|
// ============ Backup/Export Data Structures ============
|
|
|
|
// BackupData represents the full backup structure
|
|
type BackupData struct {
|
|
Version string `json:"version"`
|
|
ExportedAt int64 `json:"exportedAt"`
|
|
Users []UserBackup `json:"users,omitempty"`
|
|
Nodes []NodeBackup `json:"nodes,omitempty"`
|
|
Tunnels []TunnelBackup `json:"tunnels,omitempty"`
|
|
Forwards []ForwardBackup `json:"forwards,omitempty"`
|
|
UserTunnels []UserTunnelBackup `json:"userTunnels,omitempty"`
|
|
SpeedLimits []SpeedLimitBackup `json:"speedLimits,omitempty"`
|
|
TunnelGroups []TunnelGroupBackup `json:"tunnelGroups,omitempty"`
|
|
UserGroups []UserGroupBackup `json:"userGroups,omitempty"`
|
|
Permissions []PermissionBackup `json:"permissions,omitempty"`
|
|
Configs map[string]string `json:"configs,omitempty"`
|
|
}
|
|
|
|
type UserBackup struct {
|
|
ID int64 `json:"id"`
|
|
User string `json:"user"`
|
|
Pwd string `json:"pwd"`
|
|
RoleID int `json:"roleId"`
|
|
ExpTime int64 `json:"expTime"`
|
|
Flow int64 `json:"flow"`
|
|
InFlow int64 `json:"inFlow"`
|
|
OutFlow int64 `json:"outFlow"`
|
|
FlowResetTime int64 `json:"flowResetTime"`
|
|
Num int `json:"num"`
|
|
CreatedTime int64 `json:"createdTime"`
|
|
UpdatedTime int64 `json:"updatedTime,omitempty"`
|
|
Status int `json:"status"`
|
|
}
|
|
|
|
type NodeBackup struct {
|
|
ID int64 `json:"id"`
|
|
Name string `json:"name"`
|
|
Secret string `json:"secret"`
|
|
ServerIP string `json:"serverIp"`
|
|
ServerIPv4 string `json:"serverIpV4,omitempty"`
|
|
ServerIPv6 string `json:"serverIpV6,omitempty"`
|
|
Port string `json:"port"`
|
|
InterfaceName string `json:"interfaceName,omitempty"`
|
|
Version string `json:"version,omitempty"`
|
|
HTTP int `json:"http"`
|
|
TLS int `json:"tls"`
|
|
Socks int `json:"socks"`
|
|
CreatedTime int64 `json:"createdTime"`
|
|
UpdatedTime int64 `json:"updatedTime,omitempty"`
|
|
Status int `json:"status"`
|
|
TCPListenAddr string `json:"tcpListenAddr"`
|
|
UDPListenAddr string `json:"udpListenAddr"`
|
|
Inx int `json:"inx"`
|
|
IsRemote int `json:"isRemote"`
|
|
RemoteURL string `json:"remoteUrl,omitempty"`
|
|
RemoteToken string `json:"remoteToken,omitempty"`
|
|
RemoteConfig string `json:"remoteConfig,omitempty"`
|
|
}
|
|
|
|
type TunnelBackup struct {
|
|
ID int64 `json:"id"`
|
|
Name string `json:"name"`
|
|
TrafficRatio float64 `json:"trafficRatio"`
|
|
Type int `json:"type"`
|
|
Protocol string `json:"protocol"`
|
|
Flow int64 `json:"flow"`
|
|
CreatedTime int64 `json:"createdTime"`
|
|
UpdatedTime int64 `json:"updatedTime"`
|
|
Status int `json:"status"`
|
|
InIP string `json:"inIp,omitempty"`
|
|
Inx int `json:"inx"`
|
|
ChainTunnels []ChainTunnelBackup `json:"chainTunnels,omitempty"`
|
|
}
|
|
|
|
type ChainTunnelBackup struct {
|
|
ID int64 `json:"id"`
|
|
TunnelID int64 `json:"tunnelId"`
|
|
ChainType string `json:"chainType"`
|
|
NodeID int64 `json:"nodeId"`
|
|
Port int `json:"port,omitempty"`
|
|
Strategy string `json:"strategy,omitempty"`
|
|
Inx int `json:"inx,omitempty"`
|
|
Protocol string `json:"protocol,omitempty"`
|
|
}
|
|
|
|
type ForwardBackup struct {
|
|
ID int64 `json:"id"`
|
|
UserID int64 `json:"userId"`
|
|
UserName string `json:"userName"`
|
|
Name string `json:"name"`
|
|
TunnelID int64 `json:"tunnelId"`
|
|
RemoteAddr string `json:"remoteAddr"`
|
|
Strategy string `json:"strategy"`
|
|
InFlow int64 `json:"inFlow"`
|
|
OutFlow int64 `json:"outFlow"`
|
|
CreatedTime int64 `json:"createdTime"`
|
|
UpdatedTime int64 `json:"updatedTime"`
|
|
Status int `json:"status"`
|
|
Inx int `json:"inx"`
|
|
ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"`
|
|
}
|
|
|
|
type ForwardPortBackup struct {
|
|
NodeID int64 `json:"nodeId"`
|
|
Port int `json:"port"`
|
|
}
|
|
|
|
type UserTunnelBackup struct {
|
|
ID int64 `json:"id"`
|
|
UserID int64 `json:"userId"`
|
|
TunnelID int64 `json:"tunnelId"`
|
|
SpeedID int64 `json:"speedId,omitempty"`
|
|
Num int `json:"num"`
|
|
Flow int64 `json:"flow"`
|
|
InFlow int64 `json:"inFlow"`
|
|
OutFlow int64 `json:"outFlow"`
|
|
FlowResetTime int64 `json:"flowResetTime"`
|
|
ExpTime int64 `json:"expTime"`
|
|
Status int `json:"status"`
|
|
}
|
|
|
|
type SpeedLimitBackup struct {
|
|
ID int64 `json:"id"`
|
|
Name string `json:"name"`
|
|
Speed int64 `json:"speed"`
|
|
TunnelID int64 `json:"tunnelId"`
|
|
TunnelName string `json:"tunnelName"`
|
|
CreatedTime int64 `json:"createdTime"`
|
|
UpdatedTime int64 `json:"updatedTime,omitempty"`
|
|
Status int `json:"status"`
|
|
}
|
|
|
|
type TunnelGroupBackup struct {
|
|
ID int64 `json:"id"`
|
|
Name string `json:"name"`
|
|
CreatedTime int64 `json:"createdTime"`
|
|
UpdatedTime int64 `json:"updatedTime"`
|
|
Status int `json:"status"`
|
|
Tunnels []int64 `json:"tunnels,omitempty"`
|
|
}
|
|
|
|
type UserGroupBackup struct {
|
|
ID int64 `json:"id"`
|
|
Name string `json:"name"`
|
|
CreatedTime int64 `json:"createdTime"`
|
|
UpdatedTime int64 `json:"updatedTime"`
|
|
Status int `json:"status"`
|
|
Users []int64 `json:"users,omitempty"`
|
|
}
|
|
|
|
type PermissionBackup struct {
|
|
ID int64 `json:"id"`
|
|
UserGroupID int64 `json:"userGroupId"`
|
|
TunnelGroupID int64 `json:"tunnelGroupId"`
|
|
CreatedTime int64 `json:"createdTime"`
|
|
CreatedByGroup int `json:"createdByGroup"`
|
|
Grants []PermissionGrantBackup `json:"grants,omitempty"`
|
|
}
|
|
|
|
type PermissionGrantBackup struct {
|
|
ID int64 `json:"id"`
|
|
UserGroupID int64 `json:"userGroupId"`
|
|
TunnelGroupID int64 `json:"tunnelGroupId"`
|
|
UserTunnelID int64 `json:"userTunnelId"`
|
|
CreatedTime int64 `json:"createdTime"`
|
|
CreatedByGroup int `json:"createdByGroup"`
|
|
}
|
|
|
|
// ============ Export Methods ============
|
|
|
|
// ExportAll exports all data as BackupData
|
|
func (r *Repository) ExportAll() (*BackupData, error) {
|
|
backup := &BackupData{
|
|
Version: "1.0",
|
|
ExportedAt: unixMilliNow(),
|
|
}
|
|
|
|
// Export all data types
|
|
users, err := r.exportUsers()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export users failed: %w", err)
|
|
}
|
|
backup.Users = users
|
|
|
|
nodes, err := r.exportNodes()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export nodes failed: %w", err)
|
|
}
|
|
backup.Nodes = nodes
|
|
|
|
tunnels, err := r.exportTunnels()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export tunnels failed: %w", err)
|
|
}
|
|
backup.Tunnels = tunnels
|
|
|
|
forwards, err := r.exportForwards()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export forwards failed: %w", err)
|
|
}
|
|
backup.Forwards = forwards
|
|
|
|
userTunnels, err := r.exportUserTunnels()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export user tunnels failed: %w", err)
|
|
}
|
|
backup.UserTunnels = userTunnels
|
|
|
|
speedLimits, err := r.exportSpeedLimits()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export speed limits failed: %w", err)
|
|
}
|
|
backup.SpeedLimits = speedLimits
|
|
|
|
tunnelGroups, err := r.exportTunnelGroups()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export tunnel groups failed: %w", err)
|
|
}
|
|
backup.TunnelGroups = tunnelGroups
|
|
|
|
userGroups, err := r.exportUserGroups()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export user groups failed: %w", err)
|
|
}
|
|
backup.UserGroups = userGroups
|
|
|
|
permissions, err := r.exportPermissions()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export permissions failed: %w", err)
|
|
}
|
|
backup.Permissions = permissions
|
|
|
|
configs, err := r.ListConfigs()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export configs failed: %w", err)
|
|
}
|
|
backup.Configs = configs
|
|
|
|
return backup, nil
|
|
}
|
|
|
|
// ExportPartial exports selected data types
|
|
func (r *Repository) ExportPartial(types []string) (*BackupData, error) {
|
|
backup := &BackupData{
|
|
Version: "1.0",
|
|
ExportedAt: unixMilliNow(),
|
|
}
|
|
|
|
typeSet := make(map[string]bool)
|
|
for _, t := range types {
|
|
typeSet[t] = true
|
|
}
|
|
|
|
if typeSet["users"] {
|
|
users, err := r.exportUsers()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export users failed: %w", err)
|
|
}
|
|
backup.Users = users
|
|
}
|
|
if typeSet["nodes"] {
|
|
nodes, err := r.exportNodes()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export nodes failed: %w", err)
|
|
}
|
|
backup.Nodes = nodes
|
|
}
|
|
if typeSet["tunnels"] {
|
|
tunnels, err := r.exportTunnels()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export tunnels failed: %w", err)
|
|
}
|
|
backup.Tunnels = tunnels
|
|
}
|
|
if typeSet["forwards"] {
|
|
forwards, err := r.exportForwards()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export forwards failed: %w", err)
|
|
}
|
|
backup.Forwards = forwards
|
|
}
|
|
if typeSet["userTunnels"] {
|
|
userTunnels, err := r.exportUserTunnels()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export user tunnels failed: %w", err)
|
|
}
|
|
backup.UserTunnels = userTunnels
|
|
}
|
|
if typeSet["speedLimits"] {
|
|
speedLimits, err := r.exportSpeedLimits()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export speed limits failed: %w", err)
|
|
}
|
|
backup.SpeedLimits = speedLimits
|
|
}
|
|
if typeSet["tunnelGroups"] {
|
|
tunnelGroups, err := r.exportTunnelGroups()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export tunnel groups failed: %w", err)
|
|
}
|
|
backup.TunnelGroups = tunnelGroups
|
|
}
|
|
if typeSet["userGroups"] {
|
|
userGroups, err := r.exportUserGroups()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export user groups failed: %w", err)
|
|
}
|
|
backup.UserGroups = userGroups
|
|
}
|
|
if typeSet["permissions"] {
|
|
permissions, err := r.exportPermissions()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export permissions failed: %w", err)
|
|
}
|
|
backup.Permissions = permissions
|
|
}
|
|
if typeSet["configs"] {
|
|
configs, err := r.ListConfigs()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("export configs failed: %w", err)
|
|
}
|
|
backup.Configs = configs
|
|
}
|
|
|
|
return backup, nil
|
|
}
|
|
|
|
func (r *Repository) exportUsers() ([]UserBackup, error) {
|
|
rows, err := r.db.Query(`
|
|
SELECT id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status
|
|
FROM user ORDER BY id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var users []UserBackup
|
|
for rows.Next() {
|
|
var u UserBackup
|
|
var updatedTime sql.NullInt64
|
|
if err := rows.Scan(&u.ID, &u.User, &u.Pwd, &u.RoleID, &u.ExpTime, &u.Flow, &u.InFlow, &u.OutFlow, &u.FlowResetTime, &u.Num, &u.CreatedTime, &updatedTime, &u.Status); err != nil {
|
|
return nil, err
|
|
}
|
|
if updatedTime.Valid {
|
|
u.UpdatedTime = updatedTime.Int64
|
|
}
|
|
users = append(users, u)
|
|
}
|
|
return users, rows.Err()
|
|
}
|
|
|
|
func (r *Repository) exportNodes() ([]NodeBackup, error) {
|
|
rows, err := r.db.Query(`
|
|
SELECT id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config
|
|
FROM node ORDER BY inx ASC, id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var nodes []NodeBackup
|
|
for rows.Next() {
|
|
var n NodeBackup
|
|
var updatedTime sql.NullInt64
|
|
var serverIPv4, serverIPv6, interfaceName, version, remoteURL, remoteToken, remoteConfig sql.NullString
|
|
if err := rows.Scan(&n.ID, &n.Name, &n.Secret, &n.ServerIP, &serverIPv4, &serverIPv6, &n.Port, &interfaceName, &version, &n.HTTP, &n.TLS, &n.Socks, &n.CreatedTime, &updatedTime, &n.Status, &n.TCPListenAddr, &n.UDPListenAddr, &n.Inx, &n.IsRemote, &remoteURL, &remoteToken, &remoteConfig); err != nil {
|
|
return nil, err
|
|
}
|
|
if updatedTime.Valid {
|
|
n.UpdatedTime = updatedTime.Int64
|
|
}
|
|
if serverIPv4.Valid {
|
|
n.ServerIPv4 = serverIPv4.String
|
|
}
|
|
if serverIPv6.Valid {
|
|
n.ServerIPv6 = serverIPv6.String
|
|
}
|
|
if interfaceName.Valid {
|
|
n.InterfaceName = interfaceName.String
|
|
}
|
|
if version.Valid {
|
|
n.Version = version.String
|
|
}
|
|
if remoteURL.Valid {
|
|
n.RemoteURL = remoteURL.String
|
|
}
|
|
if remoteToken.Valid {
|
|
n.RemoteToken = remoteToken.String
|
|
}
|
|
if remoteConfig.Valid {
|
|
n.RemoteConfig = remoteConfig.String
|
|
}
|
|
nodes = append(nodes, n)
|
|
}
|
|
return nodes, rows.Err()
|
|
}
|
|
|
|
func (r *Repository) exportTunnels() ([]TunnelBackup, error) {
|
|
rows, err := r.db.Query(`
|
|
SELECT id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx
|
|
FROM tunnel ORDER BY inx ASC, id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var tunnels []TunnelBackup
|
|
for rows.Next() {
|
|
var t TunnelBackup
|
|
var protocol sql.NullString
|
|
var updatedTime sql.NullInt64
|
|
var inIP sql.NullString
|
|
var inx sql.NullInt64
|
|
if err := rows.Scan(&t.ID, &t.Name, &t.TrafficRatio, &t.Type, &protocol, &t.Flow, &t.CreatedTime, &updatedTime, &t.Status, &inIP, &inx); err != nil {
|
|
return nil, err
|
|
}
|
|
if protocol.Valid {
|
|
t.Protocol = protocol.String
|
|
}
|
|
if updatedTime.Valid {
|
|
t.UpdatedTime = updatedTime.Int64
|
|
}
|
|
if inIP.Valid {
|
|
t.InIP = inIP.String
|
|
}
|
|
if inx.Valid {
|
|
t.Inx = int(inx.Int64)
|
|
}
|
|
// Export chain tunnels
|
|
chainTunnels, err := r.exportChainTunnels(t.ID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
t.ChainTunnels = chainTunnels
|
|
tunnels = append(tunnels, t)
|
|
}
|
|
return tunnels, rows.Err()
|
|
}
|
|
|
|
func (r *Repository) exportChainTunnels(tunnelID int64) ([]ChainTunnelBackup, error) {
|
|
rows, err := r.db.Query(`
|
|
SELECT id, tunnel_id, chain_type, node_id, port, strategy, inx, protocol
|
|
FROM chain_tunnel WHERE tunnel_id = ? ORDER BY inx ASC, id ASC
|
|
`, tunnelID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var chainTunnels []ChainTunnelBackup
|
|
for rows.Next() {
|
|
var ct ChainTunnelBackup
|
|
var port sql.NullInt64
|
|
var strategy, protocol sql.NullString
|
|
var inx sql.NullInt64
|
|
if err := rows.Scan(&ct.ID, &ct.TunnelID, &ct.ChainType, &ct.NodeID, &port, &strategy, &inx, &protocol); err != nil {
|
|
return nil, err
|
|
}
|
|
if port.Valid {
|
|
ct.Port = int(port.Int64)
|
|
}
|
|
if strategy.Valid {
|
|
ct.Strategy = strategy.String
|
|
}
|
|
if inx.Valid {
|
|
ct.Inx = int(inx.Int64)
|
|
}
|
|
if protocol.Valid {
|
|
ct.Protocol = protocol.String
|
|
}
|
|
chainTunnels = append(chainTunnels, ct)
|
|
}
|
|
return chainTunnels, rows.Err()
|
|
}
|
|
|
|
func (r *Repository) exportForwards() ([]ForwardBackup, error) {
|
|
rows, err := r.db.Query(`
|
|
SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx
|
|
FROM forward ORDER BY id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var forwards []ForwardBackup
|
|
for rows.Next() {
|
|
var f ForwardBackup
|
|
var strategy sql.NullString
|
|
var updatedTime sql.NullInt64
|
|
var inx sql.NullInt64
|
|
if err := rows.Scan(&f.ID, &f.UserID, &f.UserName, &f.Name, &f.TunnelID, &f.RemoteAddr, &strategy, &f.InFlow, &f.OutFlow, &f.CreatedTime, &updatedTime, &f.Status, &inx); err != nil {
|
|
return nil, err
|
|
}
|
|
if strategy.Valid {
|
|
f.Strategy = strategy.String
|
|
}
|
|
if updatedTime.Valid {
|
|
f.UpdatedTime = updatedTime.Int64
|
|
}
|
|
if inx.Valid {
|
|
f.Inx = int(inx.Int64)
|
|
}
|
|
|
|
forwardPorts, err := r.exportForwardPorts(f.ID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
portsCopy := append([]ForwardPortBackup(nil), forwardPorts...)
|
|
f.ForwardPorts = &portsCopy
|
|
|
|
forwards = append(forwards, f)
|
|
}
|
|
return forwards, rows.Err()
|
|
}
|
|
|
|
func (r *Repository) exportForwardPorts(forwardID int64) ([]ForwardPortBackup, error) {
|
|
rows, err := r.db.Query(`
|
|
SELECT node_id, port
|
|
FROM forward_port
|
|
WHERE forward_id = ?
|
|
ORDER BY id ASC
|
|
`, forwardID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
ports := make([]ForwardPortBackup, 0)
|
|
for rows.Next() {
|
|
var fp ForwardPortBackup
|
|
if err := rows.Scan(&fp.NodeID, &fp.Port); err != nil {
|
|
return nil, err
|
|
}
|
|
ports = append(ports, fp)
|
|
}
|
|
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return ports, nil
|
|
}
|
|
|
|
func (r *Repository) exportUserTunnels() ([]UserTunnelBackup, error) {
|
|
rows, err := r.db.Query(`
|
|
SELECT id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status
|
|
FROM user_tunnel ORDER BY id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var userTunnels []UserTunnelBackup
|
|
for rows.Next() {
|
|
var ut UserTunnelBackup
|
|
var speedID sql.NullInt64
|
|
if err := rows.Scan(&ut.ID, &ut.UserID, &ut.TunnelID, &speedID, &ut.Num, &ut.Flow, &ut.InFlow, &ut.OutFlow, &ut.FlowResetTime, &ut.ExpTime, &ut.Status); err != nil {
|
|
return nil, err
|
|
}
|
|
if speedID.Valid {
|
|
ut.SpeedID = speedID.Int64
|
|
}
|
|
userTunnels = append(userTunnels, ut)
|
|
}
|
|
return userTunnels, rows.Err()
|
|
}
|
|
|
|
func (r *Repository) exportSpeedLimits() ([]SpeedLimitBackup, error) {
|
|
rows, err := r.db.Query(`
|
|
SELECT id, name, speed, tunnel_id, tunnel_name, created_time, updated_time, status
|
|
FROM speed_limit ORDER BY id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var speedLimits []SpeedLimitBackup
|
|
for rows.Next() {
|
|
var sl SpeedLimitBackup
|
|
var updatedTime sql.NullInt64
|
|
if err := rows.Scan(&sl.ID, &sl.Name, &sl.Speed, &sl.TunnelID, &sl.TunnelName, &sl.CreatedTime, &updatedTime, &sl.Status); err != nil {
|
|
return nil, err
|
|
}
|
|
if updatedTime.Valid {
|
|
sl.UpdatedTime = updatedTime.Int64
|
|
}
|
|
speedLimits = append(speedLimits, sl)
|
|
}
|
|
return speedLimits, rows.Err()
|
|
}
|
|
|
|
func (r *Repository) exportTunnelGroups() ([]TunnelGroupBackup, error) {
|
|
rows, err := r.db.Query(`
|
|
SELECT id, name, created_time, updated_time, status
|
|
FROM tunnel_group ORDER BY id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var groups []TunnelGroupBackup
|
|
for rows.Next() {
|
|
var tg TunnelGroupBackup
|
|
if err := rows.Scan(&tg.ID, &tg.Name, &tg.CreatedTime, &tg.UpdatedTime, &tg.Status); err != nil {
|
|
return nil, err
|
|
}
|
|
// Get tunnel IDs for this group
|
|
tunnelRows, err := r.db.Query(`SELECT tunnel_id FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, tg.ID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
for tunnelRows.Next() {
|
|
var tunnelID int64
|
|
if err := tunnelRows.Scan(&tunnelID); err != nil {
|
|
tunnelRows.Close()
|
|
return nil, err
|
|
}
|
|
tg.Tunnels = append(tg.Tunnels, tunnelID)
|
|
}
|
|
tunnelRows.Close()
|
|
groups = append(groups, tg)
|
|
}
|
|
return groups, rows.Err()
|
|
}
|
|
|
|
func (r *Repository) exportUserGroups() ([]UserGroupBackup, error) {
|
|
rows, err := r.db.Query(`
|
|
SELECT id, name, created_time, updated_time, status
|
|
FROM user_group ORDER BY id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var groups []UserGroupBackup
|
|
for rows.Next() {
|
|
var ug UserGroupBackup
|
|
if err := rows.Scan(&ug.ID, &ug.Name, &ug.CreatedTime, &ug.UpdatedTime, &ug.Status); err != nil {
|
|
return nil, err
|
|
}
|
|
// Get user IDs for this group
|
|
userRows, err := r.db.Query(`SELECT user_id FROM user_group_user WHERE user_group_id = ?`, ug.ID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
for userRows.Next() {
|
|
var userID int64
|
|
if err := userRows.Scan(&userID); err != nil {
|
|
userRows.Close()
|
|
return nil, err
|
|
}
|
|
ug.Users = append(ug.Users, userID)
|
|
}
|
|
userRows.Close()
|
|
groups = append(groups, ug)
|
|
}
|
|
return groups, rows.Err()
|
|
}
|
|
|
|
func (r *Repository) exportPermissions() ([]PermissionBackup, error) {
|
|
rows, err := r.db.Query(`
|
|
SELECT id, user_group_id, tunnel_group_id, created_time
|
|
FROM group_permission ORDER BY id ASC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var permissions []PermissionBackup
|
|
for rows.Next() {
|
|
var p PermissionBackup
|
|
if err := rows.Scan(&p.ID, &p.UserGroupID, &p.TunnelGroupID, &p.CreatedTime); err != nil {
|
|
return nil, err
|
|
}
|
|
p.CreatedByGroup = 0
|
|
// Get grants for this permission
|
|
grantRows, err := r.db.Query(`SELECT id, user_group_id, tunnel_group_id, user_tunnel_id, created_time, created_by_group FROM group_permission_grant WHERE user_group_id = ? AND tunnel_group_id = ?`, p.UserGroupID, p.TunnelGroupID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
for grantRows.Next() {
|
|
var g PermissionGrantBackup
|
|
if err := grantRows.Scan(&g.ID, &g.UserGroupID, &g.TunnelGroupID, &g.UserTunnelID, &g.CreatedTime, &g.CreatedByGroup); err != nil {
|
|
grantRows.Close()
|
|
return nil, err
|
|
}
|
|
p.Grants = append(p.Grants, g)
|
|
}
|
|
grantRows.Close()
|
|
permissions = append(permissions, p)
|
|
}
|
|
return permissions, rows.Err()
|
|
}
|
|
|
|
// ============ Import Methods ============
|
|
|
|
// ImportResult contains the result of an import operation
|
|
type ImportResult struct {
|
|
UsersImported int `json:"usersImported"`
|
|
NodesImported int `json:"nodesImported"`
|
|
TunnelsImported int `json:"tunnelsImported"`
|
|
ForwardsImported int `json:"forwardsImported"`
|
|
UserTunnelsImported int `json:"userTunnelsImported"`
|
|
SpeedLimitsImported int `json:"speedLimitsImported"`
|
|
TunnelGroupsImported int `json:"tunnelGroupsImported"`
|
|
UserGroupsImported int `json:"userGroupsImported"`
|
|
PermissionsImported int `json:"permissionsImported"`
|
|
ConfigsImported int `json:"configsImported"`
|
|
AutoBackup *BackupData `json:"autoBackup,omitempty"`
|
|
}
|
|
|
|
// Import imports data from BackupData with transaction support
|
|
func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult, error) {
|
|
result := &ImportResult{}
|
|
|
|
typeSet := make(map[string]bool)
|
|
for _, t := range types {
|
|
typeSet[t] = true
|
|
}
|
|
|
|
tx, err := r.db.Begin()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to begin transaction: %w", err)
|
|
}
|
|
defer func() { _ = tx.Rollback() }()
|
|
|
|
now := unixMilliNow()
|
|
|
|
if typeSet["users"] && len(backup.Users) > 0 {
|
|
count, err := r.importUsers(tx, backup.Users, now)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("import users failed: %w", err)
|
|
}
|
|
result.UsersImported = count
|
|
}
|
|
|
|
if typeSet["nodes"] && len(backup.Nodes) > 0 {
|
|
count, err := r.importNodes(tx, backup.Nodes, now)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("import nodes failed: %w", err)
|
|
}
|
|
result.NodesImported = count
|
|
}
|
|
|
|
if typeSet["tunnels"] && len(backup.Tunnels) > 0 {
|
|
count, err := r.importTunnels(tx, backup.Tunnels, now)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("import tunnels failed: %w", err)
|
|
}
|
|
result.TunnelsImported = count
|
|
}
|
|
|
|
if typeSet["forwards"] && len(backup.Forwards) > 0 {
|
|
count, err := r.importForwards(tx, backup.Forwards, now)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("import forwards failed: %w", err)
|
|
}
|
|
result.ForwardsImported = count
|
|
}
|
|
|
|
if typeSet["userTunnels"] && len(backup.UserTunnels) > 0 {
|
|
count, err := r.importUserTunnels(tx, backup.UserTunnels, now)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("import user tunnels failed: %w", err)
|
|
}
|
|
result.UserTunnelsImported = count
|
|
}
|
|
|
|
if typeSet["speedLimits"] && len(backup.SpeedLimits) > 0 {
|
|
count, err := r.importSpeedLimits(tx, backup.SpeedLimits, now)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("import speed limits failed: %w", err)
|
|
}
|
|
result.SpeedLimitsImported = count
|
|
}
|
|
|
|
if typeSet["tunnelGroups"] && len(backup.TunnelGroups) > 0 {
|
|
count, err := r.importTunnelGroups(tx, backup.TunnelGroups, now)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("import tunnel groups failed: %w", err)
|
|
}
|
|
result.TunnelGroupsImported = count
|
|
}
|
|
|
|
if typeSet["userGroups"] && len(backup.UserGroups) > 0 {
|
|
count, err := r.importUserGroups(tx, backup.UserGroups, now)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("import user groups failed: %w", err)
|
|
}
|
|
result.UserGroupsImported = count
|
|
}
|
|
|
|
if typeSet["permissions"] && len(backup.Permissions) > 0 {
|
|
count, err := r.importPermissions(tx, backup.Permissions, now)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("import permissions failed: %w", err)
|
|
}
|
|
result.PermissionsImported = count
|
|
}
|
|
|
|
if typeSet["configs"] && len(backup.Configs) > 0 {
|
|
count, err := r.importConfigs(tx, backup.Configs, now)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("import configs failed: %w", err)
|
|
}
|
|
result.ConfigsImported = count
|
|
}
|
|
|
|
if err := tx.Commit(); err != nil {
|
|
return nil, fmt.Errorf("failed to commit transaction: %w", err)
|
|
}
|
|
|
|
return result, nil
|
|
}
|
|
|
|
func (r *Repository) importUsers(db Execer, users []UserBackup, now int64) (int, error) {
|
|
count := 0
|
|
for _, u := range users {
|
|
_, err := db.Exec(`
|
|
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
ON CONFLICT(id) DO UPDATE SET
|
|
user = excluded.user,
|
|
pwd = excluded.pwd,
|
|
role_id = excluded.role_id,
|
|
exp_time = excluded.exp_time,
|
|
flow = excluded.flow,
|
|
in_flow = excluded.in_flow,
|
|
out_flow = excluded.out_flow,
|
|
flow_reset_time = excluded.flow_reset_time,
|
|
num = excluded.num,
|
|
updated_time = excluded.updated_time,
|
|
status = excluded.status
|
|
`, u.ID, u.User, u.Pwd, u.RoleID, u.ExpTime, u.Flow, u.InFlow, u.OutFlow, u.FlowResetTime, u.Num, u.CreatedTime, now, u.Status)
|
|
if err != nil {
|
|
return count, err
|
|
}
|
|
count++
|
|
}
|
|
return count, nil
|
|
}
|
|
|
|
func (r *Repository) UsernameExists(username string) (bool, error) {
|
|
var count int
|
|
err := r.db.QueryRow(`SELECT COUNT(1) FROM user WHERE user = ?`, username).Scan(&count)
|
|
if err != nil {
|
|
return false, err
|
|
}
|
|
return count > 0, nil
|
|
}
|
|
|
|
func (r *Repository) importNodes(db Execer, nodes []NodeBackup, now int64) (int, error) {
|
|
count := 0
|
|
for _, n := range nodes {
|
|
_, err := db.Exec(`
|
|
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config)
|
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
ON CONFLICT(id) DO UPDATE SET
|
|
name = excluded.name,
|
|
secret = excluded.secret,
|
|
server_ip = excluded.server_ip,
|
|
server_ip_v4 = excluded.server_ip_v4,
|
|
server_ip_v6 = excluded.server_ip_v6,
|
|
port = excluded.port,
|
|
interface_name = excluded.interface_name,
|
|
version = excluded.version,
|
|
http = excluded.http,
|
|
tls = excluded.tls,
|
|
socks = excluded.socks,
|
|
updated_time = excluded.updated_time,
|
|
status = excluded.status,
|
|
tcp_listen_addr = excluded.tcp_listen_addr,
|
|
udp_listen_addr = excluded.udp_listen_addr,
|
|
inx = excluded.inx,
|
|
is_remote = excluded.is_remote,
|
|
remote_url = excluded.remote_url,
|
|
remote_token = excluded.remote_token,
|
|
remote_config = excluded.remote_config
|
|
`, n.ID, n.Name, n.Secret, n.ServerIP, n.ServerIPv4, n.ServerIPv6, n.Port, n.InterfaceName, n.Version, n.HTTP, n.TLS, n.Socks, n.CreatedTime, now, n.Status, n.TCPListenAddr, n.UDPListenAddr, n.Inx, n.IsRemote, n.RemoteURL, n.RemoteToken, n.RemoteConfig)
|
|
if err != nil {
|
|
return count, err
|
|
}
|
|
count++
|
|
}
|
|
return count, nil
|
|
}
|
|
|
|
func (r *Repository) importTunnels(db Execer, tunnels []TunnelBackup, now int64) (int, error) {
|
|
count := 0
|
|
for _, t := range tunnels {
|
|
_, err := db.Exec(`
|
|
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
ON CONFLICT(id) DO UPDATE SET
|
|
name = excluded.name,
|
|
traffic_ratio = excluded.traffic_ratio,
|
|
type = excluded.type,
|
|
protocol = excluded.protocol,
|
|
flow = excluded.flow,
|
|
updated_time = excluded.updated_time,
|
|
status = excluded.status,
|
|
in_ip = excluded.in_ip,
|
|
inx = excluded.inx
|
|
`, t.ID, t.Name, t.TrafficRatio, t.Type, t.Protocol, t.Flow, t.CreatedTime, now, t.Status, t.InIP, t.Inx)
|
|
if err != nil {
|
|
return count, err
|
|
}
|
|
if len(t.ChainTunnels) > 0 {
|
|
for _, ct := range t.ChainTunnels {
|
|
_, err = db.Exec(`
|
|
INSERT INTO chain_tunnel(id, tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?)
|
|
ON CONFLICT(id) DO UPDATE SET
|
|
chain_type = excluded.chain_type,
|
|
node_id = excluded.node_id,
|
|
port = excluded.port,
|
|
strategy = excluded.strategy,
|
|
inx = excluded.inx,
|
|
protocol = excluded.protocol
|
|
`, ct.ID, ct.TunnelID, ct.ChainType, ct.NodeID, ct.Port, ct.Strategy, ct.Inx, ct.Protocol)
|
|
if err != nil {
|
|
return count, err
|
|
}
|
|
}
|
|
}
|
|
count++
|
|
}
|
|
return count, nil
|
|
}
|
|
|
|
func (r *Repository) importForwards(db Execer, forwards []ForwardBackup, now int64) (int, error) {
|
|
count := 0
|
|
for _, f := range forwards {
|
|
_, err := db.Exec(`
|
|
INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
ON CONFLICT(id) DO UPDATE SET
|
|
user_id = excluded.user_id,
|
|
user_name = excluded.user_name,
|
|
name = excluded.name,
|
|
tunnel_id = excluded.tunnel_id,
|
|
remote_addr = excluded.remote_addr,
|
|
strategy = excluded.strategy,
|
|
in_flow = excluded.in_flow,
|
|
out_flow = excluded.out_flow,
|
|
updated_time = excluded.updated_time,
|
|
status = excluded.status,
|
|
inx = excluded.inx
|
|
`, f.ID, f.UserID, f.UserName, f.Name, f.TunnelID, f.RemoteAddr, f.Strategy, f.InFlow, f.OutFlow, f.CreatedTime, now, f.Status, f.Inx)
|
|
if err != nil {
|
|
return count, err
|
|
}
|
|
|
|
if f.ForwardPorts != nil {
|
|
if _, err := db.Exec(`DELETE FROM forward_port WHERE forward_id = ?`, f.ID); err != nil {
|
|
return count, err
|
|
}
|
|
for _, fp := range *f.ForwardPorts {
|
|
if _, err := db.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, f.ID, fp.NodeID, fp.Port); err != nil {
|
|
return count, err
|
|
}
|
|
}
|
|
}
|
|
|
|
count++
|
|
}
|
|
return count, nil
|
|
}
|
|
|
|
func (r *Repository) importUserTunnels(db Execer, userTunnels []UserTunnelBackup, now int64) (int, error) {
|
|
count := 0
|
|
for _, ut := range userTunnels {
|
|
var speedID interface{}
|
|
if ut.SpeedID > 0 {
|
|
speedID = ut.SpeedID
|
|
}
|
|
_, err := db.Exec(`
|
|
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
ON CONFLICT(id) DO UPDATE SET
|
|
user_id = excluded.user_id,
|
|
tunnel_id = excluded.tunnel_id,
|
|
speed_id = excluded.speed_id,
|
|
num = excluded.num,
|
|
flow = excluded.flow,
|
|
in_flow = excluded.in_flow,
|
|
out_flow = excluded.out_flow,
|
|
flow_reset_time = excluded.flow_reset_time,
|
|
exp_time = excluded.exp_time,
|
|
status = excluded.status
|
|
`, ut.ID, ut.UserID, ut.TunnelID, speedID, ut.Num, ut.Flow, ut.InFlow, ut.OutFlow, ut.FlowResetTime, ut.ExpTime, ut.Status)
|
|
if err != nil {
|
|
return count, err
|
|
}
|
|
count++
|
|
}
|
|
return count, nil
|
|
}
|
|
|
|
func (r *Repository) importSpeedLimits(db Execer, speedLimits []SpeedLimitBackup, now int64) (int, error) {
|
|
count := 0
|
|
for _, sl := range speedLimits {
|
|
_, err := db.Exec(`
|
|
INSERT INTO speed_limit(id, name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?)
|
|
ON CONFLICT(id) DO UPDATE SET
|
|
name = excluded.name,
|
|
speed = excluded.speed,
|
|
tunnel_id = excluded.tunnel_id,
|
|
tunnel_name = excluded.tunnel_name,
|
|
updated_time = excluded.updated_time,
|
|
status = excluded.status
|
|
`, sl.ID, sl.Name, sl.Speed, sl.TunnelID, sl.TunnelName, sl.CreatedTime, now, sl.Status)
|
|
if err != nil {
|
|
return count, err
|
|
}
|
|
count++
|
|
}
|
|
return count, nil
|
|
}
|
|
|
|
func (r *Repository) importTunnelGroups(db Execer, tunnelGroups []TunnelGroupBackup, now int64) (int, error) {
|
|
count := 0
|
|
for _, tg := range tunnelGroups {
|
|
_, err := db.Exec(`
|
|
INSERT INTO tunnel_group(id, name, created_time, updated_time, status)
|
|
VALUES(?, ?, ?, ?, ?)
|
|
ON CONFLICT(id) DO UPDATE SET
|
|
name = excluded.name,
|
|
updated_time = excluded.updated_time,
|
|
status = excluded.status
|
|
`, tg.ID, tg.Name, tg.CreatedTime, now, tg.Status)
|
|
if err != nil {
|
|
return count, err
|
|
}
|
|
_, err = db.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, tg.ID)
|
|
if err != nil {
|
|
return count, err
|
|
}
|
|
for _, tunnelID := range tg.Tunnels {
|
|
_, err = db.Exec(`
|
|
INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time)
|
|
VALUES(?, ?, ?)
|
|
`, tg.ID, tunnelID, now)
|
|
if err != nil {
|
|
return count, err
|
|
}
|
|
}
|
|
count++
|
|
}
|
|
return count, nil
|
|
}
|
|
|
|
func (r *Repository) importUserGroups(db Execer, userGroups []UserGroupBackup, now int64) (int, error) {
|
|
count := 0
|
|
for _, ug := range userGroups {
|
|
_, err := db.Exec(`
|
|
INSERT INTO user_group(id, name, created_time, updated_time, status)
|
|
VALUES(?, ?, ?, ?, ?)
|
|
ON CONFLICT(id) DO UPDATE SET
|
|
name = excluded.name,
|
|
updated_time = excluded.updated_time,
|
|
status = excluded.status
|
|
`, ug.ID, ug.Name, ug.CreatedTime, now, ug.Status)
|
|
if err != nil {
|
|
return count, err
|
|
}
|
|
_, err = db.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, ug.ID)
|
|
if err != nil {
|
|
return count, err
|
|
}
|
|
for _, userID := range ug.Users {
|
|
_, err = db.Exec(`
|
|
INSERT INTO user_group_user(user_group_id, user_id, created_time)
|
|
VALUES(?, ?, ?)
|
|
`, ug.ID, userID, now)
|
|
if err != nil {
|
|
return count, err
|
|
}
|
|
}
|
|
count++
|
|
}
|
|
return count, nil
|
|
}
|
|
|
|
func (r *Repository) importPermissions(db Execer, permissions []PermissionBackup, now int64) (int, error) {
|
|
count := 0
|
|
for _, p := range permissions {
|
|
_, err := db.Exec(`
|
|
INSERT INTO group_permission(id, user_group_id, tunnel_group_id, created_time, created_by_group)
|
|
VALUES(?, ?, ?, ?, ?)
|
|
ON CONFLICT(id) DO UPDATE SET
|
|
user_group_id = excluded.user_group_id,
|
|
tunnel_group_id = excluded.tunnel_group_id,
|
|
created_by_group = excluded.created_by_group
|
|
`, p.ID, p.UserGroupID, p.TunnelGroupID, p.CreatedTime, p.CreatedByGroup)
|
|
if err != nil {
|
|
return count, err
|
|
}
|
|
for _, g := range p.Grants {
|
|
_, err = db.Exec(`
|
|
INSERT INTO group_permission_grant(id, user_group_id, tunnel_group_id, user_tunnel_id, created_time, created_by_group)
|
|
VALUES(?, ?, ?, ?, ?, ?)
|
|
ON CONFLICT(id) DO UPDATE SET
|
|
user_tunnel_id = excluded.user_tunnel_id,
|
|
created_by_group = excluded.created_by_group
|
|
`, g.ID, g.UserGroupID, g.TunnelGroupID, g.UserTunnelID, g.CreatedTime, g.CreatedByGroup)
|
|
if err != nil {
|
|
return count, err
|
|
}
|
|
}
|
|
count++
|
|
}
|
|
return count, nil
|
|
}
|
|
|
|
func (r *Repository) importConfigs(db Execer, configs map[string]string, now int64) (int, error) {
|
|
count := 0
|
|
for name, value := range configs {
|
|
err := r.UpsertConfig(name, value, now)
|
|
if err != nil {
|
|
return count, err
|
|
}
|
|
count++
|
|
}
|
|
return count, nil
|
|
}
|