diff --git a/internal/builders/conversation_builder.go b/internal/builders/conversation_builder.go index 7d633cb..3250aa1 100644 --- a/internal/builders/conversation_builder.go +++ b/internal/builders/conversation_builder.go @@ -101,7 +101,7 @@ func BuildMessage(item *models.Message) response.MessageResponse { } func BuildMessageWithReadStates(item *models.Message, agentReadState, customerReadState *models.ConversationReadState, aiSenderNames, userSenderNames map[int64]string, agentProfiles map[int64]*models.AgentProfile) response.MessageResponse { - content, payload := services.BuildRenderableMessage(item) + content, payload := utils.BuildRenderableMessage(item) ret := response.MessageResponse{ ID: item.ID, ConversationID: item.ConversationID, diff --git a/internal/pkg/utils/message.go b/internal/pkg/utils/message.go new file mode 100644 index 0000000..fecddcf --- /dev/null +++ b/internal/pkg/utils/message.go @@ -0,0 +1,287 @@ +package utils + +import ( + "bytes" + "cs-agent/internal/models" + "cs-agent/internal/pkg/enums" + "cs-agent/internal/repositories" + "cs-agent/internal/services/storage" + "encoding/json" + "strings" + + "github.com/microcosm-cc/bluemonday" + "github.com/mlogclub/simple/sqls" + "golang.org/x/net/html" +) + +type imMessageAssetPayload struct { + AssetID string `json:"assetId"` + Provider enums.AssetProvider `json:"provider,omitempty"` + StorageKey string `json:"storageKey,omitempty"` + Filename string `json:"filename,omitempty"` + FileSize int64 `json:"fileSize,omitempty"` + MimeType string `json:"mimeType,omitempty"` + URL string `json:"url,omitempty"` +} + +func SanitizeMessageHTML(content string) string { + policy := bluemonday.UGCPolicy() + policy.AllowElements("img") + policy.AllowAttrs("src", "alt", "title", "data-provider", "data-storage-key").OnElements("img") + policy.AllowURLSchemes("http", "https") + policy.AllowStandardURLs() + policy.AllowElements("p", "br") + return stripHTMLImageSrcIfBound(strings.TrimSpace(policy.Sanitize(content))) +} + +func BuildHTMLSummary(content string) string { + if strings.TrimSpace(content) == "" { + return "" + } + doc, err := html.Parse(strings.NewReader("
" + content + "
")) + if err != nil { + return strings.TrimSpace(content) + } + parts := make([]string, 0, 8) + var walk func(*html.Node) + walk = func(node *html.Node) { + if node == nil { + return + } + if node.Type == html.TextNode { + text := strings.TrimSpace(node.Data) + if text != "" { + parts = append(parts, text) + } + } + if node.Type == html.ElementNode && node.Data == "img" { + parts = append(parts, "[图片]") + } + for child := node.FirstChild; child != nil; child = child.NextSibling { + walk(child) + } + } + walk(doc) + return strings.TrimSpace(strings.Join(parts, " ")) +} + +func BuildRenderableMessage(item *models.Message) (content, payload string) { + if item == nil { + return "", "" + } + if item.RecalledAt != nil { + return "该消息已撤回", "" + } + if item.SendStatus == enums.IMMessageStatusRecalled { + return "该消息已撤回", "" + } + + content = item.Content + payload = item.Payload + switch item.MessageType { + case enums.IMMessageTypeImage, enums.IMMessageTypeAttachment: + payload = buildIMMessageAssetPayloadForResponse(item.Payload) + case enums.IMMessageTypeHTML: + content = BuildMessageHTMLForResponse(item.Content) + } + return content, payload +} + +func BuildMessageHTMLForResponse(content string) string { + content = strings.TrimSpace(content) + if content == "" { + return "" + } + doc, err := html.Parse(strings.NewReader("
" + content + "
")) + if err != nil { + return content + } + var walk func(*html.Node) + walk = func(node *html.Node) { + if node == nil { + return + } + if node.Type == html.ElementNode && node.Data == "img" { + provider := enums.AssetProvider(strings.TrimSpace(findHTMLAttr(node, "data-provider"))) + storageKey := strings.TrimSpace(findHTMLAttr(node, "data-storage-key")) + if provider != "" && storageKey != "" { + if storageProvider, err := storage.NewProvider(provider); err == nil { + setHTMLAttr(node, "src", storageProvider.GetSignedURL(storageKey)) + } + } + } + for child := node.FirstChild; child != nil; child = child.NextSibling { + walk(child) + } + } + walk(doc) + return renderHTMLFragment(doc) +} + +func stripHTMLImageSrcIfBound(content string) string { + content = strings.TrimSpace(content) + if content == "" { + return "" + } + doc, err := html.Parse(strings.NewReader("
" + content + "
")) + if err != nil { + return content + } + var walk func(*html.Node) + walk = func(node *html.Node) { + if node == nil { + return + } + if node.Type == html.ElementNode && node.Data == "img" { + provider := strings.TrimSpace(findHTMLAttr(node, "data-provider")) + storageKey := strings.TrimSpace(findHTMLAttr(node, "data-storage-key")) + if provider != "" && storageKey != "" { + removeHTMLAttr(node, "src") + } + } + for child := node.FirstChild; child != nil; child = child.NextSibling { + walk(child) + } + } + walk(doc) + return renderHTMLFragment(doc) +} + +func renderHTMLFragment(doc *html.Node) string { + if doc == nil { + return "" + } + root := findHTMLRoot(doc) + if root == nil { + return "" + } + var buf bytes.Buffer + for child := root.FirstChild; child != nil; child = child.NextSibling { + if err := html.Render(&buf, child); err != nil { + return "" + } + } + return strings.TrimSpace(buf.String()) +} + +func findHTMLRoot(doc *html.Node) *html.Node { + var walk func(*html.Node) *html.Node + walk = func(node *html.Node) *html.Node { + if node == nil { + return nil + } + if node.Type == html.ElementNode && node.Data == "div" { + return node + } + for child := node.FirstChild; child != nil; child = child.NextSibling { + if found := walk(child); found != nil { + return found + } + } + return nil + } + return walk(doc) +} + +func findHTMLAttr(node *html.Node, key string) string { + if node == nil { + return "" + } + for _, attr := range node.Attr { + if attr.Key == key { + return attr.Val + } + } + return "" +} + +func setHTMLAttr(node *html.Node, key, value string) { + if node == nil { + return + } + for i := range node.Attr { + if node.Attr[i].Key == key { + node.Attr[i].Val = value + return + } + } + node.Attr = append(node.Attr, html.Attribute{Key: key, Val: value}) +} + +func removeHTMLAttr(node *html.Node, key string) { + if node == nil { + return + } + dst := node.Attr[:0] + for _, attr := range node.Attr { + if attr.Key != key { + dst = append(dst, attr) + } + } + node.Attr = dst +} + +func buildIMMessageAssetPayloadForResponse(payload string) string { + assetPayload, err := parseIMMessageAssetPayload(payload) + if err != nil { + return strings.TrimSpace(payload) + } + assetPayload = hydrateIMMessageAssetPayload(assetPayload) + if assetPayload.Provider != "" && assetPayload.StorageKey != "" { + if provider, err := storage.NewProvider(assetPayload.Provider); err == nil { + assetPayload.URL = provider.GetSignedURL(assetPayload.StorageKey) + } + } + data, err := json.Marshal(assetPayload) + if err != nil { + return strings.TrimSpace(payload) + } + return string(data) +} + +func parseIMMessageAssetPayload(payload string) (*imMessageAssetPayload, error) { + payload = strings.TrimSpace(payload) + if payload == "" { + return nil, nil + } + ret := &imMessageAssetPayload{} + if err := json.Unmarshal([]byte(payload), ret); err != nil { + return nil, err + } + ret.AssetID = strings.TrimSpace(ret.AssetID) + ret.Provider = enums.AssetProvider(strings.TrimSpace(string(ret.Provider))) + ret.StorageKey = strings.TrimSpace(ret.StorageKey) + return ret, nil +} + +func hydrateIMMessageAssetPayload(payload *imMessageAssetPayload) *imMessageAssetPayload { + if payload == nil { + return nil + } + if payload.Provider != "" && payload.StorageKey != "" { + return payload + } + if payload.AssetID == "" { + return payload + } + asset := repositories.AssetRepository.GetByAssetID(sqls.DB(), payload.AssetID) + if asset == nil { + return payload + } + if payload.Provider == "" { + payload.Provider = asset.Provider + } + if payload.StorageKey == "" { + payload.StorageKey = strings.TrimSpace(asset.StorageKey) + } + if payload.Filename == "" { + payload.Filename = strings.TrimSpace(asset.Filename) + } + if payload.FileSize <= 0 { + payload.FileSize = asset.FileSize + } + if payload.MimeType == "" { + payload.MimeType = strings.TrimSpace(asset.MimeType) + } + return payload +} diff --git a/internal/services/message_render_test.go b/internal/pkg/utils/message_test.go similarity index 97% rename from internal/services/message_render_test.go rename to internal/pkg/utils/message_test.go index b1ebfe5..eab1572 100644 --- a/internal/services/message_render_test.go +++ b/internal/pkg/utils/message_test.go @@ -1,4 +1,4 @@ -package services +package utils import ( "cs-agent/internal/models" @@ -35,7 +35,7 @@ func TestBuildIMMessageAssetPayloadForResponseAddsSignedURL(t *testing.T) { func TestSanitizeMessageHTMLStripsStoredSrcForManagedImages(t *testing.T) { html := `

demo

` - got := sanitizeMessageHTML(html) + got := SanitizeMessageHTML(html) if strings.Contains(got, `src=`) { t.Fatalf("expected src removed from stored html, got: %s", got) @@ -59,7 +59,7 @@ func TestBuildMessageHTMLForResponseAddsSignedURL(t *testing.T) { }) html := `

demo

` - got := buildMessageHTMLForResponse(html) + got := BuildMessageHTMLForResponse(html) if !strings.Contains(got, `src="https://files.example.com/attachments/demo.png"`) { t.Fatalf("expected signed src in response html, got: %s", got) diff --git a/internal/services/message_service.go b/internal/services/message_service.go index 8bd752e..fad0856 100644 --- a/internal/services/message_service.go +++ b/internal/services/message_service.go @@ -1,24 +1,21 @@ package services import ( - "bytes" "cs-agent/internal/models" "cs-agent/internal/pkg/dto" "cs-agent/internal/pkg/enums" "cs-agent/internal/pkg/errorsx" "cs-agent/internal/pkg/openidentity" + "cs-agent/internal/pkg/utils" "cs-agent/internal/repositories" - "cs-agent/internal/services/storage" "log/slog" "slices" "strings" "time" - "github.com/microcosm-cc/bluemonday" "github.com/mlogclub/simple/common/strs" "github.com/mlogclub/simple/sqls" "github.com/mlogclub/simple/web/params" - "golang.org/x/net/html" ) var MessageService = newMessageService() @@ -194,11 +191,11 @@ func (s *messageService) RecallAgentMessage(messageID int64, operator *dto.AuthP message.UpdateUserName = operator.Username agentReadState, customerReadState := ConversationReadStateService.getConversationReadStates(ctx.Tx, conversation.ID) - agentUnreadCount, err := ConversationReadStateService.CountUnreadMessages(ctx, conversation.ID, readSeqNo(agentReadState), enums.IMSenderTypeCustomer) + agentUnreadCount, err := ConversationReadStateService.CountUnreadMessages(ctx, conversation.ID, s.readSeqNo(agentReadState), enums.IMSenderTypeCustomer) if err != nil { return err } - customerUnreadCount, err := ConversationReadStateService.CountUnreadMessages(ctx, conversation.ID, readSeqNo(customerReadState), enums.IMSenderTypeAgent, enums.IMSenderTypeAI) + customerUnreadCount, err := ConversationReadStateService.CountUnreadMessages(ctx, conversation.ID, s.readSeqNo(customerReadState), enums.IMSenderTypeAgent, enums.IMSenderTypeAI) if err != nil { return err } @@ -354,11 +351,11 @@ func (s *messageService) sendMessage(conversationID int64, senderType enums.IMSe } } agentReadState, customerReadState := ConversationReadStateService.getConversationReadStates(ctx.Tx, conversationID) - agentUnreadCount, err := ConversationReadStateService.CountUnreadMessages(ctx, conversationID, readSeqNo(agentReadState), enums.IMSenderTypeCustomer) + agentUnreadCount, err := ConversationReadStateService.CountUnreadMessages(ctx, conversationID, s.readSeqNo(agentReadState), enums.IMSenderTypeCustomer) if err != nil { return err } - customerUnreadCount, err := ConversationReadStateService.CountUnreadMessages(ctx, conversationID, readSeqNo(customerReadState), enums.IMSenderTypeAgent, enums.IMSenderTypeAI) + customerUnreadCount, err := ConversationReadStateService.CountUnreadMessages(ctx, conversationID, s.readSeqNo(customerReadState), enums.IMSenderTypeAgent, enums.IMSenderTypeAI) if err != nil { return err } @@ -448,7 +445,7 @@ func buildMessageSummary(messageType enums.IMMessageType, content string) string case enums.IMMessageTypeAttachment: return "[附件]" case enums.IMMessageTypeHTML: - return buildHTMLSummary(content) + return utils.BuildHTMLSummary(content) case "": return "" default: @@ -459,8 +456,8 @@ func buildMessageSummary(messageType enums.IMMessageType, content string) string func (s *messageService) normalizeMessageContent(conversationID int64, messageType enums.IMMessageType, content, payload string) (string, string, string, error) { switch messageType { case enums.IMMessageTypeHTML: - sanitized := sanitizeMessageHTML(content) - summary := buildHTMLSummary(sanitized) + sanitized := utils.SanitizeMessageHTML(content) + summary := utils.BuildHTMLSummary(sanitized) if summary == "" { return "", "", "", errorsx.InvalidParam("消息内容不能为空") } @@ -483,7 +480,7 @@ func (s *messageService) normalizeMessageContent(conversationID int64, messageTy summary = "[图片]" } content = strings.TrimSpace(asset.Filename) - return content, canonicalPayload, summary + suffixFilenameForSummary(asset.Filename), nil + return content, canonicalPayload, summary + s.suffixFilenameForSummary(asset.Filename), nil default: content = strings.TrimSpace(content) if content == "" && strings.TrimSpace(payload) == "" { @@ -516,7 +513,7 @@ func (s *messageService) ValidateConversationSender(conversationID int64, sender if operator == nil { return nil, errorsx.Unauthorized("未登录或登录已过期") } - if conversation.Status != enums.IMConversationStatusAIServing && !allowAIMessageOnPendingHandoff(conversation) { + if conversation.Status != enums.IMConversationStatusAIServing && !s.allowAIMessageOnPendingHandoff(conversation) { return nil, errorsx.Forbidden("当前会话不处于 AI 接待状态") } if conversation.CurrentAssigneeID != 0 { @@ -532,7 +529,7 @@ func (s *messageService) ValidateConversationSender(conversationID int64, sender return conversation, nil } -func allowAIMessageOnPendingHandoff(conversation *models.Conversation) bool { +func (s *messageService) allowAIMessageOnPendingHandoff(conversation *models.Conversation) bool { if conversation == nil { return false } @@ -541,7 +538,7 @@ func allowAIMessageOnPendingHandoff(conversation *models.Conversation) bool { conversation.CurrentAssigneeID == 0 } -func suffixFilenameForSummary(filename string) string { +func (s *messageService) suffixFilenameForSummary(filename string) string { filename = strings.TrimSpace(filename) if filename == "" { return "" @@ -549,206 +546,9 @@ func suffixFilenameForSummary(filename string) string { return " " + filename } -func readSeqNo(state *models.ConversationReadState) int64 { +func (s *messageService) readSeqNo(state *models.ConversationReadState) int64 { if state == nil { return 0 } return state.LastReadSeqNo } - -func sanitizeMessageHTML(content string) string { - policy := bluemonday.UGCPolicy() - policy.AllowElements("img") - policy.AllowAttrs("src", "alt", "title", "data-provider", "data-storage-key").OnElements("img") - policy.AllowURLSchemes("http", "https") - policy.AllowStandardURLs() - policy.AllowElements("p", "br") - return stripHTMLImageSrcIfBound(strings.TrimSpace(policy.Sanitize(content))) -} - -func buildHTMLSummary(content string) string { - if strings.TrimSpace(content) == "" { - return "" - } - doc, err := html.Parse(strings.NewReader("
" + content + "
")) - if err != nil { - return strings.TrimSpace(content) - } - parts := make([]string, 0, 8) - var walk func(*html.Node) - walk = func(node *html.Node) { - if node == nil { - return - } - if node.Type == html.TextNode { - text := strings.TrimSpace(node.Data) - if text != "" { - parts = append(parts, text) - } - } - if node.Type == html.ElementNode && node.Data == "img" { - parts = append(parts, "[图片]") - } - for child := node.FirstChild; child != nil; child = child.NextSibling { - walk(child) - } - } - walk(doc) - return strings.TrimSpace(strings.Join(parts, " ")) -} - -func BuildRenderableMessage(item *models.Message) (content, payload string) { - if item == nil { - return "", "" - } - if item.RecalledAt != nil { - return "该消息已撤回", "" - } - if item.SendStatus == enums.IMMessageStatusRecalled { - return "该消息已撤回", "" - } - - content = item.Content - payload = item.Payload - switch item.MessageType { - case enums.IMMessageTypeImage, enums.IMMessageTypeAttachment: - payload = buildIMMessageAssetPayloadForResponse(item.Payload) - case enums.IMMessageTypeHTML: - content = buildMessageHTMLForResponse(item.Content) - } - return content, payload -} - -func buildMessageHTMLForResponse(content string) string { - content = strings.TrimSpace(content) - if content == "" { - return "" - } - doc, err := html.Parse(strings.NewReader("
" + content + "
")) - if err != nil { - return content - } - var walk func(*html.Node) - walk = func(node *html.Node) { - if node == nil { - return - } - if node.Type == html.ElementNode && node.Data == "img" { - provider := enums.AssetProvider(strings.TrimSpace(findHTMLAttr(node, "data-provider"))) - storageKey := strings.TrimSpace(findHTMLAttr(node, "data-storage-key")) - if provider != "" && storageKey != "" { - if storageProvider, err := storage.NewProvider(provider); err == nil { - setHTMLAttr(node, "src", storageProvider.GetSignedURL(storageKey)) - } - } - } - for child := node.FirstChild; child != nil; child = child.NextSibling { - walk(child) - } - } - walk(doc) - return renderHTMLFragment(doc) -} - -func stripHTMLImageSrcIfBound(content string) string { - content = strings.TrimSpace(content) - if content == "" { - return "" - } - doc, err := html.Parse(strings.NewReader("
" + content + "
")) - if err != nil { - return content - } - var walk func(*html.Node) - walk = func(node *html.Node) { - if node == nil { - return - } - if node.Type == html.ElementNode && node.Data == "img" { - provider := strings.TrimSpace(findHTMLAttr(node, "data-provider")) - storageKey := strings.TrimSpace(findHTMLAttr(node, "data-storage-key")) - if provider != "" && storageKey != "" { - removeHTMLAttr(node, "src") - } - } - for child := node.FirstChild; child != nil; child = child.NextSibling { - walk(child) - } - } - walk(doc) - return renderHTMLFragment(doc) -} - -func renderHTMLFragment(doc *html.Node) string { - if doc == nil { - return "" - } - root := findHTMLRoot(doc) - if root == nil { - return "" - } - var buf bytes.Buffer - for child := root.FirstChild; child != nil; child = child.NextSibling { - if err := html.Render(&buf, child); err != nil { - return "" - } - } - return strings.TrimSpace(buf.String()) -} - -func findHTMLRoot(doc *html.Node) *html.Node { - var walk func(*html.Node) *html.Node - walk = func(node *html.Node) *html.Node { - if node == nil { - return nil - } - if node.Type == html.ElementNode && node.Data == "div" { - return node - } - for child := node.FirstChild; child != nil; child = child.NextSibling { - if found := walk(child); found != nil { - return found - } - } - return nil - } - return walk(doc) -} - -func findHTMLAttr(node *html.Node, key string) string { - if node == nil { - return "" - } - for _, attr := range node.Attr { - if attr.Key == key { - return attr.Val - } - } - return "" -} - -func setHTMLAttr(node *html.Node, key, value string) { - if node == nil { - return - } - for i := range node.Attr { - if node.Attr[i].Key == key { - node.Attr[i].Val = value - return - } - } - node.Attr = append(node.Attr, html.Attribute{Key: key, Val: value}) -} - -func removeHTMLAttr(node *html.Node, key string) { - if node == nil { - return - } - dst := node.Attr[:0] - for _, attr := range node.Attr { - if attr.Key != key { - dst = append(dst, attr) - } - } - node.Attr = dst -} diff --git a/internal/services/message_service_test.go b/internal/services/message_service_test.go index d1f4e6b..19ce11a 100644 --- a/internal/services/message_service_test.go +++ b/internal/services/message_service_test.go @@ -14,12 +14,12 @@ func TestAllowAIMessageOnPendingHandoff(t *testing.T) { CurrentAssigneeID: 0, HandoffAt: ptrTime(time.Now()), } - if !allowAIMessageOnPendingHandoff(conversation) { + if !MessageService.allowAIMessageOnPendingHandoff(conversation) { t.Fatalf("expected pending handoff conversation to allow ai handoff notice") } conversation.Status = enums.IMConversationStatusAIServing - if allowAIMessageOnPendingHandoff(conversation) { + if MessageService.allowAIMessageOnPendingHandoff(conversation) { t.Fatalf("expected ai serving conversation not to use pending handoff allowance") } } diff --git a/internal/services/ws_service.go b/internal/services/ws_service.go index 6b557ab..e2f9da9 100644 --- a/internal/services/ws_service.go +++ b/internal/services/ws_service.go @@ -238,7 +238,7 @@ func (s *wsService) PublishMessageCreated(conversation *models.Conversation, mes if conversation == nil || message == nil { return } - content, payload := BuildRenderableMessage(message) + content, payload := utils.BuildRenderableMessage(message) event := s.newEvent(s.conversationTopic(conversation.ID), RealtimeMessageCreatedEvent{ Payload: RealtimeMessageCreatedPayload{