liuchengsen 1 місяць тому
батько
коміт
3c6bb285a5

+ 136 - 5
backend-java/src/main/java/com/pharmacopoeia/service/RetrieverService.java

@@ -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);