Files

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()
}