feat(kb): 统一知识库并支持 search/find/open

This commit is contained in:
Wenjie Zhang 2026-05-18 14:09:42 +08:00
parent 7be27ebc7e
commit 0faecbb886
8 changed files with 516 additions and 252 deletions

View File

@ -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):

View File

@ -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"]

View File

@ -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_idfile_id 对应检索结果 metadata.file_id
resource_id 是知识库资源 ID也就是 kb_idfile_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]

View File

@ -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

View File

@ -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()

View File

@ -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)

View File

@ -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="带行号的窗口内容")

View File

@ -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())