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
+53
View File
@@ -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`)和密码认证两种方式
+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()
+39 -3
View File
@@ -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)
}