417 lines
12 KiB
Go
417 lines
12 KiB
Go
package db
|
|
|
|
import (
|
|
"database/sql"
|
|
"fmt"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/google/uuid"
|
|
|
|
"github.com/orohi/vpn-panel/internal/models"
|
|
)
|
|
|
|
func CreateProfile(db *sql.DB, p *models.ConfigProfile) error {
|
|
if p.ID == uuid.Nil {
|
|
p.ID = uuid.New()
|
|
}
|
|
now := time.Now().UTC()
|
|
p.CreatedAt = now
|
|
p.UpdatedAt = now
|
|
_, err := db.Exec(`
|
|
INSERT INTO config_profiles (id, name, description, created_at, updated_at)
|
|
VALUES ($1, $2, $3, $4, $5)`, p.ID, p.Name, p.Description, p.CreatedAt, p.UpdatedAt)
|
|
return err
|
|
}
|
|
|
|
func UpdateProfile(db *sql.DB, p *models.ConfigProfile) error {
|
|
_, err := db.Exec(`
|
|
UPDATE config_profiles SET name = $2, description = $3, updated_at = NOW() WHERE id = $1`,
|
|
p.ID, p.Name, p.Description)
|
|
return err
|
|
}
|
|
|
|
func DeleteProfile(db *sql.DB, id uuid.UUID) error {
|
|
res, err := db.Exec(`DELETE FROM config_profiles WHERE id = $1`, id)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
n, _ := res.RowsAffected()
|
|
if n == 0 {
|
|
return fmt.Errorf("profile not found")
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func GetProfile(db *sql.DB, id uuid.UUID) (*models.ConfigProfile, error) {
|
|
p := &models.ConfigProfile{}
|
|
err := db.QueryRow(`
|
|
SELECT id, name, description, created_at, updated_at,
|
|
(SELECT COUNT(*) FROM inbounds i WHERE i.profile_id = config_profiles.id),
|
|
(SELECT COUNT(*) FROM nodes n WHERE n.profile_id = config_profiles.id)
|
|
FROM config_profiles WHERE id = $1`, id).Scan(
|
|
&p.ID, &p.Name, &p.Description, &p.CreatedAt, &p.UpdatedAt,
|
|
&p.InboundCount, &p.NodeCount,
|
|
)
|
|
if err == sql.ErrNoRows {
|
|
return nil, nil
|
|
}
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
inbounds, err := ListInboundsByProfile(db, id)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
p.Inbounds = inbounds
|
|
return p, nil
|
|
}
|
|
|
|
func ListProfiles(db *sql.DB) ([]models.ConfigProfile, error) {
|
|
rows, err := db.Query(`
|
|
SELECT id, name, description, created_at, updated_at,
|
|
(SELECT COUNT(*) FROM inbounds i WHERE i.profile_id = config_profiles.id),
|
|
(SELECT COUNT(*) FROM nodes n WHERE n.profile_id = config_profiles.id)
|
|
FROM config_profiles
|
|
ORDER BY name ASC`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
var list []models.ConfigProfile
|
|
for rows.Next() {
|
|
var p models.ConfigProfile
|
|
if err := rows.Scan(
|
|
&p.ID, &p.Name, &p.Description, &p.CreatedAt, &p.UpdatedAt,
|
|
&p.InboundCount, &p.NodeCount,
|
|
); err != nil {
|
|
return nil, err
|
|
}
|
|
list = append(list, p)
|
|
}
|
|
return list, rows.Err()
|
|
}
|
|
|
|
func scanInbound(scanner interface {
|
|
Scan(dest ...any) error
|
|
}) (*models.Inbound, error) {
|
|
in := &models.Inbound{}
|
|
err := scanner.Scan(
|
|
&in.ID, &in.ProfileID, &in.Tag, &in.ProtocolID, &in.ProtocolCode, &in.ProtocolName,
|
|
&in.Port, &in.Network, &in.Security, &in.Listen, &in.Enabled, &in.SortOrder, &in.Remark,
|
|
&in.Path, &in.Host, &in.SNI, &in.Fingerprint, &in.Flow, &in.ALPN,
|
|
&in.RealityDest, &in.RealityServerNames, &in.RealityPrivateKey, &in.RealityPublicKey, &in.RealityShortID,
|
|
&in.RealityShortIDs, &in.SpiderX, &in.RealityXver,
|
|
&in.SSMethod, &in.Password,
|
|
&in.FallbackDest, &in.FallbacksJSON, &in.TLSCertPEM, &in.TLSKeyPEM,
|
|
&in.Sniffing, &in.SniffingRouteOnly,
|
|
&in.CreatedAt, &in.UpdatedAt,
|
|
)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return in, nil
|
|
}
|
|
|
|
const inboundColumns = `
|
|
i.id, i.profile_id, i.tag, i.protocol_id, p.code, p.name,
|
|
i.port, i.network, i.security, i.listen, i.enabled, i.sort_order, i.remark,
|
|
i.path, i.host, i.sni, i.fingerprint, i.flow, i.alpn,
|
|
i.reality_dest, i.reality_server_names, i.reality_private_key, i.reality_public_key, i.reality_short_id,
|
|
i.reality_short_ids, i.spider_x, i.reality_xver,
|
|
i.ss_method, i.password,
|
|
i.fallback_dest, i.fallbacks_json, i.tls_cert_pem, i.tls_key_pem,
|
|
i.sniffing, i.sniffing_route_only,
|
|
i.created_at, i.updated_at`
|
|
|
|
func ListInboundsByProfile(db *sql.DB, profileID uuid.UUID) ([]models.Inbound, error) {
|
|
rows, err := db.Query(`
|
|
SELECT `+inboundColumns+`
|
|
FROM inbounds i
|
|
JOIN protocols p ON p.id = i.protocol_id
|
|
WHERE i.profile_id = $1
|
|
ORDER BY i.sort_order ASC, i.tag ASC`, profileID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
var list []models.Inbound
|
|
for rows.Next() {
|
|
in, err := scanInbound(rows)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
list = append(list, *in)
|
|
}
|
|
return list, rows.Err()
|
|
}
|
|
|
|
func GetInbound(db *sql.DB, id uuid.UUID) (*models.Inbound, error) {
|
|
row := db.QueryRow(`
|
|
SELECT `+inboundColumns+`
|
|
FROM inbounds i
|
|
JOIN protocols p ON p.id = i.protocol_id
|
|
WHERE i.id = $1`, id)
|
|
in, err := scanInbound(row)
|
|
if err == sql.ErrNoRows {
|
|
return nil, nil
|
|
}
|
|
return in, err
|
|
}
|
|
|
|
func CreateInbound(db *sql.DB, in *models.Inbound) error {
|
|
if in.ID == uuid.Nil {
|
|
in.ID = uuid.New()
|
|
}
|
|
if in.Tag == "" {
|
|
in.Tag = strings.ToUpper(in.ProtocolCode) + "-" + fmt.Sprintf("%d", in.Port)
|
|
}
|
|
if in.Network == "" {
|
|
in.Network = "tcp"
|
|
}
|
|
if in.Security == "" {
|
|
in.Security = "none"
|
|
}
|
|
if in.Listen == "" {
|
|
in.Listen = "0.0.0.0"
|
|
}
|
|
if in.Fingerprint == "" {
|
|
in.Fingerprint = "chrome"
|
|
}
|
|
if in.SSMethod == "" {
|
|
in.SSMethod = "aes-128-gcm"
|
|
}
|
|
if in.SpiderX == "" {
|
|
in.SpiderX = "/"
|
|
}
|
|
if in.RealityShortIDs == "" && in.RealityShortID != "" {
|
|
in.RealityShortIDs = in.RealityShortID
|
|
}
|
|
now := time.Now().UTC()
|
|
in.CreatedAt = now
|
|
in.UpdatedAt = now
|
|
_, err := db.Exec(`
|
|
INSERT INTO inbounds (
|
|
id, profile_id, tag, protocol_id, port, network, security, listen, enabled, sort_order, remark,
|
|
path, host, sni, fingerprint, flow, alpn,
|
|
reality_dest, reality_server_names, reality_private_key, reality_public_key, reality_short_id,
|
|
reality_short_ids, spider_x, reality_xver,
|
|
ss_method, password,
|
|
fallback_dest, fallbacks_json, tls_cert_pem, tls_key_pem,
|
|
sniffing, sniffing_route_only, created_at, updated_at
|
|
) VALUES (
|
|
$1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,
|
|
$12,$13,$14,$15,$16,$17,
|
|
$18,$19,$20,$21,$22,
|
|
$23,$24,$25,
|
|
$26,$27,
|
|
$28,$29,$30,$31,
|
|
$32,$33,$34,$35
|
|
)`,
|
|
in.ID, in.ProfileID, in.Tag, in.ProtocolID, in.Port, in.Network, in.Security, in.Listen,
|
|
in.Enabled, in.SortOrder, in.Remark,
|
|
in.Path, in.Host, in.SNI, in.Fingerprint, in.Flow, in.ALPN,
|
|
in.RealityDest, in.RealityServerNames, in.RealityPrivateKey, in.RealityPublicKey, in.RealityShortID,
|
|
in.RealityShortIDs, in.SpiderX, in.RealityXver,
|
|
in.SSMethod, in.Password,
|
|
in.FallbackDest, in.FallbacksJSON, in.TLSCertPEM, in.TLSKeyPEM,
|
|
in.Sniffing, in.SniffingRouteOnly, in.CreatedAt, in.UpdatedAt,
|
|
)
|
|
return err
|
|
}
|
|
|
|
func UpdateInbound(db *sql.DB, in *models.Inbound) error {
|
|
if in.Fingerprint == "" {
|
|
in.Fingerprint = "chrome"
|
|
}
|
|
if in.SSMethod == "" {
|
|
in.SSMethod = "aes-128-gcm"
|
|
}
|
|
if in.SpiderX == "" {
|
|
in.SpiderX = "/"
|
|
}
|
|
if in.RealityShortIDs == "" && in.RealityShortID != "" {
|
|
in.RealityShortIDs = in.RealityShortID
|
|
}
|
|
_, err := db.Exec(`
|
|
UPDATE inbounds SET
|
|
tag = $2, protocol_id = $3, port = $4, network = $5, security = $6,
|
|
listen = $7, enabled = $8, sort_order = $9, remark = $10,
|
|
path = $11, host = $12, sni = $13, fingerprint = $14, flow = $15, alpn = $16,
|
|
reality_dest = $17, reality_server_names = $18, reality_private_key = $19,
|
|
reality_public_key = $20, reality_short_id = $21,
|
|
reality_short_ids = $22, spider_x = $23, reality_xver = $24,
|
|
ss_method = $25, password = $26,
|
|
fallback_dest = $27, fallbacks_json = $28, tls_cert_pem = $29, tls_key_pem = $30,
|
|
sniffing = $31, sniffing_route_only = $32,
|
|
updated_at = NOW()
|
|
WHERE id = $1`,
|
|
in.ID, in.Tag, in.ProtocolID, in.Port, in.Network, in.Security,
|
|
in.Listen, in.Enabled, in.SortOrder, in.Remark,
|
|
in.Path, in.Host, in.SNI, in.Fingerprint, in.Flow, in.ALPN,
|
|
in.RealityDest, in.RealityServerNames, in.RealityPrivateKey, in.RealityPublicKey, in.RealityShortID,
|
|
in.RealityShortIDs, in.SpiderX, in.RealityXver,
|
|
in.SSMethod, in.Password,
|
|
in.FallbackDest, in.FallbacksJSON, in.TLSCertPEM, in.TLSKeyPEM,
|
|
in.Sniffing, in.SniffingRouteOnly,
|
|
)
|
|
return err
|
|
}
|
|
|
|
func DeleteInbound(db *sql.DB, id uuid.UUID) error {
|
|
_, err := db.Exec(`DELETE FROM inbounds WHERE id = $1`, id)
|
|
return err
|
|
}
|
|
|
|
func ToggleInbound(db *sql.DB, id uuid.UUID) error {
|
|
_, err := db.Exec(`UPDATE inbounds SET enabled = NOT enabled, updated_at = NOW() WHERE id = $1`, id)
|
|
return err
|
|
}
|
|
|
|
// AssignProfileToNode sets profile and enables all profile inbounds on the node.
|
|
func AssignProfileToNode(db *sql.DB, nodeID uuid.UUID, profileID uuid.UUID) error {
|
|
tx, err := db.Begin()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer func() { _ = tx.Rollback() }()
|
|
|
|
if _, err := tx.Exec(`UPDATE nodes SET profile_id = $2, updated_at = NOW() WHERE id = $1`, nodeID, profileID); err != nil {
|
|
return err
|
|
}
|
|
if _, err := tx.Exec(`DELETE FROM node_inbounds WHERE node_id = $1`, nodeID); err != nil {
|
|
return err
|
|
}
|
|
if _, err := tx.Exec(`
|
|
INSERT INTO node_inbounds (node_id, inbound_id, enabled)
|
|
SELECT $1, i.id, TRUE FROM inbounds i WHERE i.profile_id = $2 AND i.enabled = TRUE`, nodeID, profileID); err != nil {
|
|
return err
|
|
}
|
|
if err := tx.Commit(); err != nil {
|
|
return err
|
|
}
|
|
return ApplyProfileInboundsToNodeProtocols(db, nodeID)
|
|
}
|
|
|
|
func ClearNodeProfile(db *sql.DB, nodeID uuid.UUID) error {
|
|
tx, err := db.Begin()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer func() { _ = tx.Rollback() }()
|
|
if _, err := tx.Exec(`UPDATE nodes SET profile_id = NULL, updated_at = NOW() WHERE id = $1`, nodeID); err != nil {
|
|
return err
|
|
}
|
|
if _, err := tx.Exec(`DELETE FROM node_inbounds WHERE node_id = $1`, nodeID); err != nil {
|
|
return err
|
|
}
|
|
return tx.Commit()
|
|
}
|
|
|
|
func SetNodeInboundEnabled(db *sql.DB, nodeID, inboundID uuid.UUID, enabled bool) error {
|
|
_, err := db.Exec(`
|
|
INSERT INTO node_inbounds (node_id, inbound_id, enabled) VALUES ($1, $2, $3)
|
|
ON CONFLICT (node_id, inbound_id) DO UPDATE SET enabled = EXCLUDED.enabled`, nodeID, inboundID, enabled)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return ApplyProfileInboundsToNodeProtocols(db, nodeID)
|
|
}
|
|
|
|
func ListNodeInbounds(db *sql.DB, nodeID uuid.UUID) ([]models.Inbound, error) {
|
|
n, err := GetNode(db, nodeID)
|
|
if err != nil || n == nil || n.ProfileID == nil {
|
|
return nil, err
|
|
}
|
|
rows, err := db.Query(`
|
|
SELECT `+inboundColumns+`, COALESCE(ni.enabled, FALSE)
|
|
FROM inbounds i
|
|
JOIN protocols p ON p.id = i.protocol_id
|
|
LEFT JOIN node_inbounds ni ON ni.inbound_id = i.id AND ni.node_id = $1
|
|
WHERE i.profile_id = $2
|
|
ORDER BY i.sort_order ASC, i.tag ASC`, nodeID, *n.ProfileID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
var list []models.Inbound
|
|
for rows.Next() {
|
|
in := models.Inbound{}
|
|
if err := rows.Scan(
|
|
&in.ID, &in.ProfileID, &in.Tag, &in.ProtocolID, &in.ProtocolCode, &in.ProtocolName,
|
|
&in.Port, &in.Network, &in.Security, &in.Listen, &in.Enabled, &in.SortOrder, &in.Remark,
|
|
&in.Path, &in.Host, &in.SNI, &in.Fingerprint, &in.Flow, &in.ALPN,
|
|
&in.RealityDest, &in.RealityServerNames, &in.RealityPrivateKey, &in.RealityPublicKey, &in.RealityShortID,
|
|
&in.RealityShortIDs, &in.SpiderX, &in.RealityXver,
|
|
&in.SSMethod, &in.Password,
|
|
&in.FallbackDest, &in.FallbacksJSON, &in.TLSCertPEM, &in.TLSKeyPEM,
|
|
&in.Sniffing, &in.SniffingRouteOnly,
|
|
&in.CreatedAt, &in.UpdatedAt, &in.NodeEnabled,
|
|
); err != nil {
|
|
return nil, err
|
|
}
|
|
list = append(list, in)
|
|
}
|
|
return list, rows.Err()
|
|
}
|
|
|
|
// ApplyProfileInboundsToNodeProtocols mirrors enabled node inbounds into node_protocols.
|
|
func ApplyProfileInboundsToNodeProtocols(db *sql.DB, nodeID uuid.UUID) error {
|
|
rows, err := db.Query(`
|
|
SELECT i.protocol_id, i.port, BOOL_OR(ni.enabled AND i.enabled) AS active
|
|
FROM inbounds i
|
|
JOIN nodes n ON n.profile_id = i.profile_id AND n.id = $1
|
|
JOIN node_inbounds ni ON ni.inbound_id = i.id AND ni.node_id = $1
|
|
GROUP BY i.protocol_id, i.port`, nodeID)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer rows.Close()
|
|
|
|
type row struct {
|
|
protocolID uuid.UUID
|
|
port int
|
|
active bool
|
|
}
|
|
var active []row
|
|
seen := map[uuid.UUID]row{}
|
|
for rows.Next() {
|
|
var r row
|
|
if err := rows.Scan(&r.protocolID, &r.port, &r.active); err != nil {
|
|
return err
|
|
}
|
|
if prev, ok := seen[r.protocolID]; ok {
|
|
if r.active {
|
|
prev.active = true
|
|
prev.port = r.port
|
|
seen[r.protocolID] = prev
|
|
}
|
|
continue
|
|
}
|
|
seen[r.protocolID] = r
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return err
|
|
}
|
|
for _, r := range seen {
|
|
active = append(active, r)
|
|
}
|
|
|
|
// Reset then apply.
|
|
if _, err := db.Exec(`
|
|
UPDATE node_protocols SET installed = FALSE, enabled = FALSE, updated_at = NOW()
|
|
WHERE node_id = $1`, nodeID); err != nil {
|
|
return err
|
|
}
|
|
for _, r := range active {
|
|
if !r.active {
|
|
continue
|
|
}
|
|
if err := UpsertNodeProtocol(db, nodeID, r.protocolID, true, true, r.port); err != nil {
|
|
return err
|
|
}
|
|
_ = SetProtocolEnabled(db, r.protocolID, true)
|
|
}
|
|
return nil
|
|
}
|