283 lines
6.6 KiB
Go
283 lines
6.6 KiB
Go
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:<deviceKey>)
|
||
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)
|
||
}
|
||
}
|
||
}
|