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
+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()