|
|
@@ -209,7 +209,7 @@ public class RetrieverService {
|
|
|
.collect(Collectors.joining(",", "[", "]"));
|
|
|
|
|
|
String drugName = extractDrugName(query);
|
|
|
- return searchWithVersion(vecStr, drugName, "2025年版", topK);
|
|
|
+ return searchWithVersion(vecStr, drugName, query, "2025年版", topK);
|
|
|
}
|
|
|
|
|
|
/** 响应式检索:供 SSE 端点使用,避免 reactor 线程中 block() */
|
|
|
@@ -221,13 +221,14 @@ public class RetrieverService {
|
|
|
.map(String::valueOf)
|
|
|
.collect(Collectors.joining(",", "[", "]"));
|
|
|
String drugName = extractDrugName(query);
|
|
|
- return searchWithVersion(vecStr, drugName, "2025年版", topK);
|
|
|
+ return searchWithVersion(vecStr, drugName, query, "2025年版", topK);
|
|
|
});
|
|
|
}
|
|
|
|
|
|
- /** 带版本过滤的向量检索。version 为 null 时不限制版本。 */
|
|
|
+ /** 带版本过滤的向量检索。version 为 null 时不限制版本。
|
|
|
+ * @param originalQuery 用户原始输入,用于向量搜索不足时的关键词兜底搜索 */
|
|
|
private List<Map<String, Object>> searchWithVersion(
|
|
|
- String vecStr, String drugName, String version, int topK) {
|
|
|
+ String vecStr, String drugName, String originalQuery, String version, int topK) {
|
|
|
|
|
|
String versionFilter = (version != null)
|
|
|
? " AND d.source_version = '" + version + "' "
|
|
|
@@ -277,7 +278,14 @@ public class RetrieverService {
|
|
|
params = new Object[]{vecStr, vecStr, topK};
|
|
|
}
|
|
|
|
|
|
- return jdbc.queryForList(sql, params);
|
|
|
+ List<Map<String, Object>> results = jdbc.queryForList(sql, params);
|
|
|
+
|
|
|
+ // 关键词文本兜底:向量结果不足时用 ILIKE 补充疾病/症状名匹配
|
|
|
+ if (results.size() < KEYWORD_FALLBACK_THRESHOLD) {
|
|
|
+ mergeKeywordFallback(originalQuery, results, topK);
|
|
|
+ }
|
|
|
+
|
|
|
+ return results;
|
|
|
}
|
|
|
|
|
|
// 剂型后缀(长后缀优先,避免"缓释胶囊"被错误截断为"缓释")
|
|
|
@@ -348,6 +356,129 @@ public class RetrieverService {
|
|
|
return "";
|
|
|
}
|
|
|
|
|
|
+ // ============================================================
|
|
|
+ // 关键词文本兜底搜索
|
|
|
+ // ============================================================
|
|
|
+
|
|
|
+ private static final Set<String> STOP_WORDS = Set.of(
|
|
|
+ "是什么", "怎么", "如何", "为什么", "告诉我", "请问",
|
|
|
+ "查询", "搜索", "查一下", "看一下", "的", "吗", "呢", "啊", "什么", "哪些"
|
|
|
+ );
|
|
|
+
|
|
|
+ private static final int KEYWORD_FALLBACK_THRESHOLD = 3;
|
|
|
+ private static final double KEYWORD_FALLBACK_SIMILARITY = 0.25;
|
|
|
+
|
|
|
+ /** 向量结果不足时,用关键词 ILIKE 搜索 drugs.name / drug_chunks.content 作为兜底 */
|
|
|
+ private void mergeKeywordFallback(String query, List<Map<String, Object>> results, int topK) {
|
|
|
+ if (query == null || query.isBlank()) return;
|
|
|
+
|
|
|
+ Set<String> keywords = extractKeywords(query);
|
|
|
+ if (keywords.isEmpty()) return;
|
|
|
+
|
|
|
+ // 已存在的 drug_id,兜底时跳过
|
|
|
+ Set<String> seen = results.stream()
|
|
|
+ .map(r -> (String) r.getOrDefault("drug_id", ""))
|
|
|
+ .filter(s -> !s.isEmpty())
|
|
|
+ .collect(Collectors.toSet());
|
|
|
+
|
|
|
+ try {
|
|
|
+ List<Map<String, Object>> fallback = searchDrugsByName(keywords, seen, topK);
|
|
|
+ if (fallback.isEmpty()) {
|
|
|
+ fallback = searchChunksByContent(keywords, seen, topK);
|
|
|
+ }
|
|
|
+ fallback.forEach(results::add);
|
|
|
+ } catch (Exception ignored) {
|
|
|
+ // 兜底失败不影响主流程
|
|
|
+ }
|
|
|
+ }
|
|
|
+
|
|
|
+ /** 搜 drugs.name ILIKE 关键词 → 取对应 drug_chunks(重要 section 优先) */
|
|
|
+ private List<Map<String, Object>> searchDrugsByName(
|
|
|
+ Set<String> keywords, Set<String> excludeDrugIds, int topK) {
|
|
|
+
|
|
|
+ List<String> kwList = keywords.stream().filter(k -> k.length() >= 2).toList();
|
|
|
+ if (kwList.isEmpty()) return List.of();
|
|
|
+
|
|
|
+ String nameOr = kwList.stream().map(k -> "d.name ILIKE ?").collect(Collectors.joining(" OR "));
|
|
|
+ List<Object> params = new ArrayList<>(kwList.stream().map(k -> "%" + k + "%").toList());
|
|
|
+ params.add(topK * 3);
|
|
|
+
|
|
|
+ String excludeClause = buildExcludeClause("d.drug_id", excludeDrugIds);
|
|
|
+
|
|
|
+ List<Map<String, Object>> drugs = jdbc.queryForList(
|
|
|
+ "SELECT DISTINCT d.drug_id, d.name, d.category, d.source_version, d.source_volume "
|
|
|
+ + "FROM drugs d WHERE d.is_active = TRUE AND (" + nameOr + ")" + excludeClause
|
|
|
+ + " ORDER BY d.name LIMIT ?",
|
|
|
+ params.toArray());
|
|
|
+
|
|
|
+ List<Map<String, Object>> chunks = new ArrayList<>();
|
|
|
+ for (Map<String, Object> drug : drugs) {
|
|
|
+ String did = (String) drug.get("drug_id");
|
|
|
+ List<Map<String, Object>> drugChunks = jdbc.queryForList(
|
|
|
+ """
|
|
|
+ SELECT c.content, c.source, c.drug_id, c.section
|
|
|
+ FROM drug_chunks c WHERE c.drug_id = ? AND c.vec IS NOT NULL
|
|
|
+ ORDER BY CASE c.section
|
|
|
+ WHEN '功能主治' THEN 1 WHEN '适应证' THEN 1 WHEN '主治' THEN 1
|
|
|
+ WHEN '用法与用量' THEN 2 WHEN '用法用量' THEN 2 WHEN '类别' THEN 3
|
|
|
+ ELSE 99 END
|
|
|
+ LIMIT 3""", did);
|
|
|
+ for (Map<String, Object> c : drugChunks) {
|
|
|
+ c.put("name", drug.get("name"));
|
|
|
+ c.put("category", drug.getOrDefault("category", ""));
|
|
|
+ c.put("source_version", drug.getOrDefault("source_version", ""));
|
|
|
+ c.put("source_volume", drug.getOrDefault("source_volume", ""));
|
|
|
+ c.put("similarity", KEYWORD_FALLBACK_SIMILARITY);
|
|
|
+ chunks.add(c);
|
|
|
+ }
|
|
|
+ }
|
|
|
+ return chunks;
|
|
|
+ }
|
|
|
+
|
|
|
+ /** 搜 drug_chunks.content ILIKE 关键词(drugs.name 匹配不到时的二层兜底) */
|
|
|
+ private List<Map<String, Object>> searchChunksByContent(
|
|
|
+ Set<String> keywords, Set<String> excludeDrugIds, int topK) {
|
|
|
+
|
|
|
+ List<String> kwList = keywords.stream().filter(k -> k.length() >= 2).toList();
|
|
|
+ if (kwList.isEmpty()) return List.of();
|
|
|
+
|
|
|
+ String contentOr = kwList.stream().map(k -> "c.content ILIKE ?").collect(Collectors.joining(" OR "));
|
|
|
+ List<Object> params = new ArrayList<>(kwList.stream().map(k -> "%" + k + "%").toList());
|
|
|
+ params.add(topK * 2);
|
|
|
+
|
|
|
+ String excludeClause = buildExcludeClause("c.drug_id", excludeDrugIds);
|
|
|
+
|
|
|
+ return jdbc.queryForList(
|
|
|
+ "SELECT c.content, c.source, c.drug_id, c.section, "
|
|
|
+ + "d.name, d.category, d.source_version, d.source_volume "
|
|
|
+ + "FROM drug_chunks c JOIN drugs d ON d.drug_id = c.drug_id "
|
|
|
+ + "WHERE c.vec IS NOT NULL AND (" + contentOr + ")" + excludeClause
|
|
|
+ + " ORDER BY d.name LIMIT ?",
|
|
|
+ params.toArray());
|
|
|
+ }
|
|
|
+
|
|
|
+ private String buildExcludeClause(String col, Set<String> ids) {
|
|
|
+ if (ids.isEmpty()) return "";
|
|
|
+ return " AND " + col + " NOT IN ('" + String.join("','", ids.stream()
|
|
|
+ .map(s -> s.replace("'", "''")).toList()) + "')";
|
|
|
+ }
|
|
|
+
|
|
|
+ /** 提取中文 2-6 字片段,过滤停用词,长片段优先 */
|
|
|
+ Set<String> extractKeywords(String query) {
|
|
|
+ Set<String> keywords = new LinkedHashSet<>();
|
|
|
+ String cleaned = query;
|
|
|
+ for (String sw : STOP_WORDS) cleaned = cleaned.replace(sw, " ");
|
|
|
+ for (int n = 6; n >= 2; n--) {
|
|
|
+ for (int i = 0; i <= cleaned.length() - n; i++) {
|
|
|
+ String seg = cleaned.substring(i, n + i).trim();
|
|
|
+ if (seg.length() >= 2 && seg.codePoints().allMatch(Character::isIdeographic)) {
|
|
|
+ keywords.add(seg);
|
|
|
+ }
|
|
|
+ }
|
|
|
+ }
|
|
|
+ return keywords;
|
|
|
+ }
|
|
|
+
|
|
|
private boolean anyMatch(String text, String... keywords) {
|
|
|
for (String kw : keywords) {
|
|
|
int idx = text.indexOf(kw);
|