ai_scheduler/internal/services/advice/wx_ws_hub.go

162 lines
4.0 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package advice
import (
"ai_scheduler/internal/biz"
"ai_scheduler/internal/config"
"ai_scheduler/internal/pkg"
"encoding/json"
"sync"
"time"
"github.com/gofiber/fiber/v2/log"
"github.com/gofiber/websocket/v2"
)
// WxWsHub WebSocket 连接中心,管理前端 WebSocket 客户端并广播实时事件
type WxWsHub struct {
mu sync.RWMutex
clients map[*wsClient]struct{}
jwtSecret string
}
type wsClient struct {
conn *websocket.Conn
appId string // 订阅的微信 appId(subscribe 消息后设置)
mu sync.Mutex
}
// NewWxWsHub 创建 WebSocket Hub(实现 biz.WsBroadcaster 接口)
func NewWxWsHub(cfg *config.Config) *WxWsHub {
return &WxWsHub{
clients: make(map[*wsClient]struct{}),
jwtSecret: cfg.JwtSecret,
}
}
// wsEvent WebSocket 推送的事件格式
type wsEvent struct {
Event string `json:"event"`
Data interface{} `json:"data"`
}
// WsHandler Fiber WebSocket 升级处理函数
// 连接流程:
// 1. 前端通过 ws://host/api/v1/advicer/ws?token=JWT 连接
// 2. Hub 验证 JWT token
// 3. 前端发送 {"action":"subscribe","appId":"xxx"} 订阅特定 appId 的事件
// 4. Hub 向该 appId 的所有已订阅客户端广播事件
func (h *WxWsHub) WsHandler(c *websocket.Conn) {
// 验证 JWT token(从 query 参数获取)
token := c.Query("token")
if token == "" {
log.Warn("ws connect: missing token")
_ = c.Close()
return
}
claims, err := pkg.ParseToken(token, h.jwtSecret)
if err != nil || claims == nil {
log.Warn("ws connect: invalid token")
_ = c.Close()
return
}
client := &wsClient{conn: c}
h.register(client)
defer h.unregister(client)
log.Infof("ws client connected: user=%s", claims.Username)
// 读取循环:处理客户端发来的消息(subscribe 等)
for {
_, msg, err := c.ReadMessage()
if err != nil {
// 连接关闭或读取错误,退出循环(defer 会清理)
break
}
h.handleClientMsg(client, msg)
}
}
// clientMsg 客户端发送的消息格式
type clientMsg struct {
Action string `json:"action"`
AppId string `json:"appId"`
}
func (h *WxWsHub) handleClientMsg(client *wsClient, raw []byte) {
var msg clientMsg
if err := json.Unmarshal(raw, &msg); err != nil {
return
}
switch msg.Action {
case "subscribe":
if msg.AppId != "" {
client.mu.Lock()
client.appId = msg.AppId
client.mu.Unlock()
log.Infof("[WS] client subscribed appId=%s", msg.AppId)
} else {
log.Warn("[WS] subscribe with empty appId")
}
case "ping":
// 客户端心跳,回复 pong
client.mu.Lock()
_ = client.conn.WriteMessage(websocket.PongMessage, []byte("pong"))
client.mu.Unlock()
}
}
func (h *WxWsHub) register(c *wsClient) {
h.mu.Lock()
h.clients[c] = struct{}{}
h.mu.Unlock()
}
func (h *WxWsHub) unregister(c *wsClient) {
h.mu.Lock()
if _, ok := h.clients[c]; ok {
delete(h.clients, c)
_ = c.conn.Close()
}
h.mu.Unlock()
log.Info("ws client disconnected")
}
// Broadcast 向指定 appId 的所有已订阅客户端广播 JSON 事件(异步非阻塞)
func (h *WxWsHub) Broadcast(appId string, event string, data interface{}) {
h.mu.RLock()
defer h.mu.RUnlock()
payload, err := json.Marshal(wsEvent{Event: event, Data: data})
if err != nil {
return
}
clientCount := 0
matchedCount := 0
for client := range h.clients {
client.mu.Lock()
clientCount++
if client.appId == appId {
matchedCount++
// 设置写超时,避免慢客户端阻塞广播
_ = client.conn.SetWriteDeadline(time.Now().Add(3 * time.Second))
_ = client.conn.WriteMessage(websocket.TextMessage, payload)
_ = client.conn.SetWriteDeadline(time.Time{})
}
client.mu.Unlock()
}
log.Infof("[WS] Broadcast event=%s appId=%s clients=%d matched=%d", event, appId, clientCount, matchedCount)
}
// BroadcastOnline 广播在线状态变化
func (h *WxWsHub) BroadcastOnline(appId string, online bool) {
h.Broadcast(appId, "online_status", map[string]interface{}{
"appId": appId,
"online": online,
})
}
// 编译期断言:WxWsHub 实现 biz.WsBroadcaster 接口
var _ biz.WsBroadcaster = (*WxWsHub)(nil)