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:
@@ -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 参数获取 token(WebSocket 无法设置 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,
|
||||
})
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user