ChatPersistenceService.java 6.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149
  1. package com.pharmacopoeia.service;
  2. import com.fasterxml.jackson.databind.ObjectMapper;
  3. import com.pharmacopoeia.entity.Conversation;
  4. import com.pharmacopoeia.entity.Message;
  5. import com.pharmacopoeia.repository.ConversationRepository;
  6. import com.pharmacopoeia.repository.MessageRepository;
  7. import org.springframework.data.domain.Page;
  8. import org.springframework.data.domain.PageRequest;
  9. import org.springframework.data.domain.Pageable;
  10. import org.springframework.stereotype.Service;
  11. import org.springframework.transaction.annotation.Transactional;
  12. import java.util.*;
  13. import java.util.stream.Collectors;
  14. @Service
  15. public class ChatPersistenceService {
  16. private final ConversationRepository conversationRepository;
  17. private final MessageRepository messageRepository;
  18. private final ObjectMapper mapper = new ObjectMapper();
  19. public ChatPersistenceService(ConversationRepository conversationRepository,
  20. MessageRepository messageRepository) {
  21. this.conversationRepository = conversationRepository;
  22. this.messageRepository = messageRepository;
  23. }
  24. @Transactional
  25. public void saveMessage(String userKey, String conversationId, String role, String content,
  26. String intent, List<Map<String, Object>> sources) {
  27. saveMessage(userKey, conversationId, role, content, intent, sources, null);
  28. }
  29. @Transactional
  30. public void saveMessage(String userKey, String conversationId, String role, String content,
  31. String intent, List<Map<String, Object>> sources,
  32. List<Map<String, Object>> brandRecommendations) {
  33. conversationRepository.findByConversationId(conversationId).orElseGet(() -> {
  34. var conv = Conversation.builder()
  35. .conversationId(conversationId)
  36. .userKey(userKey)
  37. .userId(0L)
  38. .title(role.equals("user") && content.length() > 0
  39. ? content.substring(0, Math.min(50, content.length()))
  40. : "")
  41. .build();
  42. return conversationRepository.save(conv);
  43. });
  44. var msg = Message.builder()
  45. .conversationId(conversationId)
  46. .userKey(userKey)
  47. .role(role)
  48. .content(content)
  49. .intent(intent)
  50. .build();
  51. msg.setSourcesList(sources);
  52. if (brandRecommendations != null) {
  53. msg.setBrandRecommendationsList(brandRecommendations);
  54. }
  55. messageRepository.save(msg);
  56. }
  57. public Map<String, Object> getHistory(String userKey, int page, int pageSize) {
  58. int safePage = Math.max(1, page);
  59. int safePageSize = Math.max(1, pageSize);
  60. Pageable pageable = PageRequest.of(safePage - 1, safePageSize);
  61. Page<Conversation> convPage = conversationRepository
  62. .findByUserKeyOrderByCreatedAtDesc(userKey, pageable);
  63. List<String> cids = convPage.getContent().stream()
  64. .map(Conversation::getConversationId)
  65. .collect(Collectors.toList());
  66. Map<String, Long> countMap = new HashMap<>();
  67. if (!cids.isEmpty()) {
  68. for (Object[] row : messageRepository.countByConversationIds(cids)) {
  69. countMap.put((String) row[0], (Long) row[1]);
  70. }
  71. }
  72. List<Map<String, Object>> items = new ArrayList<>();
  73. for (Conversation c : convPage.getContent()) {
  74. long count = countMap.getOrDefault(c.getConversationId(), 0L);
  75. items.add(Map.of(
  76. "conversation_id", c.getConversationId(),
  77. "title", c.getTitle() != null && !c.getTitle().isBlank() ? c.getTitle() : "新的对话",
  78. "created_at", c.getCreatedAt() != null ? c.getCreatedAt().toString() : "",
  79. "message_count", count
  80. ));
  81. }
  82. return Map.of(
  83. "items", items,
  84. "page", safePage,
  85. "page_size", safePageSize,
  86. "total", convPage.getTotalElements(),
  87. "total_pages", convPage.getTotalPages()
  88. );
  89. }
  90. public List<Map<String, Object>> getConversationDetail(String cid) {
  91. List<Map<String, Object>> result = new ArrayList<>();
  92. for (var msg : messageRepository.findByConversationIdOrderByCreatedAtAsc(cid)) {
  93. result.add(Map.of(
  94. "id", msg.getId(),
  95. "role", msg.getRole(),
  96. "content", msg.getContent(),
  97. "intent", msg.getIntent() != null ? msg.getIntent() : "",
  98. "sources", msg.getSourcesList(),
  99. "brand_recommendations", msg.getBrandRecommendationsList(),
  100. "created_at", msg.getCreatedAt() != null ? msg.getCreatedAt().toString() : ""
  101. ));
  102. }
  103. return result;
  104. }
  105. public Optional<Message> findMessageById(Long messageId) {
  106. return messageRepository.findById(messageId);
  107. }
  108. @Transactional
  109. public void updateFeedback(Long messageId, String feedback) {
  110. messageRepository.findById(messageId).ifPresent(msg -> {
  111. msg.setFeedback(feedback);
  112. messageRepository.save(msg);
  113. });
  114. }
  115. /** 获取当前用户最近 N 条消息,用于前端恢复对话 */
  116. public List<Map<String, Object>> getRecentMessages(String userKey, int limit) {
  117. List<Map<String, Object>> result = new ArrayList<>();
  118. var page = messageRepository.findRecentByUserKey(userKey, PageRequest.of(0, limit));
  119. var msgs = new ArrayList<>(page.getContent());
  120. java.util.Collections.reverse(msgs);
  121. for (var msg : msgs) {
  122. result.add(Map.of(
  123. "role", msg.getRole(),
  124. "content", msg.getContent(),
  125. "intent", msg.getIntent() != null ? msg.getIntent() : "",
  126. "sources", msg.getSourcesList(),
  127. "brand_recommendations", msg.getBrandRecommendationsList()
  128. ));
  129. }
  130. return result;
  131. }
  132. }