Files
EveryPublish/server/internal/ws/hub.go
T

230 lines
4.9 KiB
Go
Raw 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"
"errors"
"strconv"
"sync"
"time"
"github.com/coder/websocket"
"github.com/gin-gonic/gin"
"everypublish/shared/proto"
)
// parseID 解析设备ID字符串(失败返回 0)
func parseID(s string) uint64 {
id, _ := strconv.ParseUint(s, 10, 64)
return id
}
// Conn 一条已认证连接(Agent 或浏览器)
type Conn struct {
ID string
Kind string
WSID uint64
ws *websocket.Conn
send chan []byte
closeOnce sync.Once
mu sync.Mutex
closed bool
}
// close 安全关闭发送通道(仅一次)。
// 关闭在 mu 保护下进行;所有发送侧经 sendRaw 先检查 closed,
// 从根本上避免「send on closed channel」panic(关闭与发送的竞态)。
func (c *Conn) close() {
c.closeOnce.Do(func() {
c.mu.Lock()
c.closed = true
close(c.send)
c.mu.Unlock()
})
}
// sendRaw 向发送通道投递:closed 直接返回 false 不恐慌;
// block=true 时队列满将阻塞最多 1s(可靠消息用),否则满即丢弃(广播用)。
func (c *Conn) sendRaw(raw []byte, block bool) bool {
c.mu.Lock()
defer c.mu.Unlock()
if c.closed {
return false
}
if block {
select {
case c.send <- raw:
return true
case <-time.After(time.Second):
return false
}
}
select {
case c.send <- raw:
return true
default:
return false
}
}
// sendEnvelope 统一写出口:所有消息经 send 通道由 writeLoop 单协程写入(避免并发写)
func (c *Conn) sendEnvelope(msg *proto.Envelope) {
raw, err := json.Marshal(msg)
if err != nil {
return
}
c.sendRaw(raw, true)
}
// Hub 连接中枢(goroutine-safe)
type Hub struct {
mu sync.RWMutex
agents map[uint64]*Conn
browser map[string]*Conn
}
// NewHub 创建中枢
func NewHub() *Hub {
return &Hub{
agents: make(map[uint64]*Conn),
browser: make(map[string]*Conn),
}
}
// RegisterAgent 上线登记(同设备重连替换旧连接)
// 网络关闭放到锁外执行,避免阻塞中枢其它操作。
func (h *Hub) RegisterAgent(c *Conn) {
var old *Conn
h.mu.Lock()
if o, ok := h.agents[parseID(c.ID)]; ok {
old = o
}
h.agents[parseID(c.ID)] = c
h.mu.Unlock()
if old != nil {
old.close()
if old.ws != nil {
_ = old.ws.Close(websocket.StatusGoingAway, "replaced by new connection")
}
}
}
// UnregisterAgent 下线(仅当仍是当前连接时才移除,避免旧连接误删新连接)
func (h *Hub) UnregisterAgent(c *Conn) {
var toClose *Conn
h.mu.Lock()
if cur, ok := h.agents[parseID(c.ID)]; ok && cur == c {
delete(h.agents, parseID(c.ID))
toClose = cur
}
h.mu.Unlock()
if toClose != nil {
toClose.close()
}
}
// RegisterBrowser 浏览器连接登记
func (h *Hub) RegisterBrowser(c *Conn) {
h.mu.Lock()
h.browser[c.ID] = c
h.mu.Unlock()
}
// UnregisterBrowser 浏览器断开(指针比较,防串号)
func (h *Hub) UnregisterBrowser(c *Conn) {
var toClose *Conn
h.mu.Lock()
if cur, ok := h.browser[c.ID]; ok && cur == c {
delete(h.browser, c.ID)
toClose = cur
}
h.mu.Unlock()
if toClose != nil {
toClose.close()
}
}
// Clear 清空所有连接(测试/停机用)
func (h *Hub) Clear() {
h.mu.Lock()
for id, c := range h.agents {
c.close()
_ = c.ws.Close(websocket.StatusGoingAway, "hub clear")
delete(h.agents, id)
}
for id, c := range h.browser {
c.close()
_ = c.ws.Close(websocket.StatusGoingAway, "hub clear")
delete(h.browser, id)
}
h.mu.Unlock()
}
// SendToAgent 给设备发消息(1s 超时;队列满丢弃)
func (h *Hub) SendToAgent(deviceID uint64, msg *proto.Envelope) error {
h.mu.RLock()
c, ok := h.agents[deviceID]
h.mu.RUnlock()
if !ok {
return errOffline
}
raw, err := json.Marshal(msg)
if err != nil {
return err
}
if !c.sendRaw(raw, true) {
return errBusy
}
return nil
}
// KickAgent 强制断开设备连接(吊销/远程指令用)
func (h *Hub) KickAgent(deviceID uint64) {
h.mu.RLock()
c, ok := h.agents[deviceID]
h.mu.RUnlock()
if ok {
_ = c.ws.Close(websocket.StatusPolicyViolation, "device revoked")
}
}
// AgentOnline 设备是否在线
func (h *Hub) AgentOnline(deviceID uint64) bool {
h.mu.RLock()
defer h.mu.RUnlock()
_, ok := h.agents[deviceID]
return ok
}
// OnlineDevices 工作空间在线设备列表
func (h *Hub) OnlineDevices(wsid uint64) []uint64 {
h.mu.RLock()
defer h.mu.RUnlock()
ids := make([]uint64, 0)
for id, c := range h.agents {
if c.WSID == wsid {
ids = append(ids, id)
}
}
return ids
}
// BroadcastWS 向工作空间所有浏览器连接广播事件
func (h *Hub) BroadcastWS(wsid uint64, typ string, payload interface{}) {
raw, err := json.Marshal(gin.H{"type": typ, "payload": payload})
if err != nil {
return
}
h.mu.RLock()
defer h.mu.RUnlock()
for _, c := range h.browser {
if c.WSID != wsid {
continue
}
c.sendRaw(raw, false)
}
}
var errOffline = errors.New("device offline")
var errBusy = errors.New("device send queue full")