ChatController.java 38 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793
  1. package com.pharmacopoeia.controller;
  2. import com.pharmacopoeia.config.QwenProperties;
  3. import com.pharmacopoeia.dto.ChatRequest;
  4. import com.pharmacopoeia.dto.FeedbackRequest;
  5. import com.pharmacopoeia.dto.ImageChatRequest;
  6. import com.pharmacopoeia.dto.MultimodalChatRequest;
  7. import com.pharmacopoeia.service.*;
  8. import jakarta.servlet.http.HttpServletRequest;
  9. import org.springframework.http.MediaType;
  10. import org.springframework.http.ResponseEntity;
  11. import org.springframework.http.codec.ServerSentEvent;
  12. import org.springframework.security.core.context.SecurityContextHolder;
  13. import org.springframework.web.bind.annotation.*;
  14. import reactor.core.publisher.Flux;
  15. import reactor.core.publisher.Sinks;
  16. import org.springframework.jdbc.core.JdbcTemplate;
  17. import org.springframework.web.multipart.MultipartFile;
  18. import java.util.*;
  19. import java.util.stream.Collectors;
  20. @RestController
  21. @RequestMapping("/api/v1/chat")
  22. public class ChatController {
  23. // 复用 PromptService.SECTION_DISPLAY 统一权威映射,避免两处重复定义导致不一致
  24. private final RetrieverService retrieverService;
  25. private final LLMService llmService;
  26. private final PromptService promptService;
  27. private final ChatPersistenceService persistenceService;
  28. private final RerankerService rerankerService;
  29. private final QACacheService qaCache;
  30. private final JdbcTemplate jdbc;
  31. private final QwenProperties props;
  32. private final HttpServletRequest request;
  33. public ChatController(RetrieverService rs, LLMService ls, PromptService ps,
  34. ChatPersistenceService cps, RerankerService rrs,
  35. QACacheService qaCache,
  36. JdbcTemplate jdbc, QwenProperties props,
  37. HttpServletRequest request) {
  38. this.retrieverService = rs;
  39. this.llmService = ls;
  40. this.promptService = ps;
  41. this.persistenceService = cps;
  42. this.rerankerService = rrs;
  43. this.qaCache = qaCache;
  44. this.jdbc = jdbc;
  45. this.props = props;
  46. this.request = request;
  47. }
  48. @PostMapping("/ask")
  49. public ResponseEntity<Map<String, Object>> chatAsk(@RequestBody ChatRequest request) {
  50. String query = request.getMessage();
  51. String normalized = qaCache.normalize(query);
  52. String cid = request.getConversationId() != null && !request.getConversationId().isBlank()
  53. ? request.getConversationId()
  54. : UUID.randomUUID().toString();
  55. // 检查全局缓存(24 小时有效,不区分用户)
  56. Map<String, Object> cached = qaCache.get(normalized);
  57. if (cached != null) {
  58. String answer = (String) cached.get("answer");
  59. String intent = (String) cached.getOrDefault("intent", "");
  60. @SuppressWarnings("unchecked")
  61. List<Map<String, Object>> sources = (List<Map<String, Object>>) cached.getOrDefault("sources", List.of());
  62. persistenceService.saveMessage(userKey, cid, "user", query, intent, null);
  63. persistenceService.saveMessage(userKey, cid, "assistant", answer, intent, sources);
  64. return ResponseEntity.ok(Map.of(
  65. "answer", answer,
  66. "sources", sources,
  67. "intent", intent,
  68. "conversation_id", cid,
  69. "cached", true
  70. ));
  71. }
  72. // 未命中缓存:尝试抢占处理权,避免并发重复调 LLM
  73. if (!qaCache.tryMarkPending(normalized)) {
  74. Map<String, Object> waited = qaCache.waitForCache(normalized);
  75. if (waited != null) {
  76. String answer = (String) waited.get("answer");
  77. String intent = (String) waited.getOrDefault("intent", "");
  78. @SuppressWarnings("unchecked")
  79. List<Map<String, Object>> sources = (List<Map<String, Object>>) waited.getOrDefault("sources", List.of());
  80. persistenceService.saveMessage(userKey, cid, "user", query, intent, null);
  81. persistenceService.saveMessage(userKey, cid, "assistant", answer, intent, sources);
  82. return ResponseEntity.ok(Map.of(
  83. "answer", answer, "sources", sources, "intent", intent,
  84. "conversation_id", cid, "cached", true
  85. ));
  86. }
  87. }
  88. String intent = retrieverService.classifyIntent(query);
  89. List<Map<String, Object>> docs = retrieverService.search(query, intent, 20);
  90. docs = rerankerService.rerank(docs, query, 5);
  91. // 统一:LLM 回答 + 原文对照
  92. List<Map<String, String>> messages = promptService.buildPrompt(query, docs, intent);
  93. String llmAnswer = cleanAnswer(llmService.chat(messages));
  94. String answer = llmAnswer;
  95. List<Map<String, Object>> sources = buildSources(docs);
  96. persistenceService.saveMessage(userKey, cid, "user", query, intent, null);
  97. persistenceService.saveMessage(userKey, cid, "assistant", answer, intent, sources);
  98. // 写入全局缓存
  99. qaCache.put(normalized, answer, intent, sources);
  100. return ResponseEntity.ok(Map.of(
  101. "answer", answer,
  102. "sources", sources,
  103. "intent", intent,
  104. "conversation_id", cid
  105. ));
  106. }
  107. @PostMapping(value = "/stream", produces = MediaType.TEXT_EVENT_STREAM_VALUE)
  108. public Flux<ServerSentEvent<String>> chatStream(@RequestBody ChatRequest request) {
  109. final String userKey = getCurrentUserKey();
  110. String query = request.getMessage();
  111. String normalized = qaCache.normalize(query);
  112. final String cid = request.getConversationId() != null && !request.getConversationId().isBlank()
  113. ? request.getConversationId()
  114. : UUID.randomUUID().toString();
  115. // 检查全局缓存
  116. Map<String, Object> cached = qaCache.get(normalized);
  117. if (cached != null) {
  118. return streamCached(cached, cid, query, normalized);
  119. }
  120. // 未命中缓存:尝试抢占处理权,避免并发重复调 LLM
  121. if (!qaCache.tryMarkPending(normalized)) {
  122. Map<String, Object> waited = qaCache.waitForCache(normalized);
  123. if (waited != null) {
  124. return streamCached(waited, cid, query, normalized);
  125. }
  126. }
  127. final String intent = retrieverService.classifyIntent(query);
  128. return retrieverService.searchReactive(query, intent, 20)
  129. .map(docs -> rerankerService.rerank(docs, query, 5))
  130. .flatMapMany(docs -> {
  131. Sinks.Many<ServerSentEvent<String>> sink = Sinks.many().unicast().onBackpressureBuffer();
  132. sink.tryEmitNext(ServerSentEvent.<String>builder().event("intent").data(intent).build());
  133. sink.tryEmitNext(ServerSentEvent.<String>builder().event("status").data("Retrieving...").build());
  134. sink.tryEmitNext(ServerSentEvent.<String>builder().event("status").data("Matched " + docs.size() + " records, generating...").build());
  135. final List<Map<String, String>> messages = promptService.buildPrompt(query, docs, intent);
  136. StringBuilder fullAnswer = new StringBuilder();
  137. llmService.chatStream(messages)
  138. .doOnNext(token -> {
  139. fullAnswer.append(token);
  140. sink.tryEmitNext(ServerSentEvent.<String>builder().data(token).build());
  141. })
  142. .doOnComplete(() -> {
  143. final List<Map<String, Object>> sources = buildSources(docs);
  144. try {
  145. String meta = new com.fasterxml.jackson.databind.ObjectMapper().writeValueAsString(Map.of(
  146. "intent", intent,
  147. "sources", sources,
  148. "conversation_id", cid
  149. ));
  150. sink.tryEmitNext(ServerSentEvent.<String>builder().event("meta").data(meta).build());
  151. } catch (Exception ignored) {}
  152. String finalAnswer = cleanAnswer(fullAnswer.toString());
  153. persistenceService.saveMessage(userKey, cid, "user", query, intent, null);
  154. persistenceService.saveMessage(userKey, cid, "assistant", finalAnswer, intent, sources);
  155. // 写入全局缓存
  156. qaCache.put(normalized, finalAnswer, intent, sources);
  157. sink.tryEmitComplete();
  158. })
  159. .doOnError(sink::tryEmitError)
  160. .subscribe();
  161. return sink.asFlux();
  162. });
  163. }
  164. /** 将缓存命中结果以流式 SSE 形式返回 */
  165. private Flux<ServerSentEvent<String>> streamCached(Map<String, Object> cached, String cid,
  166. String query, String normalized) {
  167. String answer = (String) cached.get("answer");
  168. String intent = (String) cached.getOrDefault("intent", "");
  169. @SuppressWarnings("unchecked")
  170. List<Map<String, Object>> sources = (List<Map<String, Object>>) cached.getOrDefault("sources", List.of());
  171. persistenceService.saveMessage(userKey, cid, "user", query, intent, null);
  172. persistenceService.saveMessage(userKey, cid, "assistant", answer, intent, sources);
  173. return Flux.create(sink -> {
  174. sink.next(ServerSentEvent.<String>builder().event("intent").data(intent).build());
  175. sink.next(ServerSentEvent.<String>builder().event("status").data("命中缓存,直接返回...").build());
  176. // 将缓存答案按段落拆分发送,模拟流式体验
  177. String[] chunks = answer.split("(?<=\\n)");
  178. for (String chunk : chunks) {
  179. sink.next(ServerSentEvent.<String>builder().data(chunk).build());
  180. }
  181. try {
  182. String meta = new com.fasterxml.jackson.databind.ObjectMapper().writeValueAsString(Map.of(
  183. "intent", intent, "sources", sources, "conversation_id", cid, "cached", true
  184. ));
  185. sink.next(ServerSentEvent.<String>builder().event("meta").data(meta).build());
  186. } catch (Exception ignored) {}
  187. sink.complete();
  188. });
  189. }
  190. // ============================================================
  191. // 图片对话 API(Qwen VL 分析 + OCR → RAG 检索 → 联网搜索)
  192. // ============================================================
  193. @PostMapping("/ask-image")
  194. public ResponseEntity<Map<String, Object>> chatAskImage(@RequestBody ImageChatRequest request) {
  195. String cid = request.getConversationId() != null && !request.getConversationId().isBlank()
  196. ? request.getConversationId()
  197. : UUID.randomUUID().toString();
  198. // Step 1: Qwen VL 分析图片 + OCR 提取文字
  199. String ocrText = llmService.analyzeImage(
  200. request.getImageBase64(), request.getMimeType(),
  201. "请分析这张图片,提取其中所有文字信息(OCR),特别是药品名称、成分、用法用量等关键药学信息。简要输出即可。");
  202. // Step 2: 拼接查询 → RAG 检索
  203. String query = (!request.getMessage().isBlank())
  204. ? request.getMessage() + "\n\n(图片OCR提取内容:" + ocrText + ")"
  205. : ocrText;
  206. String intent = retrieverService.classifyIntent(query);
  207. List<Map<String, Object>> docs = retrieverService.search(query, intent, 20);
  208. docs = rerankerService.rerank(docs, query, 5);
  209. // Step 3: 构建 Prompt(含图片分析上下文)+ 联网搜索
  210. List<Map<String, String>> messages = promptService.buildPrompt(query, docs, intent);
  211. String imageContext = "\n\n【图片分析结果】\n" + ocrText + "\n";
  212. messages.get(0).put("content", messages.get(0).get("content") + imageContext);
  213. String llmAnswer = cleanAnswer(llmService.chat(messages, true));
  214. String answer = llmAnswer;
  215. List<Map<String, Object>> sources = buildSources(docs);
  216. persistenceService.saveMessage(userKey, cid, "user",
  217. request.getMessage().isBlank() ? "[图片]" : "[图片] " + request.getMessage(),
  218. intent, null);
  219. persistenceService.saveMessage(userKey, cid, "assistant", answer, intent, sources);
  220. return ResponseEntity.ok(Map.of(
  221. "answer", answer,
  222. "sources", sources,
  223. "intent", intent,
  224. "conversation_id", cid
  225. ));
  226. }
  227. @PostMapping(value = "/stream-image", produces = MediaType.TEXT_EVENT_STREAM_VALUE)
  228. public Flux<ServerSentEvent<String>> chatStreamImage(@RequestBody ImageChatRequest request) {
  229. final String cid = request.getConversationId() != null && !request.getConversationId().isBlank()
  230. ? request.getConversationId()
  231. : UUID.randomUUID().toString();
  232. Sinks.Many<ServerSentEvent<String>> sink = Sinks.many().unicast().onBackpressureBuffer();
  233. // 校验图片数据
  234. if (request.getImageBase64() == null || request.getImageBase64().isBlank()) {
  235. sink.tryEmitNext(ServerSentEvent.<String>builder().event("status").data("图片数据为空,请重新上传").build());
  236. sink.tryEmitComplete();
  237. return sink.asFlux();
  238. }
  239. sink.tryEmitNext(ServerSentEvent.<String>builder().event("status").data("正在分析图片(OCR 文字识别)...").build());
  240. // Step 1: Qwen VL 分析图片(流式),收集完整结果
  241. StringBuilder ocrBuilder = new StringBuilder();
  242. llmService.analyzeImageStream(request.getImageBase64(), request.getMimeType(),
  243. "请分析这张图片,提取其中所有文字信息(OCR),特别是药品名称、成分、用法用量等。简要输出。")
  244. .doOnNext(ocrBuilder::append)
  245. .collectList()
  246. .flatMap(tokens -> {
  247. String ocrText = ocrBuilder.toString();
  248. sink.tryEmitNext(ServerSentEvent.<String>builder().event("status").data("图片分析完成,正在检索药典知识库...").build());
  249. // Step 2: 拼接查询 → RAG(响应式)
  250. String query = (!request.getMessage().isBlank())
  251. ? request.getMessage() + "\n\n(图片OCR提取内容:" + ocrText + ")"
  252. : ocrText;
  253. final String intent = retrieverService.classifyIntent(query);
  254. sink.tryEmitNext(ServerSentEvent.<String>builder().event("intent").data(intent).build());
  255. return retrieverService.searchReactive(query, intent, 20)
  256. .map(docs -> rerankerService.rerank(docs, query, 5))
  257. .map(docs -> {
  258. sink.tryEmitNext(ServerSentEvent.<String>builder().event("status")
  259. .data("已匹配 " + docs.size() + " 条药典资料,生成回答中(已启用联网搜索)...").build());
  260. // Step 3: 构建 Prompt + 联网搜索流式生成
  261. final List<Map<String, String>> messages = promptService.buildPrompt(query, docs, intent);
  262. String imageContext = "\n\n【图片分析结果】\n" + ocrText + "\n";
  263. messages.get(0).put("content", messages.get(0).get("content") + imageContext);
  264. StringBuilder fullAnswer = new StringBuilder();
  265. llmService.chatStream(messages, true)
  266. .doOnNext(token -> {
  267. fullAnswer.append(token);
  268. sink.tryEmitNext(ServerSentEvent.<String>builder().data(token).build());
  269. })
  270. .doOnComplete(() -> {
  271. final List<Map<String, Object>> sources = buildSources(docs);
  272. try {
  273. String meta = new com.fasterxml.jackson.databind.ObjectMapper().writeValueAsString(Map.of(
  274. "intent", intent,
  275. "sources", sources,
  276. "conversation_id", cid,
  277. "ocr_text", ocrText
  278. ));
  279. sink.tryEmitNext(ServerSentEvent.<String>builder().event("meta").data(meta).build());
  280. } catch (Exception ignored) {}
  281. String finalAnswer = cleanAnswer(fullAnswer.toString());
  282. persistenceService.saveMessage(userKey, cid, "user",
  283. request.getMessage().isBlank() ? "[图片]" : "[图片] " + request.getMessage(),
  284. intent, null);
  285. persistenceService.saveMessage(userKey, cid, "assistant", finalAnswer, intent, sources);
  286. sink.tryEmitComplete();
  287. })
  288. .doOnError(sink::tryEmitError)
  289. .subscribe();
  290. return docs; // dummy return for map
  291. });
  292. })
  293. .doOnError(sink::tryEmitError)
  294. .subscribe();
  295. return sink.asFlux();
  296. }
  297. // ============================================================
  298. // 统一多模态对话 API(文本 + 图片 + 视频)
  299. // ============================================================
  300. @PostMapping("/ask-multimodal")
  301. public ResponseEntity<Map<String, Object>> chatAskMultimodal(@RequestBody MultimodalChatRequest request) {
  302. String cid = request.getConversationId() != null && !request.getConversationId().isBlank()
  303. ? request.getConversationId()
  304. : UUID.randomUUID().toString();
  305. // Step 1: 媒体分析(如果有附件)
  306. String ocrText = "";
  307. String mediaLabel = "";
  308. if (request.getMediaBase64() != null && !request.getMediaBase64().isBlank()
  309. && request.getMediaType() != null && !request.getMediaType().isBlank()) {
  310. mediaLabel = "video".equals(request.getMediaType()) ? "视频" : "图片";
  311. ocrText = llmService.analyzeMedia(
  312. request.getMediaBase64(), request.getMediaType(),
  313. request.getMediaMime(), "");
  314. }
  315. // Step 2: 拼接查询
  316. String query = request.getMessage() != null ? request.getMessage().trim() : "";
  317. if (!query.isEmpty() && !ocrText.isEmpty()) {
  318. query = query + "\n\n(" + mediaLabel + "OCR提取内容:" + ocrText + ")";
  319. } else if (!ocrText.isEmpty()) {
  320. query = ocrText;
  321. } else if (query.isEmpty()) {
  322. query = "请介绍一下自己";
  323. }
  324. // Step 3: RAG 检索
  325. String intent = retrieverService.classifyIntent(query);
  326. List<Map<String, Object>> docs = retrieverService.search(query, intent, 20);
  327. docs = rerankerService.rerank(docs, query, 5);
  328. // Step 4: 构建 Prompt + 联网搜索
  329. List<Map<String, String>> messages = promptService.buildPrompt(query, docs, intent);
  330. if (!ocrText.isEmpty()) {
  331. messages.get(0).put("content",
  332. messages.get(0).get("content") + "\n\n【" + mediaLabel + "分析结果】\n" + ocrText + "\n");
  333. }
  334. boolean enableSearch = !ocrText.isEmpty() || props.isEnableWebSearch();
  335. String llmAnswer = cleanAnswer(llmService.chat(messages, enableSearch));
  336. String answer = llmAnswer;
  337. List<Map<String, Object>> sources = buildSources(docs);
  338. String userMsg = !request.getMessage().isBlank() ? request.getMessage()
  339. : !ocrText.isEmpty() ? "[" + mediaLabel + "]" : request.getMessage();
  340. persistenceService.saveMessage(userKey, cid, "user", userMsg, intent, null);
  341. persistenceService.saveMessage(userKey, cid, "assistant", answer, intent, sources);
  342. return ResponseEntity.ok(Map.of(
  343. "answer", answer,
  344. "sources", sources,
  345. "intent", intent,
  346. "conversation_id", cid
  347. ));
  348. }
  349. @PostMapping(value = "/stream-multimodal", produces = MediaType.TEXT_EVENT_STREAM_VALUE)
  350. public Flux<ServerSentEvent<String>> chatStreamMultimodal(@RequestBody MultimodalChatRequest request) {
  351. final String cid = request.getConversationId() != null && !request.getConversationId().isBlank()
  352. ? request.getConversationId()
  353. : UUID.randomUUID().toString();
  354. final boolean hasMedia = request.getMediaBase64() != null && !request.getMediaBase64().isBlank()
  355. && request.getMediaType() != null && !request.getMediaType().isBlank();
  356. final String mediaLabel = hasMedia && "video".equals(request.getMediaType()) ? "视频" : "图片";
  357. Sinks.Many<ServerSentEvent<String>> sink = Sinks.many().unicast().onBackpressureBuffer();
  358. try {
  359. if (hasMedia) {
  360. sink.tryEmitNext(ServerSentEvent.<String>builder().event("status")
  361. .data("🔍 正在分析" + mediaLabel + "...").build());
  362. // 先发 OCR section header
  363. sink.tryEmitNext(ServerSentEvent.<String>builder()
  364. .data("【📷 " + mediaLabel + "分析】\n\n").build());
  365. StringBuilder ocrBuilder = new StringBuilder();
  366. llmService.analyzeMediaStream(request.getMediaBase64(), request.getMediaType(),
  367. request.getMediaMime(), "")
  368. .doOnNext(token -> {
  369. ocrBuilder.append(token);
  370. sink.tryEmitNext(ServerSentEvent.<String>builder().data(token).build());
  371. })
  372. .doOnComplete(() -> {
  373. sink.tryEmitNext(ServerSentEvent.<String>builder().data("\n\n").build());
  374. doStreamAnswer(sink, cid, request, ocrBuilder.toString(), mediaLabel);
  375. })
  376. .doOnError(sink::tryEmitError)
  377. .subscribe();
  378. } else {
  379. doStreamAnswer(sink, cid, request, "", "");
  380. }
  381. } catch (Exception e) {
  382. sink.tryEmitError(e);
  383. }
  384. return sink.asFlux();
  385. }
  386. // ============================================================
  387. // 文件上传 API(multipart → base64 → 复用已有对话管线)
  388. // ============================================================
  389. @PostMapping("/upload-image")
  390. public ResponseEntity<Map<String, Object>> uploadImage(
  391. @RequestParam("file") MultipartFile file,
  392. @RequestParam(defaultValue = "") String message,
  393. @RequestParam(defaultValue = "") String conversationId) {
  394. // 校验 MIME 类型
  395. Set<String> allowed = Set.of("image/jpeg", "image/png", "image/webp", "image/bmp");
  396. String contentType = file.getContentType();
  397. if (contentType == null || !allowed.contains(contentType)) {
  398. throw new IllegalArgumentException(
  399. "不支持的图片格式: " + contentType + ",支持 jpg/png/webp/bmp");
  400. }
  401. // 校验大小 ≤ 10MB
  402. if (file.getSize() > 10 * 1024 * 1024) {
  403. throw new IllegalArgumentException("图片大小不能超过 10MB");
  404. }
  405. // 转 base64 → 委托给 ask-image
  406. String base64;
  407. try {
  408. base64 = Base64.getEncoder().encodeToString(file.getBytes());
  409. } catch (Exception e) {
  410. throw new RuntimeException("读取上传文件失败", e);
  411. }
  412. ImageChatRequest req = new ImageChatRequest();
  413. req.setImageBase64(base64);
  414. req.setMimeType(contentType);
  415. req.setMessage(message);
  416. req.setConversationId(
  417. conversationId.isBlank() ? UUID.randomUUID().toString() : conversationId);
  418. return chatAskImage(req);
  419. }
  420. @PostMapping("/upload-media")
  421. public ResponseEntity<Map<String, Object>> uploadMedia(
  422. @RequestParam("file") MultipartFile file,
  423. @RequestParam(defaultValue = "") String message,
  424. @RequestParam(defaultValue = "") String conversationId) {
  425. String contentType = file.getContentType();
  426. if (contentType == null) {
  427. throw new IllegalArgumentException("无法识别的媒体类型");
  428. }
  429. String mediaType;
  430. long maxSize;
  431. if (contentType.startsWith("image/")) {
  432. mediaType = "image";
  433. maxSize = 10 * 1024 * 1024; // 10MB
  434. } else if (contentType.startsWith("video/")) {
  435. mediaType = "video";
  436. maxSize = 50 * 1024 * 1024; // 50MB
  437. } else {
  438. throw new IllegalArgumentException(
  439. "不支持的媒体格式: " + contentType + ",支持 jpg/png/webp/bmp/mp4/mov/avi/webm");
  440. }
  441. if (file.getSize() > maxSize) {
  442. throw new IllegalArgumentException(
  443. "文件大小不能超过 " + (maxSize / 1024 / 1024) + "MB");
  444. }
  445. String base64;
  446. try {
  447. base64 = Base64.getEncoder().encodeToString(file.getBytes());
  448. } catch (Exception e) {
  449. throw new RuntimeException("读取上传文件失败", e);
  450. }
  451. MultimodalChatRequest req = new MultimodalChatRequest();
  452. req.setMessage(message);
  453. req.setMediaType(mediaType);
  454. req.setMediaBase64(base64);
  455. req.setMediaMime(contentType);
  456. req.setConversationId(
  457. conversationId.isBlank() ? UUID.randomUUID().toString() : conversationId);
  458. return chatAskMultimodal(req);
  459. }
  460. /** 流式多模态:OCR 完成后,走 RAG + 生成 */
  461. private void doStreamAnswer(Sinks.Many<ServerSentEvent<String>> sink, String cid,
  462. MultimodalChatRequest request, String ocrText, String mediaLabel) {
  463. if (!ocrText.isEmpty()) {
  464. sink.tryEmitNext(ServerSentEvent.<String>builder().event("status")
  465. .data("📚 检索药典知识库...").build());
  466. }
  467. String rawQuery = request.getMessage() != null ? request.getMessage().trim() : "";
  468. if (!rawQuery.isEmpty() && !ocrText.isEmpty()) {
  469. rawQuery = rawQuery + "\n\n(" + mediaLabel + "OCR提取内容:" + ocrText + ")";
  470. } else if (!ocrText.isEmpty()) {
  471. rawQuery = ocrText;
  472. } else if (rawQuery.isEmpty()) {
  473. rawQuery = "请介绍一下自己";
  474. }
  475. final String query = rawQuery;
  476. final String intent = retrieverService.classifyIntent(query);
  477. sink.tryEmitNext(ServerSentEvent.<String>builder().event("intent").data(intent).build());
  478. final boolean enableSearch = !ocrText.isEmpty() || props.isEnableWebSearch();
  479. retrieverService.searchReactive(query, intent, 20)
  480. .map(docs -> rerankerService.rerank(docs, query, 5))
  481. .subscribe(docs -> {
  482. sink.tryEmitNext(ServerSentEvent.<String>builder().event("status")
  483. .data("已匹配 " + docs.size() + " 条药典资料,生成回答中"
  484. + (enableSearch ? "(已启用联网搜索)" : "") + "...").build());
  485. final List<Map<String, String>> messages = promptService.buildPrompt(query, docs, intent);
  486. if (!ocrText.isEmpty()) {
  487. messages.get(0).put("content",
  488. messages.get(0).get("content") + "\n\n【" + mediaLabel + "分析结果】\n" + ocrText + "\n");
  489. }
  490. // 发送回答 section header
  491. sink.tryEmitNext(ServerSentEvent.<String>builder().data("\n【📚 药典参考回答】\n\n").build());
  492. StringBuilder fullAnswer = new StringBuilder();
  493. llmService.chatStream(messages, enableSearch)
  494. .doOnNext(token -> {
  495. fullAnswer.append(token);
  496. sink.tryEmitNext(ServerSentEvent.<String>builder().data(token).build());
  497. })
  498. .doOnComplete(() -> {
  499. final List<Map<String, Object>> sources = buildSources(docs);
  500. try {
  501. String meta = new com.fasterxml.jackson.databind.ObjectMapper().writeValueAsString(Map.of(
  502. "intent", intent, "sources", sources, "conversation_id", cid,
  503. "ocr_text", ocrText));
  504. sink.tryEmitNext(ServerSentEvent.<String>builder().event("meta").data(meta).build());
  505. } catch (Exception ignored) {}
  506. String finalAnswer = cleanAnswer(fullAnswer.toString());
  507. String userMsg = !request.getMessage().isBlank() ? request.getMessage()
  508. : !ocrText.isEmpty() ? "[" + mediaLabel + "]" : "";
  509. persistenceService.saveMessage(userKey, cid, "user", userMsg, intent, null);
  510. persistenceService.saveMessage(userKey, cid, "assistant", finalAnswer, intent, sources);
  511. sink.tryEmitComplete();
  512. })
  513. .doOnError(sink::tryEmitError)
  514. .subscribe();
  515. }, sink::tryEmitError);
  516. }
  517. @GetMapping("/history")
  518. public ResponseEntity<Map<String, Object>> getHistory(
  519. @RequestParam(defaultValue = "1") int page,
  520. @RequestParam(defaultValue = "20") int pageSize) {
  521. // getHistory 现在直接返回包含 items/page/page_size/total/total_pages 的 Map
  522. var result = persistenceService.getHistory(userKey, page, pageSize);
  523. return ResponseEntity.ok(result);
  524. }
  525. @GetMapping("/history/{cid}")
  526. public ResponseEntity<Map<String, Object>> getConversationDetail(@PathVariable String cid) {
  527. var msgs = persistenceService.getConversationDetail(cid);
  528. return ResponseEntity.ok(Map.of("conversation_id", cid, "messages", msgs));
  529. }
  530. /** 返回最近 N 条消息,供前端恢复对话(微信 WebView 等 IndexedDB 不可用场景) */
  531. @GetMapping("/recent-messages")
  532. public ResponseEntity<Map<String, Object>> getRecentMessages(
  533. @RequestParam(defaultValue = "50") int limit) {
  534. var msgs = persistenceService.getRecentMessages(userKey, Math.min(limit, 200));
  535. return ResponseEntity.ok(Map.of("messages", msgs));
  536. }
  537. @PostMapping("/feedback")
  538. public ResponseEntity<Map<String, Object>> submitFeedback(@RequestBody FeedbackRequest request) {
  539. persistenceService.updateFeedback(request.getMessageId(), request.getFeedback());
  540. return ResponseEntity.ok(Map.of("status", "ok"));
  541. }
  542. @GetMapping("/admin/conversations")
  543. public ResponseEntity<Map<String, Object>> adminListConversations(
  544. @RequestParam(defaultValue = "1") int page,
  545. @RequestParam(defaultValue = "20") int pageSize,
  546. @RequestParam(required = false) String keyword) {
  547. int offset = (page - 1) * pageSize;
  548. StringBuilder sql = new StringBuilder("""
  549. SELECT DISTINCT ON (c.conversation_id)
  550. c.conversation_id, c.title, c.created_at,
  551. m.content AS last_msg, m.role
  552. FROM conversations c
  553. JOIN messages m ON m.conversation_id = c.conversation_id
  554. """);
  555. List<Object> params = new ArrayList<>();
  556. if (keyword != null && !keyword.isBlank()) {
  557. sql.append("WHERE m.content ILIKE ? ");
  558. params.add("%" + keyword + "%");
  559. }
  560. sql.append("""
  561. ORDER BY c.conversation_id, m.created_at DESC
  562. LIMIT ? OFFSET ?
  563. """);
  564. params.add(pageSize);
  565. params.add(offset);
  566. List<Map<String, Object>> items = jdbc.queryForList(
  567. sql.toString(), params.toArray());
  568. return ResponseEntity.ok(Map.of(
  569. "items", items,
  570. "page", page,
  571. "page_size", pageSize
  572. ));
  573. }
  574. private List<Map<String, Object>> buildSources(List<Map<String, Object>> docs) {
  575. Set<String> rawSeen = new HashSet<>();
  576. Set<String> seen = new HashSet<>();
  577. return docs.stream()
  578. // 第一层:按原始 name|section 去重,消除同一栏目的多个分块
  579. .filter(d -> {
  580. String key = d.getOrDefault("name", "") + "|" + d.getOrDefault("section", "");
  581. return rawSeen.add(key);
  582. })
  583. .map(d -> {
  584. String content = (String) d.getOrDefault("content", "");
  585. String drugName = (String) d.getOrDefault("name", "");
  586. String storedSection = (String) d.getOrDefault("section", "");
  587. String sourceVersion = (String) d.getOrDefault("source_version", "");
  588. String sourceVolume = (String) d.getOrDefault("source_volume", "");
  589. String category = (String) d.getOrDefault("category", "");
  590. // 优先用 DB 元数据,回退到内容解析
  591. if (drugName == null || drugName.isEmpty()) {
  592. drugName = extractDrugName(content);
  593. }
  594. String sectionDisplay = PromptService.SECTION_DISPLAY.getOrDefault(storedSection, storedSection);
  595. if (sectionDisplay == null || sectionDisplay.isEmpty()) {
  596. sectionDisplay = realSection(content, storedSection);
  597. }
  598. // 跳过"正文"类无意义栏目
  599. if ("正文".equals(sectionDisplay)) {
  600. return null;
  601. }
  602. // 构建完整来源引用
  603. StringBuilder sourceBuilder = new StringBuilder();
  604. if (!sourceVersion.isEmpty()) {
  605. sourceBuilder.append(sourceVersion);
  606. }
  607. if (!sourceVolume.isEmpty()) {
  608. if (!sourceBuilder.isEmpty()) {
  609. sourceBuilder.append(" ");
  610. }
  611. sourceBuilder.append(sourceVolume);
  612. }
  613. String fullSource = sourceBuilder.toString();
  614. content = content.replaceAll("\\s*来源:.*$", "");
  615. content = content.replaceAll("[\\r\\n]+", " ").trim();
  616. String excerpt = content.length() > 500 ? content.substring(0, 500) + "…" : content;
  617. return Map.<String, Object>of(
  618. "drug_id", d.getOrDefault("drug_id", ""),
  619. "name", drugName,
  620. "section", sectionDisplay,
  621. "category", category != null ? category : "",
  622. "source", fullSource,
  623. "excerpt", excerpt
  624. );
  625. })
  626. .filter(Objects::nonNull)
  627. // 第二层:按显示名去重,避免别名映射(如"功能""主治"→"功能与主治")导致重复
  628. .filter(m -> {
  629. String key = m.get("name") + "|" + m.get("section");
  630. return seen.add(key);
  631. })
  632. // 每种药最多 5 个栏目,总数最多 8 条
  633. .collect(Collectors.groupingBy(m -> (String) m.get("name"), LinkedHashMap::new, Collectors.toList()))
  634. .values().stream()
  635. .flatMap(list -> list.stream().limit(5))
  636. .limit(8)
  637. .collect(Collectors.toList());
  638. }
  639. private String extractDrugName(String content) {
  640. if (content == null) {
  641. return "";
  642. }
  643. int start = content.indexOf("【");
  644. int end = content.indexOf(" - ");
  645. if (start >= 0 && end > start) {
  646. return content.substring(start + 1, end);
  647. }
  648. return content.length() > 20 ? content.substring(0, 20) : content;
  649. }
  650. /** 从 content 文本中提取真实 section(兜底"正文") */
  651. private String realSection(String content, String storedSection) {
  652. if (!"正文".equals(storedSection) || content == null) {
  653. return storedSection;
  654. }
  655. int sep = content.indexOf(" - ");
  656. if (sep < 0) {
  657. return storedSection;
  658. }
  659. int end = content.indexOf("】", sep);
  660. if (end > sep) {
  661. return content.substring(sep + 3, end).trim();
  662. }
  663. return storedSection;
  664. }
  665. /** 从 SecurityContext 获取当前用户标识(JWT subject),未登录则用 IP 隔离 */
  666. private String getCurrentUserKey() {
  667. var auth = SecurityContextHolder.getContext().getAuthentication();
  668. if (auth != null && auth.isAuthenticated() && !"anonymousUser".equals(auth.getPrincipal())) {
  669. return auth.getName();
  670. }
  671. // 未登录用 IP 隔离,避免不同手机会话串了
  672. String ip = request.getRemoteAddr();
  673. String forwarded = request.getHeader("X-Forwarded-For");
  674. if (forwarded != null && !forwarded.isBlank()) {
  675. ip = forwarded.split(",")[0].trim();
  676. }
  677. return "ip:" + ip;
  678. }
  679. private String cleanAnswer(String text) {
  680. if (text == null) {
  681. return "";
  682. }
  683. // 去除多余空白行(保留单个换行),修复 Qwen 常见格式问题
  684. return text
  685. .replace("\r\n", "\n")
  686. .replaceAll("\\n{3,}", "\n\n")
  687. .trim();
  688. }
  689. }