Files
frpc-console/internal/db/db.go
T

556 lines
16 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package db
import (
"crypto/rand"
"database/sql"
"encoding/hex"
"fmt"
"io"
"log"
"os"
"sort"
"strings"
"time"
_ "modernc.org/sqlite"
)
var DB *sql.DB
const SchemaVersion = "v3"
// ================================================================
// 初始化
// ================================================================
func InitDB() error {
if err := os.MkdirAll("./data", 0755); err != nil {
return fmt.Errorf("创建数据目录失败: %w", err)
}
dbPath := "./data/frpc-console.db"
var err error
DB, err = sql.Open("sqlite", dbPath)
if err != nil {
return err
}
if err := createTables(); err != nil {
return err
}
if err := runMigrations(); err != nil {
return err
}
if err := ensureJwtSecret(); err != nil {
return err
}
log.Println("✅ 数据库初始化完成 (Schema " + SchemaVersion + ")")
return nil
}
func createTables() error {
// users 表
_, 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
}
// global_config 表
_, 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',
admin_port INTEGER NOT NULL DEFAULT 7400,
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
}
// proxies 表
_, 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
}
// app_config 表
_, 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
}
// 初始化 global_config 默认值(新数据库)
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, admin_port, 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', 7400, 'info', 3, 1, 30, 15, 70, 8, 0)
`)
if err != nil {
return err
}
log.Println("✅ 全局配置初始化完成")
}
return nil
}
// ================================================================
// v2 → v3 重型迁移:global_config 表添加 admin_port
// ================================================================
func migrateGlobalConfigToV3() error {
log.Println(" 开始 global_config 表迁移 (v2→v3)")
// 1. 检查 admin_port 列是否已存在
var hasAdminPort bool
rows, err := DB.Query("PRAGMA table_info(global_config)")
if err != nil {
return fmt.Errorf("查询表结构失败: %w", err)
}
defer rows.Close()
for rows.Next() {
var cid int
var name, ctype string
var notnull, pk int
var dflt sql.NullString
if err := rows.Scan(&cid, &name, &ctype, &notnull, &dflt, &pk); err != nil {
return err
}
if name == "admin_port" {
hasAdminPort = true
break
}
}
if hasAdminPort {
log.Println(" ✅ admin_port 列已存在,跳过迁移")
return nil
}
log.Println(" 创建 global_config_new 表...")
_, err = DB.Exec(`
CREATE TABLE global_config_new (
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',
admin_port INTEGER NOT NULL DEFAULT 7400,
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 fmt.Errorf("创建 global_config_new 表失败: %w", err)
}
log.Println(" 迁移数据 (admin_port = 7400)...")
_, err = DB.Exec(`
INSERT INTO global_config_new (
id, server_addr, server_port, token, admin_port,
log_level, log_max_days, tcp_mux, tcp_mux_keepalive,
heartbeat_interval, heartbeat_timeout, pool_count,
wire_protocol_v2, updated_at
)
SELECT
id, server_addr, server_port, token, 7400,
log_level, log_max_days, tcp_mux, tcp_mux_keepalive,
heartbeat_interval, heartbeat_timeout, pool_count,
wire_protocol_v2, updated_at
FROM global_config
`)
if err != nil {
return fmt.Errorf("复制数据失败: %w", err)
}
var oldCount, newCount int
DB.QueryRow("SELECT COUNT(*) FROM global_config").Scan(&oldCount)
DB.QueryRow("SELECT COUNT(*) FROM global_config_new").Scan(&newCount)
if oldCount != newCount {
return fmt.Errorf("数据迁移不完整: 旧表 %d 行,新表 %d 行", oldCount, newCount)
}
log.Printf(" 数据迁移验证通过: %d 行", newCount)
log.Println(" 交换表名...")
if _, err := DB.Exec("ALTER TABLE global_config RENAME TO global_config_old"); err != nil {
return fmt.Errorf("重命名旧表失败: %w", err)
}
if _, err := DB.Exec("ALTER TABLE global_config_new RENAME TO global_config"); err != nil {
DB.Exec("ALTER TABLE global_config_old RENAME TO global_config")
return fmt.Errorf("重命名新表失败: %w", err)
}
if _, err := DB.Exec("DROP TABLE global_config_old"); err != nil {
log.Printf("⚠️ 删除旧表失败(不影响使用): %v", err)
}
log.Println(" ✅ global_config 表迁移完成")
return nil
}
// ================================================================
// Schema 版本管理
// ================================================================
func getCurrentSchemaVersion() 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 {
var count int
DB.QueryRow("SELECT COUNT(*) FROM users").Scan(&count)
if count > 0 {
return "v1"
}
return SchemaVersion
}
log.Printf("⚠️ 读取 Schema 版本失败: %v", err)
return "v1"
}
return version
}
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
}
func backupDatabase() (string, error) {
src := "./data/frpc-console.db"
if _, err := os.Stat(src); os.IsNotExist(err) {
return "", nil
}
timestamp := time.Now().Format("20060102_150405")
dst := fmt.Sprintf("./data/frpc-console.db.pre-%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
}
func restoreDatabase(backupPath string) error {
srcFile, err := os.Open(backupPath)
if err != nil {
return err
}
defer srcFile.Close()
dstFile, err := os.Create("./data/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
}
// ================================================================
// runMigrations - 核心迁移入口
// ================================================================
func runMigrations() error {
currentVer := getCurrentSchemaVersion()
targetVer := SchemaVersion
log.Printf("📌 当前数据库 Schema: %s, 目标版本: %s", currentVer, targetVer)
if currentVer == targetVer {
var userCount int
err := DB.QueryRow("SELECT COUNT(*) FROM users").Scan(&userCount)
if err != nil || userCount == 0 {
log.Println(" 数据库为空或无效,无需迁移,直接初始化")
return nil
}
log.Println("✅ Schema 已是最新,数据库有效")
return nil
}
log.Printf("🔄 检测到版本变更 (%s → %s),开始迁移...", currentVer, targetVer)
backupPath, err := backupDatabase()
if err != nil {
return fmt.Errorf("备份数据库失败: %w", err)
}
if backupPath != "" {
log.Printf("📦 备份文件: %s", backupPath)
}
// ---- 阶段1: v1 → v2proxies 表迁移) ----
if currentVer == "v1" {
oldDef := getSchemaDef("v1")
newDef := getSchemaDef("v2")
if oldDef == nil || newDef == nil {
return fmt.Errorf("v1 或 v2 schema 定义不存在")
}
if !schemaVersionsEqual(oldDef, newDef) {
log.Println(" 阶段1: v1→v2 重型迁移(proxies 表结构变更)")
if err := heavyMigration(oldDef, newDef); err != nil {
if backupPath != "" {
restoreDatabase(backupPath)
}
return fmt.Errorf("v1→v2 迁移失败: %w", err)
}
} else {
log.Println(" 阶段1: v1→v2 轻量迁移(proxies 表结构无变更)")
}
currentVer = "v2"
}
// ---- 阶段2: v2 → v3global_config 表新增 admin_port ----
if currentVer == "v2" {
log.Println(" 阶段2: v2→v3 重型迁移(global_config 表新增 admin_port")
if err := migrateGlobalConfigToV3(); err != nil {
if backupPath != "" {
restoreDatabase(backupPath)
}
return fmt.Errorf("v2→v3 迁移失败: %w", err)
}
currentVer = "v3"
}
if err := setSchemaVersion(targetVer); err != nil {
return fmt.Errorf("更新 Schema 版本失败: %w", err)
}
log.Printf("✅ 迁移完成,当前 Schema: %s", targetVer)
return nil
}
// ================================================================
// 重型迁移引擎(用于 proxies 表)
// ================================================================
func heavyMigration(oldDef, newDef *SchemaVersionDef) error {
if oldDef == nil {
return fmt.Errorf("旧 Schema 定义为空")
}
oldTable := oldDef.TableName
newTable := oldTable + "_new"
createSQL := buildCreateTableSQL(newTable, newDef)
log.Printf(" 创建新表: %s", newTable)
if _, err := DB.Exec(createSQL); err != nil {
return fmt.Errorf("创建新表失败: %w", err)
}
insertSQL, err := buildInsertSQL(oldTable, newTable, oldDef, newDef)
if err != nil {
return fmt.Errorf("构建数据迁移 SQL 失败: %w", err)
}
log.Printf(" 迁移数据: %s → %s", oldTable, newTable)
if _, err := DB.Exec(insertSQL); err != nil {
return fmt.Errorf("数据迁移失败: %w", err)
}
var oldCount, newCount int
DB.QueryRow(fmt.Sprintf("SELECT COUNT(*) FROM %s", oldTable)).Scan(&oldCount)
DB.QueryRow(fmt.Sprintf("SELECT COUNT(*) FROM %s", newTable)).Scan(&newCount)
if oldCount != newCount {
return fmt.Errorf("数据迁移不完整: 旧表 %d 行,新表 %d 行", oldCount, newCount)
}
log.Printf(" 数据迁移验证通过: %d 行", newCount)
tempTable := oldTable + "_old_temp"
if _, err := DB.Exec(fmt.Sprintf("ALTER TABLE %s RENAME TO %s", oldTable, tempTable)); err != nil {
return fmt.Errorf("重命名旧表失败: %w", err)
}
if _, err := DB.Exec(fmt.Sprintf("ALTER TABLE %s RENAME TO %s", newTable, oldTable)); err != nil {
DB.Exec(fmt.Sprintf("ALTER TABLE %s RENAME TO %s", tempTable, oldTable))
return fmt.Errorf("重命名新表失败: %w", err)
}
if _, err := DB.Exec(fmt.Sprintf("DROP TABLE %s", tempTable)); err != nil {
log.Printf("⚠️ 删除临时表失败(不影响使用): %v", err)
}
log.Printf(" 表交换完成: %s (新表已生效)", oldTable)
return nil
}
func buildCreateTableSQL(tableName string, def *SchemaVersionDef) string {
var cols []string
var primaryKey string
names := make([]string, 0, len(def.Columns))
for name := range def.Columns {
names = append(names, name)
}
sort.Strings(names)
for _, name := range names {
col := def.Columns[name]
parts := []string{name, col.Type}
if col.NotNull {
parts = append(parts, "NOT NULL")
}
if col.Default != "" {
parts = append(parts, "DEFAULT "+col.Default)
}
if col.Primary {
primaryKey = "PRIMARY KEY (" + name + ")"
} else {
cols = append(cols, strings.Join(parts, " "))
}
}
if primaryKey != "" {
cols = append(cols, primaryKey)
}
return fmt.Sprintf("CREATE TABLE %s (\n %s\n)", tableName, strings.Join(cols, ",\n "))
}
func buildInsertSQL(oldTable, newTable string, oldDef, newDef *SchemaVersionDef) (string, error) {
newCols := make([]string, 0, len(newDef.Columns))
for name := range newDef.Columns {
newCols = append(newCols, name)
}
sort.Strings(newCols)
var selectParts []string
var colNames []string
for _, name := range newCols {
colNames = append(colNames, name)
if _, ok := oldDef.Columns[name]; ok {
selectParts = append(selectParts, name)
} else {
colDef := newDef.Columns[name]
if colDef.Default != "" {
selectParts = append(selectParts, colDef.Default+" AS "+name)
} else if colDef.Type == "INTEGER" {
selectParts = append(selectParts, "0 AS "+name)
} else if colDef.Type == "TEXT" {
selectParts = append(selectParts, "'' AS "+name)
} else {
selectParts = append(selectParts, "NULL AS "+name)
}
}
}
return fmt.Sprintf(
"INSERT INTO %s (%s) SELECT %s FROM %s",
newTable,
strings.Join(colNames, ", "),
strings.Join(selectParts, ", "),
oldTable,
), nil
}
// ================================================================
// 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
}