diff --git a/backend/package/yuxi/agents/middlewares/knowledge_base_middleware.py b/backend/package/yuxi/agents/middlewares/knowledge_base_middleware.py index bd1b1ec7..285e94b2 100644 --- a/backend/package/yuxi/agents/middlewares/knowledge_base_middleware.py +++ b/backend/package/yuxi/agents/middlewares/knowledge_base_middleware.py @@ -12,10 +12,12 @@ from yuxi.utils.logging_config import logger class KnowledgeBaseMiddleware(AgentMiddleware): """知识库中间件 - 提供通用知识库工具 - 提供 3 个通用工具: + 提供通用知识库工具: - list_kbs: 列出用户可访问的知识库 - get_mindmap: 获取指定知识库的思维导图 - query_kb: 在指定知识库中检索 + - find_kb_document: 在指定文件内定位关键词或正则模式 + - open_kb_document: 按 file_id 分段打开知识库文档 """ def __init__(self): diff --git a/backend/package/yuxi/agents/toolkits/kbs/__init__.py b/backend/package/yuxi/agents/toolkits/kbs/__init__.py index dc052ee1..f4c550b9 100644 --- a/backend/package/yuxi/agents/toolkits/kbs/__init__.py +++ b/backend/package/yuxi/agents/toolkits/kbs/__init__.py @@ -1,6 +1,7 @@ from .tools import ( + find_kb_document, get_common_kb_tools, open_kb_document, ) -__all__ = ["get_common_kb_tools", "open_kb_document"] +__all__ = ["find_kb_document", "get_common_kb_tools", "open_kb_document"] diff --git a/backend/package/yuxi/agents/toolkits/kbs/tools.py b/backend/package/yuxi/agents/toolkits/kbs/tools.py index 9536bc48..c19f44e3 100644 --- a/backend/package/yuxi/agents/toolkits/kbs/tools.py +++ b/backend/package/yuxi/agents/toolkits/kbs/tools.py @@ -8,6 +8,15 @@ from langgraph.prebuilt.tool_node import ToolRuntime from pydantic import BaseModel, Field from yuxi import knowledge_base +from yuxi.knowledge.base import KnowledgeBase +from yuxi.knowledge.schemas import ( + FindInputSchema, + FindOutputSchema, + OpenInputSchema, + OpenOutputSchema, + SearchInputSchema, + SearchOutputSchema, +) from yuxi.utils import logger # ========== 通用知识库工具函数 ========== @@ -62,7 +71,7 @@ async def list_kbs(dummy: str, runtime: ToolRuntime) -> str: # Now has 2 params for kb in available_kbs: name = kb.get("name", "") desc = kb.get("description") or "无描述" - kb_list.append({"name": name, "description": desc}) + kb_list.append({"resource_id": kb.get("db_id"), "name": name, "description": desc}) return kb_list @@ -137,33 +146,9 @@ async def get_mindmap(kb_name: str, runtime: ToolRuntime) -> str: return f"获取思维导图失败: {str(e)}" -class QueryKBInput(BaseModel): - """知识库检索输入模型""" - - kb_name: str = Field(description="知识库名称,用于指定要在哪个知识库中检索") - query_text: str = Field( - description=( - "查询的关键词,查询的时候,应该尽量以可能帮助回答这个问题的关键词进行查询," - "不要直接使用用户的原始输入去查询。" - ) - ) - file_name: str | None = Field( - default=None, - description=( - "(非必要不启用此参数,留空即可)当已经读取思维导图之后,可以指定文件关键词,支持模糊匹配。\n" - "仅当检索结果过多且不相关,需要进一步缩小范围时使用。" - ), - ) - - -class OpenKBDocumentInput(BaseModel): - """打开知识库文档输入模型""" - - resource_id: str = Field(description="知识库资源 ID,当前对应知识库 db_id") - file_id: str = Field(description="要打开的文档 ID,也就是知识库文件 file_id") - line: int | None = Field(default=None, ge=1, description="可选,1-based 起始行号") - offset: int | None = Field(default=None, ge=0, description="可选,0-based 起始偏移;line 优先于 offset") - window_size: int = Field(default=800, ge=1, le=2000, description="读取窗口行数,默认 800 行") +QueryKBInput = SearchInputSchema +OpenKBDocumentInput = OpenInputSchema +FindKBDocumentInput = FindInputSchema async def _resolve_visible_knowledge_bases_for_query(runtime: ToolRuntime | None) -> list[dict[str, Any]]: @@ -189,81 +174,40 @@ async def _resolve_visible_knowledge_bases_for_query(runtime: ToolRuntime | None def _find_query_target( *, - kb_name: str, + resource_id: str, retrievers: dict[str, Any], visible_kbs: list[dict[str, Any]], ) -> tuple[dict[str, Any] | None, str | None, str | None]: if not visible_kbs: return None, None, "无法获取当前会话可访问的知识库" - matched_kbs = [db for db in visible_kbs if str(db.get("name") or "").strip() == kb_name] - if not matched_kbs: - return None, None, f"知识库 '{kb_name}' 不存在或当前会话未启用" - if len(matched_kbs) > 1: - return None, None, f"知识库 '{kb_name}' 存在重名,请先调整名称后重试" + normalized_resource_id = str(resource_id or "").strip() + visible_resource_ids = {str(kb.get("db_id") or "").strip() for kb in visible_kbs} + if normalized_resource_id not in visible_resource_ids: + return None, None, f"知识库资源 '{normalized_resource_id}' 不存在或当前会话未启用" - target_db_id = str(matched_kbs[0].get("db_id") or "") - target_info = retrievers.get(target_db_id) + target_info = retrievers.get(normalized_resource_id) if target_info is None: - return None, None, f"知识库 '{kb_name}' 不存在" - return target_info, target_db_id, None - - -def _normalize_retrieval_result_metadata(result: Any, resource_id: str | None = None) -> Any: - if not isinstance(result, list): - return result - - for chunk in result: - if not isinstance(chunk, dict): - continue - - metadata = chunk.get("metadata") - if not isinstance(metadata, dict): - metadata = {} - - file_id = metadata.get("file_id") or chunk.get("file_id") or chunk.get("full_doc_id") - if file_id and not metadata.get("file_id"): - metadata["file_id"] = str(file_id) - if resource_id: - metadata["resource_id"] = resource_id - - for key in ("chunk_id", "chunk_index"): - value = metadata.get(key) if metadata.get(key) is not None else chunk.get(key) - if value is not None and metadata.get(key) is None: - metadata[key] = value - - chunk["metadata"] = metadata - - return result + return None, None, f"知识库资源 '{normalized_resource_id}' 不存在" + return target_info, normalized_resource_id, None @tool(args_schema=QueryKBInput) -async def query_kb(kb_name: str, query_text: str, file_name: str | None = None, runtime: ToolRuntime = None) -> Any: +async def query_kb(resource_id: str, query_text: str, file_name: str | None = None, runtime: ToolRuntime = None) -> Any: """在指定知识库中检索内容 - 当用户需要查询具体内容时使用此工具。根据关键词在知识库中检索相关文档片段;结果中的 - metadata.file_id 和知识库 resource_id 可用于继续调用 open_kb_document 打开原文。 - - Args: - kb_name: 知识库名称 - query_text: 查询的关键词 - file_name: (可选)文件名称过滤 - - Returns: - 检索结果 + 当用户需要查询具体内容时使用此工具。resource_id 是知识库资源 ID,也就是 kb_id;返回结果中的 + file_id 可继续用于 find_kb_document 或 open_kb_document。 """ - if not kb_name: - return "请提供知识库名称" + if not resource_id: + return "请提供 resource_id" if not query_text: return "请提供查询内容" - # 获取所有检索器 retrievers = knowledge_base.get_retrievers() - visible_kbs = await _resolve_visible_knowledge_bases_for_query(runtime) - target_info, target_db_id, target_error = _find_query_target( - kb_name=kb_name, + resource_id=resource_id, retrievers=retrievers, visible_kbs=visible_kbs, ) @@ -281,7 +225,13 @@ async def query_kb(kb_name: str, query_text: str, file_name: str | None = None, else: result = retriever(query_text, **kwargs) - return _normalize_retrieval_result_metadata(result, target_db_id) + if ( + isinstance(result, dict) + and result.get("resource_id") == target_db_id + and isinstance(result.get("results"), list) + ): + return SearchOutputSchema(**result).model_dump() + return KnowledgeBase.build_search_output(target_db_id, result) except Exception as e: logger.error(f"检索失败: {e}") @@ -294,13 +244,13 @@ async def open_kb_document( file_id: str, line: int | None = None, offset: int | None = None, - window_size: int = 800, + window_size: int = 1800, runtime: ToolRuntime = None, ) -> dict[str, Any] | str: """按行窗口打开知识库文档原文 当 query_kb 返回的片段不足以回答问题,或需要查看某个文档的上下文时使用。 - resource_id 对应知识库 db_id,file_id 对应检索结果 metadata.file_id。 + resource_id 是知识库资源 ID,也就是 kb_id;file_id 是知识库文件 ID。 """ normalized_resource_id = str(resource_id or "").strip() normalized_file_id = str(file_id or "").strip() @@ -335,20 +285,79 @@ async def open_kb_document( offset=start_offset, limit=window_size, ) - return {"resource_id": normalized_resource_id, "file_id": normalized_file_id, **window} + return OpenOutputSchema(resource_id=normalized_resource_id, file_id=normalized_file_id, **window).model_dump() except Exception as e: logger.error(f"打开知识库文档失败: {e}") return f"打开知识库文档失败: {str(e)}" +@tool(args_schema=FindKBDocumentInput) +async def find_kb_document( + resource_id: str, + file_id: str, + patterns: list[str], + use_regex: bool = False, + case_sensitive: bool = False, + max_windows: int = 5, + window_size: int = 80, + runtime: ToolRuntime = None, +) -> dict[str, Any] | str: + """在已知知识库文件内做关键词或正则定位。 + + 当 query_kb 已找到候选文件,但需要在该文件内定位术语、指标、章节或实体时使用。 + """ + normalized_resource_id = str(resource_id or "").strip() + normalized_file_id = str(file_id or "").strip() + if not normalized_resource_id: + return "请提供 resource_id" + if not normalized_file_id: + return "请提供 file_id" + if not patterns: + return "请提供 patterns" + + visible_kbs = await _resolve_visible_knowledge_bases_for_query(runtime) + if not visible_kbs: + return "无法获取当前会话可访问的知识库" + + visible_resource_ids = {str(kb.get("db_id") or "").strip() for kb in visible_kbs} + if normalized_resource_id not in visible_resource_ids: + return f"知识库资源 '{normalized_resource_id}' 不存在或当前会话未启用" + + retrievers = knowledge_base.get_retrievers() + target_info = retrievers.get(normalized_resource_id) + if target_info is None: + return f"知识库资源 '{normalized_resource_id}' 不存在" + + metadata = target_info.get("metadata") if isinstance(target_info, dict) else None + kb_type = str((metadata or {}).get("kb_type") or "").strip().lower() + if kb_type == "dify": + return "Dify 知识库为外部只读检索源,当前不支持通过 Find 检索全文" + + try: + result = await knowledge_base.find_file_content( + normalized_resource_id, + normalized_file_id, + patterns, + use_regex=use_regex, + case_sensitive=case_sensitive, + max_windows=max_windows, + window_size=window_size, + ) + return FindOutputSchema(resource_id=normalized_resource_id, file_id=normalized_file_id, **result).model_dump() + except Exception as e: + logger.error(f"知识库文档内检索失败: {e}") + return f"知识库文档内检索失败: {str(e)}" + + def get_common_kb_tools() -> list: """获取通用知识库工具列表 - 返回 4 个通用工具: + 返回 5 个通用工具: - list_kbs: 列出用户可访问的知识库 - get_mindmap: 获取指定知识库的思维导图 - query_kb: 在指定知识库中检索 + - find_kb_document: 在指定文件内定位关键词或正则模式 - open_kb_document: 按 file_id 分段打开知识库文档 """ - return [list_kbs, get_mindmap, query_kb, open_kb_document] + return [list_kbs, get_mindmap, query_kb, find_kb_document, open_kb_document] diff --git a/backend/package/yuxi/knowledge/base.py b/backend/package/yuxi/knowledge/base.py index d57a90b6..5cf12f0f 100644 --- a/backend/package/yuxi/knowledge/base.py +++ b/backend/package/yuxi/knowledge/base.py @@ -1,10 +1,12 @@ import asyncio import mimetypes import os +import re from abc import ABC, abstractmethod from typing import Any from yuxi.knowledge.chunking.ragflow_like.presets import ensure_chunk_defaults_in_additional_params +from yuxi.knowledge.schemas import FindOutputSchema, FindWindowSchema, SearchOutputSchema, SearchResultSchema from yuxi.knowledge.utils import resolve_processing_params, sanitize_processing_params from yuxi.utils import logger from yuxi.utils.datetime_utils import coerce_any_to_utc_datetime, utc_isoformat @@ -675,6 +677,110 @@ class KnowledgeBase(ABC): "content": "\n".join(f"{start + idx + 1:6d}\t{line}" for idx, line in enumerate(selected)), } + @staticmethod + def build_search_output(resource_id: str, retrieval_results: Any) -> dict[str, Any] | Any: + if not isinstance(retrieval_results, list): + return retrieval_results + + results = [] + for index, chunk in enumerate(retrieval_results): + if not isinstance(chunk, dict): + continue + + metadata = chunk.get("metadata") if isinstance(chunk.get("metadata"), dict) else {} + metadata = { + key: value + for key, value in metadata.items() + if key not in {"filepath", "parsed_path", "path", "markdown_file"} + } + file_id = metadata.get("file_id") or chunk.get("file_id") or chunk.get("full_doc_id") or "" + chunk_id = metadata.get("chunk_id") or chunk.get("chunk_id") or chunk.get("id") + chunk_index = metadata.get("chunk_index") + if chunk_index is None: + chunk_index = chunk.get("chunk_index") + if chunk_index is not None: + metadata.setdefault("chunk_index", chunk_index) + if chunk.get("score") is not None: + metadata.setdefault("score", chunk.get("score")) + if chunk.get("distance") is not None: + metadata.setdefault("distance", chunk.get("distance")) + + results.append( + SearchResultSchema( + id=str(chunk_id or f"{file_id}:{index + 1}"), + resource_id=str(resource_id), + file_id=str(file_id or ""), + content=str(chunk.get("content") or ""), + metadata=metadata, + ) + ) + + return SearchOutputSchema(resource_id=str(resource_id), results=results).model_dump() + + @staticmethod + def _build_find_file_windows( + content: str, + *, + patterns: list[str], + use_regex: bool = False, + case_sensitive: bool = False, + max_windows: int = 5, + window_size: int = 80, + ) -> dict[str, Any]: + patterns = [pattern for pattern in patterns if pattern] + if not patterns: + raise ValueError("请提供至少一个 pattern") + + lines = content.splitlines() + flags = 0 if case_sensitive else re.IGNORECASE + if use_regex: + matchers = [re.compile(pattern, flags) for pattern in patterns] + + def line_matches(line: str) -> bool: + return any(matcher.search(line) for matcher in matchers) + + else: + normalized_patterns = patterns if case_sensitive else [pattern.lower() for pattern in patterns] + + def line_matches(line: str) -> bool: + haystack = line if case_sensitive else line.lower() + return any(pattern in haystack for pattern in normalized_patterns) + + matched_indexes = [index for index, line in enumerate(lines) if line_matches(line)] + windows: list[FindWindowSchema] = [] + covered_until = -1 + normalized_window_size = min(max(int(window_size), 1), 200) + half_window = normalized_window_size // 2 + + for matched_index in matched_indexes: + if matched_index < covered_until: + continue + start = max(matched_index - half_window, 0) + end = min(start + normalized_window_size, len(lines)) + start = max(end - normalized_window_size, 0) + matched_lines = [index + 1 for index in matched_indexes if start <= index < end] + selected = lines[start:end] + windows.append( + FindWindowSchema( + start_line=start + 1 if selected else 0, + end_line=end, + matched_lines=matched_lines, + content="\n".join(f"{start + idx + 1:6d}\t{line}" for idx, line in enumerate(selected)), + ) + ) + covered_until = end + if len(windows) >= max_windows: + break + + return FindOutputSchema( + resource_id="", + file_id="", + semantic=False, + match_mode="regex" if use_regex else "keyword", + total_matches=len(matched_indexes), + windows=windows, + ).model_dump(exclude={"resource_id", "file_id"}) + async def open_file_content(self, db_id: str, file_id: str, offset: int = 0, limit: int = 800) -> dict: """按行窗口打开文件解析后的 Markdown 内容""" file_meta = self.files_meta.get(file_id) @@ -692,6 +798,39 @@ class KnowledgeBase(ABC): content = await self._read_markdown_from_minio(markdown_file) return self._build_open_file_window(content, offset=offset, limit=limit) + async def find_file_content( + self, + db_id: str, + file_id: str, + patterns: list[str], + *, + use_regex: bool = False, + case_sensitive: bool = False, + max_windows: int = 5, + window_size: int = 80, + ) -> dict: + file_meta = self.files_meta.get(file_id) + if file_meta is None: + raise Exception(f"文件不存在: {file_id}") + if file_meta.get("database_id") != db_id: + raise Exception(f"文件 {file_id} 不属于知识库 {db_id}") + if file_meta.get("is_folder"): + raise Exception(f"文件 {file_id} 是文件夹") + + markdown_file = file_meta.get("markdown_file") + if not markdown_file: + raise Exception(f"文件 {file_id} 没有解析后的 Markdown 内容") + + content = await self._read_markdown_from_minio(markdown_file) + return self._build_find_file_windows( + content, + patterns=patterns, + use_regex=use_regex, + case_sensitive=case_sensitive, + max_windows=max_windows, + window_size=window_size, + ) + @abstractmethod async def index_file(self, db_id: str, file_id: str, operator_id: str | None = None) -> dict: """ @@ -1278,7 +1417,8 @@ class KnowledgeBase(ABC): def make_retriever(db_id): async def retriever(query_text, **kwargs): - return await self.aquery(query_text, db_id, agent_call=True, **kwargs) + results = await self.aquery(query_text, db_id, agent_call=True, **kwargs) + return self.build_search_output(db_id, results) return retriever diff --git a/backend/package/yuxi/knowledge/implementations/read_only_connectors.py b/backend/package/yuxi/knowledge/implementations/read_only_connectors.py index f7f62ac8..6dd023fd 100644 --- a/backend/package/yuxi/knowledge/implementations/read_only_connectors.py +++ b/backend/package/yuxi/knowledge/implementations/read_only_connectors.py @@ -70,6 +70,20 @@ class ReadOnlyConnectors(KnowledgeBase): del offset, limit raise self._readonly_error() + async def find_file_content( + self, + db_id: str, + file_id: str, + patterns: list[str], + *, + use_regex: bool = False, + case_sensitive: bool = False, + max_windows: int = 5, + window_size: int = 80, + ) -> dict: + del db_id, file_id, patterns, use_regex, case_sensitive, max_windows, window_size + raise self._readonly_error() + async def get_file_info(self, db_id: str, file_id: str) -> dict: raise self._readonly_error() diff --git a/backend/package/yuxi/knowledge/manager.py b/backend/package/yuxi/knowledge/manager.py index 0ef1c49f..99d6e839 100644 --- a/backend/package/yuxi/knowledge/manager.py +++ b/backend/package/yuxi/knowledge/manager.py @@ -228,18 +228,6 @@ class KnowledgeBaseManager: return user_department_id in accessible_departments - async def get_databases_by_raw_id(self, user_id: int) -> dict: - """根据用户ID获取知识库列表""" - from yuxi.repositories.user_repository import UserRepository - - # 通过数据库获取用户信息 - user_repo = UserRepository() - user: User | None = await user_repo.get_by_id(id=int(user_id)) - if not user: - logger.warning(f"User not found: {uid}") - return {"databases": []} - return await self.get_databases_by_user(user) - async def get_databases_by_uid(self, uid: str) -> dict: """根据 uid 获取知识库列表""" from yuxi.repositories.user_repository import UserRepository @@ -493,6 +481,28 @@ class KnowledgeBaseManager: kb_instance = await self._get_kb_for_database(db_id) return await kb_instance.open_file_content(db_id, file_id, offset, limit) + async def find_file_content( + self, + db_id: str, + file_id: str, + patterns: list[str], + *, + use_regex: bool = False, + case_sensitive: bool = False, + max_windows: int = 5, + window_size: int = 80, + ) -> dict: + kb_instance = await self._get_kb_for_database(db_id) + return await kb_instance.find_file_content( + db_id, + file_id, + patterns, + use_regex=use_regex, + case_sensitive=case_sensitive, + max_windows=max_windows, + window_size=window_size, + ) + async def get_file_info(self, db_id: str, file_id: str) -> dict: """获取文件完整信息(基本信息+内容信息)""" kb_instance = await self._get_kb_for_database(db_id) diff --git a/backend/package/yuxi/knowledge/schemas.py b/backend/package/yuxi/knowledge/schemas.py new file mode 100644 index 00000000..bfd0da98 --- /dev/null +++ b/backend/package/yuxi/knowledge/schemas.py @@ -0,0 +1,70 @@ +from typing import Any, Literal + +from pydantic import BaseModel, Field + + +class SearchInputSchema(BaseModel): + resource_id: str = Field(description="知识库资源 ID,也就是 kb_id") + query_text: str = Field(description="检索关键词,应提炼为有助于召回答案的关键词或短语") + file_name: str | None = Field(default=None, description="可选文件名关键词过滤,非必要不要使用") + + +class SearchResultSchema(BaseModel): + id: str = Field(description="检索结果 ID,通常对应 chunk_id") + resource_id: str = Field(description="知识库资源 ID,也就是 kb_id") + file_id: str = Field(default="", description="结果所属文件 ID,可用于 Find/Open") + content: str = Field(description="chunk 内容") + metadata: dict[str, Any] = Field(default_factory=dict, description="来源、分数、chunk_index 等附加信息") + + +class SearchOutputSchema(BaseModel): + resource_id: str = Field(description="知识库资源 ID,也就是 kb_id") + results: list[SearchResultSchema] = Field(default_factory=list, description="检索结果列表") + + +class FindInputSchema(BaseModel): + resource_id: str = Field(description="知识库资源 ID,也就是 kb_id") + file_id: str = Field(description="要检索的文件 ID") + patterns: list[str] = Field(description="关键词或正则模式列表,至少提供一个") + use_regex: bool = Field(default=False, description="是否将 patterns 作为正则表达式处理") + case_sensitive: bool = Field(default=False, description="是否区分大小写") + max_windows: int = Field(default=5, ge=1, le=20, description="最多返回的上下文窗口数量") + window_size: int = Field(default=80, ge=1, le=200, description="每个上下文窗口的行数") + + +class FindWindowSchema(BaseModel): + start_line: int = Field(description="窗口起始行号,1-based") + end_line: int = Field(description="窗口结束行号,1-based") + matched_lines: list[int] = Field(default_factory=list, description="该窗口内匹配到的行号") + content: str = Field(description="带行号的窗口内容") + + +class FindOutputSchema(BaseModel): + resource_id: str = Field(description="知识库资源 ID,也就是 kb_id") + file_id: str = Field(description="文件 ID") + semantic: bool = Field(default=False, description="是否为语义查找") + match_mode: Literal["keyword", "regex"] = Field(description="匹配模式") + total_matches: int = Field(description="匹配到的行数") + windows: list[FindWindowSchema] = Field(default_factory=list, description="上下文窗口") + + +class OpenInputSchema(BaseModel): + resource_id: str = Field(description="知识库资源 ID,也就是 kb_id") + file_id: str = Field(description="要打开的文件 ID") + line: int | None = Field(default=None, ge=1, description="可选,1-based 起始行号") + offset: int | None = Field(default=None, ge=0, description="可选,0-based 起始偏移;line 优先于 offset") + window_size: int = Field(default=1800, ge=1, le=2000, description="读取窗口行数") + + +class OpenOutputSchema(BaseModel): + resource_id: str = Field(description="知识库资源 ID,也就是 kb_id") + file_id: str = Field(description="文件 ID") + start_line: int = Field(description="窗口起始行号,1-based;空结果为 0") + end_line: int = Field(description="窗口结束行号,1-based;空结果为 0") + total_lines: int = Field(description="文件总行数") + offset: int = Field(description="窗口起始偏移,0-based") + window_size: int = Field(description="本次请求的窗口行数") + has_more_before: bool = Field(description="窗口前是否还有内容") + has_more_after: bool = Field(description="窗口后是否还有内容") + next_offset: int | None = Field(default=None, description="下一窗口 offset;没有更多内容时为 null") + content: str = Field(description="带行号的窗口内容") diff --git a/backend/test/unit/toolkits/test_kbs_tools.py b/backend/test/unit/toolkits/test_kbs_tools.py index 21e12872..aa628781 100644 --- a/backend/test/unit/toolkits/test_kbs_tools.py +++ b/backend/test/unit/toolkits/test_kbs_tools.py @@ -24,6 +24,10 @@ def _query_kb_callable(): return _tool_callable(tools.query_kb) +def _find_kb_document_callable(): + return _tool_callable(tools.find_kb_document) + + def _open_kb_document_callable(): return _tool_callable(tools.open_kb_document) @@ -39,11 +43,15 @@ async def _run_query_kb(**kwargs): return await _run_tool(_query_kb_callable(), **kwargs) +async def _run_find_kb_document(**kwargs): + return await _run_tool(_find_kb_document_callable(), **kwargs) + + async def _run_open_kb_document(**kwargs): return await _run_tool(_open_kb_document_callable(), **kwargs) -def _build_test_window(content: str, offset: int = 0, limit: int = 800) -> dict: +def _build_test_window(content: str, offset: int = 0, limit: int = 1800) -> dict: lines = content.splitlines() start = min(max(offset, 0), len(lines)) selected = lines[start : start + limit] @@ -61,52 +69,55 @@ def _build_test_window(content: str, offset: int = 0, limit: int = 800) -> dict: } -@pytest.mark.asyncio -async def test_query_kb_returns_milvus_chunks_without_sandbox_paths(monkeypatch) -> None: - async def _fake_retriever(query_text: str, **kwargs): - assert query_text == "auth" - return [ - { - "content": "auth guide", - "metadata": { - "file_id": "file-1", - "source": "auth-guide.pdf", - }, - } - ] - +def _patch_retrievers(monkeypatch, *, kb_type: str = "milvus", retriever=None): monkeypatch.setattr( tools.knowledge_base, "get_retrievers", lambda: { "db-1": { "name": "FAQ", - "retriever": _fake_retriever, - "metadata": {"kb_type": "milvus"}, + "retriever": retriever or object(), + "metadata": {"kb_type": kb_type}, } }, ) - async def _fake_visible_kbs(runtime): - return [{"db_id": "db-1", "name": "FAQ"}] +async def _fake_visible_kbs(runtime): + del runtime + return [{"db_id": "db-1", "name": "FAQ"}] + + +@pytest.mark.asyncio +async def test_query_kb_returns_search_schema_without_sandbox_paths(monkeypatch) -> None: + async def _fake_retriever(query_text: str, **kwargs): + assert query_text == "auth" + assert kwargs == {} + return [ + { + "content": "auth guide", + "metadata": { + "file_id": "file-1", + "source": "auth-guide.pdf", + "filepath": "/tmp/sandbox/auth-guide.pdf", + }, + } + ] + + _patch_retrievers(monkeypatch, retriever=_fake_retriever) monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs) runtime = SimpleNamespace(context=SimpleNamespace()) - result = await _run_query_kb(kb_name="FAQ", query_text="auth", runtime=runtime) + result = await _run_query_kb(resource_id="db-1", query_text="auth", runtime=runtime) - assert result == [ - { - "content": "auth guide", - "metadata": { - "file_id": "file-1", - "source": "auth-guide.pdf", - "resource_id": "db-1", - }, - } - ] - assert "filepath" not in result[0]["metadata"] - assert "parsed_path" not in result[0]["metadata"] + assert result["resource_id"] == "db-1" + assert result["results"][0]["id"] == "file-1:1" + assert result["results"][0]["resource_id"] == "db-1" + assert result["results"][0]["file_id"] == "file-1" + assert result["results"][0]["content"] == "auth guide" + assert result["results"][0]["metadata"]["source"] == "auth-guide.pdf" + assert "filepath" not in result["results"][0]["metadata"] + assert "parsed_path" not in result["results"][0]["metadata"] @pytest.mark.asyncio @@ -119,42 +130,35 @@ async def test_query_kb_allows_dify_knowledge_base(monkeypatch) -> None: "score": 0.98, "metadata": { "file_id": "dify-doc-1", + "chunk_id": "dify-segment-1", "source": "Dify Doc", }, } ] - monkeypatch.setattr( - tools.knowledge_base, - "get_retrievers", - lambda: { - "db-1": { - "name": "FAQ", - "retriever": _fake_retriever, - "metadata": {"kb_type": "dify"}, - } - }, - ) - - async def _fake_visible_kbs(runtime): - return [{"db_id": "db-1", "name": "FAQ"}] - + _patch_retrievers(monkeypatch, kb_type="dify", retriever=_fake_retriever) monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs) runtime = SimpleNamespace(context=SimpleNamespace()) - result = await _run_query_kb(kb_name="FAQ", query_text="auth", runtime=runtime) + result = await _run_query_kb(resource_id="db-1", query_text="auth", runtime=runtime) - assert result == [ - { - "content": "auth guide", - "score": 0.98, - "metadata": { - "file_id": "dify-doc-1", - "source": "Dify Doc", + assert result == { + "resource_id": "db-1", + "results": [ + { + "id": "dify-segment-1", "resource_id": "db-1", - }, - } - ] + "file_id": "dify-doc-1", + "content": "auth guide", + "metadata": { + "file_id": "dify-doc-1", + "chunk_id": "dify-segment-1", + "source": "Dify Doc", + "score": 0.98, + }, + } + ], + } @pytest.mark.asyncio @@ -163,31 +167,17 @@ async def test_query_kb_returns_plain_result_without_path_injection(monkeypatch) assert query_text == "auth" return "Milvus context" - monkeypatch.setattr( - tools.knowledge_base, - "get_retrievers", - lambda: { - "db-1": { - "name": "FAQ", - "retriever": _fake_retriever, - "metadata": {"kb_type": "milvus"}, - } - }, - ) - - async def _fake_visible_kbs(runtime): - return [{"db_id": "db-1", "name": "FAQ"}] - + _patch_retrievers(monkeypatch, retriever=_fake_retriever) monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs) runtime = SimpleNamespace(context=SimpleNamespace()) - result = await _run_query_kb(kb_name="FAQ", query_text="auth", runtime=runtime) + result = await _run_query_kb(resource_id="db-1", query_text="auth", runtime=runtime) assert result == "Milvus context" @pytest.mark.asyncio -async def test_query_kb_normalizes_file_metadata_for_open(monkeypatch) -> None: +async def test_query_kb_maps_full_doc_id_and_chunk_metadata(monkeypatch) -> None: async def _fake_retriever(query_text: str, **kwargs): assert query_text == "auth" return [ @@ -199,59 +189,112 @@ async def test_query_kb_normalizes_file_metadata_for_open(monkeypatch) -> None: } ] - monkeypatch.setattr( - tools.knowledge_base, - "get_retrievers", - lambda: { - "db-1": { - "name": "FAQ", - "retriever": _fake_retriever, - "metadata": {"kb_type": "milvus"}, - } - }, - ) - - async def _fake_visible_kbs(runtime): - return [{"db_id": "db-1", "name": "FAQ"}] - + _patch_retrievers(monkeypatch, retriever=_fake_retriever) monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs) runtime = SimpleNamespace(context=SimpleNamespace()) - result = await _run_query_kb(kb_name="FAQ", query_text="auth", runtime=runtime) + result = await _run_query_kb(resource_id="db-1", query_text="auth", runtime=runtime) - assert result[0]["metadata"] == { - "file_id": "file-1", + assert result["results"][0] == { + "id": "chunk-1", "resource_id": "db-1", - "chunk_id": "chunk-1", - "chunk_index": 3, + "file_id": "file-1", + "content": "auth guide", + "metadata": {"chunk_index": 3}, } @pytest.mark.asyncio -async def test_open_kb_document_reads_markdown_content_by_default_window(monkeypatch) -> None: - lines = [f"line {index}" for index in range(1, 1001)] +async def test_find_kb_document_returns_context_windows(monkeypatch) -> None: + _patch_retrievers(monkeypatch) + monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs) - monkeypatch.setattr( - tools.knowledge_base, - "get_retrievers", - lambda: { - "db-1": { - "name": "FAQ", - "retriever": object(), - "metadata": {"kb_type": "milvus"}, - } - }, + async def _fake_find_file_content( + db_id: str, + file_id: str, + patterns: list[str], + *, + use_regex: bool = False, + case_sensitive: bool = False, + max_windows: int = 5, + window_size: int = 80, + ): + assert db_id == "db-1" + assert file_id == "file-1" + assert patterns == ["token"] + assert use_regex is False + assert case_sensitive is False + assert max_windows == 5 + assert window_size == 80 + return { + "semantic": False, + "match_mode": "keyword", + "total_matches": 2, + "windows": [ + { + "start_line": 1, + "end_line": 3, + "matched_lines": [2], + "content": " 1\tintro\n 2\ttoken value\n 3\toutro", + } + ], + } + + monkeypatch.setattr(tools.knowledge_base, "find_file_content", _fake_find_file_content) + + runtime = SimpleNamespace(context=SimpleNamespace()) + result = await _run_find_kb_document( + resource_id="db-1", + file_id="file-1", + patterns=["token"], + runtime=runtime, ) - async def _fake_visible_kbs(runtime): - return [{"db_id": "db-1", "name": "FAQ"}] + assert result == { + "resource_id": "db-1", + "file_id": "file-1", + "semantic": False, + "match_mode": "keyword", + "total_matches": 2, + "windows": [ + { + "start_line": 1, + "end_line": 3, + "matched_lines": [2], + "content": " 1\tintro\n 2\ttoken value\n 3\toutro", + } + ], + } - async def _fake_open_file_content(db_id: str, file_id: str, offset: int = 0, limit: int = 800): + +@pytest.mark.asyncio +async def test_find_kb_document_rejects_dify(monkeypatch) -> None: + _patch_retrievers(monkeypatch, kb_type="dify") + monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs) + + runtime = SimpleNamespace(context=SimpleNamespace()) + result = await _run_find_kb_document( + resource_id="db-1", + file_id="file-1", + patterns=["token"], + runtime=runtime, + ) + + assert "Dify 知识库" in result + + +@pytest.mark.asyncio +async def test_open_kb_document_reads_markdown_content_by_default_window(monkeypatch) -> None: + lines = [f"line {index}" for index in range(1, 2001)] + + _patch_retrievers(monkeypatch) + monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs) + + async def _fake_open_file_content(db_id: str, file_id: str, offset: int = 0, limit: int = 1800): assert db_id == "db-1" assert file_id == "file-1" return _build_test_window("\n".join(lines), offset=offset, limit=limit) - monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs) monkeypatch.setattr(tools.knowledge_base, "open_file_content", _fake_open_file_content) runtime = SimpleNamespace(context=SimpleNamespace()) @@ -260,41 +303,28 @@ async def test_open_kb_document_reads_markdown_content_by_default_window(monkeyp assert result["resource_id"] == "db-1" assert result["file_id"] == "file-1" assert result["start_line"] == 1 - assert result["end_line"] == 800 - assert result["total_lines"] == 1000 - assert result["window_size"] == 800 + assert result["end_line"] == 1800 + assert result["total_lines"] == 2000 + assert result["window_size"] == 1800 assert result["has_more_before"] is False assert result["has_more_after"] is True - assert result["next_offset"] == 800 + assert result["next_offset"] == 1800 assert " 1\tline 1" in result["content"] - assert " 800\tline 800" in result["content"] + assert " 1800\tline 1800" in result["content"] @pytest.mark.asyncio async def test_open_kb_document_prefers_line_over_offset(monkeypatch) -> None: lines = [f"line {index}" for index in range(1, 1001)] - monkeypatch.setattr( - tools.knowledge_base, - "get_retrievers", - lambda: { - "db-1": { - "name": "FAQ", - "retriever": object(), - "metadata": {"kb_type": "milvus"}, - } - }, - ) + _patch_retrievers(monkeypatch) + monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs) - async def _fake_visible_kbs(runtime): - return [{"db_id": "db-1", "name": "FAQ"}] - - async def _fake_open_file_content(db_id: str, file_id: str, offset: int = 0, limit: int = 800): + async def _fake_open_file_content(db_id: str, file_id: str, offset: int = 0, limit: int = 1800): assert db_id == "db-1" assert file_id == "file-1" return _build_test_window("\n".join(lines), offset=offset, limit=limit) - monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs) monkeypatch.setattr(tools.knowledge_base, "open_file_content", _fake_open_file_content) runtime = SimpleNamespace(context=SimpleNamespace()) @@ -319,6 +349,7 @@ async def test_open_kb_document_prefers_line_over_offset(monkeypatch) -> None: @pytest.mark.asyncio async def test_open_kb_document_rejects_invisible_resource(monkeypatch) -> None: async def _fake_visible_kbs(runtime): + del runtime return [{"db_id": "db-2", "name": "FAQ"}] monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs) @@ -331,26 +362,13 @@ async def test_open_kb_document_rejects_invisible_resource(monkeypatch) -> None: @pytest.mark.asyncio async def test_open_kb_document_requires_markdown_content(monkeypatch) -> None: - monkeypatch.setattr( - tools.knowledge_base, - "get_retrievers", - lambda: { - "db-1": { - "name": "FAQ", - "retriever": object(), - "metadata": {"kb_type": "milvus"}, - } - }, - ) + _patch_retrievers(monkeypatch) + monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs) - async def _fake_visible_kbs(runtime): - return [{"db_id": "db-1", "name": "FAQ"}] - - async def _fake_open_file_content(db_id: str, file_id: str, offset: int = 0, limit: int = 800): + async def _fake_open_file_content(db_id: str, file_id: str, offset: int = 0, limit: int = 1800): del db_id, file_id, offset, limit raise Exception("文件 file-1 没有解析后的 Markdown 内容") - monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs) monkeypatch.setattr(tools.knowledge_base, "open_file_content", _fake_open_file_content) runtime = SimpleNamespace(context=SimpleNamespace())