Files

434 lines
12 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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(&current).Error; err == nil {
session = &current
} 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)
}