diff --git a/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/config/SearchChannelProperties.java b/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/config/SearchChannelProperties.java index 877ae2f4c..2f724ecfa 100644 --- a/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/config/SearchChannelProperties.java +++ b/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/config/SearchChannelProperties.java @@ -51,6 +51,16 @@ public static class Channels { * 意图定向检索配置 */ private IntentDirected intentDirected = new IntentDirected(); + + /** + * 关键词检索配置 + */ + private Keyword keyword = new Keyword(); + + /** + * 混合检索融合配置 + */ + private Hybrid hybrid = new Hybrid(); } @Data @@ -99,4 +109,47 @@ public static class IntentDirected { */ private int topKMultiplier = 2; } + + @Data + public static class Keyword { + + /** + * 是否启用关键词检索通道 + */ + private boolean enabled = true; + + /** + * TopK 倍数,关键词检索时召回更多候选 + */ + private int topKMultiplier = 3; + + /** + * 融合时的关键词通道权重(仅 WEIGHTED_SUM 模式) + */ + private float boost = 1.0f; + } + + @Data + public static class Hybrid { + + /** + * 是否启用混合融合 + */ + private boolean enabled = true; + + /** + * 融合策略:RRF / WEIGHTED_SUM + */ + private FusionMode fusion = FusionMode.RRF; + + /** + * 向量权重(仅 WEIGHTED_SUM 模式生效) + */ + private float vectorWeight = 0.7f; + } + + public enum FusionMode { + RRF, + WEIGHTED_SUM + } } diff --git a/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/controller/RAGSettingsController.java b/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/controller/RAGSettingsController.java index 97efb5d6e..062dd888a 100644 --- a/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/controller/RAGSettingsController.java +++ b/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/controller/RAGSettingsController.java @@ -24,9 +24,13 @@ import com.nageoffer.ai.ragent.rag.config.RAGConfigProperties; import com.nageoffer.ai.ragent.rag.config.RAGDefaultProperties; import com.nageoffer.ai.ragent.rag.config.RAGRateLimitProperties; +import com.nageoffer.ai.ragent.rag.config.SearchChannelProperties; import com.nageoffer.ai.ragent.rag.controller.vo.SystemSettingsVO; import com.nageoffer.ai.ragent.rag.controller.vo.SystemSettingsVO.AISettings; +import com.nageoffer.ai.ragent.rag.controller.vo.SystemSettingsVO.ChannelConfig; +import com.nageoffer.ai.ragent.rag.controller.vo.SystemSettingsVO.ChannelSettings; import com.nageoffer.ai.ragent.rag.controller.vo.SystemSettingsVO.DefaultSettings; +import com.nageoffer.ai.ragent.rag.controller.vo.SystemSettingsVO.HybridChannelConfig; import com.nageoffer.ai.ragent.rag.controller.vo.SystemSettingsVO.MemorySettings; import lombok.RequiredArgsConstructor; import org.springframework.beans.factory.annotation.Value; @@ -51,6 +55,7 @@ public class RAGSettingsController { private final RAGRateLimitProperties ragRateLimitProperties; private final MemoryProperties memoryProperties; private final AIModelProperties aiModelProperties; + private final SearchChannelProperties searchChannelProperties; @Value("${spring.servlet.multipart.max-file-size:50MB}") private DataSize maxFileSize; @@ -83,6 +88,7 @@ public Result settings() { .build()) .build()) .memory(toMemorySettings(memoryProperties)) + .channels(toChannelSettings(searchChannelProperties)) .build()) .ai(toAISettings(aiModelProperties)) .build(); @@ -160,6 +166,33 @@ private AISettings.ModelGroup toModelGroup(AIModelProperties.ModelGroup group) { .build(); } + private ChannelSettings toChannelSettings(SearchChannelProperties props) { + SearchChannelProperties.Channels channels = props.getChannels(); + return ChannelSettings.builder() + .vectorGlobal(ChannelConfig.builder() + .enabled(channels.getVectorGlobal().isEnabled()) + .confidenceThreshold(channels.getVectorGlobal().getConfidenceThreshold()) + .singleIntentSupplementThreshold(channels.getVectorGlobal().getSingleIntentSupplementThreshold()) + .topKMultiplier(channels.getVectorGlobal().getTopKMultiplier()) + .build()) + .intentDirected(ChannelConfig.builder() + .enabled(channels.getIntentDirected().isEnabled()) + .minIntentScore(channels.getIntentDirected().getMinIntentScore()) + .topKMultiplier(channels.getIntentDirected().getTopKMultiplier()) + .build()) + .keyword(ChannelConfig.builder() + .enabled(channels.getKeyword().isEnabled()) + .topKMultiplier(channels.getKeyword().getTopKMultiplier()) + .boost(channels.getKeyword().getBoost()) + .build()) + .hybrid(HybridChannelConfig.builder() + .enabled(channels.getHybrid().isEnabled()) + .fusion(channels.getHybrid().getFusion().name()) + .vectorWeight(channels.getHybrid().getVectorWeight()) + .build()) + .build(); + } + private String maskApiKey(String apiKey) { if (!StringUtils.hasText(apiKey)) { return null; diff --git a/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/controller/vo/SystemSettingsVO.java b/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/controller/vo/SystemSettingsVO.java index 109b9db8c..af0eb6c2c 100644 --- a/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/controller/vo/SystemSettingsVO.java +++ b/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/controller/vo/SystemSettingsVO.java @@ -159,13 +159,16 @@ public static class RagSettings { private QueryRewriteSettings queryRewrite; private RateLimitSettings rateLimit; private MemorySettings memory; + private ChannelSettings channels; public RagSettings(DefaultSettings defaultConfig, QueryRewriteSettings queryRewrite, - RateLimitSettings rateLimit, MemorySettings memory) { + RateLimitSettings rateLimit, MemorySettings memory, + ChannelSettings channels) { this.defaultConfig = defaultConfig; this.queryRewrite = queryRewrite; this.rateLimit = rateLimit; this.memory = memory; + this.channels = channels; } public static RagSettingsBuilder builder() { @@ -177,6 +180,7 @@ public static class RagSettingsBuilder { private QueryRewriteSettings queryRewrite; private RateLimitSettings rateLimit; private MemorySettings memory; + private ChannelSettings channels; public RagSettingsBuilder defaultConfig(DefaultSettings defaultConfig) { this.defaultConfig = defaultConfig; @@ -198,12 +202,73 @@ public RagSettingsBuilder memory(MemorySettings memory) { return this; } + public RagSettingsBuilder channels(ChannelSettings channels) { + this.channels = channels; + return this; + } + public RagSettings build() { - return new RagSettings(defaultConfig, queryRewrite, rateLimit, memory); + return new RagSettings(defaultConfig, queryRewrite, rateLimit, memory, channels); } } } + @Setter + @Getter + public static class ChannelSettings { + private ChannelConfig vectorGlobal; + private ChannelConfig intentDirected; + private ChannelConfig keyword; + private HybridChannelConfig hybrid; + + public ChannelSettings(ChannelConfig vectorGlobal, ChannelConfig intentDirected, + ChannelConfig keyword, HybridChannelConfig hybrid) { + this.vectorGlobal = vectorGlobal; + this.intentDirected = intentDirected; + this.keyword = keyword; + this.hybrid = hybrid; + } + + public static ChannelSettingsBuilder builder() { + return new ChannelSettingsBuilder(); + } + + public static class ChannelSettingsBuilder { + private ChannelConfig vectorGlobal; + private ChannelConfig intentDirected; + private ChannelConfig keyword; + private HybridChannelConfig hybrid; + + public ChannelSettingsBuilder vectorGlobal(ChannelConfig c) { this.vectorGlobal = c; return this; } + public ChannelSettingsBuilder intentDirected(ChannelConfig c) { this.intentDirected = c; return this; } + public ChannelSettingsBuilder keyword(ChannelConfig c) { this.keyword = c; return this; } + public ChannelSettingsBuilder hybrid(HybridChannelConfig c) { this.hybrid = c; return this; } + + public ChannelSettings build() { + return new ChannelSettings(vectorGlobal, intentDirected, keyword, hybrid); + } + } + } + + @Data + @Builder + public static class ChannelConfig { + private Boolean enabled; + private Double confidenceThreshold; + private Double minIntentScore; + private Double singleIntentSupplementThreshold; + private Integer topKMultiplier; + private Float boost; + } + + @Data + @Builder + public static class HybridChannelConfig { + private Boolean enabled; + private String fusion; + private Float vectorWeight; + } + @Setter @Getter public static class QueryRewriteSettings { diff --git a/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/core/retrieve/channel/KeywordSearchChannel.java b/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/core/retrieve/channel/KeywordSearchChannel.java new file mode 100644 index 000000000..fe69c1764 --- /dev/null +++ b/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/core/retrieve/channel/KeywordSearchChannel.java @@ -0,0 +1,195 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.nageoffer.ai.ragent.rag.core.retrieve.channel; + +import com.baomidou.mybatisplus.core.toolkit.Wrappers; +import com.nageoffer.ai.ragent.framework.convention.RetrievedChunk; +import com.nageoffer.ai.ragent.knowledge.dao.entity.KnowledgeBaseDO; +import com.nageoffer.ai.ragent.knowledge.dao.mapper.KnowledgeBaseMapper; +import com.nageoffer.ai.ragent.rag.config.SearchChannelProperties; +import lombok.extern.slf4j.Slf4j; +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.jdbc.core.JdbcTemplate; +import org.springframework.stereotype.Component; + +import java.util.ArrayList; +import java.util.HashSet; +import java.util.List; +import java.util.Set; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.Executor; + +/** + * 关键词全文检索通道 + * 使用 PostgreSQL tsvector/tsquery + zhparser 中文分词进行关键词匹配, + * 弥补向量检索对专有名词、产品型号等精确关键词的召回不足。 + */ +@Slf4j +@Component +@ConditionalOnProperty(name = "rag.vector.type", havingValue = "pg") +public class KeywordSearchChannel implements SearchChannel { + + private final SearchChannelProperties properties; + private final KnowledgeBaseMapper knowledgeBaseMapper; + private final JdbcTemplate jdbcTemplate; + private final Executor innerRetrievalExecutor; + + public KeywordSearchChannel(SearchChannelProperties properties, + KnowledgeBaseMapper knowledgeBaseMapper, + JdbcTemplate jdbcTemplate, + Executor innerRetrievalExecutor) { + this.properties = properties; + this.knowledgeBaseMapper = knowledgeBaseMapper; + this.jdbcTemplate = jdbcTemplate; + this.innerRetrievalExecutor = innerRetrievalExecutor; + } + + @Override + public String getName() { + return "KeywordSearch"; + } + + @Override + public int getPriority() { + return 20; + } + + @Override + public boolean isEnabled(SearchContext context) { + return properties.getChannels().getKeyword().isEnabled() + && context.getMainQuestion() != null + && !context.getMainQuestion().isBlank(); + } + + @Override + public SearchChannelResult search(SearchContext context) { + long startTime = System.currentTimeMillis(); + + try { + String query = context.getMainQuestion(); + log.info("执行关键词检索,问题:{}", query); + + List collections = getAllKBCollections(); + if (collections.isEmpty()) { + log.warn("未找到任何 KB collection,跳过关键词检索"); + return emptyResult(startTime); + } + + int topK = context.getTopK() * properties.getChannels().getKeyword().getTopKMultiplier(); + List allChunks = retrieveFromAllCollections(query, collections, topK); + + long latency = System.currentTimeMillis() - startTime; + log.info("关键词检索完成,检索到 {} 个 Chunk,耗时 {}ms", allChunks.size(), latency); + + return SearchChannelResult.builder() + .channelType(SearchChannelType.KEYWORD_ES) + .channelName(getName()) + .chunks(allChunks) + .latencyMs(latency) + .build(); + + } catch (Exception e) { + log.error("关键词检索失败", e); + return emptyResult(startTime); + } + } + + @Override + public SearchChannelType getType() { + return SearchChannelType.KEYWORD_ES; + } + + private List getAllKBCollections() { + Set collections = new HashSet<>(); + List kbList = knowledgeBaseMapper.selectList( + Wrappers.lambdaQuery(KnowledgeBaseDO.class) + .select(KnowledgeBaseDO::getCollectionName) + .eq(KnowledgeBaseDO::getDeleted, 0) + ); + for (KnowledgeBaseDO kb : kbList) { + String name = kb.getCollectionName(); + if (name != null && !name.isBlank()) { + collections.add(name); + } + } + return new ArrayList<>(collections); + } + + private List retrieveFromAllCollections(String query, + List collections, + int topK) { + List>> futures = collections.stream() + .map(collection -> CompletableFuture.supplyAsync( + () -> searchInCollection(query, collection, topK), + innerRetrievalExecutor + )) + .toList(); + + List allChunks = new ArrayList<>(); + int success = 0; + int failure = 0; + for (int i = 0; i < futures.size(); i++) { + try { + List chunks = futures.get(i).join(); + allChunks.addAll(chunks); + success++; + } catch (Exception e) { + failure++; + log.error("关键词检索失败 - Collection: {}", collections.get(i), e); + } + } + log.info("关键词检索统计 - 总 collection 数: {}, 成功: {}, 失败: {}, Chunk 总数: {}", + collections.size(), success, failure, allChunks.size()); + return allChunks; + } + + private List searchInCollection(String query, String collectionName, int topK) { + String sanitized = query.replaceAll("[^\\w\\u4e00-\\u9fff\\s]", " "); + // 在中英文/数字交界处插入空格,避免 "Redis介绍" 被当成一个无法解析的 token + sanitized = sanitized.replaceAll( + "(?<=[\\u4e00-\\u9fff])(?=[a-zA-Z0-9])|(?<=[a-zA-Z0-9])(?=[\\u4e00-\\u9fff])", " "); + if (sanitized.isBlank()) { + return List.of(); + } + // 将空格替换为 | 实现 OR 语义:匹配任一词即可命中 + String tsQuery = sanitized.trim().replaceAll("\\s+", " | "); + + // noinspection SqlDialectInspection,SqlNoDataSourceInspection + return jdbcTemplate.query( + "SELECT id, content, ts_rank(tsv, to_tsquery('zhparser', ?)) AS score " + + "FROM t_knowledge_vector " + + "WHERE metadata->>'collection_name' = ? AND tsv @@ to_tsquery('zhparser', ?) " + + "ORDER BY score DESC LIMIT ?", + (rs, rowNum) -> RetrievedChunk.builder() + .id(rs.getString("id")) + .text(rs.getString("content")) + .score(rs.getFloat("score")) + .build(), + tsQuery, collectionName, tsQuery, topK + ); + } + + private SearchChannelResult emptyResult(long startTime) { + return SearchChannelResult.builder() + .channelType(SearchChannelType.KEYWORD_ES) + .channelName(getName()) + .chunks(List.of()) + .latencyMs(System.currentTimeMillis() - startTime) + .build(); + } +} diff --git a/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/core/retrieve/channel/fusion/FusionStrategy.java b/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/core/retrieve/channel/fusion/FusionStrategy.java new file mode 100644 index 000000000..00d5927bc --- /dev/null +++ b/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/core/retrieve/channel/fusion/FusionStrategy.java @@ -0,0 +1,39 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.nageoffer.ai.ragent.rag.core.retrieve.channel.fusion; + +import com.nageoffer.ai.ragent.framework.convention.RetrievedChunk; + +import java.util.List; +import java.util.Map; + +/** + * 多通道检索结果融合策略接口 + */ +public interface FusionStrategy { + + /** + * 融合多个通道的检索结果 + * + * @param rankedChunks 各通道的排序结果列表,key 为通道标识 + * @param weights 各通道的权重(可选,用于加权融合) + * @return 融合后的 Chunk 列表,按融合分数降序排列 + */ + List fuse(Map> rankedChunks, + Map weights); +} diff --git a/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/core/retrieve/channel/fusion/RRFFusionStrategy.java b/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/core/retrieve/channel/fusion/RRFFusionStrategy.java new file mode 100644 index 000000000..64647c7b8 --- /dev/null +++ b/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/core/retrieve/channel/fusion/RRFFusionStrategy.java @@ -0,0 +1,75 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.nageoffer.ai.ragent.rag.core.retrieve.channel.fusion; + +import com.nageoffer.ai.ragent.framework.convention.RetrievedChunk; + +import java.util.ArrayList; +import java.util.Comparator; +import java.util.HashMap; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; + +/** + * RRF (Reciprocal Rank Fusion) 融合策略 + * RRF_score(d) = Σ 1/(k + rank_i(d)) + * 不需要分数归一化,仅依赖各通道内的排名。 + */ +public class RRFFusionStrategy implements FusionStrategy { + + static final int K = 60; + + @Override + public List fuse(Map> rankedChunks, + Map weights) { + // chunkKey → RRF score + Map rrfScores = new LinkedHashMap<>(); + // chunkKey → chunk + Map chunkMap = new HashMap<>(); + + for (Map.Entry> entry : rankedChunks.entrySet()) { + List chunks = entry.getValue(); + for (int rank = 0; rank < chunks.size(); rank++) { + RetrievedChunk chunk = chunks.get(rank); + String key = chunk.getId() != null ? chunk.getId() : String.valueOf(chunk.getText().hashCode()); + + double score = 1.0 / (K + rank + 1); + rrfScores.merge(key, score, Double::sum); + + chunkMap.putIfAbsent(key, chunk); + } + } + + List fused = new ArrayList<>(); + for (Map.Entry entry : rrfScores.entrySet()) { + RetrievedChunk chunk = chunkMap.get(entry.getKey()); + if (chunk != null) { + RetrievedChunk fusedChunk = RetrievedChunk.builder() + .id(chunk.getId()) + .text(chunk.getText()) + .score(entry.getValue().floatValue()) + .build(); + fused.add(fusedChunk); + } + } + + fused.sort(Comparator.comparingDouble(RetrievedChunk::getScore).reversed()); + return fused; + } +} diff --git a/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/core/retrieve/channel/fusion/WeightedSumFusionStrategy.java b/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/core/retrieve/channel/fusion/WeightedSumFusionStrategy.java new file mode 100644 index 000000000..892d21b44 --- /dev/null +++ b/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/core/retrieve/channel/fusion/WeightedSumFusionStrategy.java @@ -0,0 +1,79 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.nageoffer.ai.ragent.rag.core.retrieve.channel.fusion; + +import com.nageoffer.ai.ragent.framework.convention.RetrievedChunk; + +import java.util.ArrayList; +import java.util.Comparator; +import java.util.HashMap; +import java.util.List; +import java.util.Map; + +/** + * 加权求和融合策略 + * weighted_score(d) = w_i * norm(score_i(d)) + * 对各通道的分数做 min-max 归一化后再加权求和。 + */ +public class WeightedSumFusionStrategy implements FusionStrategy { + + @Override + public List fuse(Map> rankedChunks, + Map weights) { + // chunkKey → weighted score + Map weightedScores = new HashMap<>(); + // chunkKey → chunk + Map chunkMap = new HashMap<>(); + + for (Map.Entry> entry : rankedChunks.entrySet()) { + String channel = entry.getKey(); + List chunks = entry.getValue(); + if (chunks.isEmpty()) continue; + + float weight = weights != null && weights.containsKey(channel) ? weights.get(channel) : 1.0f; + + // min-max normalize scores within this channel + double maxScore = chunks.stream().mapToDouble(c -> c.getScore() != null ? c.getScore() : 0).max().orElse(1.0); + double minScore = chunks.stream().mapToDouble(c -> c.getScore() != null ? c.getScore() : 0).min().orElse(0.0); + double range = maxScore - minScore; + if (range == 0) range = 1.0; + + for (RetrievedChunk chunk : chunks) { + String key = chunk.getId() != null ? chunk.getId() : String.valueOf(chunk.getText().hashCode()); + double normScore = ((chunk.getScore() != null ? chunk.getScore() : 0) - minScore) / range; + weightedScores.merge(key, weight * normScore, Double::sum); + chunkMap.putIfAbsent(key, chunk); + } + } + + List fused = new ArrayList<>(); + for (Map.Entry entry : weightedScores.entrySet()) { + RetrievedChunk chunk = chunkMap.get(entry.getKey()); + if (chunk != null) { + fused.add(RetrievedChunk.builder() + .id(chunk.getId()) + .text(chunk.getText()) + .score(entry.getValue().floatValue()) + .build()); + } + } + + fused.sort(Comparator.comparingDouble(RetrievedChunk::getScore).reversed()); + return fused; + } +} diff --git a/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/core/retrieve/postprocessor/HybridFusionPostProcessor.java b/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/core/retrieve/postprocessor/HybridFusionPostProcessor.java new file mode 100644 index 000000000..5fdd72ac4 --- /dev/null +++ b/bootstrap/src/main/java/com/nageoffer/ai/ragent/rag/core/retrieve/postprocessor/HybridFusionPostProcessor.java @@ -0,0 +1,135 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.nageoffer.ai.ragent.rag.core.retrieve.postprocessor; + +import com.nageoffer.ai.ragent.framework.convention.RetrievedChunk; +import com.nageoffer.ai.ragent.rag.config.SearchChannelProperties; +import com.nageoffer.ai.ragent.rag.config.SearchChannelProperties.FusionMode; +import com.nageoffer.ai.ragent.rag.core.retrieve.channel.SearchChannelResult; +import com.nageoffer.ai.ragent.rag.core.retrieve.channel.SearchChannelType; +import com.nageoffer.ai.ragent.rag.core.retrieve.channel.SearchContext; +import com.nageoffer.ai.ragent.rag.core.retrieve.channel.fusion.FusionStrategy; +import com.nageoffer.ai.ragent.rag.core.retrieve.channel.fusion.RRFFusionStrategy; +import com.nageoffer.ai.ragent.rag.core.retrieve.channel.fusion.WeightedSumFusionStrategy; +import lombok.extern.slf4j.Slf4j; +import org.springframework.stereotype.Component; + +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; + +/** + * 混合检索融合后置处理器 + *

+ * 将向量检索通道(VECTOR_GLOBAL / INTENT_DIRECTED)和关键词检索通道(KEYWORD_ES)的结果 + * 按照配置的融合策略(RRF 或加权求和)合并重排。 + *

+ * 执行顺序位于去重之后、Rerank 之前,确保融合后的结果能进一步由 Rerank 精排。 + */ +@Slf4j +@Component +public class HybridFusionPostProcessor implements SearchResultPostProcessor { + + private final SearchChannelProperties properties; + + public HybridFusionPostProcessor(SearchChannelProperties properties) { + this.properties = properties; + } + + @Override + public String getName() { + return "HybridFusion"; + } + + @Override + public int getOrder() { + return 5; + } + + @Override + public boolean isEnabled(SearchContext context) { + // 只在混合融合启用 且 存在关键词检索结果时才执行 + if (!properties.getChannels().getHybrid().isEnabled()) { + return false; + } + // 检查 results 中是否同时包含向量和关键词两类结果,在 process() 中判断 + return true; + } + + @Override + public List process(List chunks, + List results, + SearchContext context) { + // 按通道类型分组 + Map> vectorGroups = new LinkedHashMap<>(); + SearchChannelResult keywordResult = null; + + for (SearchChannelResult result : results) { + if (result.getChunks().isEmpty()) continue; + if (result.getChannelType() == SearchChannelType.KEYWORD_ES) { + keywordResult = result; + } else if (isVectorChannel(result.getChannelType())) { + vectorGroups.put(result.getChannelName(), result.getChunks()); + } + } + + // 没有关键词结果或没有向量结果时不需要融合,直接返回原始 chunks + if (keywordResult == null || vectorGroups.isEmpty()) { + log.debug("混合融合跳过:向量通道={},关键词通道={}", + !vectorGroups.isEmpty(), + keywordResult != null); + return chunks; + } + + // 合并所有向量通道的结果为一路 + Map> inputs = new LinkedHashMap<>(); + inputs.put("vector", chunks.stream() + .filter(c -> vectorGroups.values().stream().anyMatch(v -> v.contains(c))) + .toList()); + inputs.put("keyword", keywordResult.getChunks()); + + // 选择融合策略 + FusionStrategy strategy = createStrategy(); + + // 权重配置 + Map weights = Map.of( + "vector", properties.getChannels().getHybrid().getVectorWeight(), + "keyword", 1.0f - properties.getChannels().getHybrid().getVectorWeight() + ); + + List fused = strategy.fuse(inputs, weights); + log.info("混合融合完成:向量 {} 个 + 关键词 {} 个 → 融合后 {} 个", + inputs.get("vector").size(), + inputs.get("keyword").size(), + fused.size()); + + return fused; + } + + private boolean isVectorChannel(SearchChannelType type) { + return type == SearchChannelType.VECTOR_GLOBAL || type == SearchChannelType.INTENT_DIRECTED; + } + + private FusionStrategy createStrategy() { + FusionMode mode = properties.getChannels().getHybrid().getFusion(); + if (mode == FusionMode.WEIGHTED_SUM) { + return new WeightedSumFusionStrategy(); + } + return new RRFFusionStrategy(); + } +} diff --git a/frontend/src/pages/admin/settings/SystemSettingsPage.tsx b/frontend/src/pages/admin/settings/SystemSettingsPage.tsx index 9db4a9dba..0112ea3f9 100644 --- a/frontend/src/pages/admin/settings/SystemSettingsPage.tsx +++ b/frontend/src/pages/admin/settings/SystemSettingsPage.tsx @@ -258,6 +258,52 @@ export function SystemSettingsPage() { + + + 检索通道配置 + 多路检索通道的启用状态与参数 + + + {rag.channels && ( +

+
+

向量全局检索

+
+ } /> + + + +
+
+
+

意图定向检索

+
+ } /> + + +
+
+
+

关键词检索

+
+ } /> + + +
+
+
+

混合融合

+
+ } /> + + +
+
+
+ )} + + + Rerank 模型配置 diff --git a/frontend/src/services/settingsService.ts b/frontend/src/services/settingsService.ts index 917001244..95a5ccfbf 100644 --- a/frontend/src/services/settingsService.ts +++ b/frontend/src/services/settingsService.ts @@ -30,6 +30,37 @@ export interface SystemSettings { summaryMaxChars: number; titleMaxLength: number; }; + channels: { + vectorGlobal: { + enabled: boolean; + confidenceThreshold: number; + singleIntentSupplementThreshold: number; + topKMultiplier: number; + minIntentScore?: number; + boost?: number; + }; + intentDirected: { + enabled: boolean; + minIntentScore: number; + topKMultiplier: number; + confidenceThreshold?: number; + singleIntentSupplementThreshold?: number; + boost?: number; + }; + keyword: { + enabled: boolean; + topKMultiplier: number; + boost: number; + confidenceThreshold?: number; + minIntentScore?: number; + singleIntentSupplementThreshold?: number; + }; + hybrid: { + enabled: boolean; + fusion: string; + vectorWeight: number; + }; + }; }; ai: { providers: Record< diff --git a/resources/database/schema_pg.sql b/resources/database/schema_pg.sql index 23a4416e9..6965bd523 100644 --- a/resources/database/schema_pg.sql +++ b/resources/database/schema_pg.sql @@ -4,6 +4,24 @@ -- Enable pgvector extension CREATE EXTENSION IF NOT EXISTS vector; +-- Enable zhparser extension (requires zhparser compiled in PG) +CREATE EXTENSION IF NOT EXISTS zhparser; +DO $$ +BEGIN + IF NOT EXISTS (SELECT 1 FROM pg_ts_config WHERE cfgname = 'zhparser') THEN + CREATE TEXT SEARCH CONFIGURATION zhparser (PARSER = zhparser); + END IF; +END $$; + +-- zhparser token 类型映射(n/v/a/i/e/l → simple),否则 token 被丢弃 +ALTER TEXT SEARCH CONFIGURATION zhparser DROP MAPPING IF EXISTS FOR n; +ALTER TEXT SEARCH CONFIGURATION zhparser DROP MAPPING IF EXISTS FOR v; +ALTER TEXT SEARCH CONFIGURATION zhparser DROP MAPPING IF EXISTS FOR a; +ALTER TEXT SEARCH CONFIGURATION zhparser DROP MAPPING IF EXISTS FOR i; +ALTER TEXT SEARCH CONFIGURATION zhparser DROP MAPPING IF EXISTS FOR e; +ALTER TEXT SEARCH CONFIGURATION zhparser DROP MAPPING IF EXISTS FOR l; +ALTER TEXT SEARCH CONFIGURATION zhparser ADD MAPPING FOR n,v,a,i,e,l WITH simple; + -- ============================================ -- User & Conversation Tables -- ============================================ @@ -423,16 +441,30 @@ CREATE TABLE t_knowledge_vector ( id VARCHAR(20) PRIMARY KEY, content TEXT, metadata JSONB, - embedding vector(1536) + embedding vector(1536), + tsv tsvector ); CREATE INDEX idx_kv_metadata ON t_knowledge_vector USING gin(metadata); CREATE INDEX idx_kv_embedding ON t_knowledge_vector USING hnsw (embedding vector_cosine_ops); +CREATE INDEX idx_kv_tsv ON t_knowledge_vector USING GIN(tsv); COMMENT ON TABLE t_knowledge_vector IS '知识库向量存储表'; COMMENT ON COLUMN t_knowledge_vector.id IS '分块ID'; COMMENT ON COLUMN t_knowledge_vector.content IS '分块文本内容'; COMMENT ON COLUMN t_knowledge_vector.metadata IS '元数据'; COMMENT ON COLUMN t_knowledge_vector.embedding IS '向量'; +COMMENT ON COLUMN t_knowledge_vector.tsv IS 'tsvector 全文检索列'; + +-- 触发器:自动维护 tsv +CREATE OR REPLACE FUNCTION kv_tsv_trigger() RETURNS trigger AS $$ +BEGIN + NEW.tsv := to_tsvector('zhparser', COALESCE(NEW.content, '')); + RETURN NEW; +END; +$$ LANGUAGE plpgsql; + +CREATE TRIGGER trg_kv_tsv BEFORE INSERT OR UPDATE OF content ON t_knowledge_vector + FOR EACH ROW EXECUTE FUNCTION kv_tsv_trigger(); -- ============================================ -- Column Comments diff --git a/resources/database/upgrade_v1.2_to_v1.3.sql b/resources/database/upgrade_v1.2_to_v1.3.sql new file mode 100644 index 000000000..05a041c18 --- /dev/null +++ b/resources/database/upgrade_v1.2_to_v1.3.sql @@ -0,0 +1,41 @@ +-- ragent v1.2 -> v1.3 升级脚本 +-- t_knowledge_vector 表:新增 tsvector 列及全文检索索引,支持混合检索通道 + +-- 1. 新增 tsvector 列 +ALTER TABLE t_knowledge_vector ADD COLUMN IF NOT EXISTS tsv tsvector; + +-- 2. GIN 索引加速全文检索 +CREATE INDEX IF NOT EXISTS idx_kv_tsv ON t_knowledge_vector USING GIN(tsv); + +-- 3. 启用 zhparser 中文分词扩展 +CREATE EXTENSION IF NOT EXISTS zhparser; +DO $$ +BEGIN + IF NOT EXISTS (SELECT 1 FROM pg_ts_config WHERE cfgname = 'zhparser') THEN + CREATE TEXT SEARCH CONFIGURATION zhparser (PARSER = zhparser); + END IF; +END $$; + +-- 3a. token 类型映射:zhparser 产出的 n/v/a/i/e/l 必须映射到 simple 词典,否则 PG 丢弃这些 token +ALTER TEXT SEARCH CONFIGURATION zhparser DROP MAPPING IF EXISTS FOR n; +ALTER TEXT SEARCH CONFIGURATION zhparser DROP MAPPING IF EXISTS FOR v; +ALTER TEXT SEARCH CONFIGURATION zhparser DROP MAPPING IF EXISTS FOR a; +ALTER TEXT SEARCH CONFIGURATION zhparser DROP MAPPING IF EXISTS FOR i; +ALTER TEXT SEARCH CONFIGURATION zhparser DROP MAPPING IF EXISTS FOR e; +ALTER TEXT SEARCH CONFIGURATION zhparser DROP MAPPING IF EXISTS FOR l; +ALTER TEXT SEARCH CONFIGURATION zhparser ADD MAPPING FOR n,v,a,i,e,l WITH simple; + +-- 4. 触发器:自动维护 tsv 列(zhparser 分词) +CREATE OR REPLACE FUNCTION kv_tsv_trigger() RETURNS trigger AS $$ +BEGIN + NEW.tsv := to_tsvector('zhparser', COALESCE(NEW.content, '')); + RETURN NEW; +END; +$$ LANGUAGE plpgsql; + +DROP TRIGGER IF EXISTS trg_kv_tsv ON t_knowledge_vector; +CREATE TRIGGER trg_kv_tsv BEFORE INSERT OR UPDATE OF content ON t_knowledge_vector + FOR EACH ROW EXECUTE FUNCTION kv_tsv_trigger(); + +-- 5. 回填已有数据(zhparser 分词) +UPDATE t_knowledge_vector SET tsv = to_tsvector('zhparser', COALESCE(content, ''));