From 61189a53eca2389efefc17c52e2a1b8acc9898b8 Mon Sep 17 00:00:00 2001 From: kevin Date: Fri, 3 Jul 2026 14:12:03 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E4=BF=AE=E5=A4=8D=E6=9D=83=E9=99=90?= =?UTF-8?q?=E8=BE=B9=E7=95=8C=E4=B8=8E=E5=AE=89=E5=85=A8=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - JWT 密钥改为环境变量/配置文件/随机生成注入,移除硬编码 - AuthMiddleware 用数据库 role 覆盖 claims,降级/禁用即时生效 - 改 role/status 时主动失效目标用户会话 - /ws 支持 JWT 认证(?token=),兼容密码认证,CheckOrigin 改同源校验 - 默认管理员改用随机密码,stdout 与文件双交付 - 登录与 WS 密码认证加限流(5次/分钟) - role 字段白名单校验(仅 admin/user) - socket/目录权限可配置(sock_mode/sock_group/sock_dir_mode),默认兼容现状 - 配置文件权限收紧 0644→0600,目录 0755→0700 - 静态文件改用 http.Dir.Open + NoRoute,移除手动路径拼接 - 隧道加 SetReadLimit(1MB) 与单用户连接数上限(3) - README 补充安全配置与生产部署说明 --- README.md | 53 ++++++++++++++++++++ internal/config/config.go | 52 ++++++++++++++++--- internal/db/db.go | 44 ++++++++++++++-- internal/handler/user.go | 16 +++++- internal/middleware/auth.go | 17 ++++--- internal/middleware/ratelimit.go | 86 ++++++++++++++++++++++++++++++++ internal/router/router.go | 25 +++++----- internal/vpn/auth.go | 27 +++++++--- internal/vpn/handler.go | 34 ++++++++++++- internal/vpn/tunnel.go | 34 +++++++++++-- main.go | 42 ++++++++++++++-- 11 files changed, 382 insertions(+), 48 deletions(-) create mode 100644 internal/middleware/ratelimit.go diff --git a/README.md b/README.md index a05badb..28ebce1 100644 --- a/README.md +++ b/README.md @@ -15,3 +15,56 @@ - `pytest/` - Python 测试脚本 - `main.go` - Go 服务端入口文件 - `go.mod` / `go.sum` - Go 模块依赖管理 + +## 安全配置 + +### JWT 密钥 + +JWT 密钥按以下优先级加载: + +1. 环境变量 `LMVPN_JWT_SECRET` +2. 配置文件 `data/config.yml` 中的 `web.jwt_secret` +3. 首次启动时自动生成 32 字节随机密钥并写入配置文件 + +生产环境建议通过环境变量注入,避免密钥落盘。 + +### 默认管理员 + +首次启动时自动创建管理员账户,密码为随机生成的 16 位字符串。密码会: + +- 打印到 stdout(仅一次) +- 写入 `data/.initial_admin_password`(权限 0600) + +请登录后立即修改密码,并删除 `data/.initial_admin_password` 文件。 + +### 登录限流 + +`/api/login` 和 WebSocket 密码认证均限制每 IP 5 次/分钟。 + +### Unix Socket 权限 + +默认权限兼容反向代理(如 Caddy),可通过配置文件调整: + +```yaml +web: + sock: "/run/lmvpnweb.sock" + sock_mode: "0666" # socket 文件权限,默认 0666 + sock_group: "" # socket 文件 group,空=不修改 + sock_dir_mode: "0755" # socket 目录权限,默认 0755 +``` + +多租户/高安全场景建议收紧: + +```yaml +web: + sock: "/run/lmvpnweb.sock" + sock_mode: "0660" + sock_group: "caddy" # 将 lmvpn 进程用户加入 caddy group + sock_dir_mode: "0750" +``` + +### 生产部署 + +- 生产环境必须通过反向代理(如 Caddy/Nginx)提供 HTTPS/WSS +- 配置文件 `data/config.yml` 权限为 0600,仅限运行用户读写 +- WebSocket `/ws` 端点支持 JWT 认证(`?token=xxx`)和密码认证两种方式 diff --git a/internal/config/config.go b/internal/config/config.go index 495136b..b5e7ba7 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -1,6 +1,8 @@ package config import ( + cryptorand "crypto/rand" + "encoding/base64" "os" "path/filepath" @@ -8,8 +10,12 @@ import ( ) type WebConfig struct { - Port int `yaml:"port"` - Sock string `yaml:"sock"` + Port int `yaml:"port"` + Sock string `yaml:"sock"` + SockMode string `yaml:"sock_mode"` + SockGroup string `yaml:"sock_group"` + SockDirMode string `yaml:"sock_dir_mode"` + JWTSecret string `yaml:"jwt_secret"` } type DatabaseConfig struct { @@ -26,8 +32,10 @@ type Config struct { func defaultConfig() *Config { return &Config{ Web: WebConfig{ - Port: 8080, - Sock: "/run/lmvpnweb.sock", + Port: 8080, + Sock: "/run/lmvpnweb.sock", + SockMode: "0666", + SockDirMode: "0755", }, Database: DatabaseConfig{ Type: "sqlite", @@ -39,7 +47,7 @@ func defaultConfig() *Config { func Load(path string) (*Config, error) { dir := filepath.Dir(path) - if err := os.MkdirAll(dir, 0755); err != nil { + if err := os.MkdirAll(dir, 0700); err != nil { return nil, err } @@ -50,7 +58,9 @@ func Load(path string) (*Config, error) { if !os.IsNotExist(err) { return nil, err } - + if err := resolveJWTSecret(cfg); err != nil { + return nil, err + } if err := saveConfig(path, cfg); err != nil { return nil, err } @@ -61,16 +71,44 @@ func Load(path string) (*Config, error) { return nil, err } + if err := resolveJWTSecret(cfg); err != nil { + return nil, err + } + if err := saveConfig(path, cfg); err != nil { return nil, err } return cfg, nil } +func resolveJWTSecret(cfg *Config) error { + if envSecret := os.Getenv("LMVPN_JWT_SECRET"); envSecret != "" { + cfg.Web.JWTSecret = envSecret + return nil + } + if cfg.Web.JWTSecret != "" { + return nil + } + secret, err := generateRandomSecret(32) + if err != nil { + return err + } + cfg.Web.JWTSecret = secret + return nil +} + +func generateRandomSecret(n int) (string, error) { + b := make([]byte, n) + if _, err := cryptorand.Read(b); err != nil { + return "", err + } + return base64.StdEncoding.EncodeToString(b), nil +} + func saveConfig(path string, cfg *Config) error { data, err := yaml.Marshal(cfg) if err != nil { return err } - return os.WriteFile(path, data, 0644) + return os.WriteFile(path, data, 0600) } diff --git a/internal/db/db.go b/internal/db/db.go index 4350346..cce3e2a 100644 --- a/internal/db/db.go +++ b/internal/db/db.go @@ -1,8 +1,12 @@ package db import ( + cryptorand "crypto/rand" "fmt" "log" + "math/big" + "os" + "path/filepath" "lmvpn/internal/config" "lmvpn/internal/model" @@ -40,7 +44,7 @@ func Init(cfg *config.DatabaseConfig) error { return fmt.Errorf("数据库迁移失败: %w", err) } - if err := seedDefaultAdmin(); err != nil { + if err := seedDefaultAdmin(cfg); err != nil { return fmt.Errorf("创建默认管理员失败: %w", err) } @@ -48,14 +52,19 @@ func Init(cfg *config.DatabaseConfig) error { return nil } -func seedDefaultAdmin() error { +func seedDefaultAdmin(cfg *config.DatabaseConfig) error { var count int64 DB.Model(&model.User{}).Count(&count) if count > 0 { return nil } - hash, err := bcrypt.GenerateFromPassword([]byte("admin123"), bcrypt.DefaultCost) + password, err := generateRandomPassword(16) + if err != nil { + return err + } + + hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost) if err != nil { return err } @@ -71,6 +80,33 @@ func seedDefaultAdmin() error { return err } - log.Println("已创建默认管理员: admin / admin123") + fmt.Println("========================================") + fmt.Println("已创建默认管理员账户") + fmt.Println("用户名: admin") + fmt.Println("密码: " + password) + fmt.Println("请登录后立即修改密码!") + fmt.Println("========================================") + + dbDir := filepath.Dir(cfg.Path) + pwdFile := filepath.Join(dbDir, ".initial_admin_password") + if err := os.WriteFile(pwdFile, []byte("admin:"+password+"\n"), 0600); err != nil { + log.Printf("警告: 写入初始密码文件失败: %v", err) + } else { + log.Printf("初始密码已写入 %s,请登录后删除此文件", pwdFile) + } + return nil } + +func generateRandomPassword(length int) (string, error) { + const charset = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789" + b := make([]byte, length) + for i := range b { + n, err := cryptorand.Int(cryptorand.Reader, big.NewInt(int64(len(charset)))) + if err != nil { + return "", err + } + b[i] = charset[n.Int64()] + } + return string(b), nil +} diff --git a/internal/handler/user.go b/internal/handler/user.go index 1585d06..8ab78ea 100644 --- a/internal/handler/user.go +++ b/internal/handler/user.go @@ -45,6 +45,12 @@ func formatUser(u *model.User) userResponse { } } +var validRoles = map[string]bool{"admin": true, "user": true} + +func isValidRole(role string) bool { + return validRoles[role] +} + func GetUserCount(c *gin.Context) { var count int64 db.DB.Model(&model.User{}).Count(&count) @@ -89,6 +95,10 @@ func CreateUser(c *gin.Context) { Status: 1, } if req.Role != "" { + if !isValidRole(req.Role) { + c.JSON(http.StatusBadRequest, gin.H{"error": "角色无效,仅支持 admin 或 user"}) + return + } user.Role = req.Role } if req.Status != nil { @@ -132,6 +142,10 @@ func UpdateUser(c *gin.Context) { } if req.Role != "" { + if !isValidRole(req.Role) { + c.JSON(http.StatusBadRequest, gin.H{"error": "角色无效,仅支持 admin 或 user"}) + return + } if user.ID == currentUserID.(uint) { c.JSON(http.StatusBadRequest, gin.H{"error": "不能修改自己的角色"}) return @@ -167,7 +181,7 @@ func UpdateUser(c *gin.Context) { return } - if req.Password != "" { + if req.Password != "" || req.Role != "" || req.Status != nil { db.DB.Model(&model.Session{}).Where("user_id = ?", id).Update("invalid", true) } diff --git a/internal/middleware/auth.go b/internal/middleware/auth.go index 58ba6c1..a267219 100644 --- a/internal/middleware/auth.go +++ b/internal/middleware/auth.go @@ -12,10 +12,13 @@ import ( "github.com/golang-jwt/jwt/v5" ) -const ( - jwtSecret = "lmvpn-jwt-secret-key-2024" - tokenExpire = 24 * time.Hour -) +const tokenExpire = 24 * time.Hour + +var jwtSecret []byte + +func SetJWTSecret(secret string) { + jwtSecret = []byte(secret) +} type Claims struct { SessionID string `json:"session_id,omitempty"` @@ -38,12 +41,12 @@ func GenerateToken(sessionID string, userID uint, username, role string) (string } token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims) - return token.SignedString([]byte(jwtSecret)) + return token.SignedString(jwtSecret) } func ParseToken(tokenStr string) (*Claims, error) { token, err := jwt.ParseWithClaims(tokenStr, &Claims{}, func(token *jwt.Token) (interface{}, error) { - return []byte(jwtSecret), nil + return jwtSecret, nil }) if err != nil { return nil, err @@ -131,7 +134,7 @@ func AuthMiddleware() gin.HandlerFunc { c.Set("user_id", claims.UserID) c.Set("username", claims.Username) - c.Set("role", claims.Role) + c.Set("role", user.Role) c.Set("session_id", claims.SessionID) c.Next() } diff --git a/internal/middleware/ratelimit.go b/internal/middleware/ratelimit.go new file mode 100644 index 0000000..0d7e784 --- /dev/null +++ b/internal/middleware/ratelimit.go @@ -0,0 +1,86 @@ +package middleware + +import ( + "net/http" + "sync" + "time" + + "github.com/gin-gonic/gin" +) + +type RateLimiter struct { + mu sync.Mutex + attempts map[string][]time.Time + limit int + window time.Duration +} + +func NewRateLimiter(limit int, window time.Duration) *RateLimiter { + rl := &RateLimiter{ + attempts: make(map[string][]time.Time), + limit: limit, + window: window, + } + go rl.cleanup() + return rl +} + +func (rl *RateLimiter) Allow(key string) bool { + rl.mu.Lock() + defer rl.mu.Unlock() + + now := time.Now() + cutoff := now.Add(-rl.window) + + var recent []time.Time + for _, t := range rl.attempts[key] { + if t.After(cutoff) { + recent = append(recent, t) + } + } + + if len(recent) >= rl.limit { + rl.attempts[key] = recent + return false + } + + recent = append(recent, now) + rl.attempts[key] = recent + return true +} + +func (rl *RateLimiter) cleanup() { + ticker := time.NewTicker(rl.window) + defer ticker.Stop() + for range ticker.C { + rl.mu.Lock() + cutoff := time.Now().Add(-rl.window) + for key, times := range rl.attempts { + var recent []time.Time + for _, t := range times { + if t.After(cutoff) { + recent = append(recent, t) + } + } + if len(recent) == 0 { + delete(rl.attempts, key) + } else { + rl.attempts[key] = recent + } + } + rl.mu.Unlock() + } +} + +func LoginRateLimit() gin.HandlerFunc { + limiter := NewRateLimiter(5, time.Minute) + return func(c *gin.Context) { + key := c.ClientIP() + if !limiter.Allow(key) { + c.JSON(http.StatusTooManyRequests, gin.H{"error": "请求过于频繁,请稍后再试"}) + c.Abort() + return + } + c.Next() + } +} diff --git a/internal/router/router.go b/internal/router/router.go index d3cb7f3..959f3dc 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -2,7 +2,6 @@ package router import ( "net/http" - "os" "strings" "lmvpn/internal/handler" @@ -15,7 +14,7 @@ import ( func Setup(r *gin.Engine) { r.GET("/ws", vpn.HandleWS) - r.POST("/api/login", handler.Login) + r.POST("/api/login", middleware.LoginRateLimit(), handler.Login) auth := r.Group("/api") auth.Use(middleware.AuthMiddleware()) @@ -37,21 +36,21 @@ func Setup(r *gin.Engine) { admin.DELETE("/users/:id/sessions", handler.AdminRevokeUserSessions) } - fs := http.FileServer(http.Dir("./dist")) - r.Use(func(c *gin.Context) { - if strings.HasPrefix(c.Request.URL.Path, "/ws") { + distDir := http.Dir("./dist") + fs := http.FileServer(distDir) + r.NoRoute(func(c *gin.Context) { + path := c.Request.URL.Path + if strings.HasPrefix(path, "/api") || strings.HasPrefix(path, "/ws") { + c.JSON(http.StatusNotFound, gin.H{"error": "not found"}) return } - if strings.HasPrefix(c.Request.URL.Path, "/api") { - c.Next() + f, err := distDir.Open(path) + if err != nil { + c.Header("Content-Type", "text/html") + c.File("./dist/index.html") return } - - path := "./dist" + c.Request.URL.Path - if _, err := os.Stat(path); os.IsNotExist(err) { - c.Request.URL.Path = "/" - } + f.Close() fs.ServeHTTP(c.Writer, c.Request) - c.Abort() }) } diff --git a/internal/vpn/auth.go b/internal/vpn/auth.go index 8851996..f054567 100644 --- a/internal/vpn/auth.go +++ b/internal/vpn/auth.go @@ -2,7 +2,9 @@ package vpn import ( "encoding/json" + "time" + "lmvpn/internal/middleware" "lmvpn/internal/model" "github.com/gorilla/websocket" @@ -10,6 +12,8 @@ import ( "gorm.io/gorm" ) +var authLimiter = middleware.NewRateLimiter(5, time.Minute) + type authMessage struct { Type string `json:"type"` Username string `json:"username"` @@ -21,7 +25,7 @@ type authResponse struct { Message string `json:"message,omitempty"` } -func authenticate(conn *websocket.Conn, db *gorm.DB) (*model.User, error) { +func authenticate(conn *websocket.Conn, db *gorm.DB, clientIP string) (*model.User, error) { _, msgBytes, err := conn.ReadMessage() if err != nil { return nil, err @@ -29,23 +33,27 @@ func authenticate(conn *websocket.Conn, db *gorm.DB) (*model.User, error) { var msg authMessage if err := json.Unmarshal(msgBytes, &msg); err != nil || msg.Type != "auth" { - resp := authResponse{Type: "auth_err", Message: "消息格式错误"} - sendJSON(conn, resp) + sendJSON(conn, authResponse{Type: "auth_err", Message: "消息格式错误"}) + conn.Close() + return nil, nil + } + + key := clientIP + ":" + msg.Username + if !authLimiter.Allow(key) { + sendJSON(conn, authResponse{Type: "auth_err", Message: "认证尝试过于频繁,请稍后再试"}) conn.Close() return nil, nil } var user model.User if err := db.Where("username = ? AND status = 1", msg.Username).First(&user).Error; err != nil { - resp := authResponse{Type: "auth_err", Message: "用户名或密码错误"} - sendJSON(conn, resp) + sendJSON(conn, authResponse{Type: "auth_err", Message: "用户名或密码错误"}) conn.Close() return nil, nil } if err := bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(msg.Password)); err != nil { - resp := authResponse{Type: "auth_err", Message: "用户名或密码错误"} - sendJSON(conn, resp) + sendJSON(conn, authResponse{Type: "auth_err", Message: "用户名或密码错误"}) conn.Close() return nil, nil } @@ -60,6 +68,9 @@ func authenticate(conn *websocket.Conn, db *gorm.DB) (*model.User, error) { } func sendJSON(conn *websocket.Conn, v interface{}) error { - data, _ := json.Marshal(v) + data, err := json.Marshal(v) + if err != nil { + return err + } return conn.WriteMessage(websocket.TextMessage, data) } diff --git a/internal/vpn/handler.go b/internal/vpn/handler.go index 04961c2..56fbb86 100644 --- a/internal/vpn/handler.go +++ b/internal/vpn/handler.go @@ -3,8 +3,11 @@ package vpn import ( "log" "net/http" + "net/url" "lmvpn/internal/db" + "lmvpn/internal/middleware" + "lmvpn/internal/model" "github.com/gin-gonic/gin" "github.com/gorilla/websocket" @@ -14,18 +17,45 @@ var upgrader = websocket.Upgrader{ ReadBufferSize: 4096, WriteBufferSize: 4096, CheckOrigin: func(r *http.Request) bool { - return true + origin := r.Header.Get("Origin") + if origin == "" { + return true + } + u, err := url.Parse(origin) + if err != nil { + return false + } + return u.Host == r.Host }, } func HandleWS(c *gin.Context) { + tokenStr := c.Query("token") + conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) if err != nil { log.Printf("WebSocket 升级失败: %v", err) return } - user, err := authenticate(conn, db.DB) + if tokenStr != "" { + claims, err := middleware.ParseToken(tokenStr) + if err != nil { + sendJSON(conn, authResponse{Type: "auth_err", Message: "令牌无效或已过期"}) + conn.Close() + return + } + var u model.User + if err := db.DB.First(&u, claims.UserID).Error; err != nil || u.Status != 1 { + sendJSON(conn, authResponse{Type: "auth_err", Message: "用户不存在或已禁用"}) + conn.Close() + return + } + runTunnel(conn, &u) + return + } + + user, err := authenticate(conn, db.DB, c.ClientIP()) if err != nil { log.Printf("认证读取失败: %v", err) conn.Close() diff --git a/internal/vpn/tunnel.go b/internal/vpn/tunnel.go index 4e88ade..524b9db 100644 --- a/internal/vpn/tunnel.go +++ b/internal/vpn/tunnel.go @@ -2,6 +2,7 @@ package vpn import ( "log" + "sync" "time" "lmvpn/internal/model" @@ -10,16 +11,43 @@ import ( ) const ( - readTimeout = 60 * time.Second - writeTimeout = 10 * time.Second - pingPeriod = 30 * time.Second + readTimeout = 60 * time.Second + writeTimeout = 10 * time.Second + pingPeriod = 30 * time.Second + maxMessageSize = 1 << 20 + maxConnsPerUser = 3 +) + +var ( + activeConns = make(map[uint]int) + activeConnsMu sync.Mutex ) func runTunnel(conn *websocket.Conn, user *model.User) { defer conn.Close() + activeConnsMu.Lock() + if activeConns[user.ID] >= maxConnsPerUser { + activeConnsMu.Unlock() + sendJSON(conn, authResponse{Type: "auth_err", Message: "连接数已达上限"}) + return + } + activeConns[user.ID]++ + activeConnsMu.Unlock() + + defer func() { + activeConnsMu.Lock() + activeConns[user.ID]-- + if activeConns[user.ID] <= 0 { + delete(activeConns, user.ID) + } + activeConnsMu.Unlock() + }() + log.Printf("用户 %s 已连接", user.Username) + conn.SetReadLimit(maxMessageSize) + go func() { ticker := time.NewTicker(pingPeriod) defer ticker.Stop() diff --git a/main.go b/main.go index cab17da..f39f040 100644 --- a/main.go +++ b/main.go @@ -5,10 +5,13 @@ import ( "log" "net" "os" + "os/user" "path/filepath" + "strconv" "lmvpn/internal/config" "lmvpn/internal/db" + "lmvpn/internal/middleware" "lmvpn/internal/router" "github.com/gin-gonic/gin" @@ -20,6 +23,8 @@ func main() { log.Fatalf("加载配置失败: %v", err) } + middleware.SetJWTSecret(cfg.Web.JWTSecret) + if err := db.Init(&cfg.Database); err != nil { log.Fatalf("数据库初始化失败: %v", err) } @@ -45,15 +50,22 @@ func main() { if err := os.Remove(cfg.Web.Sock); err != nil && !os.IsNotExist(err) { log.Fatalf("删除残留 sock 文件失败: %v", err) } - if err := os.MkdirAll(filepath.Dir(cfg.Web.Sock), 0777); err != nil { + dirMode := parseFileMode(cfg.Web.SockDirMode, 0755) + if err := os.MkdirAll(filepath.Dir(cfg.Web.Sock), dirMode); err != nil { log.Fatalf("创建 sock 目录失败: %v", err) } listener, err := net.Listen("unix", cfg.Web.Sock) if err != nil { log.Fatalf("Unix socket 监听失败: %v", err) } - if err := os.Chmod(cfg.Web.Sock, 0777); err != nil { - log.Fatalf("设置 sock 权限失败: %v", err) + sockMode := parseFileMode(cfg.Web.SockMode, 0666) + if err := os.Chmod(cfg.Web.Sock, sockMode); err != nil { + log.Printf("警告: 设置 sock 权限失败: %v", err) + } + if cfg.Web.SockGroup != "" { + if err := chownGroup(cfg.Web.Sock, cfg.Web.SockGroup); err != nil { + log.Printf("警告: 设置 sock group 失败: %v", err) + } } go func() { log.Printf("Unix socket 监听 %s", cfg.Web.Sock) @@ -65,3 +77,27 @@ func main() { select {} } + +func parseFileMode(s string, defaultMode os.FileMode) os.FileMode { + if s == "" { + return defaultMode + } + m, err := strconv.ParseUint(s, 8, 32) + if err != nil { + log.Printf("警告: 解析文件权限 %q 失败,使用默认值 %o: %v", s, defaultMode, err) + return defaultMode + } + return os.FileMode(m) +} + +func chownGroup(path, group string) error { + g, err := user.LookupGroup(group) + if err != nil { + return err + } + gid, err := strconv.Atoi(g.Gid) + if err != nil { + return err + } + return os.Chown(path, -1, gid) +}