ChatController.java 7.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193
  1. package com.pharmacopoeia.controller;
  2. import com.pharmacopoeia.dto.ChatRequest;
  3. import com.pharmacopoeia.dto.FeedbackRequest;
  4. import com.pharmacopoeia.service.*;
  5. import org.springframework.http.MediaType;
  6. import org.springframework.http.ResponseEntity;
  7. import org.springframework.http.codec.ServerSentEvent;
  8. import org.springframework.web.bind.annotation.*;
  9. import reactor.core.publisher.Flux;
  10. import reactor.core.publisher.Sinks;
  11. import org.springframework.jdbc.core.JdbcTemplate;
  12. import java.util.*;
  13. import java.util.stream.Collectors;
  14. @RestController
  15. @RequestMapping("/api/v1/chat")
  16. public class ChatController {
  17. private final RetrieverService retrieverService;
  18. private final LLMService llmService;
  19. private final PromptService promptService;
  20. private final ChatPersistenceService persistenceService;
  21. private final JdbcTemplate jdbc;
  22. public ChatController(RetrieverService rs, LLMService ls, PromptService ps,
  23. ChatPersistenceService cps, JdbcTemplate jdbc) {
  24. this.retrieverService = rs;
  25. this.llmService = ls;
  26. this.promptService = ps;
  27. this.persistenceService = cps;
  28. this.jdbc = jdbc;
  29. }
  30. @PostMapping("/ask")
  31. public ResponseEntity<Map<String, Object>> chatAsk(@RequestBody ChatRequest request) {
  32. String query = request.getMessage();
  33. String cid = request.getConversationId() != null && !request.getConversationId().isBlank()
  34. ? request.getConversationId()
  35. : UUID.randomUUID().toString();
  36. String intent = retrieverService.classifyIntent(query);
  37. List<Map<String, Object>> docs = retrieverService.search(query, intent, 20);
  38. docs = rerank(docs, query, 5);
  39. List<Map<String, String>> messages = promptService.buildPrompt(query, docs, intent);
  40. String answer = llmService.chat(messages);
  41. List<Map<String, Object>> sources = buildSources(docs);
  42. persistenceService.saveMessage(cid, "user", query, intent, null);
  43. persistenceService.saveMessage(cid, "assistant", answer, intent, sources);
  44. return ResponseEntity.ok(Map.of(
  45. "answer", answer,
  46. "sources", sources,
  47. "intent", intent
  48. ));
  49. }
  50. @PostMapping(value = "/stream", produces = MediaType.TEXT_EVENT_STREAM_VALUE)
  51. public Flux<ServerSentEvent<String>> chatStream(@RequestBody ChatRequest request) {
  52. String query = request.getMessage();
  53. final String cid = request.getConversationId() != null && !request.getConversationId().isBlank()
  54. ? request.getConversationId()
  55. : UUID.randomUUID().toString();
  56. final String intent = retrieverService.classifyIntent(query);
  57. final List<Map<String, Object>> docs = rerank(retrieverService.search(query, intent, 20), query, 5);
  58. Sinks.Many<ServerSentEvent<String>> sink = Sinks.many().unicast().onBackpressureBuffer();
  59. try {
  60. sink.tryEmitNext(ServerSentEvent.<String>builder().event("intent").data(intent).build());
  61. sink.tryEmitNext(ServerSentEvent.<String>builder().event("status").data("Retrieving...").build());
  62. sink.tryEmitNext(ServerSentEvent.<String>builder().event("status").data("Matched " + docs.size() + " records, generating...").build());
  63. final List<Map<String, String>> messages = promptService.buildPrompt(query, docs, intent);
  64. StringBuilder fullAnswer = new StringBuilder();
  65. llmService.chatStream(messages)
  66. .doOnNext(token -> {
  67. fullAnswer.append(token);
  68. sink.tryEmitNext(ServerSentEvent.<String>builder().data(token).build());
  69. })
  70. .doOnComplete(() -> {
  71. final List<Map<String, Object>> sources = buildSources(docs);
  72. try {
  73. String meta = new com.fasterxml.jackson.databind.ObjectMapper().writeValueAsString(Map.of(
  74. "intent", intent,
  75. "sources", sources
  76. ));
  77. sink.tryEmitNext(ServerSentEvent.<String>builder().event("meta").data(meta).build());
  78. } catch (Exception ignored) {}
  79. persistenceService.saveMessage(cid, "user", query, intent, null);
  80. persistenceService.saveMessage(cid, "assistant", fullAnswer.toString(), intent, sources);
  81. sink.tryEmitComplete();
  82. })
  83. .doOnError(e -> sink.tryEmitError(e))
  84. .subscribe();
  85. } catch (Exception e) {
  86. sink.tryEmitError(e);
  87. }
  88. return sink.asFlux();
  89. }
  90. @GetMapping("/history")
  91. public ResponseEntity<Map<String, Object>> getHistory(
  92. @RequestParam(defaultValue = "1") int page,
  93. @RequestParam(defaultValue = "20") int pageSize) {
  94. var items = persistenceService.getHistory(page, pageSize);
  95. return ResponseEntity.ok(Map.of("items", items, "page", page, "page_size", pageSize));
  96. }
  97. @GetMapping("/history/{cid}")
  98. public ResponseEntity<Map<String, Object>> getConversationDetail(@PathVariable String cid) {
  99. var msgs = persistenceService.getConversationDetail(cid);
  100. return ResponseEntity.ok(Map.of("conversation_id", cid, "messages", msgs));
  101. }
  102. @PostMapping("/feedback")
  103. public ResponseEntity<Map<String, Object>> submitFeedback(@RequestBody FeedbackRequest request) {
  104. persistenceService.updateFeedback(request.getMessageId(), request.getFeedback());
  105. return ResponseEntity.ok(Map.of("status", "ok"));
  106. }
  107. @GetMapping("/admin/conversations")
  108. public ResponseEntity<Map<String, Object>> adminListConversations(
  109. @RequestParam(defaultValue = "1") int page,
  110. @RequestParam(defaultValue = "20") int pageSize,
  111. @RequestParam(required = false) String keyword) {
  112. int offset = (page - 1) * pageSize;
  113. StringBuilder sql = new StringBuilder("""
  114. SELECT DISTINCT ON (c.conversation_id)
  115. c.conversation_id, c.title, c.created_at,
  116. m.content AS last_msg, m.role
  117. FROM conversations c
  118. JOIN messages m ON m.conversation_id = c.conversation_id
  119. """);
  120. List<Object> params = new ArrayList<>();
  121. if (keyword != null && !keyword.isBlank()) {
  122. sql.append("WHERE m.content ILIKE ? ");
  123. params.add("%" + keyword + "%");
  124. }
  125. sql.append("""
  126. ORDER BY c.conversation_id, m.created_at DESC
  127. LIMIT ? OFFSET ?
  128. """);
  129. params.add(pageSize);
  130. params.add(offset);
  131. List<Map<String, Object>> items = jdbc.queryForList(
  132. sql.toString(), params.toArray());
  133. return ResponseEntity.ok(Map.of(
  134. "items", items,
  135. "page", page,
  136. "page_size", pageSize
  137. ));
  138. }
  139. private List<Map<String, Object>> buildSources(List<Map<String, Object>> docs) {
  140. return docs.stream()
  141. .map(d -> Map.<String, Object>of(
  142. "name", d.getOrDefault("source", ""),
  143. "section", d.getOrDefault("section", ""),
  144. "source", d.getOrDefault("source", "")
  145. ))
  146. .collect(Collectors.toList());
  147. }
  148. private List<Map<String, Object>> rerank(List<Map<String, Object>> docs, String query, int topK) {
  149. if (docs.size() <= topK) return docs;
  150. docs.sort((a, b) -> {
  151. double sa = toDouble(a.get("similarity"));
  152. double sb = toDouble(b.get("similarity"));
  153. return Double.compare(sb, sa);
  154. });
  155. return docs.subList(0, Math.min(topK, docs.size()));
  156. }
  157. private double toDouble(Object o) {
  158. if (o instanceof Number n) return n.doubleValue();
  159. return 0;
  160. }
  161. }