fix: 修复权限边界与安全问题

- 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 补充安全配置与生产部署说明
This commit is contained in:
2026-07-03 14:12:03 +08:00
parent c1d560fd4a
commit 61189a53ec
11 changed files with 382 additions and 48 deletions
+45 -7
View File
@@ -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)
}
+40 -4
View File
@@ -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
}
+15 -1
View File
@@ -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)
}
+10 -7
View File
@@ -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()
}
+86
View File
@@ -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()
}
}
+12 -13
View File
@@ -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()
})
}
+19 -8
View File
@@ -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)
}
+32 -2
View File
@@ -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()
+31 -3
View File
@@ -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()