429 lines
11 KiB
Go
429 lines
11 KiB
Go
package db
|
|
|
|
import (
|
|
"database/sql"
|
|
|
|
"github.com/google/uuid"
|
|
"github.com/lib/pq"
|
|
|
|
"github.com/orohi/vpn-panel/internal/models"
|
|
)
|
|
|
|
func ListProtocolsDetailed(db *sql.DB) ([]models.Protocol, error) {
|
|
list, err := ListProtocols(db)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if len(list) == 0 {
|
|
return list, nil
|
|
}
|
|
|
|
rows, err := db.Query(`
|
|
SELECT np.protocol_id, np.installed, np.enabled,
|
|
n.id, n.name, n.host, n.status
|
|
FROM node_protocols np
|
|
JOIN nodes n ON n.id = np.node_id
|
|
WHERE np.installed = TRUE OR np.enabled = TRUE
|
|
ORDER BY n.name ASC`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
byID := map[uuid.UUID]*models.Protocol{}
|
|
for i := range list {
|
|
byID[list[i].ID] = &list[i]
|
|
}
|
|
|
|
for rows.Next() {
|
|
var protocolID, nodeID uuid.UUID
|
|
var installed, enabled bool
|
|
var name, host, status string
|
|
if err := rows.Scan(&protocolID, &installed, &enabled, &nodeID, &name, &host, &status); err != nil {
|
|
return nil, err
|
|
}
|
|
p := byID[protocolID]
|
|
if p == nil {
|
|
continue
|
|
}
|
|
ref := models.NodeProtoRef{
|
|
NodeID: nodeID,
|
|
NodeName: name,
|
|
Host: host,
|
|
Status: models.NodeStatus(status),
|
|
}
|
|
if installed {
|
|
p.InstalledOn = append(p.InstalledOn, ref)
|
|
}
|
|
if enabled {
|
|
p.EnabledOn = append(p.EnabledOn, ref)
|
|
}
|
|
}
|
|
return list, rows.Err()
|
|
}
|
|
|
|
func ListNodeProtocols(db *sql.DB, nodeID uuid.UUID) ([]models.NodeProtocol, error) {
|
|
rows, err := db.Query(`
|
|
SELECT p.id, p.code, p.name, COALESCE(NULLIF(np.port, 0), p.port),
|
|
COALESCE(np.installed, FALSE), COALESCE(np.enabled, FALSE),
|
|
COALESCE(np.updated_at, p.updated_at)
|
|
FROM protocols p
|
|
LEFT JOIN node_protocols np ON np.protocol_id = p.id AND np.node_id = $1
|
|
ORDER BY p.sort_order ASC, p.name ASC`, nodeID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var list []models.NodeProtocol
|
|
for rows.Next() {
|
|
var np models.NodeProtocol
|
|
np.NodeID = nodeID
|
|
if err := rows.Scan(
|
|
&np.ProtocolID, &np.ProtocolCode, &np.ProtocolName, &np.Port,
|
|
&np.Installed, &np.Enabled, &np.UpdatedAt,
|
|
); err != nil {
|
|
return nil, err
|
|
}
|
|
list = append(list, np)
|
|
}
|
|
return list, rows.Err()
|
|
}
|
|
|
|
func UpsertNodeProtocol(db *sql.DB, nodeID, protocolID uuid.UUID, installed, enabled bool, port int) error {
|
|
_, err := db.Exec(`
|
|
INSERT INTO node_protocols (node_id, protocol_id, installed, enabled, port, updated_at)
|
|
VALUES ($1, $2, $3, $4, $5, NOW())
|
|
ON CONFLICT (node_id, protocol_id) DO UPDATE SET
|
|
installed = EXCLUDED.installed,
|
|
enabled = EXCLUDED.enabled,
|
|
port = EXCLUDED.port,
|
|
updated_at = NOW()`, nodeID, protocolID, installed, enabled, port)
|
|
return err
|
|
}
|
|
|
|
func SetNodeProtocolFlags(db *sql.DB, nodeID, protocolID uuid.UUID, installed, enabled bool) error {
|
|
var port int
|
|
err := db.QueryRow(`
|
|
SELECT COALESCE(NULLIF(np.port, 0), p.port)
|
|
FROM protocols p
|
|
LEFT JOIN node_protocols np ON np.protocol_id = p.id AND np.node_id = $1
|
|
WHERE p.id = $2`, nodeID, protocolID).Scan(&port)
|
|
if err == sql.ErrNoRows {
|
|
return sql.ErrNoRows
|
|
}
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return UpsertNodeProtocol(db, nodeID, protocolID, installed, enabled, port)
|
|
}
|
|
|
|
func GetProtocol(db *sql.DB, id uuid.UUID) (*models.Protocol, error) {
|
|
p := &models.Protocol{}
|
|
err := db.QueryRow(`
|
|
SELECT id, code, name, description, port, enabled, sort_order, created_at, updated_at
|
|
FROM protocols WHERE id = $1`, id).Scan(
|
|
&p.ID, &p.Code, &p.Name, &p.Description, &p.Port,
|
|
&p.Enabled, &p.SortOrder, &p.CreatedAt, &p.UpdatedAt,
|
|
)
|
|
if err == sql.ErrNoRows {
|
|
return nil, nil
|
|
}
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return p, nil
|
|
}
|
|
|
|
func EnabledProtocolCodesForNode(db *sql.DB, nodeID uuid.UUID) ([]string, error) {
|
|
rows, err := db.Query(`
|
|
SELECT p.code FROM node_protocols np
|
|
JOIN protocols p ON p.id = np.protocol_id
|
|
WHERE np.node_id = $1 AND np.enabled = TRUE
|
|
ORDER BY p.sort_order`, nodeID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
var codes []string
|
|
for rows.Next() {
|
|
var c string
|
|
if err := rows.Scan(&c); err != nil {
|
|
return nil, err
|
|
}
|
|
codes = append(codes, c)
|
|
}
|
|
return codes, rows.Err()
|
|
}
|
|
|
|
func InstalledProtocolCodesForNode(db *sql.DB, nodeID uuid.UUID) ([]string, error) {
|
|
rows, err := db.Query(`
|
|
SELECT p.code FROM node_protocols np
|
|
JOIN protocols p ON p.id = np.protocol_id
|
|
WHERE np.node_id = $1 AND np.installed = TRUE
|
|
ORDER BY p.sort_order`, nodeID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
var codes []string
|
|
for rows.Next() {
|
|
var c string
|
|
if err := rows.Scan(&c); err != nil {
|
|
return nil, err
|
|
}
|
|
codes = append(codes, c)
|
|
}
|
|
return codes, rows.Err()
|
|
}
|
|
|
|
// SyncNodeProtocolsFromAgent updates install/enable flags from agent-reported codes.
|
|
func SyncNodeProtocolsFromAgent(db *sql.DB, nodeID uuid.UUID, installed, enabled []string) error {
|
|
tx, err := db.Begin()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer func() { _ = tx.Rollback() }()
|
|
|
|
installedSet := map[string]bool{}
|
|
for _, c := range installed {
|
|
installedSet[c] = true
|
|
}
|
|
enabledSet := map[string]bool{}
|
|
for _, c := range enabled {
|
|
enabledSet[c] = true
|
|
installedSet[c] = true // enabled implies installed
|
|
}
|
|
|
|
rows, err := tx.Query(`SELECT id, code, port FROM protocols`)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer rows.Close()
|
|
|
|
type proto struct {
|
|
id uuid.UUID
|
|
code string
|
|
port int
|
|
}
|
|
var all []proto
|
|
for rows.Next() {
|
|
var p proto
|
|
if err := rows.Scan(&p.id, &p.code, &p.port); err != nil {
|
|
return err
|
|
}
|
|
all = append(all, p)
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return err
|
|
}
|
|
_ = rows.Close()
|
|
|
|
for _, p := range all {
|
|
inst := installedSet[p.code]
|
|
en := enabledSet[p.code]
|
|
if !inst && !en {
|
|
_, err = tx.Exec(`
|
|
UPDATE node_protocols SET installed = FALSE, enabled = FALSE, updated_at = NOW()
|
|
WHERE node_id = $1 AND protocol_id = $2`, nodeID, p.id)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
continue
|
|
}
|
|
_, err = tx.Exec(`
|
|
INSERT INTO node_protocols (node_id, protocol_id, installed, enabled, port, updated_at)
|
|
VALUES ($1, $2, $3, $4, $5, NOW())
|
|
ON CONFLICT (node_id, protocol_id) DO UPDATE SET
|
|
installed = EXCLUDED.installed,
|
|
enabled = EXCLUDED.enabled,
|
|
updated_at = NOW()`, nodeID, p.id, inst, en, p.port)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
}
|
|
return tx.Commit()
|
|
}
|
|
|
|
func NodeProtocolConfigPayload(db *sql.DB, nodeID uuid.UUID) (map[string]any, error) {
|
|
rows, err := db.Query(`
|
|
SELECT p.code, np.installed, np.enabled, COALESCE(NULLIF(np.port, 0), p.port)
|
|
FROM node_protocols np
|
|
JOIN protocols p ON p.id = np.protocol_id
|
|
WHERE np.node_id = $1 AND (np.installed = TRUE OR np.enabled = TRUE)
|
|
ORDER BY p.sort_order`, nodeID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var installed, enabled []string
|
|
var items []map[string]any
|
|
for rows.Next() {
|
|
var code string
|
|
var inst, en bool
|
|
var port int
|
|
if err := rows.Scan(&code, &inst, &en, &port); err != nil {
|
|
return nil, err
|
|
}
|
|
if inst {
|
|
installed = append(installed, code)
|
|
}
|
|
if en {
|
|
enabled = append(enabled, code)
|
|
}
|
|
items = append(items, map[string]any{
|
|
"code": code, "installed": inst, "enabled": en, "port": port,
|
|
})
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
if installed == nil {
|
|
installed = []string{}
|
|
}
|
|
if enabled == nil {
|
|
enabled = []string{}
|
|
}
|
|
|
|
payload := map[string]any{
|
|
"protocols": enabled,
|
|
"protocols_installed": installed,
|
|
"protocols_enabled": enabled,
|
|
"protocol_items": items,
|
|
}
|
|
|
|
n, err := GetNode(db, nodeID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if n != nil && n.ProfileID != nil {
|
|
payload["profile"] = map[string]any{
|
|
"id": n.ProfileID.String(),
|
|
"name": n.ProfileName,
|
|
}
|
|
inbounds, err := ListNodeInbounds(db, nodeID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
var inList []map[string]any
|
|
for _, in := range inbounds {
|
|
if !in.NodeEnabled || !in.Enabled {
|
|
continue
|
|
}
|
|
clients, err := ListActiveClientsForInbound(db, in.ID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
var clientList []map[string]any
|
|
for _, c := range clients {
|
|
clientList = append(clientList, map[string]any{
|
|
"id": c.ID.String(),
|
|
"uuid": c.UUID.String(),
|
|
"email": c.Email,
|
|
"username": c.Username,
|
|
})
|
|
}
|
|
if clientList == nil {
|
|
clientList = []map[string]any{}
|
|
}
|
|
inList = append(inList, map[string]any{
|
|
"id": in.ID.String(),
|
|
"tag": in.Tag,
|
|
"protocol": in.ProtocolCode,
|
|
"port": in.Port,
|
|
"network": in.Network,
|
|
"security": in.Security,
|
|
"listen": in.Listen,
|
|
"remark": in.Remark,
|
|
"path": in.Path,
|
|
"host": in.Host,
|
|
"sni": in.SNI,
|
|
"fingerprint": in.Fingerprint,
|
|
"flow": in.Flow,
|
|
"alpn": in.ALPN,
|
|
"reality_dest": in.RealityDest,
|
|
"reality_server_names": in.RealityServerNames,
|
|
"reality_private_key": in.RealityPrivateKey,
|
|
"reality_public_key": in.RealityPublicKey,
|
|
"reality_short_id": in.RealityShortID,
|
|
"reality_short_ids": in.RealityShortIDs,
|
|
"spider_x": in.SpiderX,
|
|
"reality_xver": in.RealityXver,
|
|
"ss_method": in.SSMethod,
|
|
"password": in.Password,
|
|
"fallback_dest": in.FallbackDest,
|
|
"fallbacks_json": in.FallbacksJSON,
|
|
"tls_cert_pem": in.TLSCertPEM,
|
|
"tls_key_pem": in.TLSKeyPEM,
|
|
"sniffing": in.Sniffing,
|
|
"sniffing_route_only": in.SniffingRouteOnly,
|
|
"clients": clientList,
|
|
})
|
|
}
|
|
if inList == nil {
|
|
inList = []map[string]any{}
|
|
}
|
|
payload["inbounds"] = inList
|
|
}
|
|
|
|
return payload, nil
|
|
}
|
|
|
|
// ListActiveClientsForInbound returns active non-expired clients assigned to an inbound.
|
|
func ListActiveClientsForInbound(db *sql.DB, inboundID uuid.UUID) ([]models.Client, error) {
|
|
rows, err := db.Query(`
|
|
SELECT c.id, c.username, c.email, c.uuid, c.status,
|
|
c.traffic_limit_bytes, c.traffic_used_bytes, c.expire_at, c.sub_token, c.note,
|
|
c.created_at, c.updated_at, 0
|
|
FROM clients c
|
|
JOIN client_inbounds ci ON ci.client_id = c.id AND ci.inbound_id = $1
|
|
WHERE c.status = 'active'
|
|
AND (c.expire_at IS NULL OR c.expire_at > NOW())
|
|
AND (c.traffic_limit_bytes = 0 OR c.traffic_used_bytes < c.traffic_limit_bytes)
|
|
ORDER BY c.username ASC`, inboundID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
var list []models.Client
|
|
for rows.Next() {
|
|
c, err := scanClient(rows)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
list = append(list, *c)
|
|
}
|
|
return list, rows.Err()
|
|
}
|
|
|
|
// ListOnlineNodesByInboundIDs returns online nodes that have any of the inbounds enabled.
|
|
func ListOnlineNodesByInboundIDs(db *sql.DB, inboundIDs []uuid.UUID) ([]models.Node, error) {
|
|
if len(inboundIDs) == 0 {
|
|
return nil, nil
|
|
}
|
|
ids := make([]string, len(inboundIDs))
|
|
for i, id := range inboundIDs {
|
|
ids[i] = id.String()
|
|
}
|
|
rows, err := db.Query(`
|
|
SELECT DISTINCT `+nodeColumns+`
|
|
FROM nodes n
|
|
JOIN node_inbounds ni ON ni.node_id = n.id AND ni.enabled = TRUE
|
|
LEFT JOIN config_profiles cp ON cp.id = n.profile_id
|
|
WHERE n.status = 'online' AND ni.inbound_id = ANY($1)`, pq.Array(ids))
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
var list []models.Node
|
|
for rows.Next() {
|
|
n, err := scanNode(rows)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
list = append(list, *n)
|
|
}
|
|
return list, rows.Err()
|
|
}
|