ChatController.java 39 KB

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