434 lines
12 KiB
Go
434 lines
12 KiB
Go
package ws
|
||
|
||
import (
|
||
"encoding/json"
|
||
"log"
|
||
"net/http"
|
||
"strings"
|
||
"sync"
|
||
"time"
|
||
"unicode/utf8"
|
||
|
||
"github.com/gorilla/websocket"
|
||
"kefu-cloud/server/internal/model"
|
||
)
|
||
|
||
// ProcessDraftContacts 由 handler 包注入,避免 ws ↔ handler 循环依赖。
|
||
// tenantID, customerID, sessionID, draftText
|
||
var ProcessDraftContacts func(tenantID, customerID, sessionID uint, text string)
|
||
|
||
func processDraftContacts(tenantID, customerID, sessionID uint, text string) {
|
||
if ProcessDraftContacts != nil {
|
||
ProcessDraftContacts(tenantID, customerID, sessionID, text)
|
||
}
|
||
}
|
||
|
||
var upgrader = websocket.Upgrader{
|
||
CheckOrigin: func(r *http.Request) bool { return true },
|
||
Subprotocols: []string{"kefu-v1", "kefu-visitor-v1"},
|
||
}
|
||
|
||
type Client struct {
|
||
Conn *websocket.Conn
|
||
UserID uint
|
||
TenantID uint
|
||
Role string
|
||
Kind string
|
||
SessionID *uint
|
||
Send chan []byte
|
||
}
|
||
|
||
type Event struct {
|
||
Type string `json:"type"`
|
||
SessionID uint `json:"session_id,omitempty"`
|
||
Data interface{} `json:"data,omitempty"`
|
||
Timestamp int64 `json:"timestamp"`
|
||
}
|
||
|
||
type ClientEvent struct {
|
||
Type string `json:"type"`
|
||
SessionID uint `json:"session_id"`
|
||
// Text 访客输入框草稿(type=input_draft 或 typing 附带)
|
||
Text string `json:"text,omitempty"`
|
||
}
|
||
|
||
type Hub struct {
|
||
clients map[*Client]bool
|
||
register chan *Client
|
||
unregister chan *Client
|
||
mu sync.RWMutex
|
||
}
|
||
|
||
var DefaultHub = NewHub()
|
||
|
||
func NewHub() *Hub {
|
||
return &Hub{
|
||
clients: make(map[*Client]bool),
|
||
register: make(chan *Client),
|
||
unregister: make(chan *Client),
|
||
}
|
||
}
|
||
|
||
func NewEvent(eventType string, sessionID uint, data interface{}) ([]byte, error) {
|
||
return json.Marshal(Event{
|
||
Type: eventType,
|
||
SessionID: sessionID,
|
||
Data: data,
|
||
Timestamp: time.Now().UnixMilli(),
|
||
})
|
||
}
|
||
|
||
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()
|
||
}
|
||
}
|
||
}
|
||
|
||
// DisconnectUser 断开指定用户的所有 WebSocket 连接并推送通知。
|
||
func (h *Hub) DisconnectUser(userID uint, message string) {
|
||
h.mu.Lock()
|
||
defer h.mu.Unlock()
|
||
for client := range h.clients {
|
||
if client.UserID == userID {
|
||
if message != "" {
|
||
payload, err := NewEvent("kicked", 0, map[string]string{"message": message})
|
||
if err == nil {
|
||
h.send(client, payload)
|
||
}
|
||
}
|
||
delete(h.clients, client)
|
||
close(client.Send)
|
||
if client.Conn != nil {
|
||
client.Conn.Close()
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
func roleHasPermission(tenantID uint, role, code string) bool {
|
||
var roleRecord model.Role
|
||
if err := model.DB.Where("tenant_id = ? AND code = ?", tenantID, role).First(&roleRecord).Error; err != nil {
|
||
for _, permissionCode := range model.BuiltinRolePermissionCodes()[role] {
|
||
if permissionCode == code {
|
||
return true
|
||
}
|
||
}
|
||
return false
|
||
}
|
||
var count int64
|
||
return model.DB.Table("role_permissions rp").
|
||
Joins("JOIN permissions p ON p.id = rp.permission_id").
|
||
Where("rp.role_id = ? AND p.code = ?", roleRecord.ID, code).
|
||
Count(&count).Error == nil && count > 0
|
||
}
|
||
|
||
func roleDataScope(tenantID uint, role, module string) string {
|
||
defaultScope := model.DefaultRoleDataScopes(role)[module]
|
||
var roleRecord model.Role
|
||
if err := model.DB.Where("tenant_id = ? AND code = ?", tenantID, role).First(&roleRecord).Error; err != nil {
|
||
return defaultScope
|
||
}
|
||
var scope model.RoleDataScope
|
||
if err := model.DB.Where("role_id = ? AND module = ?", roleRecord.ID, module).First(&scope).Error; err != nil {
|
||
return defaultScope
|
||
}
|
||
return scope.Scope
|
||
}
|
||
|
||
// canReceiveSessionEvent 在每次推送前重新校验账号、角色权限和会话数据范围。
|
||
func canReceiveSessionEvent(client *Client, session *model.Session) bool {
|
||
if client.Kind != "agent" {
|
||
return false
|
||
}
|
||
// 单元测试未初始化数据库时保留原有内置角色判定;生产环境始终走实时数据库校验。
|
||
if model.DB == nil {
|
||
if client.Role == "admin" || client.Role == "supervisor" {
|
||
return true
|
||
}
|
||
return session == nil || session.Status == "waiting" ||
|
||
(session.AgentID != nil && client.Role == "agent" && *session.AgentID == client.UserID)
|
||
}
|
||
var user model.User
|
||
if err := model.DB.Select("id", "tenant_id", "role", "status").First(&user, client.UserID).Error; err != nil ||
|
||
user.TenantID != client.TenantID || user.Status == "disabled" {
|
||
return false
|
||
}
|
||
permissionCode := "session.view"
|
||
module := "session"
|
||
if session != nil && (session.Status == "ended" || session.Status == "archived") {
|
||
permissionCode = "chat_history.view"
|
||
module = "chat_history"
|
||
}
|
||
if !roleHasPermission(user.TenantID, user.Role, permissionCode) {
|
||
return false
|
||
}
|
||
if session == nil || roleDataScope(user.TenantID, user.Role, module) == model.DataScopeAll {
|
||
return true
|
||
}
|
||
return session.Status == "waiting" || (session.AgentID != nil && *session.AgentID == user.ID)
|
||
}
|
||
|
||
// Stats 返回当前连接统计(总连接 / 坐席 / 访客)。
|
||
func (h *Hub) Stats() (total, agents, visitors int) {
|
||
h.mu.RLock()
|
||
defer h.mu.RUnlock()
|
||
for client := range h.clients {
|
||
total++
|
||
if client.Kind == "visitor" {
|
||
visitors++
|
||
} else {
|
||
agents++
|
||
}
|
||
}
|
||
return
|
||
}
|
||
|
||
func (h *Hub) send(client *Client, message []byte) {
|
||
select {
|
||
case client.Send <- message:
|
||
default:
|
||
}
|
||
}
|
||
|
||
// BroadcastToSession 仅把消息推送给当前客服、主管/管理员及该会话的访客。
|
||
func (h *Hub) BroadcastToSession(tenantID, sessionID uint, agentID *uint, message []byte) {
|
||
h.mu.RLock()
|
||
defer h.mu.RUnlock()
|
||
|
||
for client := range h.clients {
|
||
if client.TenantID != tenantID {
|
||
continue
|
||
}
|
||
if client.Kind == "visitor" {
|
||
if client.SessionID != nil && *client.SessionID == sessionID {
|
||
h.send(client, message)
|
||
}
|
||
continue
|
||
}
|
||
if canReceiveSessionEvent(client, &model.Session{
|
||
TenantID: tenantID, ID: sessionID, AgentID: agentID, Status: "active",
|
||
}) {
|
||
h.send(client, message)
|
||
}
|
||
}
|
||
}
|
||
|
||
// BroadcastToTenantStaff 只通知租户内的工作人员,不向访客泄露其他会话事件。
|
||
func (h *Hub) BroadcastToTenantStaff(tenantID uint, message []byte) {
|
||
h.mu.RLock()
|
||
defer h.mu.RUnlock()
|
||
|
||
var event Event
|
||
_ = json.Unmarshal(message, &event)
|
||
var session *model.Session
|
||
if event.SessionID != 0 {
|
||
var current model.Session
|
||
if err := model.DB.Where("id = ? AND tenant_id = ?", event.SessionID, tenantID).First(¤t).Error; err == nil {
|
||
session = ¤t
|
||
} else {
|
||
session = &model.Session{TenantID: tenantID, ID: event.SessionID, Status: "active"}
|
||
}
|
||
}
|
||
for client := range h.clients {
|
||
if client.TenantID == tenantID && canReceiveSessionEvent(client, session) {
|
||
h.send(client, message)
|
||
}
|
||
}
|
||
}
|
||
|
||
func (h *Hub) BroadcastToVisitor(tenantID, sessionID uint, message []byte) {
|
||
h.mu.RLock()
|
||
defer h.mu.RUnlock()
|
||
|
||
for client := range h.clients {
|
||
if client.TenantID == tenantID && client.Kind == "visitor" && client.SessionID != nil && *client.SessionID == sessionID {
|
||
h.send(client, message)
|
||
}
|
||
}
|
||
}
|
||
|
||
// BroadcastToSessionStaff 仅推送给可查看该会话的工作人员(不含访客)。
|
||
func (h *Hub) BroadcastToSessionStaff(tenantID uint, agentID *uint, message []byte) {
|
||
h.mu.RLock()
|
||
defer h.mu.RUnlock()
|
||
|
||
for client := range h.clients {
|
||
if client.TenantID != tenantID || client.Kind != "agent" {
|
||
continue
|
||
}
|
||
if canReceiveSessionEvent(client, &model.Session{TenantID: tenantID, AgentID: agentID, Status: "active"}) {
|
||
h.send(client, message)
|
||
}
|
||
}
|
||
}
|
||
|
||
const maxDraftTextRunes = 500
|
||
|
||
func handleClientEvent(client *Client, event ClientEvent) {
|
||
if event.SessionID == 0 {
|
||
return
|
||
}
|
||
// typing:仅指示;input_draft:带正文草稿
|
||
if event.Type != "typing" && event.Type != "input_draft" {
|
||
return
|
||
}
|
||
var session model.Session
|
||
if err := model.DB.Where("id = ? AND tenant_id = ?", event.SessionID, client.TenantID).First(&session).Error; err != nil {
|
||
return
|
||
}
|
||
|
||
// 访客输入中 / 草稿 → 通知坐席
|
||
if client.Kind == "visitor" {
|
||
if client.SessionID == nil || *client.SessionID != event.SessionID {
|
||
return
|
||
}
|
||
if session.Status == "ended" || session.Status == "archived" {
|
||
return
|
||
}
|
||
|
||
if event.Type == "input_draft" {
|
||
text := strings.TrimSpace(event.Text)
|
||
if utf8.RuneCountInString(text) > maxDraftTextRunes {
|
||
text = string([]rune(text)[:maxDraftTextRunes])
|
||
}
|
||
now := time.Now()
|
||
_ = model.DB.Model(&session).Updates(map[string]interface{}{
|
||
"draft_text": text,
|
||
"draft_updated_at": now,
|
||
}).Error
|
||
|
||
// 异步提取联系方式(不阻塞 WS)
|
||
go processDraftContacts(session.TenantID, session.CustomerID, session.ID, text)
|
||
|
||
payload, err := NewEvent("input_draft", session.ID, map[string]interface{}{
|
||
"from": "visitor",
|
||
"text": text,
|
||
})
|
||
if err != nil {
|
||
return
|
||
}
|
||
// 排队中也让租户坐席能看到
|
||
DefaultHub.BroadcastToTenantStaff(session.TenantID, payload)
|
||
return
|
||
}
|
||
|
||
// 兼容旧 typing(无正文)
|
||
payload, err := NewEvent("typing", session.ID, map[string]string{"from": "visitor"})
|
||
if err != nil {
|
||
return
|
||
}
|
||
if session.AgentID == nil {
|
||
DefaultHub.BroadcastToTenantStaff(session.TenantID, payload)
|
||
return
|
||
}
|
||
DefaultHub.BroadcastToSessionStaff(session.TenantID, session.AgentID, payload)
|
||
return
|
||
}
|
||
|
||
// 客服输入中 → 通知访客(不传草稿内容)
|
||
if client.Kind != "agent" || event.Type != "typing" {
|
||
return
|
||
}
|
||
if !canReceiveSessionEvent(client, &session) {
|
||
return
|
||
}
|
||
payload, err := NewEvent("typing", session.ID, map[string]string{"from": "agent"})
|
||
if err == nil {
|
||
DefaultHub.BroadcastToVisitor(session.TenantID, session.ID, payload)
|
||
}
|
||
}
|
||
|
||
func HandleWebSocket(client *Client) {
|
||
defer func() {
|
||
DefaultHub.unregister <- client
|
||
client.Conn.Close()
|
||
}()
|
||
|
||
client.Conn.SetReadLimit(4096) // 允许访客草稿正文
|
||
client.Conn.SetReadDeadline(time.Now().Add(60 * time.Second))
|
||
client.Conn.SetPongHandler(func(string) error {
|
||
client.Conn.SetReadDeadline(time.Now().Add(60 * time.Second))
|
||
return nil
|
||
})
|
||
go writePump(client)
|
||
for {
|
||
_, message, err := client.Conn.ReadMessage()
|
||
if err != nil {
|
||
return
|
||
}
|
||
var event ClientEvent
|
||
if json.Unmarshal(message, &event) == nil {
|
||
handleClientEvent(client, event)
|
||
}
|
||
}
|
||
}
|
||
|
||
func writePump(client *Client) {
|
||
ticker := time.NewTicker(30 * time.Second)
|
||
defer ticker.Stop()
|
||
for {
|
||
select {
|
||
case message, ok := <-client.Send:
|
||
if !ok {
|
||
client.Conn.WriteMessage(websocket.CloseMessage, []byte{})
|
||
return
|
||
}
|
||
client.Conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||
if err := client.Conn.WriteMessage(websocket.TextMessage, message); err != nil {
|
||
return
|
||
}
|
||
case <-ticker.C:
|
||
client.Conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||
if err := client.Conn.WriteMessage(websocket.PingMessage, nil); err != nil {
|
||
return
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
func upgrade(w http.ResponseWriter, r *http.Request, client *Client) (*Client, error) {
|
||
conn, err := upgrader.Upgrade(w, r, nil)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
client.Conn = conn
|
||
DefaultHub.register <- client
|
||
return client, nil
|
||
}
|
||
|
||
func UpgradeAgent(w http.ResponseWriter, r *http.Request, userID, tenantID uint, role string) (*Client, error) {
|
||
client := &Client{
|
||
UserID: userID,
|
||
TenantID: tenantID,
|
||
Role: role,
|
||
Kind: "agent",
|
||
Send: make(chan []byte, 256),
|
||
}
|
||
log.Printf("WebSocket 连接: user=%d tenant=%d", userID, tenantID)
|
||
return upgrade(w, r, client)
|
||
}
|
||
|
||
func UpgradeVisitor(w http.ResponseWriter, r *http.Request, tenantID, sessionID uint) (*Client, error) {
|
||
client := &Client{
|
||
TenantID: tenantID,
|
||
Kind: "visitor",
|
||
SessionID: &sessionID,
|
||
Send: make(chan []byte, 256),
|
||
}
|
||
log.Printf("访客 WebSocket 连接: session=%d tenant=%d", sessionID, tenantID)
|
||
return upgrade(w, r, client)
|
||
}
|