mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-14 20:54:30 +08:00
Generate navigation, navigation search (#18096)
### Summary
2 API
POST /api/v1/datasets/{dataset_id}/navigation
GET
/api/v1/datasets/{dataset_id}/navigation/search?q={query}&mode={mode}&top_k={topk}
2 cli
uv run --no-sync python3 admin/client/ragflow_cli.py -h 127.0.0.1 -p
9380 -t user
ragflow> GENERATE NAVIGATION OF DATASET 'frames tree';
ragflow> NAVIGATION SEARCH 'Christie introduced blockchain-based digital
passports' IN DATASET 'frames tree' MODE 'all' topk 20;
(mode: chunk, nav_cluster, nav_doc, navigation_tree, all)
This commit is contained in:
@@ -397,7 +397,13 @@ async def _ask_nav_select(tools, query: str, items: list[dict], noun: str, max_i
|
||||
name = str(it.get("name") or "").strip() or f"item-{i}"
|
||||
desc = str(it.get("description") or "").strip().replace("\n", " ")
|
||||
extra = f" [{it['doc_count']} docs]" if it.get("doc_count") else ""
|
||||
lines.append(f"[{i}] {name}{extra}: {desc[:300]}")
|
||||
kwds = it.get("keywords") or []
|
||||
tags = ", ".join(str(k) for k in kwds[:6]).strip()
|
||||
head = f" [tags: {tags}]" if tags else ""
|
||||
entities = it.get("entities") or []
|
||||
ents = ", ".join(str(e) for e in entities[:6]).strip()
|
||||
head += f" [entities: {ents}]" if ents else ""
|
||||
lines.append(f"[{i}] {name}{extra}{head}: {desc[:300]}")
|
||||
|
||||
system = _NAV_SELECT_SYSTEM.format(noun=noun)
|
||||
user = f"Question:\n{query}\n\n{noun.capitalize()} (numbered):\n" + "\n".join(lines) + "\n\nOutput JSON:"
|
||||
|
||||
@@ -68,6 +68,54 @@ _LOCK_BLOCKING_TIMEOUT_S = 5
|
||||
# Hard limit on how many sibling clusters we evaluate per KNN call
|
||||
_KNN_TOP_K = 5
|
||||
|
||||
# Fields needed to shape a nav hit for hybrid scoring / rendering.
|
||||
_NAV_SEARCH_FIELDS = [
|
||||
"id",
|
||||
"content_with_weight",
|
||||
"name",
|
||||
"doc_id",
|
||||
"type_kwd",
|
||||
"doc_ids_kwd",
|
||||
"doc_count_int",
|
||||
]
|
||||
|
||||
# Weight of the dense leg in hybrid search (1 - this = BM25 leg weight).
|
||||
_NAV_HYBRID_DENSE_W = 0.5
|
||||
|
||||
# Stop-words skipped when deriving routing tag-words from a summary.
|
||||
_NAV_STOP_WORDS = {
|
||||
"the",
|
||||
"a",
|
||||
"an",
|
||||
"and",
|
||||
"or",
|
||||
"of",
|
||||
"to",
|
||||
"in",
|
||||
"on",
|
||||
"for",
|
||||
"with",
|
||||
"at",
|
||||
"is",
|
||||
"are",
|
||||
"was",
|
||||
"were",
|
||||
"be",
|
||||
"been",
|
||||
"being",
|
||||
"this",
|
||||
"that",
|
||||
"these",
|
||||
"those",
|
||||
"it",
|
||||
"its",
|
||||
"as",
|
||||
"by",
|
||||
"from",
|
||||
"about",
|
||||
"into",
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
@@ -167,6 +215,60 @@ async def _store_search(
|
||||
return list(rows.values())
|
||||
|
||||
|
||||
async def _store_text_search(
|
||||
tenant_id: str,
|
||||
kb_id: str,
|
||||
query: str,
|
||||
fields: list[str],
|
||||
limit: int = 100,
|
||||
*,
|
||||
compile_kwd: str = _COMPILE_KWD,
|
||||
type_kwd: str = "",
|
||||
extra_filter: dict | None = None,
|
||||
) -> list[dict]:
|
||||
"""Full-text (BM25) leg over the nav rows' tokenized fields.
|
||||
|
||||
Recalls nav nodes whose ``content_ltks`` / ``content_sm_ltks`` match the
|
||||
query — the lexical half of hybrid search, complementing the KNN leg for
|
||||
exact/proper-noun recall (e.g. "关羽", "ImageNet 2012").
|
||||
|
||||
Args:
|
||||
compile_kwd: Which compile partition to search within
|
||||
(default ``_COMPILE_KWD`` = ``"dataset_nav"``).
|
||||
type_kwd: Optional type filter (``"nav_doc"``, ``"nav_cluster"``, or ``""``).
|
||||
extra_filter: Additional filter conditions merged into the query.
|
||||
"""
|
||||
from common import settings
|
||||
from common.doc_store.doc_store_base import MatchTextExpr, OrderByExpr
|
||||
|
||||
# Pre-tokenize the query with the same tokenizer used to index content_ltks.
|
||||
# ES content_ltks uses the "whitespace" analyzer (no stemming), while our
|
||||
# Python tokenizer applies stemming (e.g. "Christie" → "christi").
|
||||
# Without this, raw query terms won't match pre-stemmed tokens in the index.
|
||||
tokenized_query = _tokenize(query)
|
||||
|
||||
index = _index_name(tenant_id)
|
||||
filter_condition: dict = {"compile_kwd": [compile_kwd]}
|
||||
if type_kwd:
|
||||
filter_condition["type_kwd"] = type_kwd
|
||||
if extra_filter:
|
||||
filter_condition.update(extra_filter)
|
||||
res = await thread_pool_exec(
|
||||
settings.docStoreConn.search,
|
||||
fields,
|
||||
[],
|
||||
filter_condition,
|
||||
[MatchTextExpr(["content_ltks", "content_sm_ltks"], tokenized_query, limit)],
|
||||
OrderByExpr(),
|
||||
0,
|
||||
limit,
|
||||
index,
|
||||
[kb_id],
|
||||
)
|
||||
rows = settings.docStoreConn.get_fields(res, fields) if res else {}
|
||||
return list(rows.values())
|
||||
|
||||
|
||||
async def _store_knn(
|
||||
tenant_id: str,
|
||||
kb_id: str,
|
||||
@@ -320,8 +422,27 @@ def _make_nav_doc_row(
|
||||
depth_int: int,
|
||||
embd_mdl=None,
|
||||
embedding: list[float] | None = None,
|
||||
*,
|
||||
graph_content: str = "",
|
||||
) -> dict:
|
||||
"""Build a nav_doc ES/Infinity row dict for a single document leaf node."""
|
||||
"""Build a nav_doc ES/Infinity row dict for a single document leaf node.
|
||||
|
||||
Args:
|
||||
summary: Short human-readable title / description for display.
|
||||
graph_content: Optional full-graph text used for richer keyword /
|
||||
entity extraction. When set, keywords and entities are derived
|
||||
from graph_content instead of summary, and the text is stored in
|
||||
the ``graph_content`` payload field.
|
||||
"""
|
||||
kw_text = graph_content or summary
|
||||
payload = {
|
||||
"type": "nav_doc",
|
||||
"description": summary,
|
||||
"keywords": _nav_keywords(kw_text),
|
||||
"entities": _nav_entities(kw_text),
|
||||
}
|
||||
if graph_content:
|
||||
payload["graph_content"] = graph_content
|
||||
row: dict = {
|
||||
"id": _nav_doc_id(doc_id),
|
||||
"kb_id": kb_id,
|
||||
@@ -334,9 +455,8 @@ def _make_nav_doc_row(
|
||||
"depth_int": depth_int,
|
||||
"available_int": 0,
|
||||
}
|
||||
payload = {"type": "nav_doc", "description": summary}
|
||||
row["content_with_weight"] = json.dumps(payload, ensure_ascii=False)
|
||||
ltks = _tokenize(summary)
|
||||
ltks = _tokenize(kw_text)
|
||||
row["content_ltks"] = ltks
|
||||
row["content_sm_ltks"] = _fine_tokenize(ltks)
|
||||
if _vector_len(embedding) > 0:
|
||||
@@ -370,7 +490,12 @@ def _make_nav_cluster_row(
|
||||
"doc_count_int": len(doc_ids),
|
||||
"available_int": 0,
|
||||
}
|
||||
payload = {"type": "nav_cluster", "description": description}
|
||||
payload = {
|
||||
"type": "nav_cluster",
|
||||
"description": description,
|
||||
"keywords": _nav_keywords(description),
|
||||
"entities": _nav_entities(description),
|
||||
}
|
||||
row["content_with_weight"] = json.dumps(payload, ensure_ascii=False)
|
||||
ltks = _tokenize(description)
|
||||
row["content_ltks"] = ltks
|
||||
@@ -395,6 +520,72 @@ def _fine_tokenize(text: str) -> str:
|
||||
return rag_tokenizer.fine_grained_tokenize(text)
|
||||
|
||||
|
||||
def _nav_keywords(summary: str, max_kwds: int = 6) -> list[str]:
|
||||
"""Derive routing tags from a nav summary (no LLM call, zero cost).
|
||||
|
||||
These are the tokenized non-stop terms the model uses to fast-judge node
|
||||
relevance before reading the full description.
|
||||
"""
|
||||
from rag.nlp import rag_tokenizer
|
||||
|
||||
tokens = (rag_tokenizer.tokenize(summary or "") or "").split()
|
||||
seen: set[str] = set()
|
||||
out: list[str] = []
|
||||
for t in tokens:
|
||||
t = t.strip()
|
||||
if len(t) < 2 or t.isdigit() or t.lower() in _NAV_STOP_WORDS or t.lower() in seen:
|
||||
continue
|
||||
seen.add(t.lower())
|
||||
out.append(t)
|
||||
if len(out) >= max_kwds:
|
||||
break
|
||||
return out
|
||||
|
||||
|
||||
def _nav_entities(summary: str, max_entities: int = 6) -> list[str]:
|
||||
"""Extract likely named entities from a nav summary (no LLM call, zero cost).
|
||||
|
||||
Uses a two-pass heuristic:
|
||||
1. Regex for English capitalized multi-word sequences (proper nouns).
|
||||
2. Tokenizer-based extraction for CJK and other non-Latin text.
|
||||
|
||||
Returns up to ``max_entities`` deduplicated strings.
|
||||
"""
|
||||
text = summary or ""
|
||||
entities: list[str] = []
|
||||
seen: set[str] = set()
|
||||
|
||||
# Pass 1: English-like capitalized sequences (e.g. "New York", "Machine Learning")
|
||||
for m in re.finditer(r"\b([A-Z][a-z]+(?:\s+[A-Z][a-z]+)+)\b", text):
|
||||
ent = m.group(1).strip()
|
||||
key = ent.lower()
|
||||
if key not in _NAV_STOP_WORDS and key not in seen:
|
||||
seen.add(key)
|
||||
entities.append(ent)
|
||||
if len(entities) >= max_entities:
|
||||
return entities
|
||||
|
||||
# Pass 2: tokenizer-based for CJK / other scripts
|
||||
from rag.nlp import rag_tokenizer
|
||||
|
||||
tokens = (rag_tokenizer.tokenize(text) or "").split()
|
||||
for t in tokens:
|
||||
t = t.strip()
|
||||
if len(t) < 3 or t.isdigit():
|
||||
continue
|
||||
if t.isascii() and t[0].islower():
|
||||
continue # skip lowercase English fragments
|
||||
key = t.lower()
|
||||
if key in _NAV_STOP_WORDS or key in seen:
|
||||
continue
|
||||
seen.add(key)
|
||||
entities.append(t)
|
||||
if len(entities) >= max_entities:
|
||||
break
|
||||
|
||||
return entities
|
||||
|
||||
|
||||
def _matches_condition(row: dict, condition: dict) -> bool:
|
||||
"""Check the simple equality filters used by dataset navigation."""
|
||||
for field, expected in condition.items():
|
||||
@@ -574,17 +765,27 @@ async def upsert_dataset_nav_doc(
|
||||
tenant_id: Tenant owning the KB.
|
||||
kb_id: Knowledge base id.
|
||||
doc_id: Document id.
|
||||
summary_or_tree: A plain summary string, or a RAPTOR tree dict from
|
||||
which the root summary is extracted.
|
||||
summary_or_tree:
|
||||
- A plain summary string (fallback, no RAPTOR graph).
|
||||
- A RAPTOR tree dict (backwards-compat, root summary extracted).
|
||||
- A dict with keys ``title`` (short display label) and
|
||||
``graph_text`` (full graph entities + relations for richer
|
||||
embedding / keyword extraction).
|
||||
embd_mdl: LLMBundle for embedding (required for clustering).
|
||||
chat_mdl: LLMBundle for chat (required for LLM merge/summary).
|
||||
"""
|
||||
if not doc_id or not kb_id:
|
||||
return
|
||||
|
||||
# 1. Extract summary
|
||||
# 1. Extract display summary and optional full graph content.
|
||||
graph_content = ""
|
||||
if isinstance(summary_or_tree, dict):
|
||||
summary = _extract_root_summary_from_tree(summary_or_tree)
|
||||
if "title" in summary_or_tree and "graph_text" in summary_or_tree:
|
||||
# New extended format from generate_nav's RAPTOR graph path.
|
||||
summary = (summary_or_tree.get("title") or "").strip()
|
||||
graph_content = (summary_or_tree.get("graph_text") or "").strip()
|
||||
else:
|
||||
summary = _extract_root_summary_from_tree(summary_or_tree)
|
||||
elif isinstance(summary_or_tree, str):
|
||||
summary = summary_or_tree
|
||||
else:
|
||||
@@ -593,10 +794,10 @@ async def upsert_dataset_nav_doc(
|
||||
logging.info("dataset_nav: skipping doc=%s (kb=%s) — no summary", doc_id, kb_id)
|
||||
return
|
||||
|
||||
# 2. Embed doc summary before taking the KB lock. The result is
|
||||
# independent of the nav tree; all tree reads, deletes, and writes below
|
||||
# are serialized by the same lock.
|
||||
doc_embedding = await _embed(embd_mdl, summary) if embd_mdl else []
|
||||
# 2. Embed with the richest available text so that KNN placement
|
||||
# benefits from the full tree structure when present.
|
||||
embed_text = graph_content or summary
|
||||
doc_embedding = await _embed(embd_mdl, embed_text) if embd_mdl else []
|
||||
vec_dim = len(doc_embedding)
|
||||
|
||||
lock = RedisDistributedLock(
|
||||
@@ -665,6 +866,7 @@ async def upsert_dataset_nav_doc(
|
||||
depth,
|
||||
embd_mdl,
|
||||
doc_embedding,
|
||||
graph_content=graph_content,
|
||||
)
|
||||
await _store_upsert(tenant_id, kb_id, nav_doc_row)
|
||||
|
||||
@@ -712,6 +914,7 @@ async def upsert_dataset_nav_doc(
|
||||
new_depth,
|
||||
embd_mdl,
|
||||
doc_embedding,
|
||||
graph_content=graph_content,
|
||||
)
|
||||
await _store_upsert(tenant_id, kb_id, nav_doc_row)
|
||||
else:
|
||||
@@ -739,6 +942,7 @@ async def upsert_dataset_nav_doc(
|
||||
1,
|
||||
embd_mdl,
|
||||
doc_embedding,
|
||||
graph_content=graph_content,
|
||||
)
|
||||
await _store_upsert(tenant_id, kb_id, nav_doc_row)
|
||||
|
||||
@@ -1015,29 +1219,47 @@ async def search_dataset_nav(
|
||||
kb_id: str,
|
||||
query: str,
|
||||
embd_mdl=None,
|
||||
top_k: int = 8,
|
||||
top_k: int | None = None,
|
||||
*,
|
||||
type_kwd: str = "",
|
||||
compile_kwd: str = _COMPILE_KWD,
|
||||
) -> list[dict]:
|
||||
"""Find the nav-tree nodes most relevant to ``query`` for one KB.
|
||||
|
||||
The nav rows are ``available_int=0`` (invisible to the normal retriever), so
|
||||
this is the sanctioned read seam: a caller uses the returned document ids to
|
||||
route a scoped chunk retrieval. Returns items shaped as::
|
||||
route a scoped chunk retrieval.
|
||||
|
||||
Args:
|
||||
type_kwd: Optional type filter — ``"nav_doc"`` to restrict to document
|
||||
leaves, ``"nav_cluster"`` to restrict to clusters, or ``""`` for all.
|
||||
compile_kwd: Which compile partition to search within
|
||||
(default ``_COMPILE_KWD`` = ``"dataset_nav"``).
|
||||
|
||||
Returns items shaped as::
|
||||
|
||||
{"type": "nav_doc" | "nav_cluster",
|
||||
"doc_id": str | None, # the document, for a leaf
|
||||
"doc_ids": [str], # the documents a node covers
|
||||
"name": str, "description": str, "score": float}
|
||||
|
||||
Ranked by vector KNN over the node summaries when ``embd_mdl`` is given;
|
||||
otherwise a best-effort text-ranked scan.
|
||||
Ranked by hybrid search: KNN over the node-summary vectors fused with a
|
||||
BM25 (full-text) leg over the tokenized fields. When ``embd_mdl`` is absent
|
||||
the lexical leg alone ranks the results.
|
||||
"""
|
||||
query = (query or "").strip()
|
||||
if not query:
|
||||
return []
|
||||
|
||||
condition = {"compile_kwd": [_COMPILE_KWD]}
|
||||
rows_with_scores: list[tuple[dict, float]] = []
|
||||
condition: dict = {"compile_kwd": [compile_kwd]}
|
||||
if type_kwd:
|
||||
condition["type_kwd"] = type_kwd
|
||||
# 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] = {}
|
||||
dense_w = _NAV_HYBRID_DENSE_W if embd_mdl is not None else 0.0
|
||||
|
||||
# ── Dense leg: KNN over the node-summary vectors ──
|
||||
if embd_mdl is not None:
|
||||
try:
|
||||
vec = await _embed(embd_mdl, query)
|
||||
@@ -1046,28 +1268,39 @@ async def search_dataset_nav(
|
||||
vec = []
|
||||
if _vector_len(vec) > 0:
|
||||
try:
|
||||
rows = await _store_knn(tenant_id, kb_id, vec, len(vec), condition, top_k=top_k)
|
||||
rows = await _store_knn(tenant_id, kb_id, vec, len(vec), condition, top_k=top_k or 10000)
|
||||
vf = _vec_field(len(vec))
|
||||
rows_with_scores = [(r, _cosine_sim(vec, r.get(vf))) for r in rows]
|
||||
for r in rows:
|
||||
rk = r.get("name") or r.get("doc_id") or ""
|
||||
if not rk:
|
||||
continue
|
||||
fused.setdefault(rk, [r, 0.0])[1] += dense_w * _cosine_sim(vec, r.get(vf))
|
||||
except Exception:
|
||||
logging.exception("search_dataset_nav: knn failed for kb=%s", kb_id)
|
||||
rows_with_scores = []
|
||||
|
||||
if not rows_with_scores:
|
||||
fields = ["content_with_weight", "name", "doc_id", "type_kwd", "doc_ids_kwd", "doc_count_int"]
|
||||
try:
|
||||
rows = await _store_search(tenant_id, kb_id, condition, fields, limit=max(top_k * 20, 100))
|
||||
except Exception:
|
||||
logging.exception("search_dataset_nav: scan failed for kb=%s", kb_id)
|
||||
rows = []
|
||||
rows_with_scores = [(r, _nav_text_score(query, r)) for r in rows]
|
||||
rows_with_scores.sort(key=lambda item: item[1], reverse=True)
|
||||
# ── Lexical leg: engine BM25 over the tokenized fields ──
|
||||
text_w = 1.0 - dense_w
|
||||
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)
|
||||
except Exception:
|
||||
logging.exception("search_dataset_nav: text search failed for kb=%s", kb_id)
|
||||
text_rows = []
|
||||
for r in text_rows:
|
||||
rk = r.get("name") or r.get("doc_id") or ""
|
||||
if not rk:
|
||||
continue
|
||||
ts = _nav_text_score(query, r)
|
||||
if ts <= 0:
|
||||
continue
|
||||
fused.setdefault(rk, [r, 0.0])[1] += text_w * ts
|
||||
|
||||
# Discard zero-score rows (text match produced no relevant hits)
|
||||
rows_with_scores = [(r, s) for r, s in rows_with_scores if s > 0]
|
||||
rows_with_scores = [(r, s) for r, s in fused.values() if s > 0]
|
||||
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]
|
||||
|
||||
out: list[dict] = []
|
||||
for r, score in rows_with_scores[:top_k]:
|
||||
for r, score in rows_with_scores:
|
||||
try:
|
||||
payload = json.loads(r.get("content_with_weight") or "{}")
|
||||
except Exception:
|
||||
@@ -1088,6 +1321,11 @@ async def search_dataset_nav(
|
||||
"doc_ids": doc_ids,
|
||||
"name": name,
|
||||
"description": payload.get("description") or "",
|
||||
"keywords": _as_str_list(payload.get("keywords")),
|
||||
"entities": _as_str_list(payload.get("entities")),
|
||||
"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),
|
||||
"score": float(score or 0.0),
|
||||
}
|
||||
@@ -1095,6 +1333,233 @@ async def search_dataset_nav(
|
||||
return out
|
||||
|
||||
|
||||
async def search_nav_tree_descent(
|
||||
tenant_id: str,
|
||||
kb_id: str,
|
||||
query: str,
|
||||
embd_mdl,
|
||||
top_k: int | None = None,
|
||||
) -> list[dict]:
|
||||
"""Tree-structured hybrid search: descend from root into the most relevant branches.
|
||||
|
||||
Unlike ``search_dataset_nav`` (flat hybrid search over all nav nodes),
|
||||
this respects the parent-child hierarchy: it starts at root clusters,
|
||||
then at each level descends into the semantically closest children,
|
||||
collecting document identifiers from ``nav_doc`` leaves along the way.
|
||||
|
||||
Each level uses hybrid search (KNN + BM25 text) fused with the same
|
||||
weights as ``search_dataset_nav``, so both vector similarity *and*
|
||||
lexical matching drive the descent.
|
||||
|
||||
The search uses BFS with beam pruning — at each depth, only the
|
||||
*beam_width* most similar clusters are expanded further.
|
||||
|
||||
Returns items shaped as ``{"doc_id": str, "score": float}``.
|
||||
"""
|
||||
query = (query or "").strip()
|
||||
if not query:
|
||||
return []
|
||||
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")
|
||||
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)
|
||||
vec_dim = _vector_len(vec)
|
||||
if vec_dim == 0:
|
||||
return []
|
||||
|
||||
beam_width = 5
|
||||
dense_w = _NAV_HYBRID_DENSE_W # same weight as search_dataset_nav
|
||||
vf = _vec_field(vec_dim)
|
||||
|
||||
fields = [
|
||||
"content_with_weight",
|
||||
"name",
|
||||
"doc_id",
|
||||
"compile_kwd",
|
||||
"type_kwd",
|
||||
"parent_kwd",
|
||||
"depth_int",
|
||||
"doc_count_int",
|
||||
"doc_ids_kwd",
|
||||
vf,
|
||||
]
|
||||
|
||||
collected: list[dict] = []
|
||||
seen_docs: set[str] = set()
|
||||
seen_nodes: set[str] = set()
|
||||
|
||||
# ── Root clusters (KNN only — no text) ──
|
||||
# Root clusters at depth=0 are LLM-generated broad-topic summaries.
|
||||
# Specific query terms (e.g. "British") almost never appear in their
|
||||
# name / description, so we route purely by vector similarity at this
|
||||
# level. Text matching takes over from depth ≥ 1 where cluster
|
||||
# descriptions contain concrete terms.
|
||||
root_cond = {
|
||||
"kb_id": [kb_id],
|
||||
"compile_kwd": [_COMPILE_KWD],
|
||||
"type_kwd": ["nav_cluster"],
|
||||
"depth_int": [0],
|
||||
}
|
||||
roots_knn = await _store_knn(tenant_id, kb_id, vec, vec_dim, root_cond, top_k=beam_width * 3)
|
||||
|
||||
# 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 not roots_knn:
|
||||
all_cond = {
|
||||
"kb_id": [kb_id],
|
||||
"compile_kwd": [_COMPILE_KWD],
|
||||
"type_kwd": ["nav_cluster"],
|
||||
}
|
||||
all_clusters = await _store_search(tenant_id, kb_id, all_cond, fields, limit=10000)
|
||||
if not all_clusters:
|
||||
return []
|
||||
|
||||
min_depth = min(
|
||||
(r.get("depth_int", 0) for r in all_clusters if r.get("depth_int") is not None),
|
||||
default=0,
|
||||
)
|
||||
logging.warning(
|
||||
"search_nav_tree_descent: no root cluster at depth=0, starting beam search from depth=%d",
|
||||
min_depth,
|
||||
)
|
||||
|
||||
starters = [r for r in all_clusters if r.get("depth_int") == min_depth]
|
||||
starters.sort(key=lambda r: _cosine_sim(vec, r.get(vf)), reverse=True)
|
||||
current_level = starters[:beam_width]
|
||||
for r in current_level:
|
||||
r["_score"] = _cosine_sim(vec, r.get(vf))
|
||||
else:
|
||||
# Route down to beam_width semantically closest root clusters.
|
||||
for r in roots_knn:
|
||||
r["_score"] = _cosine_sim(vec, r.get(vf))
|
||||
roots_knn.sort(key=lambda r: r["_score"], reverse=True)
|
||||
current_level = roots_knn[:beam_width]
|
||||
|
||||
if not current_level:
|
||||
return []
|
||||
|
||||
while current_level and (top_k is None or len(collected) < top_k):
|
||||
next_level: list[dict] = []
|
||||
|
||||
for node in current_level:
|
||||
node_name = node.get("name", "")
|
||||
if node_name in seen_nodes:
|
||||
continue
|
||||
seen_nodes.add(node_name)
|
||||
parent_score = node.get("_score", 0.0)
|
||||
|
||||
child_cond: dict = {
|
||||
"kb_id": [kb_id],
|
||||
"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]},
|
||||
)
|
||||
candidates = _hybrid_fuse(vec, vf, query, children_knn, children_text, dense_w, beam_width)
|
||||
|
||||
for c in candidates:
|
||||
if top_k is not None and len(collected) >= top_k:
|
||||
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)})
|
||||
else:
|
||||
next_level.append(c)
|
||||
|
||||
if top_k is not None and len(collected) >= top_k:
|
||||
break
|
||||
|
||||
if top_k is not None and len(collected) >= top_k:
|
||||
break
|
||||
|
||||
next_level.sort(key=lambda c: c.get("_score", 0.0), reverse=True)
|
||||
current_level = next_level[:beam_width]
|
||||
|
||||
# Fallback: collect doc_ids from terminal cluster nodes.
|
||||
if not collected and current_level:
|
||||
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 top_k is not None and len(collected) >= top_k:
|
||||
break
|
||||
|
||||
return collected
|
||||
|
||||
|
||||
def _hybrid_fuse(
|
||||
vec: list[float],
|
||||
vf: str,
|
||||
query: str,
|
||||
knn_rows: list[dict],
|
||||
text_rows: list[dict],
|
||||
dense_w: float,
|
||||
top_k: int,
|
||||
) -> list[dict]:
|
||||
"""Fuse KNN and text results into a scored, deduplicated, sorted list.
|
||||
|
||||
``_store_knn`` already enforces filter conditions on the KNN leg;
|
||||
text rows come from ``_store_text_search`` which applies its own
|
||||
filters. Here we merge both legs by ``name``/``doc_id``.
|
||||
"""
|
||||
text_w = 1.0 - dense_w
|
||||
fused: dict[str, tuple[dict, float]] = {} # key → (row, score)
|
||||
|
||||
# KNN leg: raw cosine * dense_w, matching search_dataset_nav's scale so
|
||||
# both search paths score the dense leg identically.
|
||||
for r in knn_rows:
|
||||
rk = r.get("name") or r.get("doc_id") or ""
|
||||
if not rk:
|
||||
continue
|
||||
key = f"knn:{rk}"
|
||||
fused[key] = (r, _cosine_sim(vec, r.get(vf, [])) * dense_w)
|
||||
|
||||
# Text leg.
|
||||
for r in text_rows:
|
||||
rk = r.get("name") or r.get("doc_id") or ""
|
||||
if not rk:
|
||||
continue
|
||||
ts = _nav_text_score(query, r)
|
||||
if ts <= 0:
|
||||
continue
|
||||
key = f"text:{rk}"
|
||||
if key in fused:
|
||||
fused[key] = (fused[key][0], fused[key][1] + text_w * ts)
|
||||
else:
|
||||
key_knn = f"knn:{rk}"
|
||||
if key_knn in fused:
|
||||
fused[key_knn] = (fused[key_knn][0], fused[key_knn][1] + text_w * ts)
|
||||
else:
|
||||
fused[key] = (r, text_w * ts)
|
||||
|
||||
rows_with_scores = [(r, s) for r, s in fused.values() if s > 0]
|
||||
rows_with_scores.sort(key=lambda item: item[1], reverse=True)
|
||||
result = rows_with_scores[:top_k]
|
||||
for r, s in result:
|
||||
r["_score"] = s
|
||||
return [r for r, _ in result]
|
||||
|
||||
|
||||
def _as_str_list(value) -> list[str]:
|
||||
if isinstance(value, list):
|
||||
return [str(v) for v in value if v]
|
||||
@@ -1108,11 +1573,17 @@ def _nav_text_score(query: str, row: dict) -> float:
|
||||
payload = json.loads(row.get("content_with_weight") or "{}")
|
||||
except Exception:
|
||||
payload = {}
|
||||
keywords = payload.get("keywords") or []
|
||||
entities = payload.get("entities") or []
|
||||
graph_content = payload.get("graph_content") or ""
|
||||
haystack = " ".join(
|
||||
str(x or "")
|
||||
for x in (
|
||||
row.get("name"),
|
||||
payload.get("description"),
|
||||
graph_content,
|
||||
*keywords,
|
||||
*entities,
|
||||
)
|
||||
).lower()
|
||||
q_terms = set(re.findall(r"[\w]+", query.lower()))
|
||||
|
||||
@@ -14,18 +14,19 @@
|
||||
# limitations under the License.
|
||||
#
|
||||
|
||||
import re
|
||||
import copy
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
|
||||
import copy
|
||||
from elasticsearch_dsl import UpdateByQuery, Q, Search
|
||||
from elastic_transport import ConnectionTimeout
|
||||
from elasticsearch_dsl import Q, Search, UpdateByQuery
|
||||
|
||||
from common.constants import PAGERANK_FLD, TAG_FLD
|
||||
from common.decorator import singleton
|
||||
from common.doc_store.doc_store_base import MatchTextExpr, OrderByExpr, MatchExpr, MatchDenseExpr, FusionExpr
|
||||
from common.doc_store.doc_store_base import FusionExpr, MatchDenseExpr, MatchExpr, MatchTextExpr, OrderByExpr
|
||||
from common.doc_store.es_conn_base import ESConnectionBase
|
||||
from common.float_utils import get_float
|
||||
from common.constants import PAGERANK_FLD, TAG_FLD
|
||||
|
||||
ATTEMPT_TIME = 2
|
||||
MAX_RESULT_WINDOW = 10000
|
||||
@@ -227,7 +228,7 @@ class ESConnection(ESConnectionBase):
|
||||
elif isinstance(v, str) or isinstance(v, int):
|
||||
bool_query.filter.append(Q("term", **{k: v}))
|
||||
else:
|
||||
raise Exception(f"Condition `{str(k)}={str(v)}` value type is {str(type(v))}, expected to be int, str or list.")
|
||||
raise Exception(f"Condition `{k!s}={v!s}` value type is {type(v)!s}, expected to be int, str or list.")
|
||||
|
||||
s = Search()
|
||||
vector_similarity_weight = 0.5
|
||||
@@ -243,7 +244,7 @@ class ESConnection(ESConnectionBase):
|
||||
vector_similarity_weight = get_float(weights.split(",")[1])
|
||||
for m in match_expressions:
|
||||
if isinstance(m, MatchTextExpr):
|
||||
minimum_should_match = m.extra_options.get("minimum_should_match", 0.0)
|
||||
minimum_should_match = (m.extra_options or {}).get("minimum_should_match", 0.0)
|
||||
if isinstance(minimum_should_match, float):
|
||||
minimum_should_match = str(int(minimum_should_match * 100)) + "%"
|
||||
bool_query.must.append(Q("query_string", fields=m.fields, type="best_fields", query=m.matching_text, minimum_should_match=minimum_should_match, boost=1))
|
||||
@@ -254,10 +255,11 @@ class ESConnection(ESConnectionBase):
|
||||
similarity = 0.0
|
||||
if "similarity" in m.extra_options:
|
||||
similarity = m.extra_options["similarity"]
|
||||
k = min(m.topn, 10000)
|
||||
s = s.knn(
|
||||
m.vector_column_name,
|
||||
m.topn,
|
||||
m.topn * 2,
|
||||
k,
|
||||
min(k * 2, 10000),
|
||||
query_vector=list(m.embedding_data),
|
||||
filter=bool_query.to_dict(), # filter=_build_knn_filter_query(bool_query, vector_similarity_weight),
|
||||
similarity=similarity,
|
||||
@@ -310,7 +312,7 @@ class ESConnection(ESConnectionBase):
|
||||
vector_fields = [f for f in (select_fields or []) if f.endswith("_vec")]
|
||||
if vector_fields:
|
||||
q["fields"] = vector_fields
|
||||
self.logger.debug(f"ESConnection.search {str(index_names)} query: " + json.dumps(q))
|
||||
self.logger.debug(f"ESConnection.search {index_names!s} query: " + json.dumps(q))
|
||||
|
||||
for i in range(ATTEMPT_TIME):
|
||||
try:
|
||||
@@ -321,7 +323,7 @@ class ESConnection(ESConnectionBase):
|
||||
res = self._es_search_once(index_names, q, track_total_hits=True)
|
||||
if str(res.get("timed_out", "")).lower() == "true":
|
||||
raise Exception("Es Timeout.")
|
||||
self.logger.debug(f"ESConnection.search {str(index_names)} res: " + str(res))
|
||||
self.logger.debug(f"ESConnection.search {index_names!s} res: " + str(res))
|
||||
return res
|
||||
except ConnectionTimeout:
|
||||
self.logger.exception("ES request timeout")
|
||||
@@ -330,9 +332,9 @@ class ESConnection(ESConnectionBase):
|
||||
except Exception as e:
|
||||
# Only log debug for NotFoundError(accepted when metadata index doesn't exist)
|
||||
if "NotFound" in str(e):
|
||||
self.logger.debug(f"ESConnection.search {str(index_names)} query: " + str(q) + " - " + str(e))
|
||||
self.logger.debug(f"ESConnection.search {index_names!s} query: " + str(q) + " - " + str(e))
|
||||
else:
|
||||
self.logger.exception(f"ESConnection.search {str(index_names)} query: " + str(q) + str(e))
|
||||
self.logger.exception(f"ESConnection.search {index_names!s} query: " + str(q) + str(e))
|
||||
raise e
|
||||
|
||||
self.logger.error(f"ESConnection.search timeout for {ATTEMPT_TIME} times!")
|
||||
@@ -443,7 +445,7 @@ class ESConnection(ESConnectionBase):
|
||||
elif isinstance(v, str) or isinstance(v, int):
|
||||
bool_query.filter.append(Q("term", **{k: v}))
|
||||
else:
|
||||
raise Exception(f"Condition `{str(k)}={str(v)}` value type is {str(type(v))}, expected to be int, str or list.")
|
||||
raise Exception(f"Condition `{k!s}={v!s}` value type is {type(v)!s}, expected to be int, str or list.")
|
||||
scripts = []
|
||||
params = {}
|
||||
for k, v in new_value.items():
|
||||
@@ -473,7 +475,7 @@ class ESConnection(ESConnectionBase):
|
||||
scripts.append(f"ctx._source.{k}=params.pp_{k};")
|
||||
params[f"pp_{k}"] = json.dumps(v, ensure_ascii=False)
|
||||
else:
|
||||
raise Exception(f"newValue `{str(k)}={str(v)}` value type is {str(type(v))}, expected to be int, str.")
|
||||
raise Exception(f"newValue `{k!s}={v!s}` value type is {type(v)!s}, expected to be int, str.")
|
||||
ubq = UpdateByQuery(index=index_name).using(self.es).query(bool_query)
|
||||
ubq = ubq.script(source="".join(scripts), params=params)
|
||||
ubq = ubq.params(refresh=True)
|
||||
|
||||
Reference in New Issue
Block a user