package ws import ( "encoding/json" "log/slog" "net/http" "sync" "time" "github.com/gin-gonic/gin" "github.com/gorilla/websocket" "silk-server-go/internal/service" ) // client WebSocket 客户端 type client struct { conn *websocket.Conn userID string username string role string rooms map[string]bool // 订阅的房间(device:) send chan []byte } // Hub WebSocket 中心,管理客户端和房间 type Hub struct { authorizer DeviceAuthorizer allowedOrigins []string clients map[*client]bool mu sync.RWMutex tickets map[string]*wsTicket ticketsMu sync.Mutex } // NewHub 创建 Hub func NewHub(jwtSecret string, authorizer DeviceAuthorizer, allowedOrigins []string) *Hub { return &Hub{ authorizer: authorizer, allowedOrigins: allowedOrigins, clients: make(map[*client]bool), tickets: make(map[string]*wsTicket), } } func (h *Hub) originAllowed(origin string) bool { if origin == "" { return true } for _, allowed := range h.allowedOrigins { if origin == allowed { return true } } return false } // HandleTicket 签发一次性 WebSocket ticket,需在 JWT 鉴权路由后注册。 func (h *Hub) HandleTicket(c *gin.Context) { userVal, ok := c.Get("user") if !ok { c.JSON(http.StatusUnauthorized, gin.H{"error": "未认证"}) return } userMap, ok := userVal.(map[string]interface{}) if !ok { c.JSON(http.StatusUnauthorized, gin.H{"error": "用户信息无效"}) return } userID, _ := userMap["sub"].(string) username, _ := userMap["username"].(string) role, _ := userMap["role"].(string) if userID == "" { c.JSON(http.StatusUnauthorized, gin.H{"error": "用户信息无效"}) return } c.JSON(http.StatusOK, gin.H{ "ticket": h.IssueTicket(userID, username, role), "expiresIn": int(wsTicketTTL.Seconds()), }) } // HandleWebSocket 处理 WebSocket 连接(Gin handler) func (h *Hub) HandleWebSocket(c *gin.Context) { if !h.originAllowed(c.GetHeader("Origin")) { c.JSON(http.StatusForbidden, gin.H{"error": "origin not allowed"}) return } ticket := h.ticketFromQuery(c.Request.URL.RawQuery) info, ok := h.consumeTicket(ticket) if !ok { c.JSON(http.StatusUnauthorized, gin.H{"error": "invalid or expired ws ticket"}) return } // 升级为 WebSocket upgrader := websocket.Upgrader{ CheckOrigin: func(r *http.Request) bool { return true }, ReadBufferSize: 1024, WriteBufferSize: 1024, } conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) if err != nil { slog.Warn("WebSocket 升级失败", "err", err) return } cl := &client{ conn: conn, userID: info.userID, username: info.username, role: info.role, rooms: make(map[string]bool), send: make(chan []byte, 256), } h.mu.Lock() h.clients[cl] = true h.mu.Unlock() // 发送 auth.ok h.sendJSON(cl, "auth.ok", map[string]interface{}{ "ok": true, "user": map[string]interface{}{ "sub": info.userID, "username": info.username, "role": info.role, }, }) go h.readPump(cl) go h.writePump(cl) } // readPump 读取客户端消息 func (h *Hub) readPump(cl *client) { defer func() { h.mu.Lock() delete(h.clients, cl) h.mu.Unlock() cl.conn.Close() }() for { _, msg, err := cl.conn.ReadMessage() if err != nil { break } // 解析消息(兼容 {event: "...", deviceKey: "..."} 格式) var data struct { Event string `json:"event"` DeviceKey string `json:"deviceKey"` } if err := json.Unmarshal(msg, &data); err != nil { // 尝试 Socket.IO 格式 ["event", {deviceKey: "..."}] var arr []json.RawMessage if err2 := json.Unmarshal(msg, &arr); err2 == nil && len(arr) >= 2 { if len(arr[0]) > 0 { json.Unmarshal(arr[0], &data.Event) } if len(arr[1]) > 0 { var payload struct { DeviceKey string `json:"deviceKey"` } json.Unmarshal(arr[1], &payload) data.DeviceKey = payload.DeviceKey } } } switch data.Event { case "subscribe.device": if err := h.authorizeSubscription(cl.userID, data.DeviceKey); err != nil { h.sendJSON(cl, "subscribe.denied", map[string]interface{}{ "ok": false, "deviceKey": data.DeviceKey, "error": err.Error(), }) continue } room := "device:" + data.DeviceKey h.mu.Lock() cl.rooms[room] = true h.mu.Unlock() h.sendJSON(cl, "subscribed", map[string]interface{}{"ok": true, "room": room}) case "unsubscribe.device": room := "device:" + data.DeviceKey h.mu.Lock() delete(cl.rooms, room) h.mu.Unlock() h.sendJSON(cl, "unsubscribed", map[string]interface{}{"ok": true}) } } } // writePump 向客户端发送消息 func (h *Hub) writePump(cl *client) { ticker := time.NewTicker(30 * time.Second) defer func() { ticker.Stop() cl.conn.Close() }() for { select { case msg, ok := <-cl.send: if !ok { cl.conn.WriteMessage(websocket.CloseMessage, []byte{}) return } if err := cl.conn.WriteMessage(websocket.TextMessage, msg); err != nil { return } case <-ticker.C: if err := cl.conn.WriteMessage(websocket.PingMessage, nil); err != nil { return } } } } // sendJSON 向客户端发送 JSON 消息 func (h *Hub) sendJSON(cl *client, event string, data interface{}) { msg := map[string]interface{}{"event": event, "data": data} body, _ := json.Marshal(msg) select { case cl.send <- body: default: slog.Warn("WebSocket 客户端发送缓冲区满,丢弃消息") } } // BroadcastTelemetry 广播遥测数据(实现 service.EventHub 接口) func (h *Hub) BroadcastTelemetry(deviceKey string, data interface{}) { h.mu.RLock() defer h.mu.RUnlock() room := "device:" + deviceKey for cl := range h.clients { if cl.rooms[room] { h.sendJSON(cl, "telemetry", data) } } } // BroadcastAlarm 广播告警(实现 service.EventHub 接口) func (h *Hub) BroadcastAlarm(event service.AlarmEvent) { h.mu.RLock() defer h.mu.RUnlock() room := "device:" + event.DeviceKey isRecovery := event.Code == "recovery" for cl := range h.clients { if isRecovery { if event.DeviceKey != "" && cl.rooms[room] { h.sendJSON(cl, "alarm.device.recovery", event) } } else { if event.DeviceKey != "" && cl.rooms[room] { h.sendJSON(cl, "alarm.device", event) } } } } // BroadcastDeviceStatus 广播设备状态变更(实现 service.EventHub 接口) func (h *Hub) BroadcastDeviceStatus(deviceKey string, status string) { h.mu.RLock() defer h.mu.RUnlock() room := "device:" + deviceKey data := map[string]interface{}{"deviceKey": deviceKey, "status": status} for cl := range h.clients { if cl.rooms[room] { h.sendJSON(cl, "device.status.device", data) } } }