| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136 |
- package com.pharmacopoeia.service;
- import com.fasterxml.jackson.databind.ObjectMapper;
- import com.pharmacopoeia.entity.Conversation;
- import com.pharmacopoeia.entity.Message;
- import com.pharmacopoeia.repository.ConversationRepository;
- import com.pharmacopoeia.repository.MessageRepository;
- import org.springframework.data.domain.Page;
- import org.springframework.data.domain.PageRequest;
- import org.springframework.data.domain.Pageable;
- import org.springframework.stereotype.Service;
- import org.springframework.transaction.annotation.Transactional;
- import java.util.*;
- import java.util.stream.Collectors;
- @Service
- public class ChatPersistenceService {
- private final ConversationRepository conversationRepository;
- private final MessageRepository messageRepository;
- private final ObjectMapper mapper = new ObjectMapper();
- public ChatPersistenceService(ConversationRepository conversationRepository,
- MessageRepository messageRepository) {
- this.conversationRepository = conversationRepository;
- this.messageRepository = messageRepository;
- }
- @Transactional
- public void saveMessage(String userKey, String conversationId, String role, String content,
- String intent, List<Map<String, Object>> sources) {
- conversationRepository.findByConversationId(conversationId).orElseGet(() -> {
- var conv = Conversation.builder()
- .conversationId(conversationId)
- .userKey(userKey)
- .userId(0L)
- .title(role.equals("user") && content.length() > 0
- ? content.substring(0, Math.min(50, content.length()))
- : "")
- .build();
- return conversationRepository.save(conv);
- });
- var msg = Message.builder()
- .conversationId(conversationId)
- .userKey(userKey)
- .role(role)
- .content(content)
- .intent(intent)
- .build();
- msg.setSourcesList(sources);
- messageRepository.save(msg);
- }
- public Map<String, Object> getHistory(String userKey, int page, int pageSize) {
- int safePage = Math.max(1, page);
- int safePageSize = Math.max(1, pageSize);
- Pageable pageable = PageRequest.of(safePage - 1, safePageSize);
- Page<Conversation> convPage = conversationRepository
- .findByUserKeyOrderByCreatedAtDesc(userKey, pageable);
- List<String> cids = convPage.getContent().stream()
- .map(Conversation::getConversationId)
- .collect(Collectors.toList());
- Map<String, Long> countMap = new HashMap<>();
- if (!cids.isEmpty()) {
- for (Object[] row : messageRepository.countByConversationIds(cids)) {
- countMap.put((String) row[0], (Long) row[1]);
- }
- }
- List<Map<String, Object>> items = new ArrayList<>();
- for (Conversation c : convPage.getContent()) {
- long count = countMap.getOrDefault(c.getConversationId(), 0L);
- items.add(Map.of(
- "conversation_id", c.getConversationId(),
- "title", c.getTitle() != null && !c.getTitle().isBlank() ? c.getTitle() : "新的对话",
- "created_at", c.getCreatedAt() != null ? c.getCreatedAt().toString() : "",
- "message_count", count
- ));
- }
- return Map.of(
- "items", items,
- "page", safePage,
- "page_size", safePageSize,
- "total", convPage.getTotalElements(),
- "total_pages", convPage.getTotalPages()
- );
- }
- public List<Map<String, Object>> getConversationDetail(String cid) {
- List<Map<String, Object>> result = new ArrayList<>();
- for (var msg : messageRepository.findByConversationIdOrderByCreatedAtAsc(cid)) {
- result.add(Map.of(
- "id", msg.getId(),
- "role", msg.getRole(),
- "content", msg.getContent(),
- "intent", msg.getIntent() != null ? msg.getIntent() : "",
- "sources", msg.getSourcesList(),
- "created_at", msg.getCreatedAt() != null ? msg.getCreatedAt().toString() : ""
- ));
- }
- return result;
- }
- public Optional<Message> findMessageById(Long messageId) {
- return messageRepository.findById(messageId);
- }
- @Transactional
- public void updateFeedback(Long messageId, String feedback) {
- messageRepository.findById(messageId).ifPresent(msg -> {
- msg.setFeedback(feedback);
- messageRepository.save(msg);
- });
- }
- /** 获取当前用户最近 N 条消息,用于前端恢复对话 */
- public List<Map<String, Object>> getRecentMessages(String userKey, int limit) {
- List<Map<String, Object>> result = new ArrayList<>();
- var page = messageRepository.findRecentByUserKey(userKey, PageRequest.of(0, limit));
- var msgs = new ArrayList<>(page.getContent());
- java.util.Collections.reverse(msgs);
- for (var msg : msgs) {
- result.add(Map.of(
- "role", msg.getRole(),
- "content", msg.getContent(),
- "intent", msg.getIntent() != null ? msg.getIntent() : ""
- ));
- }
- return result;
- }
- }
|