本地 RAG 知识库
You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 
 
 
 

428 lines
17 KiB

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<FaqMatchResult> match(String question, List<Long> categoryIds) {
if (question == null || question.isBlank()) {
return Optional.empty();
}
String trimmedQuestion = question.trim();
// 第一级:精确匹配
Optional<FaqMatchResult> exactResult = exactMatch(trimmedQuestion, categoryIds);
if (exactResult.isPresent()) {
log.info("FAQ 精确匹配命中: question={}", trimmedQuestion);
return exactResult;
}
// 第二级:关键词匹配
Optional<FaqMatchResult> 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<FaqMatchResult> 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<FaqMatchResult> exactMatch(String question, List<Long> categoryIds) {
try {
List<Object> params = new ArrayList<>();
params.add(question);
String categorySql = buildCategoryFilter(categoryIds, params, "category_id");
List<KnowledgeFaq> 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<FaqMatchResult> keywordMatch(String question, List<Long> categoryIds) {
try {
List<String> 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<Object> 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<KnowledgeFaq> 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<String> extractKeywords(String question) {
String[] tokens = question.split("[\\s,,。?!?!、;;::\"\"''\\(\\)()\\[\\]【】]+");
List<String> 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<FaqMatchResult> semanticMatch(String question, List<Long> categoryIds) {
try {
// P2 门控:先确认作用域内确有可语义匹配的 FAQ(faq_embedding 有向量记录)。
// 为空直接返回未命中,省掉一次注定失败的语义 embedding HTTP(豆包多模态模型逐条调用较慢)。
List<Object> 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<Object> params = new ArrayList<>();
params.add(vectorStr);
String categorySql = buildCategoryFilter(categoryIds, params, "kf.category_id");
// 余弦距离查询:<=> 运算符返回余弦距离,相似度 = 1 - distance
List<Map<String, Object>> 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<String, Object> 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<Long> 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<Long> categoryIds, List<Object> params, String column) {
List<String> 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<String, Object> 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;
}
}