155 lines
3.8 KiB
Go
155 lines
3.8 KiB
Go
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)
|
||
}
|
||
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
|
||
}
|
||
|
||
for client := range h.clients {
|
||
client.mu.Lock()
|
||
if client.appId == appId {
|
||
// 设置写超时,避免慢客户端阻塞广播
|
||
_ = client.conn.SetWriteDeadline(time.Now().Add(3 * time.Second))
|
||
_ = client.conn.WriteMessage(websocket.TextMessage, payload)
|
||
_ = client.conn.SetWriteDeadline(time.Time{})
|
||
}
|
||
client.mu.Unlock()
|
||
}
|
||
}
|
||
|
||
// 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)
|