diff --git a/backend/package/yuxi/agents/toolkits/kbs/tools.py b/backend/package/yuxi/agents/toolkits/kbs/tools.py index 6a40a551..f9364d08 100644 --- a/backend/package/yuxi/agents/toolkits/kbs/tools.py +++ b/backend/package/yuxi/agents/toolkits/kbs/tools.py @@ -192,24 +192,24 @@ def _find_query_target( kb_name: str, retrievers: dict[str, Any], visible_kbs: list[dict[str, Any]], -) -> tuple[dict[str, Any] | None, str | None]: +) -> tuple[dict[str, Any] | None, str | None, str | None]: if not visible_kbs: - return None, "无法获取当前会话可访问的知识库" + 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, f"知识库 '{kb_name}' 不存在或当前会话未启用" + return None, None, f"知识库 '{kb_name}' 不存在或当前会话未启用" if len(matched_kbs) > 1: - return None, f"知识库 '{kb_name}' 存在重名,请先调整名称后重试" + return None, None, f"知识库 '{kb_name}' 存在重名,请先调整名称后重试" target_db_id = str(matched_kbs[0].get("db_id") or "") target_info = retrievers.get(target_db_id) if target_info is None: - return None, f"知识库 '{kb_name}' 不存在" - return target_info, None + return None, None, f"知识库 '{kb_name}' 不存在" + return target_info, target_db_id, None -def _normalize_retrieval_result_metadata(result: Any) -> Any: +def _normalize_retrieval_result_metadata(result: Any, resource_id: str | None = None) -> Any: if not isinstance(result, list): return result @@ -224,6 +224,8 @@ def _normalize_retrieval_result_metadata(result: Any) -> Any: 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) @@ -260,7 +262,7 @@ async def query_kb(kb_name: str, query_text: str, file_name: str | None = None, visible_kbs = await _resolve_visible_knowledge_bases_for_query(runtime) - target_info, target_error = _find_query_target( + target_info, target_db_id, target_error = _find_query_target( kb_name=kb_name, retrievers=retrievers, visible_kbs=visible_kbs, @@ -279,7 +281,7 @@ 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) + return _normalize_retrieval_result_metadata(result, target_db_id) except Exception as e: logger.error(f"检索失败: {e}")