ChatController.java 42 KB

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