重构进行中,preview通道尚未完全恢复全部功能
This commit is contained in:
@@ -0,0 +1,234 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
|
||||
"frpc-console/internal/db"
|
||||
)
|
||||
|
||||
var jwtSecretCache []byte
|
||||
|
||||
// getJwtSecret 从数据库获取 JWT 密钥
|
||||
func getJwtSecret() []byte {
|
||||
if len(jwtSecretCache) > 0 {
|
||||
return jwtSecretCache
|
||||
}
|
||||
secret, err := db.GetJwtSecret()
|
||||
if err != nil {
|
||||
// 如果数据库还没有密钥,生成一个默认的(仅用于开发)
|
||||
// 生产环境应该通过数据库初始化时生成
|
||||
return []byte("frpc-console-default-secret-key-2024")
|
||||
}
|
||||
jwtSecretCache = []byte(secret)
|
||||
return jwtSecretCache
|
||||
}
|
||||
|
||||
// Claims JWT 声明
|
||||
type Claims struct {
|
||||
Username string `json:"username"`
|
||||
jwt.RegisteredClaims
|
||||
}
|
||||
|
||||
// GenerateJWT 生成 JWT
|
||||
func GenerateJWT(username string) (string, error) {
|
||||
claims := Claims{
|
||||
Username: username,
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
ExpiresAt: jwt.NewNumericDate(time.Now().Add(7 * 24 * time.Hour)),
|
||||
IssuedAt: jwt.NewNumericDate(time.Now()),
|
||||
},
|
||||
}
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
return token.SignedString(getJwtSecret())
|
||||
}
|
||||
|
||||
// ParseJWT 解析 JWT
|
||||
func ParseJWT(tokenString string) (*Claims, error) {
|
||||
token, err := jwt.ParseWithClaims(tokenString, &Claims{}, func(token *jwt.Token) (interface{}, error) {
|
||||
return getJwtSecret(), nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if claims, ok := token.Claims.(*Claims); ok && token.Valid {
|
||||
return claims, nil
|
||||
}
|
||||
return nil, errors.New("invalid token")
|
||||
}
|
||||
|
||||
// ValidatePassword 校验密码复杂度
|
||||
// 要求: 至少8位,包含大小写字母、数字、特殊字符
|
||||
func ValidatePassword(pwd string) bool {
|
||||
if len(pwd) < 8 {
|
||||
return false
|
||||
}
|
||||
var hasUpper, hasLower, hasDigit, hasSpecial bool
|
||||
for _, ch := range pwd {
|
||||
switch {
|
||||
case ch >= 'A' && ch <= 'Z':
|
||||
hasUpper = true
|
||||
case ch >= 'a' && ch <= 'z':
|
||||
hasLower = true
|
||||
case ch >= '0' && ch <= '9':
|
||||
hasDigit = true
|
||||
case strings.ContainsAny(string(ch), "!@#$%^&*()_+-=[]{}|;:,.<>?"):
|
||||
hasSpecial = true
|
||||
}
|
||||
}
|
||||
return hasUpper && hasLower && hasDigit && hasSpecial
|
||||
}
|
||||
|
||||
// HashPassword 加密密码
|
||||
func HashPassword(password string) (string, error) {
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(hash), nil
|
||||
}
|
||||
|
||||
// CheckPassword 验证密码
|
||||
func CheckPassword(password, hash string) bool {
|
||||
err := bcrypt.CompareHashAndPassword([]byte(hash), []byte(password))
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// ================================================================
|
||||
// 登录 / 注册 / 修改密码 的业务逻辑 (供 handler 调用)
|
||||
// ================================================================
|
||||
|
||||
// LoginRequest 登录请求
|
||||
type LoginRequest struct {
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
}
|
||||
|
||||
// LoginResponse 登录响应
|
||||
type LoginResponse struct {
|
||||
Token string `json:"token"`
|
||||
}
|
||||
|
||||
// Login 登录业务逻辑
|
||||
func Login(username, password string) (string, error) {
|
||||
user, err := db.GetUserByUsername(username)
|
||||
if err != nil {
|
||||
return "", errors.New("用户不存在")
|
||||
}
|
||||
if !CheckPassword(password, user.PasswordHash) {
|
||||
return "", errors.New("密码错误")
|
||||
}
|
||||
return GenerateJWT(username)
|
||||
}
|
||||
|
||||
// RegisterRequest 注册请求
|
||||
type RegisterRequest struct {
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
}
|
||||
|
||||
// Register 注册业务逻辑
|
||||
func Register(username, password string) (string, error) {
|
||||
// 检查用户名是否已存在
|
||||
exists, err := db.UserExists(username)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if exists {
|
||||
return "", errors.New("用户名已存在")
|
||||
}
|
||||
// 密码复杂度校验
|
||||
if !ValidatePassword(password) {
|
||||
return "", errors.New("密码不符合复杂度要求")
|
||||
}
|
||||
// 加密密码
|
||||
hash, err := HashPassword(password)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
// 创建用户
|
||||
if err := db.CreateUser(username, hash); err != nil {
|
||||
return "", err
|
||||
}
|
||||
// 生成 JWT
|
||||
return GenerateJWT(username)
|
||||
}
|
||||
|
||||
// ChangePasswordRequest 修改密码请求
|
||||
type ChangePasswordRequest struct {
|
||||
OldPassword string `json:"oldPassword"`
|
||||
NewPassword string `json:"newPassword"`
|
||||
}
|
||||
|
||||
// ChangePassword 修改密码业务逻辑
|
||||
func ChangePassword(username, oldPassword, newPassword string) error {
|
||||
user, err := db.GetUserByUsername(username)
|
||||
if err != nil {
|
||||
return errors.New("用户不存在")
|
||||
}
|
||||
if !CheckPassword(oldPassword, user.PasswordHash) {
|
||||
return errors.New("原密码错误")
|
||||
}
|
||||
if !ValidatePassword(newPassword) {
|
||||
return errors.New("新密码不符合复杂度要求")
|
||||
}
|
||||
hash, err := HashPassword(newPassword)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return db.UpdateUserPassword(username, hash)
|
||||
}
|
||||
|
||||
// HasUsers 检查是否存在用户
|
||||
func HasUsers() (bool, error) {
|
||||
count, err := db.CountUsers()
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
// ================================================================
|
||||
// Gin 中间件
|
||||
// ================================================================
|
||||
|
||||
// AuthMiddleware JWT 认证中间件
|
||||
func AuthMiddleware() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
authHeader := c.GetHeader("Authorization")
|
||||
if authHeader == "" {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"code": 1, "msg": "未提供认证令牌"})
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
parts := strings.SplitN(authHeader, " ", 2)
|
||||
if len(parts) != 2 || parts[0] != "Bearer" {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"code": 1, "msg": "认证令牌格式错误"})
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
claims, err := ParseJWT(parts[1])
|
||||
if err != nil {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"code": 1, "msg": "无效的认证令牌"})
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
c.Set("username", claims.Username)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// GetUsernameFromContext 从 gin.Context 获取当前用户名
|
||||
func GetUsernameFromContext(c *gin.Context) string {
|
||||
username, exists := c.Get("username")
|
||||
if !exists {
|
||||
return ""
|
||||
}
|
||||
return username.(string)
|
||||
}
|
||||
@@ -0,0 +1,424 @@
|
||||
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 = "v2"
|
||||
|
||||
// ================================================================
|
||||
// 初始化
|
||||
// ================================================================
|
||||
|
||||
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 {
|
||||
_, 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
|
||||
}
|
||||
|
||||
_, 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
|
||||
}
|
||||
|
||||
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("✅ 全局配置初始化完成")
|
||||
}
|
||||
|
||||
return 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
|
||||
}
|
||||
|
||||
// ================================================================
|
||||
// 迁移引擎
|
||||
// ================================================================
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
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 未定义", 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
|
||||
|
||||
for name, col := range def.Columns {
|
||||
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) {
|
||||
var newCols []string
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,36 @@
|
||||
package db
|
||||
|
||||
// GlobalConfig 全局配置
|
||||
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"`
|
||||
}
|
||||
|
||||
// Proxy 隧道配置
|
||||
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"`
|
||||
}
|
||||
|
||||
// User 用户
|
||||
type User struct {
|
||||
ID int `json:"id"`
|
||||
Username string `json:"username"`
|
||||
PasswordHash string `json:"-"`
|
||||
CreatedAt string `json:"createdAt"`
|
||||
}
|
||||
@@ -0,0 +1,149 @@
|
||||
package db
|
||||
|
||||
// ================================================================
|
||||
// 全局配置
|
||||
// ================================================================
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
// ================================================================
|
||||
// 隧道管理
|
||||
// ================================================================
|
||||
|
||||
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)
|
||||
}
|
||||
return proxies, rows.Err()
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
// ================================================================
|
||||
// 用户管理
|
||||
// ================================================================
|
||||
|
||||
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 UserExists(username string) (bool, error) {
|
||||
var count int
|
||||
err := DB.QueryRow("SELECT COUNT(*) FROM users WHERE username = ?", username).Scan(&count)
|
||||
return count > 0, err
|
||||
}
|
||||
|
||||
func CreateUser(username, passwordHash string) error {
|
||||
_, err := DB.Exec(`
|
||||
INSERT INTO users (username, password_hash) VALUES (?, ?)
|
||||
`, username, passwordHash)
|
||||
return err
|
||||
}
|
||||
|
||||
func UpdateUserPassword(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
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
package db
|
||||
|
||||
// ================================================================
|
||||
// Schema 版本声明
|
||||
// ================================================================
|
||||
|
||||
type SchemaVersionDef struct {
|
||||
Version string
|
||||
TableName string
|
||||
Columns map[string]ColumnDef
|
||||
}
|
||||
|
||||
type ColumnDef struct {
|
||||
Type string
|
||||
NotNull bool
|
||||
Default string
|
||||
Primary bool
|
||||
}
|
||||
|
||||
var schemaHistory = []SchemaVersionDef{
|
||||
// v1: 初始版本
|
||||
{
|
||||
Version: "v1",
|
||||
TableName: "proxies",
|
||||
Columns: map[string]ColumnDef{
|
||||
"id": {Type: "INTEGER", Primary: true},
|
||||
"name": {Type: "TEXT", NotNull: true},
|
||||
"type": {Type: "TEXT", NotNull: true, Default: "'tcp'"},
|
||||
"local_ip": {Type: "TEXT", NotNull: true},
|
||||
"local_port": {Type: "INTEGER", NotNull: true},
|
||||
"remote_port": {Type: "INTEGER", NotNull: true},
|
||||
"enabled": {Type: "INTEGER", NotNull: true, Default: "1"},
|
||||
"created_at": {Type: "DATETIME", Default: "CURRENT_TIMESTAMP"},
|
||||
"updated_at": {Type: "DATETIME", Default: "CURRENT_TIMESTAMP"},
|
||||
},
|
||||
},
|
||||
// v2: 当前版本
|
||||
{
|
||||
Version: "v2",
|
||||
TableName: "proxies",
|
||||
Columns: map[string]ColumnDef{
|
||||
"id": {Type: "INTEGER", Primary: true},
|
||||
"name": {Type: "TEXT", NotNull: true},
|
||||
"type": {Type: "TEXT", NotNull: true, Default: "'tcp'"},
|
||||
"local_ip": {Type: "TEXT", NotNull: true},
|
||||
"local_port": {Type: "INTEGER", NotNull: true},
|
||||
"remote_port": {Type: "INTEGER", NotNull: true},
|
||||
"enabled": {Type: "INTEGER", NotNull: true, Default: "1"},
|
||||
"created_at": {Type: "DATETIME", Default: "CURRENT_TIMESTAMP"},
|
||||
"updated_at": {Type: "DATETIME", Default: "CURRENT_TIMESTAMP"},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
func getSchemaDef(version string) *SchemaVersionDef {
|
||||
for _, def := range schemaHistory {
|
||||
if def.Version == version {
|
||||
return &def
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func getLatestSchemaDef() *SchemaVersionDef {
|
||||
if len(schemaHistory) == 0 {
|
||||
return nil
|
||||
}
|
||||
return &schemaHistory[len(schemaHistory)-1]
|
||||
}
|
||||
|
||||
func schemaVersionsEqual(v1, v2 *SchemaVersionDef) bool {
|
||||
if v1 == nil || v2 == nil {
|
||||
return false
|
||||
}
|
||||
if v1.TableName != v2.TableName {
|
||||
return false
|
||||
}
|
||||
if len(v1.Columns) != len(v2.Columns) {
|
||||
return false
|
||||
}
|
||||
for name, col1 := range v1.Columns {
|
||||
col2, ok := v2.Columns[name]
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
if col1.Type != col2.Type || col1.NotNull != col2.NotNull || col1.Default != col2.Default {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
package frp
|
||||
|
||||
import (
|
||||
"embed"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"sync"
|
||||
)
|
||||
|
||||
//go:embed bin/*
|
||||
var embeddedFrpc embed.FS
|
||||
|
||||
var (
|
||||
cachedFrpcPath string
|
||||
frpcPathMutex sync.Mutex
|
||||
)
|
||||
|
||||
// GetFrpcPath 获取 frpc 二进制路径
|
||||
// 优先级: 本地缓存 > 内嵌二进制 > 系统 PATH
|
||||
func GetFrpcPath() (string, error) {
|
||||
frpcPathMutex.Lock()
|
||||
defer frpcPathMutex.Unlock()
|
||||
|
||||
if cachedFrpcPath != "" {
|
||||
if _, err := os.Stat(cachedFrpcPath); err == nil {
|
||||
return cachedFrpcPath, nil
|
||||
}
|
||||
cachedFrpcPath = ""
|
||||
}
|
||||
|
||||
var fileName string
|
||||
switch {
|
||||
case runtime.GOOS == "windows" && runtime.GOARCH == "amd64":
|
||||
fileName = "frpc_windows_amd64.exe"
|
||||
case runtime.GOOS == "linux" && runtime.GOARCH == "amd64":
|
||||
fileName = "frpc_linux_amd64"
|
||||
case runtime.GOOS == "linux" && runtime.GOARCH == "arm64":
|
||||
fileName = "frpc_linux_arm64"
|
||||
case runtime.GOOS == "linux" && runtime.GOARCH == "arm":
|
||||
fileName = "frpc_linux_arm_hf"
|
||||
default:
|
||||
path, err := exec.LookPath("frpc")
|
||||
if err == nil {
|
||||
cachedFrpcPath = path
|
||||
return path, nil
|
||||
}
|
||||
return "", fmt.Errorf("不支持的平台: %s/%s", runtime.GOOS, runtime.GOARCH)
|
||||
}
|
||||
|
||||
// 尝试从本地 bin 目录加载
|
||||
localPath := filepath.Join(".", "bin", fileName)
|
||||
if _, err := os.Stat(localPath); err == nil {
|
||||
cachedFrpcPath = localPath
|
||||
return localPath, nil
|
||||
}
|
||||
|
||||
// 尝试从 embed 提取到临时目录
|
||||
data, err := embeddedFrpc.ReadFile("bin/" + fileName)
|
||||
if err == nil {
|
||||
tmpPath := filepath.Join(os.TempDir(), "frpc")
|
||||
if runtime.GOOS == "windows" {
|
||||
tmpPath += ".exe"
|
||||
}
|
||||
if err := os.WriteFile(tmpPath, data, 0755); err == nil {
|
||||
cachedFrpcPath = tmpPath
|
||||
return tmpPath, nil
|
||||
}
|
||||
if _, statErr := os.Stat(tmpPath); statErr == nil {
|
||||
cachedFrpcPath = tmpPath
|
||||
return tmpPath, nil
|
||||
}
|
||||
}
|
||||
|
||||
// 最后尝试从系统 PATH 查找
|
||||
path, err := exec.LookPath("frpc")
|
||||
if err == nil {
|
||||
cachedFrpcPath = path
|
||||
return path, nil
|
||||
}
|
||||
|
||||
return "", fmt.Errorf("未找到 frpc 文件")
|
||||
}
|
||||
@@ -0,0 +1,285 @@
|
||||
package frp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
"os"
|
||||
"os/exec"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
"syscall"
|
||||
|
||||
"frpc-console/internal/process"
|
||||
)
|
||||
|
||||
// ================================================================
|
||||
// 兼容层:保持对外接口不变
|
||||
// 这些函数供 api 和外部调用,实际委托给 process.Manager
|
||||
// ================================================================
|
||||
|
||||
// IsRunning 检查 frpc 是否在运行
|
||||
func IsRunning() bool {
|
||||
pm := process.GetGlobalManager()
|
||||
if pm != nil {
|
||||
status, err := pm.Status()
|
||||
if err != nil {
|
||||
log.Printf("[WARN] ProcessManager.Status() 失败: %v,降级到 PID 文件", err)
|
||||
return isRunningLegacy()
|
||||
}
|
||||
return status.State == "running"
|
||||
}
|
||||
return isRunningLegacy()
|
||||
}
|
||||
|
||||
// Start 启动 frpc (幂等)
|
||||
func Start() error {
|
||||
pm := process.GetGlobalManager()
|
||||
if pm != nil {
|
||||
log.Println("[INFO] 使用 ProcessManager 启动 frpc")
|
||||
return pm.Start(context.Background())
|
||||
}
|
||||
log.Println("[WARN] ProcessManager 未初始化,使用兼容模式启动 frpc")
|
||||
return startLegacy()
|
||||
}
|
||||
|
||||
// Stop 停止 frpc (幂等)
|
||||
func Stop() error {
|
||||
pm := process.GetGlobalManager()
|
||||
if pm != nil {
|
||||
log.Println("[INFO] 使用 ProcessManager 停止 frpc")
|
||||
return pm.Stop(context.Background())
|
||||
}
|
||||
log.Println("[WARN] ProcessManager 未初始化,使用兼容模式停止 frpc")
|
||||
return stopLegacy()
|
||||
}
|
||||
|
||||
// Restart 重启 frpc (原子操作)
|
||||
func Restart() error {
|
||||
pm := process.GetGlobalManager()
|
||||
if pm != nil {
|
||||
log.Println("[INFO] 使用 ProcessManager 重启 frpc")
|
||||
return pm.Restart(context.Background())
|
||||
}
|
||||
log.Println("[WARN] ProcessManager 未初始化,使用兼容模式重启 frpc")
|
||||
if err := stopLegacy(); err != nil {
|
||||
return err
|
||||
}
|
||||
return startLegacy()
|
||||
}
|
||||
|
||||
// GetStatus 获取 frpc 详细状态 (供 API 调用)
|
||||
func GetStatus() (map[string]interface{}, error) {
|
||||
pm := process.GetGlobalManager()
|
||||
if pm != nil {
|
||||
status, err := pm.Status()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return map[string]interface{}{
|
||||
"state": status.State,
|
||||
"pid": status.PID,
|
||||
"port": status.Port,
|
||||
}, nil
|
||||
}
|
||||
|
||||
running := isRunningLegacy()
|
||||
return map[string]interface{}{
|
||||
"state": map[bool]string{true: "running", false: "stopped"}[running],
|
||||
"pid": 0,
|
||||
"port": 0,
|
||||
"legacy": true,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Reload 热加载 frpc 配置
|
||||
func Reload() error {
|
||||
pm := process.GetGlobalManager()
|
||||
if pm == nil {
|
||||
return reloadLegacy()
|
||||
}
|
||||
|
||||
status, err := pm.Status()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if status.State != "running" {
|
||||
return pm.Start(context.Background())
|
||||
}
|
||||
|
||||
frpcPath, err := GetFrpcPath()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
cmd := exec.Command(frpcPath, "reload", "-c", "./data/frpc.toml")
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
log.Printf("⚠️ 热加载失败 (%v),降级为重启 frpc", err)
|
||||
log.Printf(" reload 输出: %s", string(output))
|
||||
if err := pm.Restart(context.Background()); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
log.Printf("✅ frpc 热加载成功: %s", string(output))
|
||||
return nil
|
||||
}
|
||||
|
||||
// ================================================================
|
||||
// 旧版实现 (降级方案)
|
||||
// ================================================================
|
||||
|
||||
func isRunningLegacy() bool {
|
||||
pidData, err := os.ReadFile("./data/frpc.pid")
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
pid, err := strconv.Atoi(strings.TrimSpace(string(pidData)))
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
|
||||
if runtime.GOOS == "windows" {
|
||||
cmd := exec.Command("tasklist", "/FI", "PID eq", strconv.Itoa(pid))
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(string(output), strconv.Itoa(pid))
|
||||
}
|
||||
|
||||
proc, err := os.FindProcess(pid)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return proc.Signal(syscall.Signal(0)) == nil
|
||||
}
|
||||
|
||||
func startLegacy() error {
|
||||
frpcPath, err := GetFrpcPath()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := os.MkdirAll("./data", 0755); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if _, err := os.Stat("./data/frpc.toml"); os.IsNotExist(err) {
|
||||
if err := GenerateConfig(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
if isRunningLegacy() {
|
||||
return nil
|
||||
}
|
||||
|
||||
os.Remove("./data/frpc.pid")
|
||||
|
||||
cmd := exec.Command(frpcPath, "-c", "./data/frpc.toml")
|
||||
setWindowHide(cmd)
|
||||
setSysProcAttr(cmd)
|
||||
|
||||
logFile, err := os.OpenFile("./data/frpc.log", os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0644)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cmd.Stdout = logFile
|
||||
cmd.Stderr = logFile
|
||||
|
||||
if err := cmd.Start(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
go func() {
|
||||
if err := cmd.Wait(); err != nil {
|
||||
log.Printf("frpc 子进程退出: %v", err)
|
||||
}
|
||||
os.Remove("./data/frpc.pid")
|
||||
}()
|
||||
|
||||
return os.WriteFile("./data/frpc.pid", []byte(strconv.Itoa(cmd.Process.Pid)), 0644)
|
||||
}
|
||||
|
||||
func stopLegacy() error {
|
||||
if runtime.GOOS == "windows" {
|
||||
cmd := exec.Command("taskkill", "/F", "/IM", "frpc.exe")
|
||||
if err := cmd.Run(); err != nil && !strings.Contains(err.Error(), "not found") {
|
||||
return err
|
||||
}
|
||||
os.Remove("./data/frpc.pid")
|
||||
return nil
|
||||
}
|
||||
|
||||
pidData, err := os.ReadFile("./data/frpc.pid")
|
||||
if err != nil {
|
||||
cmd := exec.Command("pkill", "-f", "frpc")
|
||||
if err := cmd.Run(); err != nil && !strings.Contains(err.Error(), "no process") {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
pid, _ := strconv.Atoi(strings.TrimSpace(string(pidData)))
|
||||
proc, err := os.FindProcess(pid)
|
||||
if err != nil {
|
||||
os.Remove("./data/frpc.pid")
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := proc.Kill(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
os.Remove("./data/frpc.pid")
|
||||
return nil
|
||||
}
|
||||
|
||||
func reloadLegacy() error {
|
||||
if !isRunningLegacy() {
|
||||
return startLegacy()
|
||||
}
|
||||
|
||||
frpcPath, err := GetFrpcPath()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
cmd := exec.Command(frpcPath, "reload", "-c", "./data/frpc.toml")
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
log.Printf("⚠️ 热加载失败 (%v),降级为重启 frpc", err)
|
||||
log.Printf(" reload 输出: %s", string(output))
|
||||
if stopErr := stopLegacy(); stopErr != nil {
|
||||
return stopErr
|
||||
}
|
||||
return startLegacy()
|
||||
}
|
||||
|
||||
log.Printf("✅ frpc 热加载成功 (兼容模式): %s", string(output))
|
||||
return nil
|
||||
}
|
||||
|
||||
// ================================================================
|
||||
// 辅助函数 (平台相关)
|
||||
// ================================================================
|
||||
|
||||
func setWindowHide(cmd *exec.Cmd) {
|
||||
if runtime.GOOS == "windows" {
|
||||
cmd.SysProcAttr = &syscall.SysProcAttr{
|
||||
HideWindow: true,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func setSysProcAttr(cmd *exec.Cmd) {
|
||||
if runtime.GOOS != "windows" {
|
||||
cmd.SysProcAttr = &syscall.SysProcAttr{
|
||||
Setpgid: true,
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user