package main import ( "crypto/rand" "database/sql" "encoding/hex" "log" "strings" _ "modernc.org/sqlite" ) var DB *sql.DB // ===== 全局配置表(A表)===== 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"` // frp v2 协议开关 (v2.0 启用) } // ===== 隧道表(B表)===== 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"` } // ===== 用户表(C表)===== 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 } // ---- 用户表 ---- _, 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 } // ---- 迁移:自动补全缺失的字段 ---- 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, 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 } // ---- 应用配置表 ---- _, 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 } if err := ensureJwtSecret(); err != nil { return err } 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", } func migrateGlobalConfigTable() error { // 1. 获取当前表的列名 currentCols, err := getCurrentColumns("global_config") if err != nil { return err } // 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 } added = append(added, colName) } } if len(added) > 0 { log.Printf("✅ 数据库迁移完成,新增字段: %v", added) } else { log.Println("✅ 数据库 Schema 已是最新") } return nil } // getCurrentColumns 查询指定表的所有列名 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, ¬Null, &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 }