package com.wok.supportbot.service; import com.fasterxml.jackson.core.type.TypeReference; import com.fasterxml.jackson.databind.ObjectMapper; import com.wok.supportbot.config.EmbeddingModelFactory; import com.wok.supportbot.dao.KnowledgeFaqMapper; import com.wok.supportbot.entity.KnowledgeFaq; import com.wok.supportbot.rag.CategoryFilter; import lombok.AllArgsConstructor; import lombok.Data; import lombok.extern.slf4j.Slf4j; import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.embedding.EmbeddingRequest; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.beans.factory.annotation.Value; import org.springframework.jdbc.core.JdbcTemplate; import org.springframework.stereotype.Service; import java.util.*; import java.util.concurrent.CompletableFuture; /** * FAQ 三级匹配引擎 * 1. 精确匹配:问题文本完全一致 * 2. 关键词匹配:similar_questions 字段包含用户问题中的关键词 * 3. 语义匹配:基于向量余弦距离的相似度匹配 */ @Service @Slf4j public class FaqMatchEngine { @Autowired private KnowledgeFaqMapper faqMapper; @Autowired private JdbcTemplate jdbcTemplate; @Autowired private EmbeddingModelFactory embeddingModelFactory; @Autowired private CategoryFilter categoryFilter; @Autowired private PipelineToggleService pipelineToggleService; /** 向量维度,与 PgVectorStore 保持一致 */ @Value("${knowledge.vector.dimension:1024}") private int vectorDimension; /** 语义匹配相似度阈值 */ @Value("${knowledge.faq.semantic-threshold:0.85}") private double semanticThreshold; private final ObjectMapper objectMapper = new ObjectMapper(); // ==================== 匹配结果内部类 ==================== /** * FAQ 匹配结果 */ @Data @AllArgsConstructor public static class FaqMatchResult { /** 匹配到的 FAQ */ private KnowledgeFaq faq; /** 匹配类型: EXACT / KEYWORD / SEMANTIC */ private String matchType; /** 匹配分数 (0.0 ~ 1.0) */ private double score; } // ==================== 核心匹配方法 ==================== /** * 对用户问题进行三级匹配 * * @param question 用户问题 * @param categoryIds 角色授权分类 ID 列表(null/空表示不限制),未设置分类的 FAQ 一并可见 * @return 匹配结果(可能为空) */ public Optional match(String question, List categoryIds) { if (question == null || question.isBlank()) { return Optional.empty(); } String trimmedQuestion = question.trim(); // 第一级:精确匹配 Optional exactResult = exactMatch(trimmedQuestion, categoryIds); if (exactResult.isPresent()) { log.info("FAQ 精确匹配命中: question={}", trimmedQuestion); return exactResult; } // 第二级:关键词匹配 Optional keywordResult = keywordMatch(trimmedQuestion, categoryIds); if (keywordResult.isPresent()) { log.info("FAQ 关键词匹配命中: question={}", trimmedQuestion); return keywordResult; } // 第三级:语义匹配 // 性能开关:pipeline_faq_semantic_enabled 关闭时跳过,只保留前两级(精确/关键词,均为本地 SQL,无网络调用) if (!pipelineToggleService.isFaqSemanticEnabled()) { log.info("FAQ 语义匹配已被后台关闭,跳过第三级: question={}", trimmedQuestion); return Optional.empty(); } Optional semanticResult = semanticMatch(trimmedQuestion, categoryIds); if (semanticResult.isPresent()) { log.info("FAQ 语义匹配命中: question={}, score={}", trimmedQuestion, semanticResult.get().getScore()); return semanticResult; } log.debug("FAQ 未匹配: question={}", trimmedQuestion); return Optional.empty(); } // ==================== 第一级:精确匹配 ==================== /** * 精确匹配:问题文本完全一致 */ private Optional exactMatch(String question, List categoryIds) { try { List params = new ArrayList<>(); params.add(question); String categorySql = buildCategoryFilter(categoryIds, params, "category_id"); List results = jdbcTemplate.query( "SELECT * FROM knowledge_faq WHERE question = ? AND status = 'ENABLED' AND is_delete = false" + categorySql + " ORDER BY priority DESC LIMIT 1", (rs, rowNum) -> mapRowToFaq(rs), params.toArray() ); if (!results.isEmpty()) { KnowledgeFaq faq = results.get(0); incrementHitCount(faq.getId()); return Optional.of(new FaqMatchResult(faq, "EXACT", 1.0)); } } catch (Exception e) { log.error("FAQ 精确匹配异常", e); } return Optional.empty(); } // ==================== 第二级:关键词匹配 ==================== /** * 关键词匹配:从用户问题中提取关键词,检查 similar_questions 是否包含 */ private Optional keywordMatch(String question, List categoryIds) { try { List keywords = extractKeywords(question); if (keywords.isEmpty()) { return Optional.empty(); } // 构建 ILIKE 条件 StringBuilder sqlBuilder = new StringBuilder( "SELECT * FROM knowledge_faq WHERE status = 'ENABLED' AND is_delete = false AND (" ); List params = new ArrayList<>(); for (int i = 0; i < keywords.size(); i++) { if (i > 0) { sqlBuilder.append(" OR "); } sqlBuilder.append("similar_questions ILIKE ?"); params.add("%" + keywords.get(i) + "%"); } sqlBuilder.append(")"); sqlBuilder.append(buildCategoryFilter(categoryIds, params, "category_id")); sqlBuilder.append(" ORDER BY priority DESC LIMIT 5"); List results = jdbcTemplate.query( sqlBuilder.toString(), (rs, rowNum) -> mapRowToFaq(rs), params.toArray() ); if (!results.isEmpty()) { KnowledgeFaq faq = results.get(0); return Optional.of(new FaqMatchResult(faq, "KEYWORD", 0.9)); } } catch (Exception e) { log.error("FAQ 关键词匹配异常", e); } return Optional.empty(); } /** * 简单分词:按空格、标点切分,过滤短词 */ private List extractKeywords(String question) { String[] tokens = question.split("[\\s,,。?!?!、;;::\"\"''\\(\\)()\\[\\]【】]+"); List keywords = new ArrayList<>(); for (String token : tokens) { String trimmed = token.trim(); // 过滤长度小于2的token(避免单字匹配过于宽泛) if (trimmed.length() >= 2) { keywords.add(trimmed); } } return keywords; } // ==================== 第三级:语义匹配 ==================== /** * 语义匹配:计算问题向量,在 faq_embedding 表中做余弦距离查询 */ private Optional semanticMatch(String question, List categoryIds) { try { // P2 门控:先确认作用域内确有可语义匹配的 FAQ(faq_embedding 有向量记录)。 // 为空直接返回未命中,省掉一次注定失败的语义 embedding HTTP(豆包多模态模型逐条调用较慢)。 List gateParams = new ArrayList<>(); String gateSql = "SELECT 1 FROM faq_embedding fe " + "JOIN knowledge_faq kf ON fe.faq_id = kf.id " + "WHERE kf.status = 'ENABLED' AND kf.is_delete = false" + buildCategoryFilter(categoryIds, gateParams, "kf.category_id") + " LIMIT 1"; if (jdbcTemplate.queryForList(gateSql, gateParams.toArray()).isEmpty()) { log.debug("FAQ 语义候选为空,跳过语义 embedding: question={}", question); return Optional.empty(); } EmbeddingModel embeddingModel = embeddingModelFactory.getEmbeddingModel(); float[] embedding = embeddingModel.call(new EmbeddingRequest(List.of(question), null)) .getResult().getOutput(); // 将向量转为 PGVector 格式字符串 String vectorStr = toPgVectorFormat(embedding); List params = new ArrayList<>(); params.add(vectorStr); String categorySql = buildCategoryFilter(categoryIds, params, "kf.category_id"); // 余弦距离查询:<=> 运算符返回余弦距离,相似度 = 1 - distance List> results = jdbcTemplate.queryForList( "SELECT fe.faq_id, fe.embedding <=> ?::vector AS distance, " + "kf.id, kf.question, kf.answer, kf.similar_questions, kf.category, kf.category_id, " + "kf.status, kf.priority, kf.hit_count, kf.source, kf.create_time, kf.update_time, kf.is_delete " + "FROM faq_embedding fe " + "JOIN knowledge_faq kf ON fe.faq_id = kf.id " + "WHERE kf.status = 'ENABLED' AND kf.is_delete = false" + categorySql + " ORDER BY distance ASC LIMIT 5", params.toArray() ); if (!results.isEmpty()) { Map topResult = results.get(0); double distance = ((Number) topResult.get("distance")).doubleValue(); double similarity = 1.0 - distance; if (similarity >= semanticThreshold) { KnowledgeFaq faq = mapResultToFaq(topResult); incrementHitCount(faq.getId()); return Optional.of(new FaqMatchResult(faq, "SEMANTIC", similarity)); } } } catch (Exception e) { log.error("FAQ 语义匹配异常", e); } return Optional.empty(); } // ==================== 向量化方法 ==================== /** * 计算并保存单条 FAQ 的向量嵌入 * * @param faqId FAQ ID * @param question FAQ 问题文本 */ public void computeAndSaveEmbedding(Long faqId, String question) { try { EmbeddingModel embeddingModel = embeddingModelFactory.getEmbeddingModel(); float[] embedding = embeddingModel.call(new EmbeddingRequest(List.of(question), null)) .getResult().getOutput(); String vectorStr = toPgVectorFormat(embedding); // 获取当前模型名称用于记录 String modelName = "unknown"; try { var config = embeddingModel.getClass().getSimpleName(); modelName = config; } catch (Exception ignored) { } // 先删除旧记录,再插入新记录 jdbcTemplate.update("DELETE FROM faq_embedding WHERE faq_id = ?", faqId); jdbcTemplate.update( "INSERT INTO faq_embedding (faq_id, embedding, model_name) VALUES (?, ?::vector, ?)", faqId, vectorStr, modelName ); log.info("FAQ 向量已保存: faqId={}, dimension={}", faqId, embedding.length); } catch (Exception e) { log.error("FAQ 向量计算/保存失败: faqId={}", faqId, e); } } /** * 批量异步计算 FAQ 向量嵌入 * * @param faqIds FAQ ID 列表 */ public void batchComputeEmbeddings(List faqIds) { CompletableFuture.runAsync(() -> { log.info("开始批量计算 FAQ 向量: count={}", faqIds.size()); int success = 0; int fail = 0; for (Long faqId : faqIds) { try { // 查询 FAQ 问题文本 String question = jdbcTemplate.queryForObject( "SELECT question FROM knowledge_faq WHERE id = ? AND is_delete = false", String.class, faqId ); if (question != null) { computeAndSaveEmbedding(faqId, question); success++; } } catch (Exception e) { log.error("批量计算 FAQ 向量失败: faqId={}", faqId, e); fail++; } } log.info("批量计算 FAQ 向量完成: success={}, fail={}", success, fail); }); } // ==================== 工具方法 ==================== /** * 构建 FAQ 分类过滤 SQL 片段。 * 规则:授权分类(category_id IN ...)与未设置分类(category_id IS NULL 或 0)均可见, * 其余分类被隔离;分类列表为空时不加过滤(检索全部)。 * * @param categoryIds 角色授权分类 ID 列表,可为 null/空 * @param params 参数集合,本方法追加分类占位符参数 * @param column 分类列名(单表为 category_id,联表为 kf.category_id) * @return SQL 片段(分类为空时返回空串) */ private String buildCategoryFilter(List categoryIds, List params, String column) { List ids = categoryFilter.normalize(categoryIds); if (ids.isEmpty()) { return ""; } StringBuilder sb = new StringBuilder(" AND (").append(column).append(" IN ("); for (int i = 0; i < ids.size(); i++) { if (i > 0) { sb.append(", "); } sb.append("?"); params.add(Long.valueOf(ids.get(i))); } sb.append(") OR ").append(column).append(" IS NULL OR ").append(column).append(" = 0)"); return sb.toString(); } /** * 将 float[] 转为 PGVector 格式字符串: [0.1,0.2,0.3] */ private String toPgVectorFormat(float[] embedding) { StringBuilder sb = new StringBuilder("["); for (int i = 0; i < embedding.length; i++) { if (i > 0) { sb.append(","); } sb.append(embedding[i]); } sb.append("]"); return sb.toString(); } /** * 增加 FAQ 命中次数 */ private void incrementHitCount(Long faqId) { try { jdbcTemplate.update("UPDATE knowledge_faq SET hit_count = hit_count + 1 WHERE id = ?", faqId); } catch (Exception e) { log.warn("更新 FAQ 命中次数失败: faqId={}", faqId, e); } } /** * 从 ResultSet 映射为 KnowledgeFaq 实体 */ private KnowledgeFaq mapRowToFaq(java.sql.ResultSet rs) throws java.sql.SQLException { KnowledgeFaq faq = new KnowledgeFaq(); faq.setId(rs.getLong("id")); faq.setQuestion(rs.getString("question")); faq.setAnswer(rs.getString("answer")); faq.setSimilarQuestions(rs.getString("similar_questions")); faq.setCategory(rs.getString("category")); faq.setCategoryId(rs.getObject("category_id") != null ? rs.getLong("category_id") : null); faq.setStatus(rs.getString("status")); faq.setPriority(rs.getInt("priority")); faq.setHitCount(rs.getLong("hit_count")); faq.setSource(rs.getString("source")); faq.setCreateTime(rs.getTimestamp("create_time")); faq.setUpdateTime(rs.getTimestamp("update_time")); faq.setDelete(rs.getBoolean("is_delete")); return faq; } /** * 从 Map 结果映射为 KnowledgeFaq 实体(语义匹配使用) */ private KnowledgeFaq mapResultToFaq(Map result) { KnowledgeFaq faq = new KnowledgeFaq(); faq.setId(((Number) result.get("id")).longValue()); faq.setQuestion((String) result.get("question")); faq.setAnswer((String) result.get("answer")); faq.setSimilarQuestions((String) result.get("similar_questions")); faq.setCategory((String) result.get("category")); faq.setCategoryId(result.get("category_id") != null ? ((Number) result.get("category_id")).longValue() : null); faq.setStatus((String) result.get("status")); faq.setPriority(((Number) result.get("priority")).intValue()); faq.setHitCount(((Number) result.get("hit_count")).longValue()); faq.setSource((String) result.get("source")); faq.setCreateTime(result.get("create_time") instanceof java.sql.Timestamp ts ? new Date(ts.getTime()) : null); faq.setUpdateTime(result.get("update_time") instanceof java.sql.Timestamp ts ? new Date(ts.getTime()) : null); faq.setDelete(Boolean.TRUE.equals(result.get("is_delete"))); return faq; } }