修复会话安全与实时消息
This commit is contained in:
@@ -1,11 +1,18 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
"kefu-sys/server/internal/middleware"
|
||||
"kefu-sys/server/internal/model"
|
||||
"kefu-sys/server/internal/ws"
|
||||
)
|
||||
|
||||
type WidgetHandler struct{}
|
||||
@@ -13,33 +20,52 @@ type WidgetHandler struct{}
|
||||
func NewWidgetHandler() *WidgetHandler { return &WidgetHandler{} }
|
||||
|
||||
type WidgetInitReq struct {
|
||||
ChannelKey string `json:"channel_key" form:"channel_key"`
|
||||
ChannelKey string `json:"channel_key" form:"channel_key"`
|
||||
VisitorName string `json:"visitor_name"`
|
||||
}
|
||||
|
||||
type WidgetMessageReq struct {
|
||||
SessionID uint `json:"session_id" binding:"required"`
|
||||
Content string `json:"content" binding:"required"`
|
||||
Type string `json:"type"`
|
||||
SessionID uint `json:"session_id" binding:"required"`
|
||||
Content string `json:"content" binding:"required"`
|
||||
Type string `json:"type"`
|
||||
VisitorToken string `json:"visitor_token"`
|
||||
}
|
||||
|
||||
type WidgetRatingReq struct {
|
||||
SessionID uint `json:"session_id" binding:"required"`
|
||||
Score int `json:"score" binding:"required"`
|
||||
Text string `json:"text"`
|
||||
VisitorToken string `json:"visitor_token"`
|
||||
}
|
||||
|
||||
var channelKeyPattern = regexp.MustCompile(`^[A-Za-z0-9_-]{3,64}$`)
|
||||
|
||||
func (h *WidgetHandler) Init(c *gin.Context) {
|
||||
var req WidgetInitReq
|
||||
if err := c.ShouldBindQuery(&req); err != nil && c.ShouldBindJSON(&req) != nil {
|
||||
if err := c.ShouldBindQuery(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"code": 400, "message": "参数错误"})
|
||||
return
|
||||
}
|
||||
if c.Request.Method == http.MethodPost && strings.HasPrefix(c.GetHeader("Content-Type"), "application/json") {
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"code": 400, "message": "参数错误"})
|
||||
return
|
||||
}
|
||||
}
|
||||
if req.ChannelKey == "" {
|
||||
req.ChannelKey = c.Query("channel_key")
|
||||
}
|
||||
if !channelKeyPattern.MatchString(req.ChannelKey) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"code": 400, "message": "渠道标识无效"})
|
||||
return
|
||||
}
|
||||
|
||||
var channel model.Channel
|
||||
if err := model.DB.Where("script_code LIKE ?", "%"+req.ChannelKey+"%").Or("script_code LIKE ?", "%"+req.ChannelKey+"%").First(&channel).Error; err != nil {
|
||||
// 如果没有匹配的渠道,使用第一个启用的渠道
|
||||
if err := model.DB.Where("type = ? AND status = ?", "web", "enabled").First(&channel).Error; err != nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"code": 404, "message": "渠道不存在"})
|
||||
return
|
||||
}
|
||||
if err := model.DB.
|
||||
Where("type = ? AND status = ? AND script_code LIKE ?", "web", "enabled", "%data-id=\""+req.ChannelKey+"\"%").
|
||||
First(&channel).Error; err != nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"code": 404, "message": "渠道不存在或已关闭"})
|
||||
return
|
||||
}
|
||||
|
||||
name := req.VisitorName
|
||||
@@ -60,27 +86,61 @@ func (h *WidgetHandler) Init(c *gin.Context) {
|
||||
model.DB.Create(&customer)
|
||||
}
|
||||
|
||||
visitorToken, visitorTokenHash, err := model.NewVisitorToken()
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"code": 500, "message": "创建访客会话失败"})
|
||||
return
|
||||
}
|
||||
|
||||
// 创建会话
|
||||
session := model.Session{
|
||||
TenantID: channel.TenantID,
|
||||
ChannelID: channel.ID,
|
||||
CustomerID: customer.ID,
|
||||
Status: "waiting",
|
||||
Priority: "normal",
|
||||
TenantID: channel.TenantID,
|
||||
ChannelID: channel.ID,
|
||||
CustomerID: customer.ID,
|
||||
VisitorTokenHash: visitorTokenHash,
|
||||
Status: "waiting",
|
||||
Priority: "normal",
|
||||
}
|
||||
if err := model.DB.Create(&session).Error; err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"code": 500, "message": "创建会话失败"})
|
||||
return
|
||||
}
|
||||
if payload, err := ws.NewEvent("session_created", session.ID, session); err == nil {
|
||||
ws.DefaultHub.BroadcastToTenantStaff(session.TenantID, payload)
|
||||
}
|
||||
model.DB.Create(&session)
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"code": 0,
|
||||
"data": gin.H{
|
||||
"session_id": session.ID,
|
||||
"customer_id": customer.ID,
|
||||
"channel_id": channel.ID,
|
||||
"tenant_id": channel.TenantID,
|
||||
"session_id": session.ID,
|
||||
"customer_id": customer.ID,
|
||||
"channel_id": channel.ID,
|
||||
"tenant_id": channel.TenantID,
|
||||
"visitor_token": visitorToken,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
func visitorTokenFromRequest(c *gin.Context, bodyToken string) string {
|
||||
if token := c.GetHeader("X-Visitor-Token"); token != "" {
|
||||
return token
|
||||
}
|
||||
return bodyToken
|
||||
}
|
||||
|
||||
func loadVisitorSession(c *gin.Context, sessionID uint, token string) (*model.Session, bool) {
|
||||
var session model.Session
|
||||
if err := model.DB.First(&session, sessionID).Error; err != nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"code": 404, "message": "会话不存在"})
|
||||
return nil, false
|
||||
}
|
||||
if !model.VerifyVisitorToken(&session, token) {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"code": 401, "message": "访客会话凭证无效"})
|
||||
return nil, false
|
||||
}
|
||||
return &session, true
|
||||
}
|
||||
|
||||
func (h *WidgetHandler) SendMessage(c *gin.Context) {
|
||||
var req WidgetMessageReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
@@ -91,20 +151,24 @@ func (h *WidgetHandler) SendMessage(c *gin.Context) {
|
||||
req.Type = "text"
|
||||
}
|
||||
|
||||
var session model.Session
|
||||
if err := model.DB.First(&session, req.SessionID).Error; err != nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"code": 404, "message": "会话不存在"})
|
||||
session, ok := loadVisitorSession(c, req.SessionID, visitorTokenFromRequest(c, req.VisitorToken))
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if session.Status == "ended" || session.Status == "archived" {
|
||||
c.JSON(http.StatusConflict, gin.H{"code": 409, "message": "会话已结束,请重新发起咨询"})
|
||||
return
|
||||
}
|
||||
|
||||
// 如果是等待中的会话,更新为活跃
|
||||
if session.Status == "waiting" {
|
||||
model.DB.Model(&session).Update("status", "active")
|
||||
if err := model.DB.Model(session).Update("status", "active").Error; err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"code": 500, "message": "更新会话失败"})
|
||||
return
|
||||
}
|
||||
session.Status = "active"
|
||||
}
|
||||
|
||||
var maxSeq int
|
||||
model.DB.Model(&model.Message{}).Where("session_id = ?", session.ID).Select("COALESCE(MAX(seq), 0)").Scan(&maxSeq)
|
||||
|
||||
custID := session.CustomerID
|
||||
msg := model.Message{
|
||||
SessionID: session.ID,
|
||||
@@ -112,23 +176,110 @@ func (h *WidgetHandler) SendMessage(c *gin.Context) {
|
||||
SenderID: &custID,
|
||||
Content: req.Content,
|
||||
Type: req.Type,
|
||||
Seq: maxSeq + 1,
|
||||
SentAt: time.Now(),
|
||||
}
|
||||
|
||||
if err := model.DB.Create(&msg).Error; err != nil {
|
||||
if err := model.CreateMessage(&msg); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"code": 500, "message": "发送失败"})
|
||||
return
|
||||
}
|
||||
if payload, err := ws.NewEvent("message", session.ID, msg); err == nil {
|
||||
ws.DefaultHub.BroadcastToSession(session.TenantID, session.ID, session.AgentID, payload)
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{"code": 0, "data": msg})
|
||||
}
|
||||
|
||||
func visitorTokenFromWebSocket(c *gin.Context) string {
|
||||
protocols := strings.Split(c.GetHeader("Sec-WebSocket-Protocol"), ",")
|
||||
if len(protocols) == 2 && strings.TrimSpace(protocols[0]) == "kefu-visitor-v1" {
|
||||
return strings.TrimSpace(protocols[1])
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *WidgetHandler) Connect(c *gin.Context) {
|
||||
sessionID, err := strconv.ParseUint(c.Query("session_id"), 10, 64)
|
||||
if err != nil || sessionID == 0 {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"code": 400, "message": "会话参数错误"})
|
||||
return
|
||||
}
|
||||
session, ok := loadVisitorSession(c, uint(sessionID), visitorTokenFromWebSocket(c))
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
client, err := ws.UpgradeVisitor(c.Writer, c.Request, session.TenantID, session.ID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"code": 500, "message": "升级连接失败"})
|
||||
return
|
||||
}
|
||||
ws.HandleWebSocket(client)
|
||||
}
|
||||
|
||||
func (h *WidgetHandler) SubmitRating(c *gin.Context) {
|
||||
var req WidgetRatingReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"code": 400, "message": "参数错误"})
|
||||
return
|
||||
}
|
||||
if req.Score < 1 || req.Score > 5 {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"code": 400, "message": "评分应为 1 至 5 星"})
|
||||
return
|
||||
}
|
||||
session, ok := loadVisitorSession(c, req.SessionID, visitorTokenFromRequest(c, req.VisitorToken))
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if session.Status != "ended" {
|
||||
c.JSON(http.StatusConflict, gin.H{"code": 409, "message": "会话结束后才能评价"})
|
||||
return
|
||||
}
|
||||
if session.SatisfactionScore != nil {
|
||||
c.JSON(http.StatusConflict, gin.H{"code": 409, "message": "该会话已评价"})
|
||||
return
|
||||
}
|
||||
|
||||
if err := model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
result := tx.Model(&model.Session{}).Where("id = ? AND satisfaction_score IS NULL", session.ID).
|
||||
Updates(map[string]interface{}{"satisfaction_score": req.Score, "satisfaction_text": req.Text})
|
||||
if result.Error != nil {
|
||||
return result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
return tx.Model(&model.Customer{}).Where("id = ? AND tenant_id = ?", session.CustomerID, session.TenantID).
|
||||
Updates(map[string]interface{}{
|
||||
"satisfaction_sum": gorm.Expr("satisfaction_sum + ?", req.Score),
|
||||
"satisfaction_count": gorm.Expr("satisfaction_count + 1"),
|
||||
}).Error
|
||||
}); err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
c.JSON(http.StatusConflict, gin.H{"code": 409, "message": "该会话已评价"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"code": 500, "message": "提交评价失败"})
|
||||
return
|
||||
}
|
||||
|
||||
middleware.JSON(c, gin.H{"message": "感谢您的评价"})
|
||||
}
|
||||
|
||||
func (h *WidgetHandler) GetMessages(c *gin.Context) {
|
||||
sessionID := c.Query("session_id")
|
||||
sessionID, err := strconv.ParseUint(c.Query("session_id"), 10, 64)
|
||||
if err != nil || sessionID == 0 {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"code": 400, "message": "会话参数错误"})
|
||||
return
|
||||
}
|
||||
if _, ok := loadVisitorSession(c, uint(sessionID), visitorTokenFromRequest(c, "")); !ok {
|
||||
return
|
||||
}
|
||||
|
||||
var messages []model.Message
|
||||
model.DB.Where("session_id = ?", sessionID).Order("seq asc").Find(&messages)
|
||||
if err := model.DB.Where("session_id = ?", sessionID).Order("seq asc").Find(&messages).Error; err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"code": 500, "message": "查询消息失败"})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{"code": 0, "data": messages})
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user