|
@@ -1,8 +1,13 @@
|
|
|
package com.pharmacopoeia.service;
|
|
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.jdbc.core.JdbcTemplate;
|
|
|
import org.springframework.stereotype.Service;
|
|
import org.springframework.stereotype.Service;
|
|
|
|
|
|
|
|
|
|
+import java.nio.charset.StandardCharsets;
|
|
|
|
|
+import java.security.MessageDigest;
|
|
|
|
|
+import java.time.Duration;
|
|
|
import java.util.*;
|
|
import java.util.*;
|
|
|
import java.util.stream.Collectors;
|
|
import java.util.stream.Collectors;
|
|
|
|
|
|
|
@@ -11,15 +16,87 @@ public class RetrieverService {
|
|
|
|
|
|
|
|
private final JdbcTemplate jdbc;
|
|
private final JdbcTemplate jdbc;
|
|
|
private final LLMService llmService;
|
|
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(
|
|
private static final List<String> NEGATION_PATTERNS = List.of(
|
|
|
"不是", "没有", "并非", "算不上", "怎么会是", "不可能", "不会",
|
|
"不是", "没有", "并非", "算不上", "怎么会是", "不可能", "不会",
|
|
|
"没得", "没", "无", "不属", "不属于", "不是什么", "这不是"
|
|
"没得", "没", "无", "不属", "不属于", "不是什么", "这不是"
|
|
|
);
|
|
);
|
|
|
|
|
|
|
|
- public RetrieverService(JdbcTemplate jdbc, LLMService llmService) {
|
|
|
|
|
|
|
+ public RetrieverService(JdbcTemplate jdbc, LLMService llmService,
|
|
|
|
|
+ StringRedisTemplate redis) {
|
|
|
this.jdbc = jdbc;
|
|
this.jdbc = jdbc;
|
|
|
this.llmService = llmService;
|
|
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 中关键词之前是否存在否定词,避免误匹配 */
|
|
/** 检测 query 中关键词之前是否存在否定词,避免误匹配 */
|
|
@@ -69,46 +146,70 @@ public class RetrieverService {
|
|
|
return "drug_query";
|
|
return "drug_query";
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
|
|
+ /** 检索主入口:2025 优先,2025 结果不足时降级到全部版本(含 2020 补充) */
|
|
|
public List<Map<String, Object>> search(String query, String intent, int topK) {
|
|
public List<Map<String, Object>> search(String query, String intent, int topK) {
|
|
|
- List<Float> vec = llmService.embed(query);
|
|
|
|
|
|
|
+ List<Float> vec = embedWithCache(query);
|
|
|
String vecStr = vec.stream()
|
|
String vecStr = vec.stream()
|
|
|
.map(String::valueOf)
|
|
.map(String::valueOf)
|
|
|
.collect(Collectors.joining(",", "[", "]"));
|
|
.collect(Collectors.joining(",", "[", "]"));
|
|
|
|
|
|
|
|
- // 尝试提取药品名,用于精确定位
|
|
|
|
|
String drugName = extractDrugName(query);
|
|
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;
|
|
String sql;
|
|
|
Object[] params;
|
|
Object[] params;
|
|
|
if (!drugName.isEmpty()) {
|
|
if (!drugName.isEmpty()) {
|
|
|
- // 先精确匹配,没结果再前缀匹配(如"布洛芬"→先查"布洛芬",没有再查"布洛芬片/胶囊等")
|
|
|
|
|
|
|
+ // 精确匹配
|
|
|
sql = """
|
|
sql = """
|
|
|
SELECT c.content, c.source, c.drug_id, c.section,
|
|
SELECT c.content, c.source, c.drug_id, c.section,
|
|
|
d.name, d.category, d.source_version, d.source_volume,
|
|
d.name, d.category, d.source_version, d.source_volume,
|
|
|
1 - (c.vec <=> ?::vector) AS similarity
|
|
1 - (c.vec <=> ?::vector) AS similarity
|
|
|
FROM drug_chunks c
|
|
FROM drug_chunks c
|
|
|
JOIN drugs d ON d.drug_id = c.drug_id
|
|
JOIN drugs d ON d.drug_id = c.drug_id
|
|
|
- WHERE c.vec IS NOT NULL AND d.name = ?
|
|
|
|
|
|
|
+ WHERE c.vec IS NOT NULL AND d.name = ?""" + versionFilter + """
|
|
|
ORDER BY c.vec <=> ?::vector
|
|
ORDER BY c.vec <=> ?::vector
|
|
|
LIMIT ?
|
|
LIMIT ?
|
|
|
""";
|
|
""";
|
|
|
params = new Object[]{vecStr, drugName, vecStr, topK};
|
|
params = new Object[]{vecStr, drugName, vecStr, topK};
|
|
|
List<Map<String, Object>> results = jdbc.queryForList(sql, params);
|
|
List<Map<String, Object>> results = jdbc.queryForList(sql, params);
|
|
|
- if (results.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 LIKE ?
|
|
|
|
|
- ORDER BY c.vec <=> ?::vector
|
|
|
|
|
- LIMIT ?
|
|
|
|
|
- """;
|
|
|
|
|
- params = new Object[]{vecStr, drugName + "%", vecStr, topK};
|
|
|
|
|
- }
|
|
|
|
|
- return 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 {
|
|
} else {
|
|
|
sql = """
|
|
sql = """
|
|
|
SELECT c.content, c.source, c.drug_id, c.section,
|
|
SELECT c.content, c.source, c.drug_id, c.section,
|
|
@@ -116,7 +217,7 @@ public class RetrieverService {
|
|
|
1 - (c.vec <=> ?::vector) AS similarity
|
|
1 - (c.vec <=> ?::vector) AS similarity
|
|
|
FROM drug_chunks c
|
|
FROM drug_chunks c
|
|
|
JOIN drugs d ON d.drug_id = c.drug_id
|
|
JOIN drugs d ON d.drug_id = c.drug_id
|
|
|
- WHERE c.vec IS NOT NULL
|
|
|
|
|
|
|
+ WHERE c.vec IS NOT NULL""" + versionFilter + """
|
|
|
ORDER BY c.vec <=> ?::vector
|
|
ORDER BY c.vec <=> ?::vector
|
|
|
LIMIT ?
|
|
LIMIT ?
|
|
|
""";
|
|
""";
|
|
@@ -126,6 +227,34 @@ public class RetrieverService {
|
|
|
return jdbc.queryForList(sql, params);
|
|
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(
|
|
private static final List<String> FORMULATION_SUFFIXES = List.of(
|
|
|
"缓释胶囊", "缓释片", "肠溶胶囊", "肠溶片", "分散片", "咀嚼片",
|
|
"缓释胶囊", "缓释片", "肠溶胶囊", "肠溶片", "分散片", "咀嚼片",
|