hub.go 重写用send chan+单写协程模式避免并发写竞争 handler.go 加DEV_MODE dev-token支持+broadcast端点 config.go 加DevMode/RedisURL字段
171 lines
4.4 KiB
Go
171 lines
4.4 KiB
Go
package ws
|
||
|
||
import (
|
||
"encoding/json"
|
||
"errors"
|
||
"net/http"
|
||
"strings"
|
||
|
||
"github.com/edu-cloud/push-gateway/internal/hub"
|
||
"github.com/gin-gonic/gin"
|
||
"github.com/golang-jwt/jwt/v5"
|
||
"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
|
||
}
|
||
|
||
// 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) {
|
||
userID, err := h.authenticate(c)
|
||
if err != nil {
|
||
c.JSON(http.StatusUnauthorized, gin.H{"error": err.Error()})
|
||
return
|
||
}
|
||
|
||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||
if err != nil {
|
||
return
|
||
}
|
||
defer conn.Close()
|
||
|
||
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 {
|
||
_, msg, err := conn.ReadMessage()
|
||
if err != nil {
|
||
break
|
||
}
|
||
if strings.ToLower(string(msg)) == "ping" {
|
||
connPtr.Send([]byte("pong"))
|
||
}
|
||
}
|
||
}
|
||
|
||
// PushHandler 接收来自 Msg 服务的定向推送请求
|
||
func (h *Handler) PushHandler(c *gin.Context) {
|
||
var req struct {
|
||
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
|
||
}
|
||
|
||
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,
|
||
})
|
||
}
|