feat(push-gateway): 修复WebSocket并发写竞争并添加DEV_MODE鉴权

hub.go 重写用send chan+单写协程模式避免并发写竞争

handler.go 加DEV_MODE dev-token支持+broadcast端点

config.go 加DevMode/RedisURL字段
This commit is contained in:
SpecialX
2026-07-09 09:09:13 +08:00
parent 416e1bc0b2
commit dfb6d2bfc1
4 changed files with 179 additions and 60 deletions

View File

@@ -1,6 +1,8 @@
package ws
import (
"encoding/json"
"errors"
"net/http"
"strings"
@@ -10,50 +12,36 @@ import (
"github.com/gorilla/websocket"
)
var (
errMissingToken = errors.New("missing token")
errInvalidToken = errors.New("invalid token")
errInvalidClaims = errors.New("invalid claims")
errMissingUserID = errors.New("missing user id")
)
var upgrader = websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool {
return true // P5 骨架,生产环境需校验 origin
},
}
// Handler 处理 WebSocket 升级与内部推送 API
type Handler struct {
hub *hub.Hub
jwtSecret string
devMode bool
}
func NewHandler(h *hub.Hub, jwtSecret string) *Handler {
return &Handler{hub: h, jwtSecret: jwtSecret}
// NewHandler 创建 Handler 实例
func NewHandler(h *hub.Hub, jwtSecret string, devMode bool) *Handler {
return &Handler{hub: h, jwtSecret: jwtSecret, devMode: devMode}
}
// HandleWebSocket 升级 HTTP 为 WebSocket保持长连接并处理心跳
func (h *Handler) HandleWebSocket(c *gin.Context) {
// 从 query 参数获取 tokenWebSocket 无法设置 Authorization 头)
tokenStr := c.Query("token")
if tokenStr == "" {
c.JSON(http.StatusUnauthorized, gin.H{"error": "missing token"})
return
}
token, err := jwt.Parse(tokenStr, func(t *jwt.Token) (interface{}, error) {
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, jwt.ErrSignatureInvalid
}
return []byte(h.jwtSecret), nil
})
if err != nil || !token.Valid {
c.JSON(http.StatusUnauthorized, gin.H{"error": "invalid token"})
return
}
claims, ok := token.Claims.(jwt.MapClaims)
if !ok {
c.JSON(http.StatusUnauthorized, gin.H{"error": "invalid claims"})
return
}
userID, ok := claims["sub"].(string)
if !ok {
c.JSON(http.StatusUnauthorized, gin.H{"error": "missing user id"})
userID, err := h.authenticate(c)
if err != nil {
c.JSON(http.StatusUnauthorized, gin.H{"error": err.Error()})
return
}
@@ -63,8 +51,17 @@ func (h *Handler) HandleWebSocket(c *gin.Context) {
}
defer conn.Close()
h.hub.Register(userID, conn)
defer h.hub.Unregister(userID, conn)
connPtr := h.hub.Register(userID, conn)
defer h.hub.Unregister(connPtr)
// 写协程:从 send chan 读取并写入 WebSocket避免并发写
go func() {
for msg := range connPtr.Outgoing() {
if err := conn.WriteMessage(websocket.TextMessage, msg); err != nil {
return
}
}
}()
// 读取循环(保持连接,处理心跳)
for {
@@ -72,28 +69,102 @@ func (h *Handler) HandleWebSocket(c *gin.Context) {
if err != nil {
break
}
// 处理心跳 ping
if strings.ToLower(string(msg)) == "ping" {
conn.WriteMessage(websocket.TextMessage, []byte("pong"))
connPtr.Send([]byte("pong"))
}
}
}
// PushHandler 接收来自 Msg 服务的推送请求
// PushHandler 接收来自 Msg 服务的定向推送请求
func (h *Handler) PushHandler(c *gin.Context) {
var req struct {
UserID string `json:"user_id"`
Message string `json:"message"`
UserID string `json:"userId"`
Event string `json:"event"`
Data map[string]any `json:"data"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": gin.H{"code": "INVALID_REQUEST", "message": err.Error()}})
return
}
if err := h.hub.SendToUser(req.UserID, []byte(req.Message)); err != nil {
message, err := buildMessage(req.Event, req.Data)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": gin.H{"code": "INVALID_PAYLOAD", "message": err.Error()}})
return
}
if err := h.hub.SendToUser(req.UserID, message); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": gin.H{"code": "PUSH_FAILED", "message": err.Error()}})
return
}
c.JSON(http.StatusOK, gin.H{"success": true})
}
// BroadcastHandler 接收来自 Msg 服务的广播请求
func (h *Handler) BroadcastHandler(c *gin.Context) {
var req struct {
Event string `json:"event"`
Data map[string]any `json:"data"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": gin.H{"code": "INVALID_REQUEST", "message": err.Error()}})
return
}
message, err := buildMessage(req.Event, req.Data)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": gin.H{"code": "INVALID_PAYLOAD", "message": err.Error()}})
return
}
h.hub.Broadcast(message)
c.JSON(http.StatusOK, gin.H{"success": true})
}
// authenticate 校验 WebSocket 连接的 JWT 或 dev token
func (h *Handler) authenticate(c *gin.Context) (string, error) {
tokenStr := c.Query("token")
if tokenStr == "" {
if auth := c.GetHeader("Authorization"); strings.HasPrefix(auth, "Bearer ") {
tokenStr = strings.TrimPrefix(auth, "Bearer ")
}
}
if tokenStr == "" {
return "", errMissingToken
}
// DEV_MODE 下接受 dev-token便于本地联调
if h.devMode && tokenStr == "dev-token" {
return "dev-user", nil
}
token, err := jwt.Parse(tokenStr, func(t *jwt.Token) (any, error) {
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, jwt.ErrSignatureInvalid
}
return []byte(h.jwtSecret), nil
})
if err != nil || !token.Valid {
return "", errInvalidToken
}
claims, ok := token.Claims.(jwt.MapClaims)
if !ok {
return "", errInvalidClaims
}
userID, ok := claims["sub"].(string)
if !ok {
return "", errMissingUserID
}
return userID, nil
}
// buildMessage 构造推送 JSON 消息体
func buildMessage(event string, data map[string]any) ([]byte, error) {
return json.Marshal(map[string]any{
"event": event,
"data": data,
})
}