Files
kefu_cloud/server/internal/ws/ws.go
T

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
}