Files
frpc-console/db.go
T

589 lines
16 KiB
Go

package main
import (
"crypto/rand"
"database/sql"
"encoding/hex"
"fmt"
"io"
"log"
"os"
"strings"
"time"
_ "modernc.org/sqlite"
)
var DB *sql.DB
// ============================================================
// 数据库版本常量
// ============================================================
const (
SchemaVersion = "2.0.0" // 当前数据库 Schema 版本,与项目版本同步
)
// ============================================================
// 数据模型
// ============================================================
// GlobalConfig 全局配置表
type GlobalConfig struct {
ID int `json:"id"`
ServerAddr string `json:"serverAddr"`
ServerPort int `json:"serverPort"`
Token string `json:"token"`
LogLevel string `json:"logLevel"`
LogMaxDays int `json:"logMaxDays"`
TcpMux bool `json:"tcpMux"`
TcpMuxKeepalive int `json:"tcpMuxKeepalive"`
HeartbeatInterval int `json:"heartbeatInterval"`
HeartbeatTimeout int `json:"heartbeatTimeout"`
PoolCount int `json:"poolCount"`
WireProtocolV2 bool `json:"wireProtocolV2"` // v2.0 正式启用
}
// Proxy 隧道表
type Proxy struct {
ID int `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
LocalIP string `json:"localIP"`
LocalPort int `json:"localPort"`
RemotePort int `json:"remotePort"`
Enabled bool `json:"enabled"`
}
// User 用户表
type User struct {
ID int `json:"id"`
Username string `json:"username"`
PasswordHash string `json:"-"`
CreatedAt string `json:"createdAt"`
}
// ============================================================
// 数据库初始化
// ============================================================
func InitDB() error {
var err error
DB, err = sql.Open("sqlite", "./frpc-console.db")
if err != nil {
return err
}
// ---- 创建所有表 ----
if err := createTables(); err != nil {
return err
}
// ---- 执行版本迁移 ----
if err := runMigrations(); err != nil {
return err
}
// ---- 确保 JWT 密钥存在 ----
if err := ensureJwtSecret(); err != nil {
return err
}
log.Println("✅ 数据库初始化完成 (Schema v" + SchemaVersion + ")")
return nil
}
// ============================================================
// 建表
// ============================================================
func createTables() error {
// 用户表
_, err := DB.Exec(`
CREATE TABLE IF NOT EXISTS users (
id INTEGER PRIMARY KEY AUTOINCREMENT,
username TEXT UNIQUE NOT NULL,
password_hash TEXT NOT NULL,
created_at DATETIME DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME DEFAULT CURRENT_TIMESTAMP
)
`)
if err != nil {
return err
}
// 全局配置表
_, err = DB.Exec(`
CREATE TABLE IF NOT EXISTS global_config (
id INTEGER PRIMARY KEY CHECK (id = 1),
server_addr TEXT NOT NULL DEFAULT 'frp.example.com',
server_port INTEGER NOT NULL DEFAULT 7000,
token TEXT NOT NULL DEFAULT 'CHANGE_ME',
log_level TEXT NOT NULL DEFAULT 'info',
log_max_days INTEGER NOT NULL DEFAULT 3,
tcp_mux INTEGER NOT NULL DEFAULT 1,
tcp_mux_keepalive INTEGER NOT NULL DEFAULT 30,
heartbeat_interval INTEGER NOT NULL DEFAULT 15,
heartbeat_timeout INTEGER NOT NULL DEFAULT 70,
pool_count INTEGER NOT NULL DEFAULT 8,
wire_protocol_v2 INTEGER NOT NULL DEFAULT 0,
updated_at DATETIME DEFAULT CURRENT_TIMESTAMP
)
`)
if err != nil {
return err
}
// 隧道表
_, err = DB.Exec(`
CREATE TABLE IF NOT EXISTS proxies (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT UNIQUE NOT NULL,
type TEXT NOT NULL DEFAULT 'tcp',
local_ip TEXT NOT NULL,
local_port INTEGER NOT NULL,
remote_port INTEGER NOT NULL,
enabled INTEGER NOT NULL DEFAULT 1,
created_at DATETIME DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME DEFAULT CURRENT_TIMESTAMP
)
`)
if err != nil {
return err
}
// 应用配置表(存储 JWT 密钥、Schema 版本等)
_, err = DB.Exec(`
CREATE TABLE IF NOT EXISTS app_config (
key TEXT PRIMARY KEY,
value TEXT NOT NULL,
updated_at DATETIME DEFAULT CURRENT_TIMESTAMP
)
`)
if err != nil {
return err
}
// 初始化默认配置(仅当表为空时)
var count int
DB.QueryRow("SELECT COUNT(*) FROM global_config").Scan(&count)
if count == 0 {
_, err = DB.Exec(`
INSERT INTO global_config (
id, server_addr, server_port, token, log_level, log_max_days,
tcp_mux, tcp_mux_keepalive, heartbeat_interval, heartbeat_timeout, pool_count,
wire_protocol_v2
) VALUES (1, 'frp.example.com', 7000, 'CHANGE_ME', 'info', 3, 1, 30, 15, 70, 8, 0)
`)
if err != nil {
return err
}
log.Println("✅ 全局配置初始化完成")
}
return nil
}
// ============================================================
// 迁移引擎
// ============================================================
// getSchemaVersion 读取当前数据库的 Schema 版本
func getSchemaVersion() string {
var version string
err := DB.QueryRow("SELECT value FROM app_config WHERE key = 'schema_version'").Scan(&version)
if err != nil {
if err == sql.ErrNoRows {
// 没有版本记录 → 首次启动或 v1.x 升级
// 检查是否已有数据(通过 users 表判断)
var count int
DB.QueryRow("SELECT COUNT(*) FROM users").Scan(&count)
if count > 0 {
// 有用户数据 → 这是 v1.x 升级,标记为 1.5
return "1.5.0"
}
// 全新安装 → 直接标记为当前版本
return SchemaVersion
}
log.Printf("⚠️ 读取 Schema 版本失败: %v", err)
return "1.5.0" // 保守降级
}
return version
}
// setSchemaVersion 更新 Schema 版本
func setSchemaVersion(version string) error {
_, err := DB.Exec(`
INSERT INTO app_config (key, value) VALUES ('schema_version', ?)
ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = CURRENT_TIMESTAMP
`, version)
return err
}
// backupDatabase 备份数据库文件
func backupDatabase() (string, error) {
src := "./frpc-console.db"
if _, err := os.Stat(src); os.IsNotExist(err) {
return "", nil // 数据库不存在,无需备份
}
timestamp := time.Now().Format("20060102_150405")
dst := fmt.Sprintf("./frpc-console.db.pre-v%s.%s", SchemaVersion, timestamp)
srcFile, err := os.Open(src)
if err != nil {
return "", fmt.Errorf("打开源数据库失败: %w", err)
}
defer srcFile.Close()
dstFile, err := os.Create(dst)
if err != nil {
return "", fmt.Errorf("创建备份文件失败: %w", err)
}
defer dstFile.Close()
if _, err := io.Copy(dstFile, srcFile); err != nil {
return "", fmt.Errorf("复制数据库失败: %w", err)
}
log.Printf("✅ 数据库备份完成: %s", dst)
return dst, nil
}
// runMigrations 执行版本迁移
func runMigrations() error {
currentVersion := getSchemaVersion()
log.Printf("📌 当前数据库 Schema: %s, 目标版本: %s", currentVersion, SchemaVersion)
if currentVersion == SchemaVersion {
log.Println("✅ Schema 已是最新,无需迁移")
return nil
}
log.Printf("🔄 检测到版本变更 (%s → %s),开始迁移...", currentVersion, SchemaVersion)
// ---- 1. 备份数据库 ----
backupPath, err := backupDatabase()
if err != nil {
return fmt.Errorf("备份数据库失败: %w", err)
}
if backupPath != "" {
log.Printf("📦 备份文件: %s", backupPath)
} else {
log.Println("📦 数据库为空,跳过备份")
}
// ---- 2. 执行迁移 ----
// 按照版本号逐个升级
migrations := []struct {
from string
upgrade func() error
}{
{"1.5.0", migrateFrom1_5_0},
{"2.0.0", migrateFrom2_0_0}, // 预留,实际无操作
}
applied := 0
for _, m := range migrations {
if currentVersion == m.from {
log.Printf(" 执行迁移: %s → %s", m.from, SchemaVersion)
if err := m.upgrade(); err != nil {
// 迁移失败,尝试恢复备份
if backupPath != "" {
log.Printf("❌ 迁移失败,尝试恢复备份: %s", backupPath)
if restoreErr := restoreDatabase(backupPath); restoreErr != nil {
log.Printf("⚠️ 恢复备份失败: %v", restoreErr)
}
}
return fmt.Errorf("迁移失败: %w", err)
}
applied++
currentVersion = SchemaVersion
break
}
}
// ---- 3. 更新 Schema 版本 ----
if err := setSchemaVersion(SchemaVersion); err != nil {
return fmt.Errorf("更新 Schema 版本失败: %w", err)
}
log.Printf("✅ 迁移完成,应用了 %d 个迁移", applied)
return nil
}
// restoreDatabase 从备份恢复数据库
func restoreDatabase(backupPath string) error {
srcFile, err := os.Open(backupPath)
if err != nil {
return err
}
defer srcFile.Close()
dstFile, err := os.Create("./frpc-console.db")
if err != nil {
return err
}
defer dstFile.Close()
if _, err := io.Copy(dstFile, srcFile); err != nil {
return err
}
log.Printf("✅ 数据库已从备份恢复: %s", backupPath)
return nil
}
// ============================================================
// 迁移函数(各版本)
// ============================================================
// migrateFrom1_5_0: v1.5 → v2.0
// v1.5 已经有 wire_protocol_v2 字段(灰标占位),v2.0 无需新增字段
// 但需要确保字段存在(兼容从 v1.0 直接升级的场景)
func migrateFrom1_5_0() error {
log.Println(" 迁移: v1.5.0 → v2.0.0")
// 检查并补全 wire_protocol_v2 字段(兼容从 v1.0 直接升级的场景)
cols, err := getCurrentColumns("global_config")
if err != nil {
return fmt.Errorf("获取列信息失败: %w", err)
}
if !contains(cols, "wire_protocol_v2") {
log.Println(" 添加字段: wire_protocol_v2")
_, err := DB.Exec("ALTER TABLE global_config ADD COLUMN wire_protocol_v2 INTEGER NOT NULL DEFAULT 0")
if err != nil {
return fmt.Errorf("添加 wire_protocol_v2 字段失败: %w", err)
}
}
log.Println(" ✅ v1.5.0 → v2.0.0 迁移完成")
return nil
}
// migrateFrom2_0_0: 预留,v2.0 → 未来版本
func migrateFrom2_0_0() error {
log.Println(" v2.0.0 已是当前版本,无需迁移")
return nil
}
// ============================================================
// 辅助函数
// ============================================================
func getCurrentColumns(tableName string) ([]string, error) {
rows, err := DB.Query("PRAGMA table_info(" + tableName + ")")
if err != nil {
return nil, err
}
defer rows.Close()
var cols []string
for rows.Next() {
var (
cid int
name string
typ string
notNull int
dfltVal sql.NullString
pk int
)
if err := rows.Scan(&cid, &name, &typ, &notNull, &dfltVal, &pk); err != nil {
return nil, err
}
cols = append(cols, name)
}
return cols, rows.Err()
}
func contains(slice []string, item string) bool {
for _, s := range slice {
if strings.EqualFold(s, item) {
return true
}
}
return false
}
// ============================================================
// JWT 密钥管理
// ============================================================
func ensureJwtSecret() error {
var value string
err := DB.QueryRow("SELECT value FROM app_config WHERE key = 'jwt_secret'").Scan(&value)
if err == nil && value != "" {
return nil
}
bytes := make([]byte, 32)
if _, err := rand.Read(bytes); err != nil {
return err
}
secret := hex.EncodeToString(bytes)
_, err = DB.Exec(`
INSERT INTO app_config (key, value) VALUES ('jwt_secret', ?)
`, secret)
if err != nil {
return err
}
log.Printf("✅ JWT 密钥已生成")
return nil
}
func GetJwtSecret() (string, error) {
var secret string
err := DB.QueryRow("SELECT value FROM app_config WHERE key = 'jwt_secret'").Scan(&secret)
if err != nil {
return "", err
}
return secret, nil
}
// ============================================================
// 全局配置 CRUD
// ============================================================
func GetGlobalConfig() (*GlobalConfig, error) {
var cfg GlobalConfig
err := DB.QueryRow(`
SELECT id, server_addr, server_port, token, log_level, log_max_days,
tcp_mux, tcp_mux_keepalive, heartbeat_interval, heartbeat_timeout, pool_count,
wire_protocol_v2
FROM global_config WHERE id = 1
`).Scan(
&cfg.ID, &cfg.ServerAddr, &cfg.ServerPort, &cfg.Token,
&cfg.LogLevel, &cfg.LogMaxDays, &cfg.TcpMux, &cfg.TcpMuxKeepalive,
&cfg.HeartbeatInterval, &cfg.HeartbeatTimeout, &cfg.PoolCount,
&cfg.WireProtocolV2,
)
if err != nil {
return nil, err
}
cfg.TcpMux = true
return &cfg, nil
}
func UpdateGlobalConfig(cfg *GlobalConfig) error {
_, err := DB.Exec(`
UPDATE global_config SET
server_addr = ?, server_port = ?, token = ?, log_level = ?, log_max_days = ?,
tcp_mux = 1,
tcp_mux_keepalive = ?, heartbeat_interval = ?, heartbeat_timeout = ?, pool_count = ?,
wire_protocol_v2 = ?,
updated_at = CURRENT_TIMESTAMP
WHERE id = 1
`, cfg.ServerAddr, cfg.ServerPort, cfg.Token, cfg.LogLevel, cfg.LogMaxDays,
cfg.TcpMuxKeepalive, cfg.HeartbeatInterval, cfg.HeartbeatTimeout, cfg.PoolCount,
cfg.WireProtocolV2)
return err
}
// ============================================================
// 隧道 CRUD
// ============================================================
func GetProxies() ([]Proxy, error) {
rows, err := DB.Query(`
SELECT id, name, type, local_ip, local_port, remote_port, enabled
FROM proxies ORDER BY id
`)
if err != nil {
return nil, err
}
defer rows.Close()
var proxies []Proxy
for rows.Next() {
var p Proxy
err := rows.Scan(&p.ID, &p.Name, &p.Type, &p.LocalIP, &p.LocalPort, &p.RemotePort, &p.Enabled)
if err != nil {
return nil, err
}
proxies = append(proxies, p)
}
if err = rows.Err(); err != nil {
return nil, err
}
return proxies, nil
}
func GetProxy(id int) (*Proxy, error) {
var p Proxy
err := DB.QueryRow(`
SELECT id, name, type, local_ip, local_port, remote_port, enabled
FROM proxies WHERE id = ?
`, id).Scan(&p.ID, &p.Name, &p.Type, &p.LocalIP, &p.LocalPort, &p.RemotePort, &p.Enabled)
if err != nil {
return nil, err
}
return &p, nil
}
func CreateProxy(p *Proxy) error {
result, err := DB.Exec(`
INSERT INTO proxies (name, type, local_ip, local_port, remote_port, enabled)
VALUES (?, ?, ?, ?, ?, ?)
`, p.Name, p.Type, p.LocalIP, p.LocalPort, p.RemotePort, p.Enabled)
if err != nil {
return err
}
id, _ := result.LastInsertId()
p.ID = int(id)
return nil
}
func UpdateProxy(p *Proxy) error {
_, err := DB.Exec(`
UPDATE proxies SET
name = ?, type = ?, local_ip = ?, local_port = ?, remote_port = ?, enabled = ?,
updated_at = CURRENT_TIMESTAMP
WHERE id = ?
`, p.Name, p.Type, p.LocalIP, p.LocalPort, p.RemotePort, p.Enabled, p.ID)
return err
}
func DeleteProxy(id int) error {
_, err := DB.Exec("DELETE FROM proxies WHERE id = ?", id)
return err
}
// ============================================================
// 用户 CRUD
// ============================================================
func GetUserByUsername(username string) (*User, error) {
var u User
err := DB.QueryRow(`
SELECT id, username, password_hash, created_at
FROM users WHERE username = ?
`, username).Scan(&u.ID, &u.Username, &u.PasswordHash, &u.CreatedAt)
if err != nil {
return nil, err
}
return &u, nil
}
func CreateUser(username, passwordHash string) error {
_, err := DB.Exec(`
INSERT INTO users (username, password_hash) VALUES (?, ?)
`, username, passwordHash)
return err
}
func UpdatePassword(username, passwordHash string) error {
_, err := DB.Exec(`
UPDATE users SET password_hash = ?, updated_at = CURRENT_TIMESTAMP
WHERE username = ?
`, passwordHash, username)
return err
}
func CountUsers() (int, error) {
var count int
err := DB.QueryRow("SELECT COUNT(*) FROM users").Scan(&count)
return count, err
}