ChatPersistenceService.java 5.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136
  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. conversationRepository.findByConversationId(conversationId).orElseGet(() -> {
  28. var conv = Conversation.builder()
  29. .conversationId(conversationId)
  30. .userKey(userKey)
  31. .userId(0L)
  32. .title(role.equals("user") && content.length() > 0
  33. ? content.substring(0, Math.min(50, content.length()))
  34. : "")
  35. .build();
  36. return conversationRepository.save(conv);
  37. });
  38. var msg = Message.builder()
  39. .conversationId(conversationId)
  40. .userKey(userKey)
  41. .role(role)
  42. .content(content)
  43. .intent(intent)
  44. .build();
  45. msg.setSourcesList(sources);
  46. messageRepository.save(msg);
  47. }
  48. public Map<String, Object> getHistory(String userKey, int page, int pageSize) {
  49. int safePage = Math.max(1, page);
  50. int safePageSize = Math.max(1, pageSize);
  51. Pageable pageable = PageRequest.of(safePage - 1, safePageSize);
  52. Page<Conversation> convPage = conversationRepository
  53. .findByUserKeyOrderByCreatedAtDesc(userKey, pageable);
  54. List<String> cids = convPage.getContent().stream()
  55. .map(Conversation::getConversationId)
  56. .collect(Collectors.toList());
  57. Map<String, Long> countMap = new HashMap<>();
  58. if (!cids.isEmpty()) {
  59. for (Object[] row : messageRepository.countByConversationIds(cids)) {
  60. countMap.put((String) row[0], (Long) row[1]);
  61. }
  62. }
  63. List<Map<String, Object>> items = new ArrayList<>();
  64. for (Conversation c : convPage.getContent()) {
  65. long count = countMap.getOrDefault(c.getConversationId(), 0L);
  66. items.add(Map.of(
  67. "conversation_id", c.getConversationId(),
  68. "title", c.getTitle() != null && !c.getTitle().isBlank() ? c.getTitle() : "新的对话",
  69. "created_at", c.getCreatedAt() != null ? c.getCreatedAt().toString() : "",
  70. "message_count", count
  71. ));
  72. }
  73. return Map.of(
  74. "items", items,
  75. "page", safePage,
  76. "page_size", safePageSize,
  77. "total", convPage.getTotalElements(),
  78. "total_pages", convPage.getTotalPages()
  79. );
  80. }
  81. public List<Map<String, Object>> getConversationDetail(String cid) {
  82. List<Map<String, Object>> result = new ArrayList<>();
  83. for (var msg : messageRepository.findByConversationIdOrderByCreatedAtAsc(cid)) {
  84. result.add(Map.of(
  85. "id", msg.getId(),
  86. "role", msg.getRole(),
  87. "content", msg.getContent(),
  88. "intent", msg.getIntent() != null ? msg.getIntent() : "",
  89. "sources", msg.getSourcesList(),
  90. "created_at", msg.getCreatedAt() != null ? msg.getCreatedAt().toString() : ""
  91. ));
  92. }
  93. return result;
  94. }
  95. public Optional<Message> findMessageById(Long messageId) {
  96. return messageRepository.findById(messageId);
  97. }
  98. @Transactional
  99. public void updateFeedback(Long messageId, String feedback) {
  100. messageRepository.findById(messageId).ifPresent(msg -> {
  101. msg.setFeedback(feedback);
  102. messageRepository.save(msg);
  103. });
  104. }
  105. /** 获取当前用户最近 N 条消息,用于前端恢复对话 */
  106. public List<Map<String, Object>> getRecentMessages(String userKey, int limit) {
  107. List<Map<String, Object>> result = new ArrayList<>();
  108. var page = messageRepository.findRecentByUserKey(userKey, PageRequest.of(0, limit));
  109. var msgs = new ArrayList<>(page.getContent());
  110. java.util.Collections.reverse(msgs);
  111. for (var msg : msgs) {
  112. result.add(Map.of(
  113. "role", msg.getRole(),
  114. "content", msg.getContent(),
  115. "intent", msg.getIntent() != null ? msg.getIntent() : ""
  116. ));
  117. }
  118. return result;
  119. }
  120. }