160 lines
3.1 KiB
Go
160 lines
3.1 KiB
Go
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
|
|
}
|