第一份2.0-LTS代码上传
This commit is contained in:
@@ -4,15 +4,31 @@ import (
|
||||
"crypto/rand"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
var DB *sql.DB
|
||||
|
||||
// ===== 全局配置表(A表)=====
|
||||
// ============================================================
|
||||
// 数据库版本常量
|
||||
// ============================================================
|
||||
|
||||
const (
|
||||
SchemaVersion = "2.0.0" // 当前数据库 Schema 版本,与项目版本同步
|
||||
)
|
||||
|
||||
// ============================================================
|
||||
// 数据模型
|
||||
// ============================================================
|
||||
|
||||
// GlobalConfig 全局配置表
|
||||
type GlobalConfig struct {
|
||||
ID int `json:"id"`
|
||||
ServerAddr string `json:"serverAddr"`
|
||||
@@ -25,10 +41,10 @@ type GlobalConfig struct {
|
||||
HeartbeatInterval int `json:"heartbeatInterval"`
|
||||
HeartbeatTimeout int `json:"heartbeatTimeout"`
|
||||
PoolCount int `json:"poolCount"`
|
||||
WireProtocolV2 bool `json:"wireProtocolV2"` // frp v2 协议开关 (v2.0 启用)
|
||||
WireProtocolV2 bool `json:"wireProtocolV2"` // v2.0 正式启用
|
||||
}
|
||||
|
||||
// ===== 隧道表(B表)=====
|
||||
// Proxy 隧道表
|
||||
type Proxy struct {
|
||||
ID int `json:"id"`
|
||||
Name string `json:"name"`
|
||||
@@ -39,7 +55,7 @@ type Proxy struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
// ===== 用户表(C表)=====
|
||||
// User 用户表
|
||||
type User struct {
|
||||
ID int `json:"id"`
|
||||
Username string `json:"username"`
|
||||
@@ -58,8 +74,32 @@ func InitDB() error {
|
||||
return err
|
||||
}
|
||||
|
||||
// ---- 用户表 ----
|
||||
_, err = DB.Exec(`
|
||||
// ---- 创建所有表 ----
|
||||
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,
|
||||
@@ -72,7 +112,7 @@ func InitDB() error {
|
||||
return err
|
||||
}
|
||||
|
||||
// ---- 全局配置表 ----
|
||||
// 全局配置表
|
||||
_, err = DB.Exec(`
|
||||
CREATE TABLE IF NOT EXISTS global_config (
|
||||
id INTEGER PRIMARY KEY CHECK (id = 1),
|
||||
@@ -94,29 +134,7 @@ func InitDB() error {
|
||||
return err
|
||||
}
|
||||
|
||||
// ---- 迁移:自动补全缺失的字段 ----
|
||||
if err := migrateGlobalConfigTable(); err != nil {
|
||||
log.Printf("⚠️ 迁移警告: %v", 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("✅ 全局配置初始化完成")
|
||||
}
|
||||
|
||||
// ---- 隧道表 ----
|
||||
// 隧道表
|
||||
_, err = DB.Exec(`
|
||||
CREATE TABLE IF NOT EXISTS proxies (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
@@ -134,7 +152,7 @@ func InitDB() error {
|
||||
return err
|
||||
}
|
||||
|
||||
// ---- 应用配置表 ----
|
||||
// 应用配置表(存储 JWT 密钥、Schema 版本等)
|
||||
_, err = DB.Exec(`
|
||||
CREATE TABLE IF NOT EXISTS app_config (
|
||||
key TEXT PRIMARY KEY,
|
||||
@@ -146,66 +164,214 @@ func InitDB() error {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := ensureJwtSecret(); 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("✅ 全局配置初始化完成")
|
||||
}
|
||||
|
||||
log.Println("✅ 数据库初始化完成")
|
||||
return nil
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// Schema 迁移:自动补全缺失的字段
|
||||
// 迁移引擎
|
||||
// ============================================================
|
||||
|
||||
// expectedColumns 定义所有期望的列,格式: 列名 => 完整定义
|
||||
// 新增字段时,在这里加一行即可
|
||||
var expectedColumns = map[string]string{
|
||||
"id": "id INTEGER PRIMARY KEY CHECK (id = 1)",
|
||||
"server_addr": "server_addr TEXT NOT NULL DEFAULT 'frp.example.com'",
|
||||
"server_port": "server_port INTEGER NOT NULL DEFAULT 7000",
|
||||
"token": "token TEXT NOT NULL DEFAULT 'CHANGE_ME'",
|
||||
"log_level": "log_level TEXT NOT NULL DEFAULT 'info'",
|
||||
"log_max_days": "log_max_days INTEGER NOT NULL DEFAULT 3",
|
||||
"tcp_mux": "tcp_mux INTEGER NOT NULL DEFAULT 1",
|
||||
"tcp_mux_keepalive": "tcp_mux_keepalive INTEGER NOT NULL DEFAULT 30",
|
||||
"heartbeat_interval": "heartbeat_interval INTEGER NOT NULL DEFAULT 15",
|
||||
"heartbeat_timeout": "heartbeat_timeout INTEGER NOT NULL DEFAULT 70",
|
||||
"pool_count": "pool_count INTEGER NOT NULL DEFAULT 8",
|
||||
"wire_protocol_v2": "wire_protocol_v2 INTEGER NOT NULL DEFAULT 0", // v1.5 新增
|
||||
"updated_at": "updated_at DATETIME DEFAULT CURRENT_TIMESTAMP",
|
||||
// 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
|
||||
}
|
||||
|
||||
func migrateGlobalConfigTable() error {
|
||||
// 1. 获取当前表的列名
|
||||
currentCols, err := getCurrentColumns("global_config")
|
||||
if err != nil {
|
||||
return err
|
||||
// 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 // 数据库不存在,无需备份
|
||||
}
|
||||
|
||||
// 2. 比对并补全缺失的列
|
||||
var added []string
|
||||
for colName, colDef := range expectedColumns {
|
||||
if !contains(currentCols, colName) {
|
||||
alterSQL := "ALTER TABLE global_config ADD COLUMN " + colDef
|
||||
if _, err := DB.Exec(alterSQL); err != nil {
|
||||
log.Printf("⚠️ 添加列 %s 失败: %v", colName, err)
|
||||
continue
|
||||
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)
|
||||
}
|
||||
added = append(added, colName)
|
||||
applied++
|
||||
currentVersion = SchemaVersion
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if len(added) > 0 {
|
||||
log.Printf("✅ 数据库迁移完成,新增字段: %v", added)
|
||||
} else {
|
||||
log.Println("✅ 数据库 Schema 已是最新")
|
||||
// ---- 3. 更新 Schema 版本 ----
|
||||
if err := setSchemaVersion(SchemaVersion); err != nil {
|
||||
return fmt.Errorf("更新 Schema 版本失败: %w", err)
|
||||
}
|
||||
|
||||
log.Printf("✅ 迁移完成,应用了 %d 个迁移", applied)
|
||||
return nil
|
||||
}
|
||||
|
||||
// getCurrentColumns 查询指定表的所有列名
|
||||
// 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 {
|
||||
@@ -263,7 +429,7 @@ func ensureJwtSecret() error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
log.Printf("✅ JWT 密钥已生成并保存到数据库")
|
||||
log.Printf("✅ JWT 密钥已生成")
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user