message_service.go 3.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113
  1. package service
  2. import (
  3. "errors"
  4. "log"
  5. "github.com/2930134478/AI-CS/backend/models"
  6. "github.com/2930134478/AI-CS/backend/repository"
  7. "gorm.io/gorm"
  8. )
  9. // ErrConversationClosed indicates operations are attempted on a closed conversation.
  10. var (
  11. // ErrConversationClosed 表示会话已关闭,不能继续发送消息。
  12. ErrConversationClosed = errors.New("conversation is closed")
  13. // ErrConversationNotFound 表示未找到指定的会话记录。
  14. ErrConversationNotFound = gorm.ErrRecordNotFound
  15. )
  16. // MessageService 负责消息领域的业务处理。
  17. type MessageService struct {
  18. conversations *repository.ConversationRepository
  19. messages *repository.MessageRepository
  20. hub BroadcastHub
  21. }
  22. // NewMessageService 创建 MessageService 实例。
  23. func NewMessageService(
  24. conversations *repository.ConversationRepository,
  25. messages *repository.MessageRepository,
  26. hub BroadcastHub,
  27. ) *MessageService {
  28. return &MessageService{
  29. conversations: conversations,
  30. messages: messages,
  31. hub: hub,
  32. }
  33. }
  34. // CreateMessage 创建消息并通过 WebSocket 广播。
  35. func (s *MessageService) CreateMessage(input CreateMessageInput) (*models.Message, error) {
  36. conv, err := s.conversations.GetByID(input.ConversationID)
  37. if err != nil {
  38. return nil, err
  39. }
  40. if conv.Status == "closed" {
  41. return nil, ErrConversationClosed
  42. }
  43. if input.SenderIsAgent && input.SenderID == 0 {
  44. return nil, errors.New("sender_id is required for agent messages")
  45. }
  46. message := &models.Message{
  47. ConversationID: input.ConversationID,
  48. SenderID: input.SenderID,
  49. SenderIsAgent: input.SenderIsAgent,
  50. Content: input.Content,
  51. MessageType: "user_message",
  52. IsRead: false,
  53. }
  54. if err := s.messages.Create(message); err != nil {
  55. return nil, err
  56. }
  57. if err := s.conversations.UpdateFields(conv.ID, map[string]interface{}{
  58. "updated_at": message.CreatedAt,
  59. }); err != nil {
  60. return nil, err
  61. }
  62. if s.hub != nil {
  63. s.hub.BroadcastMessage(message.ConversationID, "new_message", message)
  64. } else {
  65. log.Printf("⚠️ WebSocket Hub 为空,无法广播消息: 消息ID=%d, 对话ID=%d", message.ID, message.ConversationID)
  66. }
  67. return message, nil
  68. }
  69. // ListMessages 返回会话内的全部消息。
  70. func (s *MessageService) ListMessages(conversationID uint) ([]models.Message, error) {
  71. return s.messages.ListByConversationID(conversationID)
  72. }
  73. // MarkMessagesRead 将消息标记为已读并通知监听方。
  74. func (s *MessageService) MarkMessagesRead(conversationID uint, readerIsAgent bool) (*MarkMessagesReadResult, error) {
  75. messageIDs, unreadRemaining, readAt, err := s.messages.MarkMessagesRead(conversationID, !readerIsAgent)
  76. if err != nil {
  77. return nil, err
  78. }
  79. result := &MarkMessagesReadResult{
  80. ConversationID: conversationID,
  81. MessageIDs: messageIDs,
  82. UnreadCount: unreadRemaining,
  83. ReadAt: readAt,
  84. }
  85. if s.hub != nil && len(messageIDs) > 0 {
  86. s.hub.BroadcastMessage(conversationID, "messages_read", map[string]interface{}{
  87. "message_ids": messageIDs,
  88. "reader_is_agent": readerIsAgent,
  89. "read_at": readAt,
  90. "unread_count": unreadRemaining,
  91. "conversation_id": conversationID, // 确保 payload 中也包含 conversation_id
  92. })
  93. }
  94. return result, nil
  95. }