package common import ( "encoding/json" "log" "net/http" "sync" "time" "github.com/gorilla/websocket" ) var upgrader = websocket.Upgrader{ CheckOrigin: func(r *http.Request) bool { return true // 允许所有来源 }, } // WebSocketConnection WebSocket连接信息 type WebSocketConnection struct { ID string Conn *websocket.Conn LastSendTime time.Time LastDataHash string IsFirstSend bool SendChan chan []byte mu sync.Mutex } // WebSocketManager WebSocket管理器 type WebSocketManager struct { connections map[string]*WebSocketConnection mutex sync.RWMutex broadcast chan []byte register chan *WebSocketConnection unregister chan *WebSocketConnection } // 全局WebSocket管理器实例 var GlobalWSManager = &WebSocketManager{ connections: make(map[string]*WebSocketConnection), broadcast: make(chan []byte, 100), register: make(chan *WebSocketConnection, 10), unregister: make(chan *WebSocketConnection, 10), } // Start 启动WebSocket管理器 func (m *WebSocketManager) Start() { go m.handleMessages() go m.cleanupStaleConnections() } // handleMessages 处理连接消息 func (m *WebSocketManager) handleMessages() { for { select { case conn := <-m.register: m.mutex.Lock() m.connections[conn.ID] = conn m.mutex.Unlock() log.Printf("WebSocket客户端已连接: %s", conn.ID) case conn := <-m.unregister: m.mutex.Lock() if _, exists := m.connections[conn.ID]; exists { delete(m.connections, conn.ID) conn.Conn.Close() close(conn.SendChan) } m.mutex.Unlock() log.Printf("WebSocket客户端已断开: %s", conn.ID) case message := <-m.broadcast: m.mutex.RLock() for _, conn := range m.connections { select { case conn.SendChan <- message: conn.LastSendTime = time.Now() default: log.Printf("发送队列已满,关闭连接: %s", conn.ID) go func(c *WebSocketConnection) { m.unregister <- c }(conn) } } m.mutex.RUnlock() } } } // cleanupStaleConnections 清理超时连接 func (m *WebSocketManager) cleanupStaleConnections() { ticker := time.NewTicker(30 * time.Second) defer ticker.Stop() for range ticker.C { m.mutex.Lock() now := time.Now() for id, conn := range m.connections { // 如果连接超过5分钟没有发送数据,则关闭 if now.Sub(conn.LastSendTime) > 5*time.Minute { log.Printf("清理超时WebSocket连接: %s", id) delete(m.connections, id) conn.Conn.Close() close(conn.SendChan) } } m.mutex.Unlock() } } // Register 注册新连接 func (m *WebSocketManager) Register(id string, conn *websocket.Conn) *WebSocketConnection { wsConn := &WebSocketConnection{ ID: id, Conn: conn, LastSendTime: time.Now(), IsFirstSend: true, SendChan: make(chan []byte, 10), } // 启动发送协程 go wsConn.writePump() m.register <- wsConn return wsConn } // Unregister 注销连接 func (m *WebSocketManager) Unregister(id string) { m.mutex.RLock() if conn, exists := m.connections[id]; exists { m.unregister <- conn } m.mutex.RUnlock() } // GetConnection 获取连接 func (m *WebSocketManager) GetConnection(id string) (*WebSocketConnection, bool) { m.mutex.RLock() defer m.mutex.RUnlock() conn, exists := m.connections[id] return conn, exists } // Broadcast 广播消息到所有连接 func (m *WebSocketManager) Broadcast(data interface{}) error { message, err := json.Marshal(data) if err != nil { return err } m.broadcast <- message return nil } // BroadcastToClient 发送消息到指定客户端 func (m *WebSocketManager) BroadcastToClient(id string, data interface{}) error { m.mutex.RLock() conn, exists := m.connections[id] m.mutex.RUnlock() if !exists { return nil } message, err := json.Marshal(data) if err != nil { return err } select { case conn.SendChan <- message: conn.LastSendTime = time.Now() default: return nil } return nil } // writePump 写入数据到WebSocket连接 func (c *WebSocketConnection) writePump() { defer func() { c.Conn.Close() GlobalWSManager.Unregister(c.ID) }() for { select { case message, ok := <-c.SendChan: if !ok { return } c.mu.Lock() err := c.Conn.WriteMessage(websocket.TextMessage, message) c.mu.Unlock() if err != nil { log.Printf("WebSocket发送消息失败 %s: %v", c.ID, err) return } } } } // UpdateLastDataHash 更新最后发送的数据哈希 func (c *WebSocketConnection) UpdateLastDataHash(hash string) { c.mu.Lock() defer c.mu.Unlock() c.LastDataHash = hash } // GetLastDataHash 获取最后发送的数据哈希 func (c *WebSocketConnection) GetLastDataHash() string { c.mu.Lock() defer c.mu.Unlock() return c.LastDataHash } // SetIsFirstSend 设置是否首次发送 func (c *WebSocketConnection) SetIsFirstSend(isFirst bool) { c.mu.Lock() defer c.mu.Unlock() c.IsFirstSend = isFirst } // GetIsFirstSend 获取是否首次发送 func (c *WebSocketConnection) GetIsFirstSend() bool { c.mu.Lock() defer c.mu.Unlock() return c.IsFirstSend }