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:
@@ -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`)和密码认证两种方式
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user