Add dataset navigation search for agentic RAG (#18218)

### Summary

Add dataset navigation search for agentic RAG.

dataset_navigation_search() calls
dataset_api_service.search_dataset_layers with mode as "nav_doc"
This commit is contained in:
qinling0210
2026-08-13 17:35:24 +08:00
committed by GitHub
parent 7427d94d50
commit a6dcc0f665
13 changed files with 489 additions and 91 deletions

View File

@@ -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/<dataset_id>/navigation/search?q=<query>&mode=<mode>&top_k=20
GET /api/v1/datasets/<dataset_id>/navigation/search?q=<query>&mode=<mode>&top_k=20&doc_ids=<id>,<id>
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": <mode>, "total": <n>,
"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")

View File

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