Files
2026-08-14 00:45:22 +08:00

283 lines
6.6 KiB
Go
Raw Permalink 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 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)
}
}
}