diff --git a/backend/package/yuxi/knowledge/implementations/milvus.py b/backend/package/yuxi/knowledge/implementations/milvus.py index b5f1c283..56dbd67b 100644 --- a/backend/package/yuxi/knowledge/implementations/milvus.py +++ b/backend/package/yuxi/knowledge/implementations/milvus.py @@ -1,12 +1,23 @@ import asyncio import os -import re import time import traceback from functools import partial from typing import Any -from pymilvus import Collection, CollectionSchema, DataType, FieldSchema, connections, db, utility +from pymilvus import ( + AnnSearchRequest, + Collection, + CollectionSchema, + DataType, + FieldSchema, + Function, + FunctionType, + WeightedRanker, + connections, + db, + utility, +) from yuxi import config from yuxi.knowledge.base import FileStatus, KnowledgeBase @@ -19,6 +30,8 @@ from yuxi.utils import hashstr, logger from yuxi.utils.datetime_utils import utc_isoformat MILVUS_AVAILABLE = True +CONTENT_SPARSE_FIELD = "content_sparse" +CONTENT_ANALYZER_PARAMS = {"type": "chinese"} class MilvusKB(KnowledgeBase): @@ -118,6 +131,11 @@ class MilvusKB(KnowledgeBase): utility.drop_collection(collection_name, using=self.connection_alias) return self._create_new_collection(collection_name, embed_info, db_id) + if not self._collection_supports_bm25(collection): + logger.warning(f"Collection {collection_name} schema does not support BM25, recreating") + utility.drop_collection(collection_name, using=self.connection_alias) + return self._create_new_collection(collection_name, embed_info, db_id) + logger.info(f"Retrieved existing collection: {collection_name}") return collection else: @@ -140,16 +158,31 @@ class MilvusKB(KnowledgeBase): # 定义集合Schema fields = [ FieldSchema(name="id", dtype=DataType.VARCHAR, max_length=100, is_primary=True), - FieldSchema(name="content", dtype=DataType.VARCHAR, max_length=65535), + FieldSchema( + name="content", + dtype=DataType.VARCHAR, + max_length=65535, + enable_analyzer=True, + analyzer_params=CONTENT_ANALYZER_PARAMS, + ), FieldSchema(name="source", dtype=DataType.VARCHAR, max_length=500), FieldSchema(name="chunk_id", dtype=DataType.VARCHAR, max_length=100), FieldSchema(name="file_id", dtype=DataType.VARCHAR, max_length=100), FieldSchema(name="chunk_index", dtype=DataType.INT64), FieldSchema(name="embedding", dtype=DataType.FLOAT_VECTOR, dim=embedding_dim), + FieldSchema(name=CONTENT_SPARSE_FIELD, dtype=DataType.SPARSE_FLOAT_VECTOR), ] + bm25_function = Function( + name="content_bm25", + input_field_names=["content"], + output_field_names=[CONTENT_SPARSE_FIELD], + function_type=FunctionType.BM25, + ) schema = CollectionSchema( - fields=fields, description=f"Knowledge base collection for {db_id} using {model_name}" + fields=fields, + description=f"Knowledge base collection for {db_id} using {model_name}", + functions=[bm25_function], ) # 创建集合 @@ -158,11 +191,38 @@ class MilvusKB(KnowledgeBase): # 创建索引 index_params = {"metric_type": "COSINE", "index_type": "IVF_FLAT", "params": {"nlist": 1024}} collection.create_index("embedding", index_params) + sparse_index_params = { + "metric_type": "BM25", + "index_type": "SPARSE_INVERTED_INDEX", + "params": {"inverted_index_algo": "DAAT_MAXSCORE"}, + } + collection.create_index(CONTENT_SPARSE_FIELD, sparse_index_params) logger.info(f"Created new Milvus collection: {collection_name} '{model_name=}', {embedding_dim=}") return collection + def _collection_supports_bm25(self, collection: Collection) -> bool: + """检查集合是否具备 Milvus 内置 BM25 所需的 schema。""" + fields = {field.name: field for field in collection.schema.fields} + content_field = fields.get("content") + sparse_field = fields.get(CONTENT_SPARSE_FIELD) + if not content_field or content_field.dtype != DataType.VARCHAR: + return False + if content_field.params.get("enable_analyzer") is not True: + return False + if not sparse_field or sparse_field.dtype != DataType.SPARSE_FLOAT_VECTOR: + return False + + for function in collection.schema.functions: + if ( + function.type == FunctionType.BM25 + and function.input_field_names == ["content"] + and function.output_field_names == [CONTENT_SPARSE_FIELD] + ): + return True + return False + async def _initialize_kb_instance(self, instance: Any) -> None: """初始化 Milvus 集合(加载到内存)""" try: @@ -468,6 +528,28 @@ class MilvusKB(KnowledgeBase): return processed_items_info + def _build_chunk_from_hit( + self, + hit: Any, + score: float, + include_distances: bool, + score_field: str | None = None, + ) -> dict: + """将 Milvus Hit 转成知识库统一返回的 Chunk 结构。""" + entity = hit.entity + metadata = { + "source": entity.get("source", "未知来源"), + "chunk_id": entity.get("chunk_id"), + "file_id": entity.get("file_id"), + "chunk_index": entity.get("chunk_index"), + } + chunk = {"content": entity.get("content", ""), "metadata": metadata, "score": float(score or 0.0)} + if score_field: + chunk[score_field] = float(score or 0.0) + if include_distances: + chunk["distance"] = hit.distance + return chunk + async def aquery(self, query_text: str, db_id: str, agent_call: bool = False, **kwargs) -> list[dict]: """异步查询知识库""" collection = await self._get_milvus_collection(db_id) @@ -508,8 +590,9 @@ class MilvusKB(KnowledgeBase): file_expr = f'source like "{safe_file_name}"' logger.debug(f"Using filter expression: {file_expr}") - vector_results: list[dict] = [] - if search_mode in {"vector", "hybrid"}: + output_fields = ["content", "source", "chunk_id", "file_id", "chunk_index"] + retrieved_chunks: list[dict] = [] + if search_mode == "vector": embed_info = self.databases_meta[db_id].get("embed_info", {}) embedding_function = self._get_embedding_function(embed_info) query_embedding = embedding_function([query_text]) @@ -522,7 +605,7 @@ class MilvusKB(KnowledgeBase): param=search_params, limit=recall_top_k, expr=file_expr, - output_fields=["content", "source", "chunk_id", "file_id", "chunk_index"], + output_fields=output_fields, ) if results and len(results) > 0 and len(results[0]) > 0: @@ -531,104 +614,80 @@ class MilvusKB(KnowledgeBase): if similarity < similarity_threshold: continue - entity = hit.entity - metadata = { - "source": entity.get("source", "未知来源"), - "chunk_id": entity.get("chunk_id"), - "file_id": entity.get("file_id"), - "chunk_index": entity.get("chunk_index"), - } - - chunk = {"content": entity.get("content", ""), "metadata": metadata, "score": similarity} - if include_distances: - chunk["distance"] = hit.distance - vector_results.append(chunk) + retrieved_chunks.append(self._build_chunk_from_hit(hit, similarity, include_distances)) logger.debug( - f"Milvus vector query response: {len(vector_results)} chunks found (after similarity filtering)" + f"Milvus vector query response: {len(retrieved_chunks)} chunks found (after similarity filtering)" ) - keyword_results: list[dict] = [] - if search_mode in {"keyword", "hybrid"}: - keyword_top_k = int(merged_kwargs.get("keyword_top_k", final_top_k)) - keyword_top_k = max(keyword_top_k, 1) - raw_keywords = re.split(r"[\s,,;;]+", str(query_text)) - keywords = [kw.strip() for kw in raw_keywords if kw and kw.strip()] - - if keywords: - keyword_clauses = [] - for kw in keywords: - safe_kw = kw.replace('"', '\\"') - keyword_clauses.append(f'content like "%{safe_kw}%"') - - keyword_expr = " or ".join(keyword_clauses) - if file_expr: - keyword_expr = f"({keyword_expr}) and ({file_expr})" - - results = collection.query( - expr=keyword_expr, - output_fields=["content", "source", "chunk_id", "file_id", "chunk_index"], - limit=keyword_top_k, - ) - - keyword_scores = [] - for result in results or []: - content = result.get("content", "") - text_lower = content.lower() - match_count = sum(text_lower.count(kw.lower()) for kw in keywords if kw) - if match_count <= 0: - continue - - metadata = { - "source": result.get("source", "未知来源"), - "chunk_id": result.get("chunk_id"), - "file_id": result.get("file_id"), - "chunk_index": result.get("chunk_index"), - } - keyword_scores.append((result, metadata, match_count)) - - if keyword_scores: - max_count = max(item[2] for item in keyword_scores) - for result, metadata, match_count in keyword_scores: - score = match_count / max_count if max_count > 0 else 0.0 - keyword_results.append( - { - "content": result.get("content", ""), - "metadata": metadata, - "score": score, - "keyword_score": score, - } - ) - keyword_results.sort(key=lambda item: item.get("score", 0.0), reverse=True) - - logger.debug(f"Milvus keyword query response: {len(keyword_results)} chunks found") - - if search_mode == "vector": - retrieved_chunks = vector_results elif search_mode == "keyword": - retrieved_chunks = keyword_results + bm25_top_k = int(merged_kwargs.get("bm25_top_k", recall_top_k)) + bm25_top_k = max(bm25_top_k, 1) + bm25_drop_ratio_search = float(merged_kwargs.get("bm25_drop_ratio_search", 0.0)) + bm25_search_params = { + "metric_type": "BM25", + "params": {"drop_ratio_search": bm25_drop_ratio_search}, + } + + results = collection.search( + data=[query_text], + anns_field=CONTENT_SPARSE_FIELD, + param=bm25_search_params, + limit=bm25_top_k, + expr=file_expr, + output_fields=output_fields, + ) + + if results and len(results) > 0 and len(results[0]) > 0: + for hit in results[0]: + retrieved_chunks.append( + self._build_chunk_from_hit(hit, hit.distance, include_distances, score_field="bm25_score") + ) + + logger.debug(f"Milvus BM25 query response: {len(retrieved_chunks)} chunks found") else: - merged: dict[str, dict] = {} - for item in vector_results: - chunk_id = item.get("metadata", {}).get("chunk_id") - if chunk_id: - merged[chunk_id] = item - else: - merged[id(item)] = item + embed_info = self.databases_meta[db_id].get("embed_info", {}) + embedding_function = self._get_embedding_function(embed_info) + query_embedding = embedding_function([query_text]) + bm25_top_k = int(merged_kwargs.get("bm25_top_k", recall_top_k)) + bm25_top_k = max(bm25_top_k, 1) + bm25_drop_ratio_search = float(merged_kwargs.get("bm25_drop_ratio_search", 0.0)) + vector_weight = float(merged_kwargs.get("vector_weight", 0.7)) + bm25_weight = float(merged_kwargs.get("bm25_weight", 0.3)) - for item in keyword_results: - chunk_id = item.get("metadata", {}).get("chunk_id") - if chunk_id in merged: - existing = merged[chunk_id] - keyword_score = item.get("keyword_score", item.get("score", 0.0)) - existing_score = existing.get("score", 0.0) - existing["score"] = max(existing_score, keyword_score) - existing["keyword_score"] = keyword_score - else: - merged[chunk_id or id(item)] = item + vector_request = AnnSearchRequest( + data=query_embedding, + anns_field="embedding", + param={"metric_type": metric_type, "params": {"nprobe": 10}}, + limit=recall_top_k, + expr=file_expr, + ) + bm25_request = AnnSearchRequest( + data=[query_text], + anns_field=CONTENT_SPARSE_FIELD, + param={ + "metric_type": "BM25", + "params": {"drop_ratio_search": bm25_drop_ratio_search}, + }, + limit=bm25_top_k, + expr=file_expr, + ) + results = collection.hybrid_search( + reqs=[vector_request, bm25_request], + rerank=WeightedRanker(vector_weight, bm25_weight), + limit=recall_top_k, + output_fields=output_fields, + ) + if results and len(results) > 0 and len(results[0]) > 0: + for hit in results[0]: + score = float(hit.distance or 0.0) + if score < similarity_threshold: + continue + retrieved_chunks.append( + self._build_chunk_from_hit(hit, score, include_distances, score_field="hybrid_score") + ) - retrieved_chunks = list(merged.values()) - retrieved_chunks.sort(key=lambda item: item.get("score", 0.0), reverse=True) + logger.debug(f"Milvus hybrid query response: {len(retrieved_chunks)} chunks found") if not retrieved_chunks: return [] @@ -804,8 +863,8 @@ class MilvusKB(KnowledgeBase): "default": "vector", "options": [ {"value": "vector", "label": "向量检索", "description": "仅使用向量相似度检索"}, - {"value": "keyword", "label": "关键词检索", "description": "仅使用关键词匹配检索"}, - {"value": "hybrid", "label": "混合检索", "description": "向量检索与关键词检索融合"}, + {"value": "keyword", "label": "BM25 全文检索", "description": "仅使用 Milvus BM25 检索"}, + {"value": "hybrid", "label": "混合检索", "description": "Milvus 向量检索与 BM25 融合检索"}, ], "description": "选择检索模式", }, @@ -829,13 +888,43 @@ class MilvusKB(KnowledgeBase): "description": "过滤相似度低于此值的结果", }, { - "key": "keyword_top_k", - "label": "使用关键词召回的最大 Chunk 数", + "key": "bm25_top_k", + "label": "BM25 召回数量", "type": "number", "default": 50, "min": 1, "max": 200, - "description": "关键词/混合检索时的候选数量", + "description": "BM25 全文检索和混合检索中的 BM25 候选数量", + }, + { + "key": "vector_weight", + "label": "向量检索权重", + "type": "number", + "default": 0.7, + "min": 0.0, + "max": 1.0, + "step": 0.1, + "description": "混合检索中向量召回结果的融合权重", + }, + { + "key": "bm25_weight", + "label": "BM25 权重", + "type": "number", + "default": 0.3, + "min": 0.0, + "max": 1.0, + "step": 0.1, + "description": "混合检索中 BM25 召回结果的融合权重", + }, + { + "key": "bm25_drop_ratio_search", + "label": "BM25 稀疏项丢弃比例", + "type": "number", + "default": 0.0, + "min": 0.0, + "max": 1.0, + "step": 0.1, + "description": "BM25 检索时丢弃低分稀疏项的比例,数值越大检索越快但可能降低召回", }, { "key": "include_distances", @@ -881,7 +970,7 @@ class MilvusKB(KnowledgeBase): "default": 50, "min": 10, "max": 200, - "description": "向量检索时保留的候选数量(启用重排序时有效)", + "description": "向量检索或混合检索保留的候选数量(启用重排序时有效)", }, ] diff --git a/backend/test/unit/plugins/test_milvus_kb.py b/backend/test/unit/plugins/test_milvus_kb.py new file mode 100644 index 00000000..d52ccc59 --- /dev/null +++ b/backend/test/unit/plugins/test_milvus_kb.py @@ -0,0 +1,160 @@ +from pymilvus import CollectionSchema, DataType, FieldSchema, Function, FunctionType + +from yuxi.knowledge.implementations.milvus import CONTENT_ANALYZER_PARAMS, CONTENT_SPARSE_FIELD, MilvusKB + + +class FakeHit: + def __init__(self, content: str, distance: float): + self.distance = distance + self.entity = { + "content": content, + "source": "demo.md", + "chunk_id": "chunk-1", + "file_id": "file-1", + "chunk_index": 0, + } + + +class FakeCollection: + def __init__(self, distance: float = 0.8): + self.search_calls = [] + self.hybrid_calls = [] + self.distance = distance + + def search(self, **kwargs): + self.search_calls.append(kwargs) + return [[FakeHit("BM25 result", self.distance)]] + + def hybrid_search(self, **kwargs): + self.hybrid_calls.append(kwargs) + return [[FakeHit("Hybrid result", self.distance)]] + + +def make_kb(collection: FakeCollection) -> MilvusKB: + kb = MilvusKB.__new__(MilvusKB) + kb.databases_meta = {"db": {"embed_info": {}}} + kb._get_query_params = lambda db_id: {} + kb._get_embedding_function = lambda embed_info: lambda texts: [[0.1, 0.2] for _ in texts] + + async def get_collection(db_id: str): + return collection + + kb._get_milvus_collection = get_collection + return kb + + +async def test_keyword_mode_uses_milvus_bm25_search(): + collection = FakeCollection() + kb = make_kb(collection) + + chunks = await kb.aquery( + "alpha beta", + "db", + search_mode="keyword", + bm25_top_k=7, + bm25_drop_ratio_search=0.2, + ) + + assert chunks[0]["content"] == "BM25 result" + assert chunks[0]["bm25_score"] == 0.8 + search_call = collection.search_calls[0] + assert search_call["data"] == ["alpha beta"] + assert search_call["anns_field"] == CONTENT_SPARSE_FIELD + assert search_call["param"] == { + "metric_type": "BM25", + "params": {"drop_ratio_search": 0.2}, + } + assert search_call["limit"] == 7 + + +async def test_hybrid_mode_uses_milvus_native_hybrid_search(): + collection = FakeCollection() + kb = make_kb(collection) + + chunks = await kb.aquery( + "hybrid query", + "db", + search_mode="hybrid", + final_top_k=3, + bm25_top_k=8, + vector_weight=0.6, + bm25_weight=0.4, + ) + + assert chunks[0]["content"] == "Hybrid result" + assert chunks[0]["hybrid_score"] == 0.8 + hybrid_call = collection.hybrid_calls[0] + assert hybrid_call["limit"] == 3 + assert hybrid_call["rerank"]._weights == [0.6, 0.4] + + vector_request, bm25_request = hybrid_call["reqs"] + assert vector_request.anns_field == "embedding" + assert vector_request.data == [[0.1, 0.2]] + assert bm25_request.anns_field == CONTENT_SPARSE_FIELD + assert bm25_request.data == ["hybrid query"] + assert bm25_request.limit == 8 + assert bm25_request.param["metric_type"] == "BM25" + + +async def test_hybrid_mode_filters_scores_below_similarity_threshold(): + collection = FakeCollection(distance=0.1) + kb = make_kb(collection) + + chunks = await kb.aquery( + "hybrid query", + "db", + search_mode="hybrid", + final_top_k=3, + similarity_threshold=0.2, + ) + + assert chunks == [] + + +def test_query_params_config_uses_bm25_parameters(): + kb = MilvusKB.__new__(MilvusKB) + + config = kb.get_query_params_config("db") + + option_keys = {option["key"] for option in config["options"]} + assert "keyword_top_k" not in option_keys + assert { + "bm25_top_k", + "vector_weight", + "bm25_weight", + "bm25_drop_ratio_search", + } <= option_keys + + search_mode = next(option for option in config["options"] if option["key"] == "search_mode") + descriptions = {option["value"]: option["description"] for option in search_mode["options"]} + assert "BM25" in descriptions["keyword"] + assert "BM25" in descriptions["hybrid"] + + +def test_collection_supports_bm25_requires_analyzed_content_sparse_field_and_function(): + kb = MilvusKB.__new__(MilvusKB) + schema = CollectionSchema( + fields=[ + FieldSchema(name="id", dtype=DataType.VARCHAR, max_length=100, is_primary=True), + FieldSchema( + name="content", + dtype=DataType.VARCHAR, + max_length=65535, + enable_analyzer=True, + analyzer_params=CONTENT_ANALYZER_PARAMS, + ), + FieldSchema(name=CONTENT_SPARSE_FIELD, dtype=DataType.SPARSE_FLOAT_VECTOR), + ], + functions=[ + Function( + name="content_bm25", + input_field_names=["content"], + output_field_names=[CONTENT_SPARSE_FIELD], + function_type=FunctionType.BM25, + ) + ], + ) + + collection = type("Collection", (), {"schema": schema})() + + assert kb._collection_supports_bm25(collection) diff --git a/docs/develop-guides/roadmap.md b/docs/develop-guides/roadmap.md index 5a996eed..2e0e7a68 100644 --- a/docs/develop-guides/roadmap.md +++ b/docs/develop-guides/roadmap.md @@ -46,6 +46,7 @@ - 调整部门删除语义:删除部门时不再要求用户数为 0,而是将部门下用户迁移到默认部门,同时清理部门级配置和部门 API Key,保证测试部门、撤换部门等场景可直接删除,并补充对应集成测试覆盖该链路 - 重构 MCP 运行时配置加载模型:移除 `MCP_SERVERS` 作为运行正确性前提的设计,改为每次直接从数据库读取最新 MCP 配置,并用 `server_name:config_hash` 作为本地工具缓存 key;同时将内置 MCP 初始化职责收敛为仅同步数据库默认项,前端 MCP 选项改为直接使用实时资源列表,解决 `api`/`worker` 分进程下的配置不一致与缓存失效问题 - 为知识库检索工具补充 `metadata.filepath` 注入:在 `query_kb` 统一出口基于会话可见知识库构建 `file_id -> /home/gem/kbs/...` 映射并回填检索结果,注入逻辑复用知识库只读后端命名规则;并将工具调用范围收敛为 Milvus(仅支持 Milvus chunks 列表且要求显式 `file_id`),不再兼容无显式 `file_id` 的推断注入,新增单测覆盖该约束 +- 调整 Milvus 混合检索实现:集合 schema 增加 Milvus 内置 BM25 稀疏向量字段、BM25 函数和中文 analyzer 配置,`keyword` 模式改为 BM25 全文检索,`hybrid` 模式改为 Milvus 原生向量 + BM25 混合检索,并同步更新检索参数说明。 - 修复前端依赖安全告警:通过 `pnpm.overrides` 将传递依赖 `flatted` 锁定到 `3.4.2`、`lodash-es` 锁定到 `4.18.1`,并同步更新 `pnpm-lock.yaml` 以消除 DriftGuard 报告的高危 CVE ---