548 lines
14 KiB
Go
548 lines
14 KiB
Go
package main
|
||
|
||
import (
|
||
"crypto/rand"
|
||
"database/sql"
|
||
"encoding/hex"
|
||
"fmt"
|
||
"io"
|
||
"log"
|
||
"os"
|
||
"sort"
|
||
"strings"
|
||
"time"
|
||
|
||
_ "modernc.org/sqlite"
|
||
)
|
||
|
||
var DB *sql.DB
|
||
|
||
const SchemaVersion = "v2"
|
||
|
||
// GlobalConfig 对应 frps 配置
|
||
type GlobalConfig struct {
|
||
ID int `json:"id"`
|
||
BindPort int `json:"bindPort"`
|
||
Token string `json:"token"`
|
||
LogLevel string `json:"logLevel"`
|
||
LogMaxDays int `json:"logMaxDays"`
|
||
TcpMux bool `json:"tcpMux"`
|
||
}
|
||
|
||
// User 保持不变
|
||
type User struct {
|
||
ID int `json:"id"`
|
||
Username string `json:"username"`
|
||
PasswordHash string `json:"-"`
|
||
CreatedAt string `json:"createdAt"`
|
||
}
|
||
|
||
// Client 记录(从 Dashboard 同步,可选)
|
||
type Client struct {
|
||
ID int `json:"id"`
|
||
Name string `json:"name"`
|
||
OS string `json:"os"`
|
||
Arch string `json:"arch"`
|
||
Version string `json:"version"`
|
||
Status string `json:"status"`
|
||
ConnTime string `json:"connTime"`
|
||
UpdatedAt string `json:"updatedAt"`
|
||
}
|
||
|
||
func InitDB() error {
|
||
// 确保 data 目录存在
|
||
if err := os.MkdirAll("./data", 0755); err != nil {
|
||
return fmt.Errorf("创建数据目录失败: %w", err)
|
||
}
|
||
|
||
dbPath := "./data/frps-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 表(frps 专用)
|
||
_, err = DB.Exec(`
|
||
CREATE TABLE IF NOT EXISTS global_config (
|
||
id INTEGER PRIMARY KEY CHECK (id = 1),
|
||
bind_port INTEGER NOT NULL DEFAULT 9358,
|
||
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,
|
||
updated_at DATETIME DEFAULT CURRENT_TIMESTAMP
|
||
)
|
||
`)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
|
||
// clients 表(用于存储从 Dashboard 拉取的客户端,可选)
|
||
_, err = DB.Exec(`
|
||
CREATE TABLE IF NOT EXISTS clients (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
name TEXT,
|
||
os TEXT,
|
||
arch TEXT,
|
||
version TEXT,
|
||
status TEXT,
|
||
conn_time DATETIME,
|
||
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
|
||
}
|
||
|
||
// 初始化默认配置
|
||
var count int
|
||
DB.QueryRow("SELECT COUNT(*) FROM global_config").Scan(&count)
|
||
if count == 0 {
|
||
_, err = DB.Exec(`
|
||
INSERT INTO global_config (id, bind_port, token, log_level, log_max_days, tcp_mux)
|
||
VALUES (1, 9358, 'CHANGE_ME', 'info', 3, 1)
|
||
`)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
log.Println("✅ 全局配置初始化完成")
|
||
}
|
||
|
||
return nil
|
||
}
|
||
|
||
// ========== 迁移引擎(复用 frps 的 db-history.go) ==========
|
||
|
||
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/frps-console.db"
|
||
if _, err := os.Stat(src); os.IsNotExist(err) {
|
||
return "", nil
|
||
}
|
||
|
||
timestamp := time.Now().Format("20060102_150405")
|
||
dst := fmt.Sprintf("./data/frps-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/frps-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
|
||
}
|
||
|
||
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)
|
||
}
|
||
|
||
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 无变更)")
|
||
var userCount int
|
||
err := DB.QueryRow("SELECT COUNT(*) FROM users").Scan(&userCount)
|
||
if err != nil || userCount == 0 {
|
||
log.Println(" 数据库为空或无效,跳过迁移,直接初始化")
|
||
return nil
|
||
}
|
||
log.Println(" 数据库有效,继续使用")
|
||
} 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)
|
||
}
|
||
}
|
||
|
||
if err := setSchemaVersion(targetVer); err != nil {
|
||
return fmt.Errorf("更新 Schema 版本失败: %w", err)
|
||
}
|
||
|
||
log.Printf("✅ 迁移完成,当前 Schema: %s", targetVer)
|
||
return nil
|
||
}
|
||
|
||
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
|
||
}
|
||
|
||
// ========== GlobalConfig CRUD ==========
|
||
|
||
func GetGlobalConfig() (*GlobalConfig, error) {
|
||
var cfg GlobalConfig
|
||
err := DB.QueryRow(`
|
||
SELECT id, bind_port, token, log_level, log_max_days, tcp_mux
|
||
FROM global_config WHERE id = 1
|
||
`).Scan(
|
||
&cfg.ID, &cfg.BindPort, &cfg.Token,
|
||
&cfg.LogLevel, &cfg.LogMaxDays,
|
||
&cfg.TcpMux,
|
||
)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
cfg.TcpMux = true
|
||
return &cfg, nil
|
||
}
|
||
|
||
func UpdateGlobalConfig(cfg *GlobalConfig) error {
|
||
_, err := DB.Exec(`
|
||
UPDATE global_config SET
|
||
bind_port = ?, token = ?, log_level = ?, log_max_days = ?,
|
||
tcp_mux = 1,
|
||
updated_at = CURRENT_TIMESTAMP
|
||
WHERE id = 1
|
||
`, cfg.BindPort, cfg.Token, cfg.LogLevel, cfg.LogMaxDays)
|
||
return err
|
||
}
|
||
|
||
// ========== Clients CRUD(用于同步 Dashboard) ==========
|
||
|
||
func SaveClient(client *Client) error {
|
||
_, err := DB.Exec(`
|
||
INSERT INTO clients (name, os, arch, version, status, conn_time)
|
||
VALUES (?, ?, ?, ?, ?, ?)
|
||
`, client.Name, client.OS, client.Arch, client.Version, client.Status, client.ConnTime)
|
||
return err
|
||
}
|
||
|
||
func GetClients() ([]Client, error) {
|
||
rows, err := DB.Query(`
|
||
SELECT id, name, os, arch, version, status, conn_time, updated_at
|
||
FROM clients ORDER BY id DESC LIMIT 100
|
||
`)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer rows.Close()
|
||
|
||
var clients []Client
|
||
for rows.Next() {
|
||
var c Client
|
||
err := rows.Scan(&c.ID, &c.Name, &c.OS, &c.Arch, &c.Version, &c.Status, &c.ConnTime, &c.UpdatedAt)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
clients = append(clients, c)
|
||
}
|
||
return clients, rows.Err()
|
||
}
|
||
|
||
// ========== User 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
|
||
}
|