255 lines
5.9 KiB
Go
255 lines
5.9 KiB
Go
package ws
|
||
|
||
import (
|
||
"encoding/json"
|
||
"log/slog"
|
||
"net/http"
|
||
"strings"
|
||
"sync"
|
||
"time"
|
||
|
||
"github.com/gin-gonic/gin"
|
||
"github.com/golang-jwt/jwt/v5"
|
||
"github.com/gorilla/websocket"
|
||
|
||
"silk-server-go/internal/service"
|
||
)
|
||
|
||
var upgrader = websocket.Upgrader{
|
||
CheckOrigin: func(r *http.Request) bool { return true },
|
||
ReadBufferSize: 1024,
|
||
WriteBufferSize: 1024,
|
||
}
|
||
|
||
// client WebSocket 客户端
|
||
type client struct {
|
||
conn *websocket.Conn
|
||
rooms map[string]bool // 订阅的房间(device:<deviceKey>)
|
||
send chan []byte
|
||
}
|
||
|
||
// Hub WebSocket 中心,管理客户端和房间
|
||
type Hub struct {
|
||
jwtSecret string
|
||
clients map[*client]bool
|
||
mu sync.RWMutex
|
||
}
|
||
|
||
// NewHub 创建 Hub
|
||
func NewHub(jwtSecret string) *Hub {
|
||
return &Hub{
|
||
jwtSecret: jwtSecret,
|
||
clients: make(map[*client]bool),
|
||
}
|
||
}
|
||
|
||
// HandleWebSocket 处理 WebSocket 连接(Gin handler)
|
||
func (h *Hub) HandleWebSocket(c *gin.Context) {
|
||
// 验证 JWT(从 query.auth.token / query.token / header.Authorization 获取)
|
||
tokenStr := ""
|
||
if t := c.Query("auth.token"); t != "" {
|
||
tokenStr = t
|
||
} else if t := c.Query("token"); t != "" {
|
||
tokenStr = t
|
||
} else if auth := c.GetHeader("Authorization"); strings.HasPrefix(auth, "Bearer ") {
|
||
tokenStr = auth[7:]
|
||
}
|
||
|
||
if tokenStr == "" {
|
||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||
if err == nil {
|
||
conn.WriteJSON(map[string]interface{}{"event": "auth.fail", "data": map[string]bool{"ok": false}})
|
||
conn.Close()
|
||
}
|
||
return
|
||
}
|
||
|
||
// 验证 token
|
||
claims := jwt.MapClaims{}
|
||
token, err := jwt.ParseWithClaims(tokenStr, claims, func(t *jwt.Token) (interface{}, error) {
|
||
return []byte(h.jwtSecret), nil
|
||
}, jwt.WithValidMethods([]string{"HS256"}))
|
||
|
||
if err != nil || !token.Valid {
|
||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||
if err == nil {
|
||
conn.WriteJSON(map[string]interface{}{"event": "auth.fail", "data": map[string]bool{"ok": false}})
|
||
conn.Close()
|
||
}
|
||
return
|
||
}
|
||
|
||
// 升级为 WebSocket
|
||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||
if err != nil {
|
||
slog.Warn("WebSocket 升级失败", "err", err)
|
||
return
|
||
}
|
||
|
||
cl := &client{
|
||
conn: conn,
|
||
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": claims,
|
||
})
|
||
|
||
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":
|
||
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)
|
||
}
|
||
// 全局推送
|
||
h.sendJSON(cl, "telemetry.all", 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 {
|
||
// 恢复通知
|
||
h.sendJSON(cl, "alarm.recovery", event)
|
||
if event.DeviceKey != "" && cl.rooms[room] {
|
||
h.sendJSON(cl, "alarm.device.recovery", event)
|
||
}
|
||
} else {
|
||
// 告警触发
|
||
h.sendJSON(cl, "alarm", event)
|
||
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 {
|
||
h.sendJSON(cl, "device.status", data)
|
||
if cl.rooms[room] {
|
||
h.sendJSON(cl, "device.status.device", data)
|
||
}
|
||
}
|
||
}
|