161 lines
6.3 KiB
Python
161 lines
6.3 KiB
Python
"""Knowledge search: keyword / vector / hybrid (phase 1 + 2).
|
|
|
|
- ``keyword`` — title>tags>summary>body scoring over the DB mirror (always works).
|
|
- ``vector`` — cosine over per-chunk embeddings (requires a configured
|
|
embedding endpoint; otherwise transparently falls back to ``keyword``).
|
|
- ``hybrid`` — min-max-normalized blend of keyword + vector.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import re
|
|
from typing import TYPE_CHECKING, Any
|
|
|
|
from deerflow.knowledge.embeddings import cosine_similarity
|
|
from deerflow.knowledge.repository import KnowledgeRepository
|
|
|
|
if TYPE_CHECKING:
|
|
from deerflow.knowledge.embeddings import EmbeddingClient
|
|
|
|
_TOKEN_RE = re.compile(r"[\w一-鿿]+", re.UNICODE)
|
|
|
|
|
|
def _tokens(text: str) -> list[str]:
|
|
return [t.lower() for t in _TOKEN_RE.findall(text or "")]
|
|
|
|
|
|
def _snippet(content: str, terms: list[str], *, width: int = 160) -> str:
|
|
body = re.sub(r"\s+", " ", content or "").strip()
|
|
low = body.lower()
|
|
for term in terms:
|
|
idx = low.find(term)
|
|
if idx >= 0:
|
|
start = max(0, idx - width // 3)
|
|
end = min(len(body), start + width)
|
|
prefix = "…" if start > 0 else ""
|
|
suffix = "…" if end < len(body) else ""
|
|
return f"{prefix}{body[start:end]}{suffix}"
|
|
return body[:width] + ("…" if len(body) > width else "")
|
|
|
|
|
|
async def _keyword_scores(repo: KnowledgeRepository, query: str, *, candidate_cap: int) -> dict[str, tuple[float, dict[str, Any]]]:
|
|
terms = _tokens(query)
|
|
if not terms:
|
|
return {}
|
|
candidates, _total = await repo.list_notes(limit=candidate_cap, offset=0)
|
|
scored: dict[str, tuple[float, dict[str, Any]]] = {}
|
|
for note in candidates:
|
|
title = (note.get("title") or "").lower()
|
|
summary = (note.get("summary") or "").lower()
|
|
content = (note.get("content_md") or "").lower()
|
|
tags = " ".join(note.get("tags") or []).lower()
|
|
score = 0.0
|
|
for term in terms:
|
|
if term in title:
|
|
score += 3.0
|
|
if term in tags:
|
|
score += 2.0
|
|
if term in summary:
|
|
score += 2.0
|
|
score += min(content.count(term), 5) * 0.5
|
|
if score > 0:
|
|
scored[note["id"]] = (score, note)
|
|
return scored
|
|
|
|
|
|
async def _vector_scores(
|
|
repo: KnowledgeRepository,
|
|
query: str,
|
|
embedding_client: EmbeddingClient,
|
|
) -> dict[str, tuple[float, dict[str, Any]]]:
|
|
qvec = await embedding_client.embed_query(query)
|
|
if not qvec:
|
|
return {}
|
|
embeddings = await repo.all_embeddings()
|
|
best: dict[str, tuple[float, str, str]] = {} # note_id -> (score, title, snippet)
|
|
for emb in embeddings:
|
|
sim = cosine_similarity(qvec, emb.get("vector") or [])
|
|
if sim <= 0:
|
|
continue
|
|
nid = emb["note_id"]
|
|
if nid not in best or sim > best[nid][0]:
|
|
best[nid] = (sim, emb.get("title") or "", emb.get("content") or "")
|
|
return {nid: (score, {"id": nid, "title": title, "_chunk": chunk}) for nid, (score, title, chunk) in best.items()}
|
|
|
|
|
|
def _normalize(scores: dict[str, float]) -> dict[str, float]:
|
|
if not scores:
|
|
return {}
|
|
mx = max(scores.values())
|
|
if mx <= 0:
|
|
return {k: 0.0 for k in scores}
|
|
return {k: v / mx for k, v in scores.items()}
|
|
|
|
|
|
async def search_notes(
|
|
repo: KnowledgeRepository,
|
|
query: str,
|
|
*,
|
|
mode: str = "keyword",
|
|
limit: int = 8,
|
|
candidate_cap: int = 500,
|
|
embedding_client: EmbeddingClient | None = None,
|
|
) -> list[dict[str, Any]]:
|
|
"""Score notes against ``query`` and return the top ``limit`` matches."""
|
|
terms = _tokens(query)
|
|
if not terms:
|
|
return []
|
|
|
|
want_vector = mode in ("vector", "hybrid") and embedding_client is not None and embedding_client.is_enabled()
|
|
keyword = await _keyword_scores(repo, query, candidate_cap=candidate_cap) if mode != "vector" or not want_vector else {}
|
|
vector = await _vector_scores(repo, query, embedding_client) if want_vector else {}
|
|
|
|
# Pure vector (when available); else keyword; hybrid blends both.
|
|
if mode == "vector" and want_vector:
|
|
combined = {nid: (sc, meta) for nid, (sc, meta) in vector.items()}
|
|
ranked = sorted(combined.items(), key=lambda kv: kv[1][0], reverse=True)[:limit]
|
|
return await _materialize(repo, ranked, terms)
|
|
|
|
if mode == "hybrid" and want_vector:
|
|
kw_norm = _normalize({k: v[0] for k, v in keyword.items()})
|
|
vec_norm = _normalize({k: v[0] for k, v in vector.items()})
|
|
ids = set(kw_norm) | set(vec_norm)
|
|
blended: dict[str, tuple[float, dict[str, Any]]] = {}
|
|
for nid in ids:
|
|
score = 0.5 * kw_norm.get(nid, 0.0) + 0.5 * vec_norm.get(nid, 0.0)
|
|
meta = keyword.get(nid, (0, None))[1] or vector.get(nid, (0, {}))[1]
|
|
blended[nid] = (score, meta)
|
|
ranked = sorted(blended.items(), key=lambda kv: kv[1][0], reverse=True)[:limit]
|
|
return await _materialize(repo, ranked, terms)
|
|
|
|
# keyword (default, and fallback when vector unavailable)
|
|
ranked = sorted(keyword.items(), key=lambda kv: kv[1][0], reverse=True)[:limit]
|
|
if not ranked and vector: # mode=vector requested but keyword empty: use vector
|
|
ranked = sorted(vector.items(), key=lambda kv: kv[1][0], reverse=True)[:limit]
|
|
return await _materialize(repo, ranked, terms)
|
|
|
|
|
|
async def _materialize(repo: KnowledgeRepository, ranked: list[tuple[str, tuple[float, dict[str, Any]]]], terms: list[str]) -> list[dict[str, Any]]:
|
|
max_score = ranked[0][1][0] if ranked else 1.0
|
|
results: list[dict[str, Any]] = []
|
|
for nid, (score, meta) in ranked:
|
|
note = meta if (meta and meta.get("content_md")) else await repo.get_note(nid, with_sources=False)
|
|
if note is None:
|
|
continue
|
|
snippet_src = note.get("content_md") or meta.get("_chunk") or note.get("summary") or ""
|
|
results.append(
|
|
{
|
|
"note_id": nid,
|
|
"chunk_id": None,
|
|
"title": note.get("title") or meta.get("title"),
|
|
"summary": note.get("summary"),
|
|
"snippet": _snippet(snippet_src, terms),
|
|
"score": round(score / max_score, 4) if max_score else 0.0,
|
|
"tags": note.get("tags") or [],
|
|
"source_type": note.get("source_type"),
|
|
"updated_at": note.get("updated_at"),
|
|
"sources": [],
|
|
}
|
|
)
|
|
return results
|