package ws import ( "encoding/json" "log" "net/http" "sync" "time" "github.com/gorilla/websocket" ) var upgrader = websocket.Upgrader{ CheckOrigin: func(r *http.Request) bool { return true }, } type Client struct { Conn *websocket.Conn UserID uint TenantID uint Role string Send chan []byte } type Message struct { Type string `json:"type"` SessionID uint `json:"session_id"` Content string `json:"content,omitempty"` FromID uint `json:"from_id,omitempty"` FromName string `json:"from_name,omitempty"` TenantID uint `json:"tenant_id"` Seq int `json:"seq,omitempty"` Timestamp int64 `json:"timestamp"` } type Hub struct { clients map[*Client]bool broadcast chan []byte register chan *Client unregister chan *Client mu sync.RWMutex } var DefaultHub = NewHub() func NewHub() *Hub { return &Hub{ clients: make(map[*Client]bool), broadcast: make(chan []byte, 256), register: make(chan *Client), unregister: make(chan *Client), } } func (h *Hub) Run() { for { select { case client := <-h.register: h.mu.Lock() h.clients[client] = true h.mu.Unlock() case client := <-h.unregister: h.mu.Lock() if _, ok := h.clients[client]; ok { delete(h.clients, client) close(client.Send) } h.mu.Unlock() case msg := <-h.broadcast: h.mu.RLock() for client := range h.clients { select { case client.Send <- msg: default: close(client.Send) delete(h.clients, client) } } h.mu.RUnlock() } } } func (h *Hub) BroadcastToTenant(tenantID uint, msg []byte) { h.mu.RLock() defer h.mu.RUnlock() for client := range h.clients { if client.TenantID == tenantID { select { case client.Send <- msg: default: } } } } func HandleWebSocket(c *Client) { conn := c.Conn go func() { ticker := time.NewTicker(30 * time.Second) defer ticker.Stop() for { select { case msg, ok := <-c.Send: if !ok { conn.WriteMessage(websocket.CloseMessage, []byte{}) return } conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) if err := conn.WriteMessage(websocket.TextMessage, msg); err != nil { return } case <-ticker.C: conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil { return } } } }() for { _, msgBytes, err := conn.ReadMessage() if err != nil { break } var msg Message if err := json.Unmarshal(msgBytes, &msg); err != nil { continue } msg.TenantID = c.TenantID msg.FromID = c.UserID msg.Timestamp = time.Now().UnixMilli() reply, _ := json.Marshal(msg) DefaultHub.BroadcastToTenant(c.TenantID, reply) } } func Upgrade(w http.ResponseWriter, r *http.Request, userID, tenantID uint, role string) (*Client, error) { conn, err := upgrader.Upgrade(w, r, nil) if err != nil { return nil, err } client := &Client{ Conn: conn, UserID: userID, TenantID: tenantID, Role: role, Send: make(chan []byte, 256), } DefaultHub.register <- client log.Printf("WebSocket 连接: user=%d tenant=%d", userID, tenantID) return client, nil }