| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335 |
- package com.pharmacopoeia.service;
- import org.springframework.dao.DataAccessException;
- import org.springframework.data.redis.core.StringRedisTemplate;
- import org.springframework.jdbc.core.JdbcTemplate;
- import org.springframework.stereotype.Service;
- import java.nio.charset.StandardCharsets;
- import java.security.MessageDigest;
- import java.time.Duration;
- import java.util.*;
- import java.util.stream.Collectors;
- @Service
- public class RetrieverService {
- private final JdbcTemplate jdbc;
- private final LLMService llmService;
- private final StringRedisTemplate redis;
- private final boolean redisAvailable;
- // Embedding 缓存 key 前缀 + TTL(7 天)
- private static final String EMBED_CACHE_PREFIX = "aiyaodian:embed:";
- private static final Duration EMBED_CACHE_TTL = Duration.ofDays(7);
- private static final List<String> NEGATION_PATTERNS = List.of(
- "不是", "没有", "并非", "算不上", "怎么会是", "不可能", "不会",
- "没得", "没", "无", "不属", "不属于", "不是什么", "这不是"
- );
- public RetrieverService(JdbcTemplate jdbc, LLMService llmService,
- StringRedisTemplate redis) {
- this.jdbc = jdbc;
- this.llmService = llmService;
- this.redis = redis;
- // 探测 Redis 是否可用,不可用时降级为每次调用远程 API
- this.redisAvailable = pingRedis();
- }
- /** 启动时探测 Redis,失败则降级为无缓存模式,避免阻塞主流程 */
- private boolean pingRedis() {
- try {
- return "PONG".equals(redis.getConnectionFactory().getConnection().ping());
- } catch (Exception e) {
- return false;
- }
- }
- /** 计算 query 的 SHA-256 作为缓存 key */
- private String cacheKey(String text) {
- try {
- MessageDigest md = MessageDigest.getInstance("SHA-256");
- byte[] hash = md.digest(text.getBytes(StandardCharsets.UTF_8));
- StringBuilder sb = new StringBuilder(2 * hash.length);
- for (byte b : hash) sb.append(String.format("%02x", b));
- return EMBED_CACHE_PREFIX + sb;
- } catch (Exception e) {
- return EMBED_CACHE_PREFIX + text.hashCode();
- }
- }
- /** 将 List<Float> 序列化为逗号分隔字符串,便于存入 Redis */
- private String encodeVec(List<Float> vec) {
- StringBuilder sb = new StringBuilder(vec.size() * 8);
- for (int i = 0; i < vec.size(); i++) {
- if (i > 0) sb.append(',');
- sb.append(vec.get(i));
- }
- return sb.toString();
- }
- /** 反序列化 */
- private List<Float> decodeVec(String s) {
- if (s == null || s.isEmpty()) return List.of();
- String[] parts = s.split(",");
- List<Float> vec = new ArrayList<>(parts.length);
- for (String p : parts) vec.add(Float.parseFloat(p));
- return vec;
- }
- /** 带缓存的 Embedding 查询:相同 query 命中缓存则跳过远程 API 调用 */
- private List<Float> embedWithCache(String query) {
- if (!redisAvailable) return llmService.embed(query);
- String key = cacheKey(query);
- try {
- String cached = redis.opsForValue().get(key);
- if (cached != null) return decodeVec(cached);
- } catch (DataAccessException ignored) {
- // Redis 异常时降级到远程调用
- }
- List<Float> vec = llmService.embed(query);
- try {
- redis.opsForValue().set(key, encodeVec(vec), EMBED_CACHE_TTL);
- } catch (DataAccessException ignored) {
- // 缓存写入失败不影响主流程
- }
- return vec;
- }
- /** 检测 query 中关键词之前是否存在否定词,避免误匹配 */
- private boolean hasNegation(String text, String keyword) {
- int idx = text.indexOf(keyword);
- if (idx < 0) return false;
- String prefix = text.substring(0, idx);
- for (String neg : NEGATION_PATTERNS) {
- if (prefix.endsWith(neg) || prefix.contains(neg)) return true;
- }
- return false;
- }
- public String classifyIntent(String query) {
- String q = query.trim();
- // 1. 用药安全/用法用量(优先级最高,避免"过敏"与症状类冲突)
- if (anyMatch(q, "怎么吃", "吃多少", "怎么用", "怎么服用", "孕妇", "儿童用量",
- "副作用多大", "伤肝", "伤肾", "安全吗", "副作用", "不良反应",
- "禁忌", "过敏", "能不能", "可以吗", "用法", "用量", "剂量",
- "用药指导", "一天几次", "一次多少", "饭前", "饭后", "空腹",
- "能不能一起吃", "相互作用", "过量", "停用", "停药", "忌口",
- "饮酒", "肝功能", "肾功能")) {
- return "usage_guide";
- }
- // 2. 考试辅导(高优先级,关键词明确)
- if (anyMatch(q, "执业药师", "考点", "历年真题", "考试大纲", "高频考点",
- "药物化学", "药剂学", "药理学", "药分", "药物分析")) {
- return "exam_tutor";
- }
- // 3. 法规条款(关键词明确)
- if (anyMatch(q, "凡例", "通则规定", "制剂通则", "一般规定", "通则")) {
- return "regulation";
- }
- // 4. 症状用药建议(安全类关键词已在上方处理,"过敏"不会落到这里)
- if (anyMatch(q, "发烧", "咳嗽", "感冒", "腹泻", "头疼", "头痛", "嗓子疼",
- "吃了什么药", "吃什么药", "该吃", "推荐用药", "推荐下用药",
- "体温", "多少度", "退烧", "止痛", "止泻", "鼻塞", "流鼻涕",
- "头晕", "乏力", "呕吐", "腹痛", "咽痛", "打喷嚏")) {
- return "symptom_advice";
- }
- // 5. 兜底:药品查询
- return "drug_query";
- }
- /** 检索主入口:2025 优先,2025 结果不足时降级到全部版本(含 2020 补充) */
- public List<Map<String, Object>> search(String query, String intent, int topK) {
- List<Float> vec = embedWithCache(query);
- String vecStr = vec.stream()
- .map(String::valueOf)
- .collect(Collectors.joining(",", "[", "]"));
- String drugName = extractDrugName(query);
- // ============================================================
- // 第一轮:限定 2025 年版
- // ============================================================
- List<Map<String, Object>> results = searchWithVersion(vecStr, drugName, "2025年版", topK);
- // ============================================================
- // 如果 2025 结果不足(<3 条 或 最高相似度 <0.4),降级到全部版本
- // ============================================================
- if (results.size() < 3 || maxSimilarity(results) < 0.4) {
- List<Map<String, Object>> fallback = searchWithVersion(vecStr, drugName, null, topK);
- // 合并:2025 结果排在前面,2020 补充排在后面,去重
- results = mergeResults(results, fallback, topK);
- }
- return results;
- }
- /** 带版本过滤的向量检索。version 为 null 时不限制版本。 */
- private List<Map<String, Object>> searchWithVersion(
- String vecStr, String drugName, String version, int topK) {
- String versionFilter = (version != null)
- ? " AND d.source_version = '" + version + "'"
- : "";
- String sql;
- Object[] params;
- if (!drugName.isEmpty()) {
- // 精确匹配
- sql = """
- SELECT c.content, c.source, c.drug_id, c.section,
- d.name, d.category, d.source_version, d.source_volume,
- 1 - (c.vec <=> ?::vector) AS similarity
- FROM drug_chunks c
- JOIN drugs d ON d.drug_id = c.drug_id
- WHERE c.vec IS NOT NULL AND d.name = ?""" + versionFilter + """
- ORDER BY c.vec <=> ?::vector
- LIMIT ?
- """;
- params = new Object[]{vecStr, drugName, vecStr, topK};
- List<Map<String, Object>> results = jdbc.queryForList(sql, params);
- if (!results.isEmpty()) return results;
- // 前缀匹配兜底
- sql = """
- SELECT c.content, c.source, c.drug_id, c.section,
- d.name, d.category, d.source_version, d.source_volume,
- 1 - (c.vec <=> ?::vector) AS similarity
- FROM drug_chunks c
- JOIN drugs d ON d.drug_id = c.drug_id
- WHERE c.vec IS NOT NULL AND d.name LIKE ?""" + versionFilter + """
- ORDER BY c.vec <=> ?::vector
- LIMIT ?
- """;
- params = new Object[]{vecStr, drugName + "%", vecStr, topK};
- } else {
- sql = """
- SELECT c.content, c.source, c.drug_id, c.section,
- d.name, d.category, d.source_version, d.source_volume,
- 1 - (c.vec <=> ?::vector) AS similarity
- FROM drug_chunks c
- JOIN drugs d ON d.drug_id = c.drug_id
- WHERE c.vec IS NOT NULL""" + versionFilter + """
- ORDER BY c.vec <=> ?::vector
- LIMIT ?
- """;
- params = new Object[]{vecStr, vecStr, topK};
- }
- return jdbc.queryForList(sql, params);
- }
- /** 合并两轮结果:2025 在前,去重,不超过 topK */
- private List<Map<String, Object>> mergeResults(
- List<Map<String, Object>> first, List<Map<String, Object>> second, int topK) {
- Set<String> seen = new HashSet<>();
- List<Map<String, Object>> merged = new ArrayList<>();
- for (Map<String, Object> r : first) {
- String key = (String) r.getOrDefault("drug_id", "") + "|" + r.getOrDefault("section", "");
- if (seen.add(key)) merged.add(r);
- }
- for (Map<String, Object> r : second) {
- String key = (String) r.getOrDefault("drug_id", "") + "|" + r.getOrDefault("section", "");
- if (seen.add(key)) merged.add(r);
- }
- return merged.subList(0, Math.min(topK, merged.size()));
- }
- private double maxSimilarity(List<Map<String, Object>> results) {
- return results.stream()
- .mapToDouble(r -> toDouble(r.get("similarity")))
- .max().orElse(0.0);
- }
- private double toDouble(Object o) {
- if (o instanceof Number n) return n.doubleValue();
- return 0.0;
- }
- // 剂型后缀(长后缀优先,避免"缓释胶囊"被错误截断为"缓释")
- private static final List<String> FORMULATION_SUFFIXES = List.of(
- "缓释胶囊", "缓释片", "肠溶胶囊", "肠溶片", "分散片", "咀嚼片",
- "口服混悬液", "口服液", "混悬液", "滴眼液", "注射液",
- "缓释", "肠溶", "胶囊", "颗粒", "糖浆", "软膏", "栓剂",
- "片", "剂", "栓"
- );
- /** 从 query 中提取已知药品名:查 drugs 表,支持剂型后缀剥离和模糊匹配 */
- private String extractDrugName(String query) {
- String cleaned = query.trim();
- // 第一步:去掉尾部常见修饰词
- String[] querySuffixes = {
- "的用法与用量", "的用法用量", "用法与用量", "用法用量", "的用量", "的用法",
- "的副作用", "不良反应", "的禁忌", "禁忌", "的注意事项", "注意事项",
- "是什么", "说明书", "怎么用", "怎么吃", "的用量", "用量", "的剂量", "剂量"
- };
- for (String s : querySuffixes) {
- if (cleaned.endsWith(s)) {
- cleaned = cleaned.substring(0, cleaned.length() - s.length()).trim();
- break;
- }
- }
- // 去掉问句前缀
- cleaned = cleaned.replaceAll("^(什么是|怎么|如何|告诉我|请问|查询|搜索|查一下)", "").trim();
- if (cleaned.length() < 2) return "";
- // 第二步:精确匹配原始 query(含剂型名如"布洛芬缓释胶囊")
- String exact = tryExactMatch(cleaned);
- if (!exact.isEmpty()) return exact;
- // 第三步:逐步剥剂型后缀再试("布洛芬缓释胶囊"→"布洛芬")
- for (String suffix : FORMULATION_SUFFIXES) {
- if (cleaned.endsWith(suffix)) {
- String base = cleaned.substring(0, cleaned.length() - suffix.length()).trim();
- if (base.length() >= 2) {
- String match = tryExactMatch(base);
- if (!match.isEmpty()) return match;
- }
- }
- }
- // 第四步:ILIKE 模糊匹配兜底
- return tryFuzzyMatch(cleaned);
- }
- private String tryExactMatch(String name) {
- try {
- List<String> matches = jdbc.queryForList(
- "SELECT name FROM drugs WHERE name = ? AND is_active = TRUE LIMIT 1",
- String.class, name);
- if (!matches.isEmpty()) return matches.get(0);
- } catch (Exception ignored) {}
- return "";
- }
- private String tryFuzzyMatch(String name) {
- try {
- List<String> matches = jdbc.queryForList(
- "SELECT name FROM drugs WHERE name ILIKE ? AND is_active = TRUE ORDER BY name LIMIT 1",
- String.class, "%" + name + "%");
- if (!matches.isEmpty()) return matches.get(0);
- } catch (Exception ignored) {}
- return "";
- }
- private boolean anyMatch(String text, String... keywords) {
- for (String kw : keywords) {
- int idx = text.indexOf(kw);
- if (idx >= 0 && !hasNegation(text, kw)) {
- return true;
- }
- }
- return false;
- }
- }
|