179 lines
4.8 KiB
Go
179 lines
4.8 KiB
Go
package db
|
||
|
||
import (
|
||
"database/sql"
|
||
"fmt"
|
||
"log"
|
||
"time"
|
||
|
||
"github.com/google/uuid"
|
||
_ "github.com/lib/pq"
|
||
"golang.org/x/crypto/bcrypt"
|
||
|
||
"github.com/orohi/vpn-panel/internal/config"
|
||
"github.com/orohi/vpn-panel/internal/models"
|
||
)
|
||
|
||
func Connect(cfg *config.Config) (*sql.DB, error) {
|
||
var db *sql.DB
|
||
var err error
|
||
|
||
for i := 1; i <= 30; i++ {
|
||
db, err = sql.Open("postgres", cfg.DatabaseURL)
|
||
if err == nil {
|
||
err = db.Ping()
|
||
}
|
||
if err == nil {
|
||
db.SetMaxOpenConns(20)
|
||
db.SetMaxIdleConns(5)
|
||
db.SetConnMaxLifetime(time.Hour)
|
||
log.Println("connected to postgres")
|
||
return db, nil
|
||
}
|
||
log.Printf("waiting for postgres (%d/30): %v", i, err)
|
||
time.Sleep(2 * time.Second)
|
||
}
|
||
return nil, fmt.Errorf("postgres unavailable: %w", err)
|
||
}
|
||
|
||
func Migrate(db *sql.DB) error {
|
||
schema := `
|
||
CREATE EXTENSION IF NOT EXISTS "pgcrypto";
|
||
|
||
CREATE TABLE IF NOT EXISTS users (
|
||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||
email TEXT NOT NULL UNIQUE,
|
||
password_hash TEXT NOT NULL,
|
||
name TEXT NOT NULL DEFAULT '',
|
||
role TEXT NOT NULL DEFAULT 'admin',
|
||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
|
||
);
|
||
|
||
CREATE TABLE IF NOT EXISTS protocols (
|
||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||
code TEXT NOT NULL UNIQUE,
|
||
name TEXT NOT NULL,
|
||
description TEXT NOT NULL DEFAULT '',
|
||
port INT NOT NULL DEFAULT 0,
|
||
enabled BOOLEAN NOT NULL DEFAULT TRUE,
|
||
sort_order INT NOT NULL DEFAULT 0,
|
||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
|
||
);
|
||
`
|
||
_, err := db.Exec(schema)
|
||
return err
|
||
}
|
||
|
||
func SeedAdmin(db *sql.DB, cfg *config.Config) error {
|
||
var exists bool
|
||
err := db.QueryRow(`SELECT EXISTS(SELECT 1 FROM users WHERE email = $1)`, cfg.AdminEmail).Scan(&exists)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
if exists {
|
||
return nil
|
||
}
|
||
|
||
hash, err := bcrypt.GenerateFromPassword([]byte(cfg.AdminPassword), bcrypt.DefaultCost)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
|
||
_, err = db.Exec(
|
||
`INSERT INTO users (email, password_hash, name, role) VALUES ($1, $2, $3, 'admin')`,
|
||
cfg.AdminEmail, string(hash), cfg.AdminName,
|
||
)
|
||
if err == nil {
|
||
log.Printf("admin seeded: %s", cfg.AdminEmail)
|
||
}
|
||
return err
|
||
}
|
||
|
||
func SeedProtocols(db *sql.DB) error {
|
||
defaults := []struct {
|
||
Code, Name, Description string
|
||
Port, Sort int
|
||
}{
|
||
{"wireguard", "WireGuard", "Современный быстрый VPN-протокол на основе Noise", 51820, 1},
|
||
{"openvpn", "OpenVPN", "Классический VPN через UDP/TCP с TLS", 1194, 2},
|
||
{"vless", "VLESS", "Лёгкий протокол Xray без шифрования на уровне протокола", 443, 3},
|
||
{"vmess", "VMess", "Протокол V2Ray/Xray с обфускацией трафика", 443, 4},
|
||
{"trojan", "Trojan", "Трафик маскируется под HTTPS", 443, 5},
|
||
{"shadowsocks", "Shadowsocks", "SOCKS5-прокси с шифрованием AEAD", 8388, 6},
|
||
{"hysteria2", "Hysteria2", "UDP-протокол на базе QUIC для нестабильных сетей", 443, 7},
|
||
}
|
||
|
||
for _, p := range defaults {
|
||
_, err := db.Exec(`
|
||
INSERT INTO protocols (code, name, description, port, enabled, sort_order)
|
||
VALUES ($1, $2, $3, $4, TRUE, $5)
|
||
ON CONFLICT (code) DO NOTHING`,
|
||
p.Code, p.Name, p.Description, p.Port, p.Sort,
|
||
)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
}
|
||
return nil
|
||
}
|
||
|
||
func GetUserByEmail(db *sql.DB, email string) (*models.User, error) {
|
||
u := &models.User{}
|
||
err := db.QueryRow(`
|
||
SELECT id, email, password_hash, name, role, created_at
|
||
FROM users WHERE email = $1`, email).Scan(
|
||
&u.ID, &u.Email, &u.PasswordHash, &u.Name, &u.Role, &u.CreatedAt,
|
||
)
|
||
if err == sql.ErrNoRows {
|
||
return nil, nil
|
||
}
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return u, nil
|
||
}
|
||
|
||
func ListProtocols(db *sql.DB) ([]models.Protocol, error) {
|
||
rows, err := db.Query(`
|
||
SELECT id, code, name, description, port, enabled, sort_order, created_at, updated_at
|
||
FROM protocols ORDER BY sort_order ASC, name ASC`)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer rows.Close()
|
||
|
||
var list []models.Protocol
|
||
for rows.Next() {
|
||
var p models.Protocol
|
||
if err := rows.Scan(
|
||
&p.ID, &p.Code, &p.Name, &p.Description, &p.Port,
|
||
&p.Enabled, &p.SortOrder, &p.CreatedAt, &p.UpdatedAt,
|
||
); err != nil {
|
||
return nil, err
|
||
}
|
||
list = append(list, p)
|
||
}
|
||
return list, rows.Err()
|
||
}
|
||
|
||
func ToggleProtocol(db *sql.DB, id uuid.UUID) error {
|
||
_, err := db.Exec(`
|
||
UPDATE protocols SET enabled = NOT enabled, updated_at = NOW() WHERE id = $1`, id)
|
||
return err
|
||
}
|
||
|
||
func GetStats(db *sql.DB) (models.DashboardStats, error) {
|
||
var s models.DashboardStats
|
||
err := db.QueryRow(`SELECT COUNT(*) FROM protocols`).Scan(&s.ProtocolsTotal)
|
||
if err != nil {
|
||
return s, err
|
||
}
|
||
err = db.QueryRow(`SELECT COUNT(*) FROM protocols WHERE enabled = TRUE`).Scan(&s.ProtocolsEnabled)
|
||
if err != nil {
|
||
return s, err
|
||
}
|
||
err = db.QueryRow(`SELECT COUNT(*) FROM users WHERE role = 'admin'`).Scan(&s.AdminsTotal)
|
||
return s, err
|
||
}
|