2.2测试代码引入中

This commit is contained in:
2026-07-28 18:43:56 +08:00
parent d2b3dd8803
commit c2516dc8ff
4 changed files with 311 additions and 146 deletions
+169 -140
View File
@@ -8,6 +8,7 @@ import (
"io"
"log"
"os"
"sort"
"strings"
"time"
@@ -21,7 +22,7 @@ var DB *sql.DB
// ============================================================
const (
SchemaVersion = "2.0.0" // 当前数据库 Schema 版本,与项目版本同步
SchemaVersion = "v2" // 当前数据库 Schema 版本(与 schemaHistory 中的版本对应)
)
// ============================================================
@@ -41,10 +42,10 @@ type GlobalConfig struct {
HeartbeatInterval int `json:"heartbeatInterval"`
HeartbeatTimeout int `json:"heartbeatTimeout"`
PoolCount int `json:"poolCount"`
WireProtocolV2 bool `json:"wireProtocolV2"` // v2.0 正式启用
WireProtocolV2 bool `json:"wireProtocolV2"` // v2 协议全局开关
}
// Proxy 隧道表
// Proxy 隧道表(不含 wire_protocol_v2
type Proxy struct {
ID int `json:"id"`
Name string `json:"name"`
@@ -74,22 +75,19 @@ func InitDB() error {
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 + ")")
log.Println("✅ 数据库初始化完成 (Schema " + SchemaVersion + ")")
return nil
}
@@ -152,7 +150,7 @@ func createTables() error {
return err
}
// 应用配置表(存储 JWT 密钥、Schema 版本等)
// 应用配置表
_, err = DB.Exec(`
CREATE TABLE IF NOT EXISTS app_config (
key TEXT PRIMARY KEY,
@@ -164,7 +162,7 @@ func createTables() error {
return err
}
// 初始化默认配置(仅当表为空时)
// 初始化默认配置
var count int
DB.QueryRow("SELECT COUNT(*) FROM global_config").Scan(&count)
if count == 0 {
@@ -188,25 +186,21 @@ func createTables() error {
// 迁移引擎
// ============================================================
// getSchemaVersion 读取当前数据库的 Schema 版本
func getSchemaVersion() string {
// getCurrentSchemaVersion 读取当前数据库的 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 {
// 没有版本记录 → 首次启动或 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 "v1"
}
// 全新安装 → 直接标记为当前版本
return SchemaVersion
}
log.Printf("⚠️ 读取 Schema 版本失败: %v", err)
return "1.5.0" // 保守降级
return "v1"
}
return version
}
@@ -224,11 +218,11 @@ func setSchemaVersion(version string) error {
func backupDatabase() (string, error) {
src := "./frpc-console.db"
if _, err := os.Stat(src); os.IsNotExist(err) {
return "", nil // 数据库不存在,无需备份
return "", nil
}
timestamp := time.Now().Format("20060102_150405")
dst := fmt.Sprintf("./frpc-console.db.pre-v%s.%s", SchemaVersion, timestamp)
dst := fmt.Sprintf("./frpc-console.db.pre-%s.%s", SchemaVersion, timestamp)
srcFile, err := os.Open(src)
if err != nil {
@@ -250,68 +244,6 @@ func backupDatabase() (string, error) {
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)
@@ -334,76 +266,176 @@ func restoreDatabase(backupPath string) error {
return nil
}
// ============================================================
// 迁移函数(各版本)
// ============================================================
// runMigrations 执行迁移
func runMigrations() error {
currentVer := getCurrentSchemaVersion()
targetVer := SchemaVersion
// 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")
log.Printf("📌 当前数据库 Schema: %s, 目标版本: %s", currentVer, targetVer)
// 检查并补全 wire_protocol_v2 字段(兼容从 v1.0 直接升级的场景)
cols, err := getCurrentColumns("global_config")
if err != nil {
return fmt.Errorf("获取列信息失败: %w", err)
if currentVer == targetVer {
log.Println("✅ Schema 已是最新,无需迁移")
return nil
}
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.Printf("🔄 检测到版本变更 (%s → %s),开始迁移...", currentVer, targetVer)
backupPath, err := backupDatabase()
if err != nil {
return fmt.Errorf("备份数据库失败: %w", err)
}
if backupPath != "" {
log.Printf("📦 备份文件: %s", backupPath)
}
currentSchema := getSchemaDef(currentVer)
targetSchema := getSchemaDef(targetVer)
if targetSchema == nil {
return fmt.Errorf("目标 Schema 版本 %s 未在 schemaHistory 中定义", targetVer)
}
if currentSchema == nil || schemaVersionsEqual(currentSchema, targetSchema) {
log.Println(" 迁移类型: 轻量复制(Schema 无变更)")
} else {
log.Println(" 迁移类型: 重型迁移(Schema 有变更,新建表 + 搬数据)")
if err := heavyMigration(currentSchema, targetSchema); err != nil {
if backupPath != "" {
log.Printf("❌ 迁移失败,尝试恢复备份: %s", backupPath)
if restoreErr := restoreDatabase(backupPath); restoreErr != nil {
log.Printf("⚠️ 恢复备份失败: %v", restoreErr)
}
}
return fmt.Errorf("重型迁移失败: %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
if err := setSchemaVersion(targetVer); err != nil {
return fmt.Errorf("更新 Schema 版本失败: %w", err)
}
defer rows.Close()
log.Printf("✅ 迁移完成,当前 Schema: %s", targetVer)
return nil
}
// heavyMigration 重型迁移
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
}
// buildCreateTableSQL 根据 Schema 定义生成 CREATE TABLE 语句
func buildCreateTableSQL(tableName string, def *SchemaVersionDef) string {
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)
var primaryKey string
names := make([]string, 0, len(def.Columns))
for name := range def.Columns {
names = append(names, name)
}
return cols, rows.Err()
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 contains(slice []string, item string) bool {
for _, s := range slice {
if strings.EqualFold(s, item) {
return true
// buildInsertSQL 构建 INSERT INTO new_table SELECT ... FROM old_table
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 false
return fmt.Sprintf(
"INSERT INTO %s (%s) SELECT %s FROM %s",
newTable,
strings.Join(colNames, ", "),
strings.Join(selectParts, ", "),
oldTable,
), nil
}
// ============================================================
@@ -504,10 +536,7 @@ func GetProxies() ([]Proxy, error) {
}
proxies = append(proxies, p)
}
if err = rows.Err(); err != nil {
return nil, err
}
return proxies, nil
return proxies, rows.Err()
}
func GetProxy(id int) (*Proxy, error) {