diff --git a/api/apps/restful_apis/dataset_api.py b/api/apps/restful_apis/dataset_api.py index b19d74dd36..bc740316e3 100644 --- a/api/apps/restful_apis/dataset_api.py +++ b/api/apps/restful_apis/dataset_api.py @@ -629,7 +629,13 @@ async def list_wiki_pages(tenant_id, dataset_id): ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -661,7 +667,13 @@ async def list_wiki_topics(tenant_id, dataset_id): ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -706,7 +718,13 @@ async def get_wiki_graph(tenant_id, dataset_id): ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -753,7 +771,13 @@ async def get_dataset_structure(tenant_id, dataset_id): ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -821,7 +845,13 @@ async def get_wiki_alteration(tenant_id, dataset_id): success, result = await dataset_api_service.get_structure_alteration(dataset_id, tenant_id, kind) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -845,7 +875,13 @@ async def clear_wiki(tenant_id, dataset_id): ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -872,7 +908,13 @@ async def get_wiki_page(tenant_id, dataset_id, page_type, slug): ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -908,7 +950,13 @@ async def get_skill_tree(tenant_id, dataset_id): ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -930,7 +978,13 @@ async def delete_all_skills(tenant_id, dataset_id): ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -953,7 +1007,13 @@ async def get_skill_page(tenant_id, dataset_id, skill_kwd): ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -975,7 +1035,13 @@ async def list_dataset_nav(tenant_id, dataset_id): ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -987,7 +1053,7 @@ async def list_dataset_nav(tenant_id, dataset_id): async def search_dataset_nav(tenant_id, dataset_id): """Unified navigation search across different knowledge layers. - GET /api/v1/datasets//navigation/search?q=&mode=&top_k=20 + GET /api/v1/datasets//navigation/search?q=&mode=&top_k=20&doc_ids=, Modes: - nav_doc: navigation tree document leaves (default) @@ -996,6 +1062,9 @@ async def search_dataset_nav(tenant_id, dataset_id): - chunk: raw document chunks (deduplicated by doc_id) - all: union of all modes above + ``doc_ids`` (comma-separated, optional) restricts the search to those + documents; omitted → all documents of the dataset. + Success: {"code": 0, "data": {"mode": , "total": , "items": [{"doc_id": str, "score": float}, ...]}} """ @@ -1010,6 +1079,7 @@ async def search_dataset_nav(tenant_id, dataset_id): top_k = max(1, int(top_k_raw)) except (ValueError, TypeError): return get_error_data_result(message="top_k must be a positive integer") + doc_scope = [d.strip() for d in request.args.get("doc_ids", "").split(",") if d.strip()] or None try: success, result = await dataset_api_service.search_dataset_layers( dataset_id, @@ -1017,10 +1087,17 @@ async def search_dataset_nav(tenant_id, dataset_id): q, mode, top_k=top_k, + doc_scope=doc_scope, ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -1043,7 +1120,13 @@ async def list_dataset_nav_children(tenant_id, dataset_id, name): ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -1065,7 +1148,13 @@ async def delete_dataset_nav(tenant_id, dataset_id): ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -1088,7 +1177,13 @@ async def delete_dataset_nav_node(tenant_id, dataset_id, name): ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -1120,7 +1215,13 @@ async def generate_dataset_nav(tenant_id, dataset_id): ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -1143,7 +1244,13 @@ async def delete_skill_page(tenant_id, dataset_id, skill_kwd): ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") @@ -1201,7 +1308,13 @@ async def update_wiki_page(tenant_id, dataset_id, page_type, slug): ) if success: return get_result(data=result) - return get_result(data=False, message=result, code=RetCode.AUTHENTICATION_ERROR) + if isinstance(result, dict): + msg = result.get("error", result) + code = result.get("code", RetCode.SERVER_ERROR) + else: + msg = result + code = RetCode.SERVER_ERROR + return get_result(data=False, message=msg, code=code) except Exception as e: logging.exception(e) return get_error_data_result(message="Internal server error") diff --git a/api/apps/services/dataset_api_service.py b/api/apps/services/dataset_api_service.py index ed80c4a430..5722bbc8e9 100644 --- a/api/apps/services/dataset_api_service.py +++ b/api/apps/services/dataset_api_service.py @@ -30,7 +30,7 @@ from api.db.services.tenant_model_service import TenantModelService from api.db.services.user_service import TenantService, UserService, UserTenantService from api.utils.api_utils import deep_merge, get_parser_config, remap_dictionary_keys, verify_embedding_availability from common import settings -from common.constants import PAGERANK_FLD, FileSource, LLMType, StatusEnum +from common.constants import PAGERANK_FLD, FileSource, LLMType, RetCode, StatusEnum from common.misc_utils import thread_pool_exec, thread_pool_exec_long_time from rag.advanced_rag.knowlege_compile.wiki import WIKI_PAGE_COMPILE_KWD @@ -3630,6 +3630,7 @@ async def search_dataset_layers( mode: str, *, top_k: int | None = None, + doc_scope: list[str] | None = None, ) -> tuple[bool, dict]: """Unified search across different knowledge layers of a dataset. @@ -3640,15 +3641,20 @@ async def search_dataset_layers( - nav_cluster: navigation tree cluster nodes - navigation_tree: tree-structured BFS beam descent - all: union of all modes, deduplicated by doc_id with best score + doc_scope: Optional set of documents to restrict the search to. None or + empty means all documents of the dataset. Forwarded to every mode: + nav modes filter the compiled nav rows by doc, ``chunk`` restricts + the retriever via ``doc_ids``. Applied query-time so scoped rows + are never dropped by the ``top_k`` truncation. Items are shaped as ``{"doc_id": str, "score": float}``. """ from rag.advanced_rag.knowlege_compile.dataset_nav import search_dataset_nav if not KnowledgebaseService.accessible(dataset_id, tenant_id): - return False, "no authorization" + return False, {"error": "no authorization", "code": RetCode.PERMISSION_ERROR} if mode not in _LAYERS_HANDLERS: - return False, f"unknown mode: {mode}, expected one of {list(_LAYERS_HANDLERS.keys())}" + return False, {"error": f"unknown mode: {mode}, expected one of {list(_LAYERS_HANDLERS.keys())}", "code": RetCode.ARGUMENT_ERROR} _, kb = KnowledgebaseService.get_by_id(dataset_id) try: @@ -3670,21 +3676,27 @@ async def search_dataset_layers( logging.exception("Full traceback for LLMBundle(EMBEDDING) failure") embd_mdl = None + logging.debug( + "search_dataset_layers: dispatching scoped mode=%s for dataset=%s, scoped_docs=%d", + mode, + dataset_id, + len([d for d in (doc_scope or []) if str(d).strip()]), + ) if mode == "nav_doc": - return await _search_layers_nav_docs(tenant_id, dataset_id, query, top_k, embd_mdl, search_dataset_nav) + return await _search_layers_nav_docs(tenant_id, dataset_id, query, top_k, embd_mdl, search_dataset_nav, doc_scope=doc_scope) elif mode == "nav_cluster": - return await _search_layers_nav_clusters(tenant_id, dataset_id, query, top_k, embd_mdl, search_dataset_nav) + return await _search_layers_nav_clusters(tenant_id, dataset_id, query, top_k, embd_mdl, search_dataset_nav, doc_scope=doc_scope) elif mode == "navigation_tree": - return await _search_layers_navigation_tree(tenant_id, dataset_id, query, top_k, embd_mdl, search_dataset_nav) + return await _search_layers_navigation_tree(tenant_id, dataset_id, query, top_k, embd_mdl, doc_scope=doc_scope) elif mode == "chunk": - return await _search_layers_chunks(tenant_id, dataset_id, query, top_k, embd_mdl, kb) + return await _search_layers_chunks(tenant_id, dataset_id, query, top_k, embd_mdl, kb, doc_scope=doc_scope) elif mode == "all": - return await _search_layers_all(tenant_id, dataset_id, query, top_k, embd_mdl, kb, search_dataset_nav) + return await _search_layers_all(tenant_id, dataset_id, query, top_k, embd_mdl, kb, search_dataset_nav, doc_scope=doc_scope) else: - return False, f"unknown mode: {mode}" + return False, {"error": f"unknown mode: {mode}", "code": RetCode.ARGUMENT_ERROR} -async def _search_layers_nav_docs(tenant_id, dataset_id, query, top_k, embd_mdl, search_fn): +async def _search_layers_nav_docs(tenant_id, dataset_id, query, top_k, embd_mdl, search_fn, *, doc_scope=None): items = await _nav_search_result( tenant_id, dataset_id, @@ -3693,11 +3705,12 @@ async def _search_layers_nav_docs(tenant_id, dataset_id, query, top_k, embd_mdl, embd_mdl, search_fn, type_kwd="nav_doc", + doc_scope=doc_scope, ) return True, {"mode": "nav_doc", "total": len(items), "items": items} -async def _search_layers_nav_clusters(tenant_id, dataset_id, query, top_k, embd_mdl, search_fn): +async def _search_layers_nav_clusters(tenant_id, dataset_id, query, top_k, embd_mdl, search_fn, *, doc_scope=None): items = await _nav_search_result( tenant_id, dataset_id, @@ -3706,11 +3719,12 @@ async def _search_layers_nav_clusters(tenant_id, dataset_id, query, top_k, embd_ embd_mdl, search_fn, type_kwd="nav_cluster", + doc_scope=doc_scope, ) return True, {"mode": "nav_cluster", "total": len(items), "items": items} -async def _search_layers_navigation_tree(tenant_id, dataset_id, query, top_k, embd_mdl, search_fn): +async def _search_layers_navigation_tree(tenant_id, dataset_id, query, top_k, embd_mdl, *, doc_scope=None): from rag.advanced_rag.knowlege_compile.dataset_nav import search_nav_tree_descent items = await search_nav_tree_descent( @@ -3719,32 +3733,48 @@ async def _search_layers_navigation_tree(tenant_id, dataset_id, query, top_k, em query, embd_mdl, top_k=top_k, + doc_scope=doc_scope, ) return True, {"mode": "navigation_tree", "total": len(items), "items": items} async def _nav_search_result(tenant_id, dataset_id, query, top_k, embd_mdl, search_fn, **kwargs): + doc_scope = kwargs.pop("doc_scope", None) results = await search_fn( tenant_id, dataset_id, query, embd_mdl=embd_mdl, top_k=top_k, + doc_scope=doc_scope, **kwargs, ) + scope_set = {str(d).strip() for d in doc_scope if str(d).strip()} if doc_scope else None items: list[dict] = [] for r in results: - doc_id = (r.get("doc_id") or "").strip() if isinstance(r.get("doc_id"), str) else (r.get("doc_ids") or [None])[0] + raw_doc_id = r.get("doc_id") + if isinstance(raw_doc_id, str) and raw_doc_id.strip(): + doc_id = raw_doc_id.strip() + else: + # Cluster row: pick the first doc_id that's in scope (if a scope + # is set) so we never leak an out-of-scope doc_id into the results. + doc_ids = r.get("doc_ids") or [] + if scope_set: + doc_id = next((str(d).strip() for d in doc_ids if str(d).strip() in scope_set), "") + else: + doc_id = str(doc_ids[0]).strip() if doc_ids else "" + if not doc_id: + continue items.append( { - "doc_id": str(doc_id) if doc_id else "", + "doc_id": doc_id, "score": round(float(r.get("score", 0.0)), 4), } ) return items -async def _search_layers_chunks(tenant_id, dataset_id, query, top_k, embd_mdl, kb): +async def _search_layers_chunks(tenant_id, dataset_id, query, top_k, embd_mdl, kb, *, doc_scope=None): from common import settings tenant_ids = [tenant_id] @@ -3752,6 +3782,8 @@ async def _search_layers_chunks(tenant_id, dataset_id, query, top_k, embd_mdl, k kwargs = {} if top_k is not None: kwargs["top"] = top_k + if doc_scope: + kwargs["doc_ids"] = [str(d) for d in doc_scope if str(d).strip()] fetch_k = max(top_k, 10) * 3 if top_k is not None else 1024 try: @@ -3767,7 +3799,7 @@ async def _search_layers_chunks(tenant_id, dataset_id, query, top_k, embd_mdl, k **kwargs, ) except Exception: - return False, "chunk retrieval failed" + return False, {"error": "chunk retrieval failed", "code": RetCode.SERVER_ERROR} doc_scores: dict[str, float] = {} for c in ranks.get("chunks", []): @@ -3787,15 +3819,15 @@ async def _search_layers_chunks(tenant_id, dataset_id, query, top_k, embd_mdl, k return True, {"mode": "chunk", "total": len(items), "items": items} -async def _search_layers_all(tenant_id, dataset_id, query, top_k, embd_mdl, kb, search_fn): +async def _search_layers_all(tenant_id, dataset_id, query, top_k, embd_mdl, kb, search_fn, *, doc_scope=None): """Run all modes and return the union of doc_ids, with best score per doc.""" import asyncio as _asyncio result_lists = await _asyncio.gather( - _search_layers_nav_docs(tenant_id, dataset_id, query, top_k, embd_mdl, search_fn), - _search_layers_nav_clusters(tenant_id, dataset_id, query, top_k, embd_mdl, search_fn), - _search_layers_navigation_tree(tenant_id, dataset_id, query, top_k, embd_mdl, search_fn), - _search_layers_chunks(tenant_id, dataset_id, query, top_k, embd_mdl, kb), + _search_layers_nav_docs(tenant_id, dataset_id, query, top_k, embd_mdl, search_fn, doc_scope=doc_scope), + _search_layers_nav_clusters(tenant_id, dataset_id, query, top_k, embd_mdl, search_fn, doc_scope=doc_scope), + _search_layers_navigation_tree(tenant_id, dataset_id, query, top_k, embd_mdl, doc_scope=doc_scope), + _search_layers_chunks(tenant_id, dataset_id, query, top_k, embd_mdl, kb, doc_scope=doc_scope), return_exceptions=True, ) diff --git a/rag/advanced_rag/harness/config.py b/rag/advanced_rag/harness/config.py index 962ab2918e..5dec62b581 100644 --- a/rag/advanced_rag/harness/config.py +++ b/rag/advanced_rag/harness/config.py @@ -56,7 +56,7 @@ THINKING_MODES: dict[str, ExecutionStrategy] = { "web_search", "bm25_search", "ontology_navigate", - "dataset_navigation_by_tree", + "dataset_navigation_search", "graph_explore", "inspector_open_context", "inspector_compare", @@ -86,7 +86,7 @@ THINKING_MODES: dict[str, ExecutionStrategy] = { "web_search", "structured_query", "ontology_navigate", - "dataset_navigation_by_tree", + "dataset_navigation_search", "mindmap_navigate", "graph_explore", "wiki_query", diff --git a/rag/advanced_rag/harness/pipeline.py b/rag/advanced_rag/harness/pipeline.py index 302a95e312..3da9f154c0 100644 --- a/rag/advanced_rag/harness/pipeline.py +++ b/rag/advanced_rag/harness/pipeline.py @@ -1,19 +1,19 @@ """Pipeline — unified tool execution dispatcher.""" -import time import logging +import time from typing import Any -from rag.advanced_rag.harness.types import ToolResult from rag.advanced_rag.harness.tools.registry import TOOL_REGISTRY +from rag.advanced_rag.harness.types import ToolResult _LOG = logging.getLogger(__name__) # Tools that retrieve *within* a set of documents. When a routing tool -# (``dataset_navigation_by_tree``) has produced a relevant-document set, these +# (``dataset_navigation_search``) has produced a relevant-document set, these # inherit it as their ``doc_scope`` unless the caller passed one explicitly, so # a follow-up search stays within the routed docs instead of re-scanning the KB. -_DOC_SCOPE_CONSUMERS = {"ontology_navigate", "mindmap_navigate", "graph_explore", "hybrid_search", "vector_search", "bm25_search", "structured_query", "dataset_navigation_by_tree"} +_DOC_SCOPE_CONSUMERS = {"ontology_navigate", "mindmap_navigate", "graph_explore", "hybrid_search", "vector_search", "bm25_search", "structured_query", "dataset_navigation_search"} class Pipeline: @@ -43,7 +43,7 @@ class Pipeline: return ToolResult(chunks=[], metadata={}, error=f"Tool {tool_name} has no executor") # Downstream scoping: a within-document tool inherits the doc IDs a prior - # router (dataset_navigation_by_tree) produced, unless the caller passed + # router (dataset_navigation_search) produced, unless the caller passed # an explicit doc_scope. if tool_name in _DOC_SCOPE_CONSUMERS and self._routed_docs and not kwargs.get("doc_scope"): kwargs["doc_scope"] = list(self._routed_docs) @@ -54,7 +54,7 @@ class Pipeline: elapsed = time.time() - start self.trace.append({"tool": tool_name, "args": kwargs, "elapsed": elapsed, "success": True}) result = self._normalize(raw) - # A routing tool (e.g. dataset_navigation_by_tree) yields the relevant + # A routing tool (e.g. dataset_navigation_search) yields the relevant # document IDs; remember them so the scope-consuming tools above can # inherit them on later turns. if result.docs: @@ -130,7 +130,7 @@ class Pipeline: ) if isinstance(raw, list): # A list of doc-id strings is a document-routing result (e.g. - # dataset_navigation_by_tree); a list of dicts is chunks. + # dataset_navigation_search); a list of dicts is chunks. if raw and all(isinstance(x, str) for x in raw): return ToolResult(docs=list(raw), metadata={}) return ToolResult(chunks=raw, metadata={}) diff --git a/rag/advanced_rag/harness/prompts/research_agent_prompt.py b/rag/advanced_rag/harness/prompts/research_agent_prompt.py index 6bd81eca3c..74868e29c4 100644 --- a/rag/advanced_rag/harness/prompts/research_agent_prompt.py +++ b/rag/advanced_rag/harness/prompts/research_agent_prompt.py @@ -16,7 +16,7 @@ Current phase: {phase} Phase hint: {phase_hint} Rules: -1. Go coarse-to-fine. First narrow the corpus with navigation tools (dataset_navigation_by_tree, +1. Go coarse-to-fine. First narrow the corpus with navigation tools (dataset_navigation_search, then ontology_navigate / mindmap_navigate). 2. After a navigation tool returns passages, judge whether they already answer the task. If they do, call generate_report immediately — do NOT search further. diff --git a/rag/advanced_rag/harness/tools/__init__.py b/rag/advanced_rag/harness/tools/__init__.py index 16050e654e..c7402d25c3 100644 --- a/rag/advanced_rag/harness/tools/__init__.py +++ b/rag/advanced_rag/harness/tools/__init__.py @@ -1,11 +1,10 @@ """Tool system: register all tools with the registry on import.""" -from rag.advanced_rag.harness.tools.registry import register_tool, _search_schema, _navigate_schema, _inspector_schema +from rag.advanced_rag.harness.tools.registry import _inspector_schema, _navigate_schema, _search_schema, register_tool # Register tools - # Search tools -from rag.advanced_rag.harness.tools.search import hybrid_search, vector_search, bm25_search, web_search, structured_query +from rag.advanced_rag.harness.tools.search import bm25_search, hybrid_search, structured_query, vector_search, web_search register_tool("hybrid_search", _search_schema("hybrid_search", "Embedding + Keywords search"), hybrid_search) register_tool("vector_search", _search_schema("vector_search", "Embedding search"), vector_search) @@ -14,7 +13,7 @@ register_tool("web_search", _search_schema("web_search", "Internet search"), web register_tool("structured_query", _search_schema("structured_query", "SQL search"), structured_query) # Navigation tools (require compilation) -from rag.advanced_rag.harness.tools.navigation import ontology_navigate, dataset_navigation_by_tree, mindmap_navigate +from rag.advanced_rag.harness.tools.navigation import dataset_navigation_search, mindmap_navigate, ontology_navigate # ontology_navigate covers both the tree/TOC outline and the page index. register_tool( @@ -28,12 +27,12 @@ register_tool( "mindmap_navigate", _navigate_schema("mindmap_navigate", "Get question-related chunks from the document's mindmap"), mindmap_navigate, requires_compilation=True, compilation_type="mindmap" ) register_tool( - "dataset_navigation_by_tree", + "dataset_navigation_search", _navigate_schema( - "dataset_navigation_by_tree", - "Find the documents most relevant to the question via the dataset map; returns their document IDs (later navigation/exploration tools are automatically scoped to them).", + "dataset_navigation_search", + "Find the documents most relevant to the question by hybrid-searching the dataset's navigation-tree document leaves (nav_doc layer) — only documents that have a compiled nav-tree node are reachable. Returns the most relevant document IDs.", ), - dataset_navigation_by_tree, + dataset_navigation_search, requires_compilation=True, compilation_type="tree", ) @@ -45,7 +44,7 @@ register_tool("graph_explore", _search_schema("graph_explore", "Knowledge graph register_tool("wiki_query", _search_schema("wiki_query", "Wiki search"), wiki_query, requires_compilation=True, compilation_type="wiki") # Inspector tools -from rag.advanced_rag.harness.tools.inspector import open_context, compare_sources, grep_within, request_adjacent +from rag.advanced_rag.harness.tools.inspector import compare_sources, grep_within, open_context, request_adjacent register_tool( "inspector_open_context", diff --git a/rag/advanced_rag/harness/tools/gating.py b/rag/advanced_rag/harness/tools/gating.py index 3b5abbcaa3..7b8c1afff2 100644 --- a/rag/advanced_rag/harness/tools/gating.py +++ b/rag/advanced_rag/harness/tools/gating.py @@ -1,8 +1,7 @@ """Tool selection gating: phase-based filtering and fallback chain.""" -from rag.advanced_rag.harness.types import OrchestratorContext from rag.advanced_rag.harness.tools.registry import TOOL_REGISTRY - +from rag.advanced_rag.harness.types import OrchestratorContext # Search phase definitions @@ -10,7 +9,7 @@ SEARCH_PHASES = { "locate": { "goal": "Locate documents or regions that may contain the answer.", "tools_priority": [ - "dataset_navigation_by_tree", + "dataset_navigation_search", "ontology_navigate", "mindmap_navigate", "hybrid_search", @@ -77,7 +76,7 @@ def tool_fits_context(tool_name: str, context: OrchestratorContext, has_routed_s return False if tool_name in {"ontology_navigate", "mindmap_navigate"} and not has_routed_scope: return False - if tool_name == "dataset_navigation_by_tree" and not context.current_claim: + if tool_name == "dataset_navigation_search" and not context.current_claim: return False if tool_name == "graph_explore" and not context.last_entity: return False diff --git a/rag/advanced_rag/harness/tools/navigation.py b/rag/advanced_rag/harness/tools/navigation.py index f56915930e..7dcbec2942 100644 --- a/rag/advanced_rag/harness/tools/navigation.py +++ b/rag/advanced_rag/harness/tools/navigation.py @@ -629,6 +629,72 @@ async def dataset_navigation_by_tree(tools, topic: str, keywords: str = "", doc_ return routed[:_NAV_MAX_DOCS] +# ── Dataset document search (hybrid, no LLM) ──────────────────────────────── + +_NAV_SEARCH_MAX_DOCS = 12 # documents the hybrid search routes to +_NAV_MIN_DOC_SCORE = 0.2 # drop docs below this score + + +async def dataset_navigation_search(tools, topic: str, keywords: str = "", doc_scope: list[str] | None = None) -> list[str]: + """Return the ``doc_id``s most relevant to the question / keywords by + searching the dataset's navigation-tree document leaves (``nav_doc`` layer). + + Runs ``search_dataset_layers`` with ``mode="nav_doc"``: a direct hybrid + search over the nav-tree doc leaves, so it only sees documents that have a + compiled nav-tree node. Faster, no LLM cost, but less precise for ambiguous + queries. + + Returns the routed ``doc_id`` list (capped at ``_NAV_SEARCH_MAX_DOCS``), or + ``[]`` when no question/keywords are given or the search returns nothing. + This function only routes — it does not retrieve. + """ + query = " ".join(part for part in ((topic or "").strip(), (keywords or "").strip()) if part).strip() + if not query: + return [] + if hasattr(tools, "scoped_doc_ids"): + doc_scope = tools.scoped_doc_ids(doc_scope) + + _LOG.info('[Dataset navigation search] Nav-tree doc search for "%s"', query) + + from api.apps.services import dataset_api_service + + kbs = getattr(tools, "kbs", []) or [] + allowed_docs = set(doc_scope or []) + + candidates: dict[str, float] = {} + for kb in kbs: + # ``doc_scope`` is forwarded query-time (search_dataset_layers applies + # it as a store filter on every mode), so the top_k truncation never + # drops scoped docs — no enlarged pool is needed. + try: + ok, result = await dataset_api_service.search_dataset_layers( + kb.id, + kb.tenant_id, + query, + "nav_doc", + top_k=_NAV_SEARCH_MAX_DOCS, + doc_scope=list(allowed_docs) or None, + ) + except Exception: + _LOG.exception("[Dataset navigation search] search_dataset_layers failed for kb=%s", kb.id) + continue + if not ok or not isinstance(result, dict): + continue + for item in result.get("items", []): + score = float(item.get("score", 0.0)) + if score < _NAV_MIN_DOC_SCORE: + continue + did = str(item.get("doc_id") or "").strip() + if not did: + continue + candidates[did] = max(candidates.get(did, float("-inf")), score) + + routed = [did for did, _ in sorted(candidates.items(), key=lambda pair: pair[1], reverse=True)[:_NAV_SEARCH_MAX_DOCS]] + + _LOG.info("[Dataset navigation search] Routed to %d document(s) (min_score=%.1f).", len(routed), _NAV_MIN_DOC_SCORE) + return routed[:_NAV_SEARCH_MAX_DOCS] + + # ── Knowledge-graph exploration ───────────────────────────────────────────── # # Unlike catalog/mindmap (which read the merged "graph" JSON of one doc), the KG diff --git a/rag/advanced_rag/harness/types.py b/rag/advanced_rag/harness/types.py index 2245f8865e..bc6dfa27c5 100644 --- a/rag/advanced_rag/harness/types.py +++ b/rag/advanced_rag/harness/types.py @@ -5,7 +5,6 @@ from __future__ import annotations from dataclasses import dataclass, field from typing import Literal - # ═══════════════════════════════════════════════════════════════ # Route # ═══════════════════════════════════════════════════════════════ @@ -153,7 +152,7 @@ class SufficiencyVerdict: @dataclass class ToolResult: chunks: list[dict] = field(default_factory=list) - # Doc-id list from routing tools (e.g. dataset_navigation_by_tree) that + # Doc-id list from routing tools (e.g. dataset_navigation_search) that # narrow the corpus to the relevant documents instead of returning chunks. docs: list[str] | None = None metadata: dict = field(default_factory=dict) diff --git a/rag/advanced_rag/knowlege_compile/dataset_nav.py b/rag/advanced_rag/knowlege_compile/dataset_nav.py index e1d19f78b7..2ec4d74356 100644 --- a/rag/advanced_rag/knowlege_compile/dataset_nav.py +++ b/rag/advanced_rag/knowlege_compile/dataset_nav.py @@ -599,6 +599,24 @@ def _matches_condition(row: dict, condition: dict) -> bool: return True +def _in_nav_scope(row: dict, allowed_docs: set[str] | None) -> bool: + """Whether a nav row belongs to the given doc scope. + + A ``nav_doc`` leaf belongs when its ``doc_id`` is in the set; a + ``nav_cluster`` row belongs when it covers at least one scoped document + (its ``doc_ids_kwd`` intersects the set). An empty/None scope allows + everything, preserving the existing unscoped behavior. + """ + if not allowed_docs: + return True + # A cluster row carries doc_id == kb_id (not a document) but lists the + # documents it covers in doc_ids_kwd, so a non-empty doc_ids_kwd is the + # reliable discriminator. A nav_doc leaf never sets doc_ids_kwd. + if row.get("doc_ids_kwd"): + return bool(set(_as_str_list(row.get("doc_ids_kwd"))) & allowed_docs) + return str(row.get("doc_id") or "").strip() in allowed_docs + + # --------------------------------------------------------------------------- # Incremental clustering core # --------------------------------------------------------------------------- @@ -1223,6 +1241,7 @@ async def search_dataset_nav( *, type_kwd: str = "", compile_kwd: str = _COMPILE_KWD, + doc_scope: list[str] | None = None, ) -> list[dict]: """Find the nav-tree nodes most relevant to ``query`` for one KB. @@ -1235,6 +1254,11 @@ async def search_dataset_nav( leaves, ``"nav_cluster"`` to restrict to clusters, or ``""`` for all. compile_kwd: Which compile partition to search within (default ``_COMPILE_KWD`` = ``"dataset_nav"``). + doc_scope: Optional set of documents to restrict results to. Applied + as a query-time store filter on the type-specific key (``doc_id`` + for ``nav_doc`` leaves, ``doc_ids_kwd`` for ``nav_cluster`` rows) + and re-checked in memory before the ``top_k`` truncation, so + scoped rows are never dropped by unscoped ones. Returns items shaped as:: @@ -1250,10 +1274,26 @@ async def search_dataset_nav( query = (query or "").strip() if not query: return [] + allowed_docs = {str(d).strip() for d in (doc_scope or []) if str(d).strip()} + logging.debug( + "search_dataset_nav: flat-search scope normalized for kb=%s, scoped_docs=%d", + kb_id, + len(allowed_docs), + ) condition: dict = {"compile_kwd": [compile_kwd]} if type_kwd: condition["type_kwd"] = type_kwd + if allowed_docs: + # Type-specific scope filter: nav_doc leaves match on doc_id, cluster + # rows on doc_ids_kwd. Only set when type_kwd pins a single type — + # a mixed query can't express both under `_matches_condition`'s + # AND-across-fields semantics, so mixed scoping relies on the + # in-memory `_in_nav_scope` check below instead. + if type_kwd == "nav_doc": + condition["doc_id"] = sorted(allowed_docs) + elif type_kwd == "nav_cluster": + condition["doc_ids_kwd"] = sorted(allowed_docs) # name -> [row, fused_score]; `name` uniquely identifies a nav node, so it # is the dedup key when the dense and lexical legs return the same node. fused: dict[str, list] = {} @@ -1280,8 +1320,23 @@ async def search_dataset_nav( # ── Lexical leg: engine BM25 over the tokenized fields ── text_w = 1.0 - dense_w + text_filter = None + if allowed_docs: + if type_kwd == "nav_doc": + text_filter = {"doc_id": sorted(allowed_docs)} + elif type_kwd == "nav_cluster": + text_filter = {"doc_ids_kwd": sorted(allowed_docs)} try: - text_rows = await _store_text_search(tenant_id, kb_id, query, _NAV_SEARCH_FIELDS, limit=max((top_k or 0) * 3, 20) if top_k else 10000, compile_kwd=compile_kwd, type_kwd=type_kwd) + text_rows = await _store_text_search( + tenant_id, + kb_id, + query, + _NAV_SEARCH_FIELDS, + limit=max((top_k or 0) * 3, 20) if top_k else 10000, + compile_kwd=compile_kwd, + type_kwd=type_kwd, + extra_filter=text_filter, + ) except Exception: logging.exception("search_dataset_nav: text search failed for kb=%s", kb_id) text_rows = [] @@ -1294,7 +1349,7 @@ async def search_dataset_nav( continue fused.setdefault(rk, [r, 0.0])[1] += text_w * ts - rows_with_scores = [(r, s) for r, s in fused.values() if s > 0] + rows_with_scores = [(r, s) for r, s in fused.values() if s > 0 and _in_nav_scope(r, allowed_docs)] rows_with_scores.sort(key=lambda item: item[1], reverse=True) if top_k is not None: rows_with_scores = rows_with_scores[:top_k] @@ -1310,6 +1365,16 @@ async def search_dataset_nav( if typ == "nav_cluster": doc_id = None doc_ids = _as_str_list(r.get("doc_ids_kwd")) + if allowed_docs: + # A scoped search must only surface in-scope documents: cut the + # cluster's coverage list down to the scoped set so no out-of- + # scope doc appears under a cluster that merely overlaps the + # scope. + doc_ids = [d for d in doc_ids if d in allowed_docs] + # Once the coverage list is scope-filtered, the reported count must + # match the returned doc_ids (the raw doc_count_int would otherwise + # include documents excluded by the scope). + scoped_cluster = bool(allowed_docs) else: # Leaf: ``name`` == the document id (see ``_make_nav_doc_row``). doc_id = r.get("doc_id") or name @@ -1326,7 +1391,7 @@ async def search_dataset_nav( "graph_content": payload.get("graph_content") or "", "doc_title": payload.get("doc_title") or "", "source_type": payload.get("source_type") or "", - "doc_count": int(r.get("doc_count_int") or len(doc_ids) or 0), + "doc_count": len(doc_ids) if (typ == "nav_cluster" and scoped_cluster) else int(r.get("doc_count_int") or len(doc_ids) or 0), "score": float(score or 0.0), } ) @@ -1339,6 +1404,7 @@ async def search_nav_tree_descent( query: str, embd_mdl, top_k: int | None = None, + doc_scope: list[str] | None = None, ) -> list[dict]: """Tree-structured hybrid search: descend from root into the most relevant branches. @@ -1354,18 +1420,38 @@ async def search_nav_tree_descent( The search uses BFS with beam pruning — at each depth, only the *beam_width* most similar clusters are expanded further. + ``doc_scope`` restricts the search to a specific set of documents: both + the KNN legs (exact if the engine misses the terms filter, ``_store_knn`` + falls back to a full in-memory match) and the text legs apply the scope + as a query-time filter, so the top-N truncation never discards scoped + documents that rank behind unscoped ones. + Returns items shaped as ``{"doc_id": str, "score": float}``. """ query = (query or "").strip() if not query: return [] + allowed_docs = {str(d).strip() for d in (doc_scope or []) if str(d).strip()} + logging.debug( + "search_nav_tree_descent: tree-search scope normalized for kb=%s, scoped_docs=%d", + kb_id, + len(allowed_docs), + ) if embd_mdl is None: logging.warning( "search_nav_tree_descent: embd_mdl is None — falling back to text-only flat search for kb=%s query=%.80s", kb_id, query, ) - raw = await search_dataset_nav(tenant_id, kb_id, query, embd_mdl=None, top_k=top_k, type_kwd="nav_doc") + raw = await search_dataset_nav( + tenant_id, + kb_id, + query, + embd_mdl=None, + top_k=top_k, + type_kwd="nav_doc", + doc_scope=list(allowed_docs) or None, + ) return [{"doc_id": r.get("doc_id", ""), "score": r.get("score", 0.0)} for r in raw if r.get("doc_id")] vec = await _embed(embd_mdl, query) @@ -1406,18 +1492,25 @@ async def search_nav_tree_descent( "type_kwd": ["nav_cluster"], "depth_int": [0], } + if allowed_docs: + root_cond["doc_ids_kwd"] = sorted(allowed_docs) roots_knn = await _store_knn(tenant_id, kb_id, vec, vec_dim, root_cond, top_k=beam_width * 3) + roots_knn = [r for r in roots_knn if _in_nav_scope(r, allowed_docs)] - # If no root cluster (depth=0) exists — the dataset may have been - # compiled without one — scan all nav_clusters to find the lowest - # available depth and start beam search there. + # If no root cluster (depth=0) exists — or none overlaps ``doc_scope`` — + # the dataset may have been compiled without a root, so scan all + # nav_clusters to find the lowest available depth and start beam search + # there. if not roots_knn: all_cond = { "kb_id": [kb_id], "compile_kwd": [_COMPILE_KWD], "type_kwd": ["nav_cluster"], } + if allowed_docs: + all_cond["doc_ids_kwd"] = sorted(allowed_docs) all_clusters = await _store_search(tenant_id, kb_id, all_cond, fields, limit=10000) + all_clusters = [r for r in all_clusters if _in_nav_scope(r, allowed_docs)] if not all_clusters: return [] @@ -1460,15 +1553,52 @@ async def search_nav_tree_descent( "compile_kwd": [_COMPILE_KWD], "parent_kwd": [node_name], } - children_knn = await _store_knn(tenant_id, kb_id, vec, vec_dim, child_cond, top_k=beam_width * 3) - children_text = await _store_text_search( - tenant_id, - kb_id, - query, - fields, - limit=beam_width * 3, - extra_filter={"parent_kwd": [node_name], "kb_id": [kb_id]}, - ) + if allowed_docs: + # Children are a mix of nav_doc leaves and nav_cluster rows, + # which scope on different keys, so fetch each type with its + # own store filter and merge. Rows are re-checked with + # `_in_nav_scope` so a doc filter the engine ignores can't + # let an unscoped row crowd out a scoped one in the fuse. + child_doc_cond = { + **child_cond, + "type_kwd": ["nav_doc"], + "doc_id": sorted(allowed_docs), + } + child_cluster_cond = { + **child_cond, + "type_kwd": ["nav_cluster"], + "doc_ids_kwd": sorted(allowed_docs), + } + children_knn = await _store_knn(tenant_id, kb_id, vec, vec_dim, child_doc_cond, top_k=beam_width * 3) + children_knn += await _store_knn(tenant_id, kb_id, vec, vec_dim, child_cluster_cond, top_k=beam_width * 3) + children_knn = [r for r in children_knn if _in_nav_scope(r, allowed_docs)] + children_text = await _store_text_search( + tenant_id, + kb_id, + query, + fields, + limit=beam_width * 3, + extra_filter={"parent_kwd": [node_name], "kb_id": [kb_id], "doc_id": sorted(allowed_docs)}, + ) + children_text += await _store_text_search( + tenant_id, + kb_id, + query, + fields, + limit=beam_width * 3, + extra_filter={"parent_kwd": [node_name], "kb_id": [kb_id], "doc_ids_kwd": sorted(allowed_docs)}, + ) + children_text = [r for r in children_text if _in_nav_scope(r, allowed_docs)] + else: + children_knn = await _store_knn(tenant_id, kb_id, vec, vec_dim, child_cond, top_k=beam_width * 3) + children_text = await _store_text_search( + tenant_id, + kb_id, + query, + fields, + limit=beam_width * 3, + extra_filter={"parent_kwd": [node_name], "kb_id": [kb_id]}, + ) candidates = _hybrid_fuse(vec, vf, query, children_knn, children_text, dense_w, beam_width) for c in candidates: @@ -1476,9 +1606,12 @@ async def search_nav_tree_descent( break if c.get("type_kwd") == "nav_doc": doc_id = (c.get("doc_id") or "").strip() - if doc_id and doc_id not in seen_docs: - seen_docs.add(doc_id) - collected.append({"doc_id": doc_id, "score": round(c["_score"] or parent_score, 4)}) + if not doc_id or doc_id in seen_docs: + continue + if allowed_docs and doc_id not in allowed_docs: + continue + seen_docs.add(doc_id) + collected.append({"doc_id": doc_id, "score": round(c["_score"] or parent_score, 4)}) else: next_level.append(c) @@ -1496,11 +1629,14 @@ async def search_nav_tree_descent( for node in current_level: for did in node.get("doc_ids_kwd") or []: did_str = str(did).strip() - if did_str and did_str not in seen_docs: - seen_docs.add(did_str) - collected.append({"doc_id": did_str, "score": round(node.get("_score", 0.0), 4)}) - if top_k is not None and len(collected) >= top_k: - break + if not did_str or did_str in seen_docs: + continue + if allowed_docs and did_str not in allowed_docs: + continue + seen_docs.add(did_str) + collected.append({"doc_id": did_str, "score": round(node.get("_score", 0.0), 4)}) + if top_k is not None and len(collected) >= top_k: + break if top_k is not None and len(collected) >= top_k: break diff --git a/test/unit_test/api/apps/services/test_dataset_api_service_list_datasets.py b/test/unit_test/api/apps/services/test_dataset_api_service_list_datasets.py index c350893c22..1430200945 100644 --- a/test/unit_test/api/apps/services/test_dataset_api_service_list_datasets.py +++ b/test/unit_test/api/apps/services/test_dataset_api_service_list_datasets.py @@ -32,7 +32,6 @@ from unittest.mock import MagicMock import pytest - pytestmark = pytest.mark.p2 @@ -97,6 +96,7 @@ def _load_list_datasets_module(monkeypatch, *, kbs, parsing_status_by_kb): FileSource=SimpleNamespace(KNOWLEDGEBASE="knowledgebase"), PipelineTaskType=SimpleNamespace(), StatusEnum=SimpleNamespace(), + RetCode=SimpleNamespace(), ModelTypeBinary=_StubModelTypeBinary, ) _stub( @@ -175,6 +175,8 @@ def _load_list_datasets_module(monkeypatch, *, kbs, parsing_status_by_kb): thread_pool_exec=MagicMock(), thread_pool_exec_long_time=MagicMock(), ) + _stub(monkeypatch, "rag.advanced_rag", __path__=[]) + _stub(monkeypatch, "rag.advanced_rag.knowlege_compile", __path__=[]) _stub( monkeypatch, "rag.advanced_rag.knowlege_compile.wiki", diff --git a/test/unit_test/api/apps/services/test_delete_datasets.py b/test/unit_test/api/apps/services/test_delete_datasets.py index 1869f7bb3e..873791779f 100644 --- a/test/unit_test/api/apps/services/test_delete_datasets.py +++ b/test/unit_test/api/apps/services/test_delete_datasets.py @@ -17,8 +17,8 @@ import importlib.util import sys -from pathlib import Path from enum import IntEnum +from pathlib import Path from types import ModuleType, SimpleNamespace from unittest.mock import MagicMock @@ -155,8 +155,16 @@ def _load_delete_datasets_module(monkeypatch, *, f2d_rows, file_filter_delete): ), StatusEnum=SimpleNamespace(), LLMType=SimpleNamespace(), + RetCode=SimpleNamespace(), ModelTypeBinary=_StubModelTypeBinary, ) + _stub(monkeypatch, "rag.advanced_rag", __path__=[]) + _stub(monkeypatch, "rag.advanced_rag.knowlege_compile", __path__=[]) + _stub( + monkeypatch, + "rag.advanced_rag.knowlege_compile.wiki", + WIKI_PAGE_COMPILE_KWD="wiki", + ) _stub( monkeypatch, "rag.nlp.search", diff --git a/test/unit_test/rag/advanced_rag/test_dataset_navigation_search.py b/test/unit_test/rag/advanced_rag/test_dataset_navigation_search.py new file mode 100644 index 0000000000..b9c0236128 --- /dev/null +++ b/test/unit_test/rag/advanced_rag/test_dataset_navigation_search.py @@ -0,0 +1,44 @@ +import sys +from types import ModuleType, SimpleNamespace + +import pytest + +from rag.advanced_rag.harness.tools.navigation import ( + _NAV_SEARCH_MAX_DOCS, + dataset_navigation_search, +) + + +@pytest.mark.asyncio +async def test_dataset_navigation_search_ranks_across_all_bound_kbs(monkeypatch): + kb1 = SimpleNamespace(id="kb-1", tenant_id="tenant-1") + kb2 = SimpleNamespace(id="kb-2", tenant_id="tenant-2") + tools = SimpleNamespace(kbs=[kb1, kb2], scoped_doc_ids=lambda doc_scope: doc_scope) + + kb1_items = [{"doc_id": f"kb1-doc-{i}", "score": 0.30 + i * 0.01} for i in range(_NAV_SEARCH_MAX_DOCS)] + kb2_items = [{"doc_id": "kb2-best", "score": 0.99}] + + async def fake_search_dataset_layers(kb_id, tenant_id, query, mode, top_k, doc_scope): + assert query == "topic keywords" + assert mode == "nav_doc" + assert top_k == _NAV_SEARCH_MAX_DOCS + assert doc_scope is None + if kb_id == kb1.id: + return True, {"items": kb1_items} + if kb_id == kb2.id: + return True, {"items": kb2_items} + raise AssertionError(f"unexpected kb_id: {kb_id}") + + dataset_api_service = ModuleType("dataset_api_service") + dataset_api_service.search_dataset_layers = fake_search_dataset_layers + services_module = ModuleType("api.apps.services") + services_module.dataset_api_service = dataset_api_service + + monkeypatch.setitem(sys.modules, "api.apps.services", services_module) + monkeypatch.setitem(sys.modules, "api.apps.services.dataset_api_service", dataset_api_service) + + routed = await dataset_navigation_search(tools, "topic", "keywords") + + assert len(routed) == _NAV_SEARCH_MAX_DOCS + assert routed[0] == "kb2-best" + assert "kb1-doc-0" not in routed