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