From 610b60e20401937fa0550dfaa632c0d0b47ee25a Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Tue, 30 Dec 2025 20:10:42 +0800 Subject: [PATCH] =?UTF-8?q?feat(knowledge):=20=E6=99=BA=E8=83=BD=E4=BD=93?= =?UTF-8?q?=E6=9F=A5=E8=AF=A2=E7=9F=A5=E8=AF=86=E5=BA=93=E6=97=B6=EF=BC=8C?= =?UTF-8?q?=E6=94=AF=E6=8C=81=E5=9F=BA=E4=BA=8E=E6=96=87=E4=BB=B6=E5=90=8D?= =?UTF-8?q?=E7=9A=84=E6=A8=A1=E7=B3=8A=E8=BF=87=E6=BB=A4=E5=8A=9F=E8=83=BD?= =?UTF-8?q?=EF=BC=8C=E4=B8=8D=E6=94=AF=E6=8C=81=20LightRAG?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 实现知识库检索时可按文件名进行模糊匹配过滤 在Milvus知识库中支持文件名的like表达式过滤 前端展示添加文件名显示 添加相关测试用例验证过滤功能 --- src/agents/common/tools.py | 23 ++- src/knowledge/base.py | 4 +- src/knowledge/implementations/milvus.py | 15 ++ test/test_milvus_filter.py | 137 ++++++++++++++++++ .../ToolCallingResult/BaseToolCall.vue | 1 - .../tools/KnowledgeBaseTool.vue | 6 + 6 files changed, 179 insertions(+), 7 deletions(-) create mode 100644 test/test_milvus_filter.py diff --git a/src/agents/common/tools.py b/src/agents/common/tools.py index 1dfaae01..1689fd1e 100644 --- a/src/agents/common/tools.py +++ b/src/agents/common/tools.py @@ -122,6 +122,11 @@ class KnowledgeRetrieverModel(BaseModel): ) +class CommonKnowledgeRetriever(KnowledgeRetrieverModel): + """Common knowledge retriever model.""" + file_name: str = Field(description="限定文件名称,当操作类型为 'search' 时,可以指定文件名称,支持模糊匹配") + + def get_kb_based_tools(db_names: list[str] | None = None) -> list: """获取所有知识库基于的工具""" # 获取所有知识库 @@ -132,7 +137,9 @@ def get_kb_based_tools(db_names: list[str] | None = None) -> list: def _create_retriever_wrapper(db_id: str, retriever_info: dict[str, Any]): """创建检索器包装函数的工厂函数,避免闭包变量捕获问题""" - async def async_retriever_wrapper(query_text: str, operation: str = "search") -> Any: + async def async_retriever_wrapper( + query_text: str, operation: str = "search", file_name: str | None = None + ) -> Any: """异步检索器包装函数,支持检索和获取思维导图""" # 获取思维导图 @@ -173,10 +180,14 @@ def get_kb_based_tools(db_names: list[str] | None = None) -> list: retriever = retriever_info["retriever"] try: logger.debug(f"Retrieving from database {db_id} with query: {query_text}") + kwargs = {} + if file_name: + kwargs["file_name"] = file_name + if asyncio.iscoroutinefunction(retriever): - result = await retriever(query_text) + result = await retriever(query_text, **kwargs) else: - result = retriever(query_text) + result = retriever(query_text, **kwargs) logger.debug(f"Retrieved {len(result) if isinstance(result, list) else 'N/A'} results from {db_id}") return result except Exception as e: @@ -207,12 +218,16 @@ def get_kb_based_tools(db_names: list[str] | None = None) -> list: safename = retrieve_info["name"].replace(" ", "_")[:20] + args_schema = KnowledgeRetrieverModel + if retrieve_info["metadata"]["kb_type"] in ["milvus"]: + args_schema = CommonKnowledgeRetriever + # 使用 StructuredTool.from_function 创建异步工具 tool = StructuredTool.from_function( coroutine=retriever_wrapper, name=safename, description=description, - args_schema=KnowledgeRetrieverModel, + args_schema=args_schema, metadata=retrieve_info["metadata"] | {"tag": ["knowledgebase"]}, ) diff --git a/src/knowledge/base.py b/src/knowledge/base.py index 4ced893d..0c950ca7 100644 --- a/src/knowledge/base.py +++ b/src/knowledge/base.py @@ -549,8 +549,8 @@ class KnowledgeBase(ABC): for db_id, meta in self.databases_meta.items(): def make_retriever(db_id): - async def retriever(query_text): - return await self.aquery(query_text, db_id) + async def retriever(query_text, **kwargs): + return await self.aquery(query_text, db_id, **kwargs) return retriever diff --git a/src/knowledge/implementations/milvus.py b/src/knowledge/implementations/milvus.py index 85688bd1..a1d81c57 100644 --- a/src/knowledge/implementations/milvus.py +++ b/src/knowledge/implementations/milvus.py @@ -458,11 +458,26 @@ class MilvusKB(KnowledgeBase): query_embedding = embedding_function([query_text]) search_params = {"metric_type": metric_type, "params": {"nprobe": 10}} + + # 构建过滤表达式 + expr = None + if file_name := kwargs.get("file_name"): + # 使用 like 支持模糊匹配 + # 注意:需要转义双引号以防止注入 + safe_file_name = file_name.replace('"', '\\"') + # 如果没有提供通配符,默认前后添加 % + if "%" not in safe_file_name: + expr = f'source like "%{safe_file_name}%"' + else: + expr = f'source like "{safe_file_name}"' + logger.debug(f"Using filter expression: {expr}") + results = collection.search( data=query_embedding, anns_field="embedding", param=search_params, limit=recall_top_k, + expr=expr, output_fields=["content", "source", "chunk_id", "file_id", "chunk_index"], ) diff --git a/test/test_milvus_filter.py b/test/test_milvus_filter.py new file mode 100644 index 00000000..ca8144d6 --- /dev/null +++ b/test/test_milvus_filter.py @@ -0,0 +1,137 @@ +import asyncio +import os +import shutil +from unittest.mock import MagicMock, patch + +from src.knowledge import knowledge_base +from src.utils import logger + +# Mock Embedding Model +class MockEmbeddingModel: + async def abatch_encode(self, texts, batch_size=None): + # Return dummy vectors of dim 4 + return [[0.1, 0.2, 0.3, 0.4] for _ in texts] + + def batch_encode(self, texts, batch_size=None): + return [[0.1, 0.2, 0.3, 0.4] for _ in texts] + +# Test function +async def test_milvus_filter(): + logger.info("Starting Milvus Filter Test") + + # Check if Milvus is available (pymilvus installed and connection works) + try: + from pymilvus import connections, utility + # Assuming Milvus is running at default location + connections.connect(alias="default", uri=os.getenv("MILVUS_URI", "http://localhost:19530")) + logger.info("Connected to Milvus") + except Exception as e: + logger.warning(f"Milvus not available or connection failed: {e}") + # Proceeding might fail, but let's try. + + db_id = "test_milvus_filter_db" + file1 = "test_file_A.txt" + file2 = "test_file_B.txt" + + # Patch embedding model + with patch("src.models.embed.select_embedding_model", return_value=MockEmbeddingModel()): + + try: + # Cleanup if exists + if db_id in knowledge_base.global_databases_meta: + await knowledge_base.delete_database(db_id) + + # Create DB + logger.info("Creating database...") + # explicitly set dimension to 4 to match mock + await knowledge_base.create_database( + database_name="Test Milvus Filter", + description="Test DB", + kb_type="milvus", + embed_info={"name": "mock-embedding", "dimension": 4, "model_id": "mock"} + ) + + # Get actual db_id + target_db = next((db for db in knowledge_base.get_databases()["databases"] if db["name"] == "Test Milvus Filter"), None) + if not target_db: + logger.error("Failed to create DB") + return + + db_id = target_db["db_id"] + logger.info(f"DB created with ID: {db_id}") + + # Create dummy files + + with open(file1, "w") as f: + f.write("Apple content.") + with open(file2, "w") as f: + f.write("Banana content.") + + # Add content + logger.info("Adding content...") + await knowledge_base.add_content(db_id, [os.path.abspath(file1), os.path.abspath(file2)]) + + # Wait for data to be visible + logger.info("Waiting for data to be visible...") + await asyncio.sleep(2) + + # Query without filter + logger.info("Querying without filter...") + results = await knowledge_base.aquery("content", db_id) + logger.info(f"No filter results: {len(results)}") + + # Verify we have chunks from both files + sources = [r['metadata']['source'] for r in results] + logger.info(f"Sources: {sources}") + + # Query with filter A (Partial Match) + logger.info("Querying with filter A (file_A)...") + results_a = await knowledge_base.aquery("content", db_id, file_name="file_A") + logger.info(f"Filter A results: {len(results_a)}") + + if len(results_a) == 0: + logger.error("FAIL: Filter A returned 0 results") + + for r in results_a: + source = r['metadata']['source'] + logger.info(f" - {source}") + if "test_file_A.txt" not in source: + logger.error(f"FAIL: Expected test_file_A.txt, got {source}") + raise AssertionError("Filter A failed") + + # Query with wildcard filter + logger.info("Querying with wildcard filter (%B.txt)...") + results_b = await knowledge_base.aquery("content", db_id, file_name="%B.txt") + logger.info(f"Filter B results: {len(results_b)}") + if len(results_b) == 0: + logger.error("FAIL: Wildcard filter returned 0 results") + + for r in results_b: + source = r['metadata']['source'] + logger.info(f" - {source}") + if "test_file_B.txt" not in source: + logger.error(f"FAIL: Expected test_file_B.txt, got {source}") + raise AssertionError("Filter B failed") + + if len(results_a) > 0 and len(results_b) > 0: + logger.info("Test passed!") + else: + logger.error("Test failed: No results found for one or more queries") + + except Exception as e: + logger.error(f"Test failed with exception: {e}") + raise + finally: + # Cleanup + logger.info("Cleaning up...") + try: + await knowledge_base.delete_database(db_id) + except Exception: + pass + if os.path.exists(file1): + os.remove(file1) + if os.path.exists(file2): + os.remove(file2) + +if __name__ == "__main__": + asyncio.run(test_milvus_filter()) diff --git a/web/src/components/ToolCallingResult/BaseToolCall.vue b/web/src/components/ToolCallingResult/BaseToolCall.vue index cd69805f..d6c4c91d 100644 --- a/web/src/components/ToolCallingResult/BaseToolCall.vue +++ b/web/src/components/ToolCallingResult/BaseToolCall.vue @@ -307,7 +307,6 @@ const formatResultData = (data) => { text-overflow: ellipsis; white-space: nowrap; min-width: 0; - flex: 1; } :deep(.tag) { diff --git a/web/src/components/ToolCallingResult/tools/KnowledgeBaseTool.vue b/web/src/components/ToolCallingResult/tools/KnowledgeBaseTool.vue index bd07cdf0..c1b579a7 100644 --- a/web/src/components/ToolCallingResult/tools/KnowledgeBaseTool.vue +++ b/web/src/components/ToolCallingResult/tools/KnowledgeBaseTool.vue @@ -5,6 +5,8 @@ {{ operationLabel }} | {{ queryText }} + | + 文件: {{ fileName }}