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" ) var upgrader = websocket.Upgrader{ ReadBufferSize: 4096, WriteBufferSize: 4096, CheckOrigin: func(r *http.Request) bool { 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") clientIP := c.ClientIP() conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) if err != nil { log.Printf("[WS] upgrade failed clientIP=%s err=%v", clientIP, err) return } if tokenStr != "" { log.Printf("[WS] upgrade clientIP=%s auth=jwt", clientIP) 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 } log.Printf("[WS] upgrade clientIP=%s auth=password", clientIP) user, err := authenticate(conn, db.DB, clientIP) if err != nil { log.Printf("[WS] auth read failed clientIP=%s err=%v", clientIP, err) conn.Close() return } if user == nil { return } runTunnel(conn, user) }