nanobot: 2026-09-10 12:33:37
This commit is contained in:
308
skills/wiki/scripts/wiki_search.py
Normal file
308
skills/wiki/scripts/wiki_search.py
Normal file
@@ -0,0 +1,308 @@
|
||||
#!/usr/bin/env -S uv run --script
|
||||
# /// script
|
||||
# requires-python = ">=3.11"
|
||||
# dependencies = ["requests", "pyyaml", "markdown-it-py", "sqlite-vec==0.1.6"]
|
||||
# ///
|
||||
"""Query side of the wiki skill — three layers, cheapest first.
|
||||
|
||||
grep live ripgrep over the files on disk. No index, never stale, and not limited
|
||||
to the indexed extensions, so it is the layer that covers source code.
|
||||
toc directory -> file -> title + tags, read straight from the `files` catalog.
|
||||
search FTS5 (BM25) and vec0 (KNN) over the same `chunks` rows, merged with RRF.
|
||||
|
||||
Both halves rank the same unit, which is what makes the merge meaningful. The
|
||||
parameters below are module constants on purpose: `--limit` is the only knob worth
|
||||
exposing, and a config key that never changes is a config key that rots.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import shutil
|
||||
import sqlite3
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import wiki_store as store
|
||||
from wiki_config import (
|
||||
GIT_KIND,
|
||||
ConfigError,
|
||||
SourceConfig,
|
||||
WikiConfig,
|
||||
db_path,
|
||||
load_config,
|
||||
source_root,
|
||||
)
|
||||
from wiki_embed import (
|
||||
EmbeddingUnavailable,
|
||||
IndexIdentityMismatch,
|
||||
OllamaEmbedder,
|
||||
require_matching_index,
|
||||
)
|
||||
|
||||
CANDIDATE_LIMIT = 50 # KNN k, and the LIMIT for BM25
|
||||
RRF_K = 60 # Cormack et al. 2009
|
||||
DEFAULT_LIMIT = 10
|
||||
PREFIX_MIN_LENGTH = 3 # a one- or two-character prefix matches too widely to carry signal
|
||||
EXCERPT_CHARS = 320
|
||||
GREP_TIMEOUT = 30
|
||||
GREP_MAX_PER_FILE = 5
|
||||
|
||||
|
||||
def rrf_merge(ranked_lists: list[list[int]]) -> list[tuple[int, float]]:
|
||||
"""score(id) = sum over lists of 1 / (RRF_K + rank). No weights: both halves count equally."""
|
||||
scores: dict[int, float] = {}
|
||||
for ids in ranked_lists:
|
||||
for rank, chunk_id in enumerate(ids, start=1):
|
||||
scores[chunk_id] = scores.get(chunk_id, 0.0) + 1.0 / (RRF_K + rank)
|
||||
return sorted(scores.items(), key=lambda item: (-item[1], item[0]))
|
||||
|
||||
|
||||
def fts_match_expression(query: str) -> str:
|
||||
"""Turn free text into an FTS5 OR-query of quoted prefix terms.
|
||||
|
||||
Terms are quoted so reserved words (`and`, `not`, `near`) and punctuation cannot be
|
||||
read as operators. The trailing `*` covers Czech inflection, which `unicode61` does
|
||||
not stem — `záloh*` finds záloha/zálohování/zálohy.
|
||||
"""
|
||||
terms = []
|
||||
for word in _tokenize(query):
|
||||
wildcard = "*" if len(word) >= PREFIX_MIN_LENGTH else ""
|
||||
terms.append(f'"{word}"{wildcard}')
|
||||
return " OR ".join(terms)
|
||||
|
||||
|
||||
def _tokenize(text: str) -> list[str]:
|
||||
return "".join(char if char.isalnum() else " " for char in text).split()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# search
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def run_search(config: WikiConfig, query: str, limit: int) -> int:
|
||||
expression = fts_match_expression(query)
|
||||
if not expression:
|
||||
print("(empty query)")
|
||||
return 0
|
||||
|
||||
with store.connection(db_path()) as conn:
|
||||
try:
|
||||
require_matching_index(store.read_meta(conn), config.embedding)
|
||||
except IndexIdentityMismatch as exc:
|
||||
print(str(exc), file=sys.stderr)
|
||||
return 1
|
||||
|
||||
bm25_ids = store.bm25_ranked_ids(conn, expression, CANDIDATE_LIMIT)
|
||||
|
||||
vector_ids: list[int] = []
|
||||
degraded: str | None = None
|
||||
try:
|
||||
vector = OllamaEmbedder(config.embedding).embed_query(query)
|
||||
vector_ids = store.knn_ranked_ids(conn, vector, CANDIDATE_LIMIT)
|
||||
except EmbeddingUnavailable as exc:
|
||||
degraded = f"note: embeddings unavailable ({exc}) — FTS-only results"
|
||||
|
||||
pending = store.index_stats(conn)["pending"]
|
||||
merged = rrf_merge([ids for ids in (bm25_ids, vector_ids) if ids])[:limit]
|
||||
rows = store.fetch_chunks(conn, [chunk_id for chunk_id, _ in merged])
|
||||
|
||||
if degraded:
|
||||
print(degraded)
|
||||
elif pending:
|
||||
print(f"note: {pending} chunks still awaiting vectors — semantic half is incomplete")
|
||||
|
||||
if not merged:
|
||||
print("(no matches)")
|
||||
return 0
|
||||
|
||||
bm25_rank = {chunk_id: rank for rank, chunk_id in enumerate(bm25_ids, start=1)}
|
||||
vector_rank = {chunk_id: rank for rank, chunk_id in enumerate(vector_ids, start=1)}
|
||||
for position, (chunk_id, score) in enumerate(merged, start=1):
|
||||
row = rows.get(chunk_id)
|
||||
if row is None:
|
||||
continue
|
||||
origin = _origin_label(bm25_rank.get(chunk_id), vector_rank.get(chunk_id))
|
||||
print(f"{position}. {row['source_id']}:{row['path']} (rrf {score:.4f}, {origin})")
|
||||
print(f" {row['breadcrumb']}")
|
||||
print(_indent(_excerpt(row["text"], row["breadcrumb"])))
|
||||
print()
|
||||
return 0
|
||||
|
||||
|
||||
def _origin_label(bm25: int | None, vector: int | None) -> str:
|
||||
parts = []
|
||||
if bm25 is not None:
|
||||
parts.append(f"bm25 #{bm25}")
|
||||
if vector is not None:
|
||||
parts.append(f"vec #{vector}")
|
||||
return ", ".join(parts)
|
||||
|
||||
|
||||
def _excerpt(text: str, breadcrumb: str) -> str:
|
||||
body = text[len(breadcrumb) :].lstrip("\n") if text.startswith(breadcrumb) else text
|
||||
body = " ".join(body.split())
|
||||
return body[:EXCERPT_CHARS] + ("…" if len(body) > EXCERPT_CHARS else "")
|
||||
|
||||
|
||||
def _indent(text: str) -> str:
|
||||
return "\n".join(f" {line}" for line in text.splitlines())
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# toc
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def run_toc(source_id: str | None, tag: str | None) -> int:
|
||||
with store.connection(db_path()) as conn:
|
||||
rows = store.list_toc_files(conn, source_id=source_id, tag=tag)
|
||||
if not rows:
|
||||
print("(no indexed files match)")
|
||||
return 0
|
||||
|
||||
by_source: dict[str, list[sqlite3.Row]] = {}
|
||||
for row in rows:
|
||||
by_source.setdefault(row["source_id"], []).append(row)
|
||||
|
||||
for source, files in by_source.items():
|
||||
print(f"{source} ({len(files)} files)")
|
||||
print()
|
||||
directory = None
|
||||
for row in files:
|
||||
path = row["path"]
|
||||
parent, _, name = path.rpartition("/")
|
||||
if parent != directory:
|
||||
directory = parent
|
||||
print(f"{parent}/" if parent else "./")
|
||||
print(f" {name:<28} {row['title'] or '':<32} {_tag_label(row['tags'])}".rstrip())
|
||||
print()
|
||||
return 0
|
||||
|
||||
|
||||
def _tag_label(raw: str | None) -> str:
|
||||
try:
|
||||
tags = json.loads(raw) if raw else []
|
||||
except json.JSONDecodeError:
|
||||
return ""
|
||||
return f"[{', '.join(tags)}]" if tags else ""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# grep
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def grep_roots(source: SourceConfig) -> list[Path]:
|
||||
"""Directories ripgrep should walk for one source.
|
||||
|
||||
A git clone is searched whole. A workspace source is bounded by the literal prefix
|
||||
of each `paths` glob, so grep stays inside what the source claims without being
|
||||
narrowed to the indexed extensions.
|
||||
"""
|
||||
root = source_root(source)
|
||||
if source.kind == GIT_KIND:
|
||||
return [root] if root.is_dir() else []
|
||||
roots = []
|
||||
for pattern in source.paths:
|
||||
prefix = _literal_prefix(pattern)
|
||||
candidate = root / prefix if prefix else root
|
||||
if candidate.is_dir() and candidate not in roots:
|
||||
roots.append(candidate)
|
||||
return roots
|
||||
|
||||
|
||||
def _literal_prefix(pattern: str) -> str:
|
||||
parts = []
|
||||
for segment in pattern.split("/"):
|
||||
if any(char in segment for char in "*?["):
|
||||
break
|
||||
parts.append(segment)
|
||||
return "/".join(parts)
|
||||
|
||||
|
||||
def run_grep(config: WikiConfig, pattern: str, source_id: str | None) -> int:
|
||||
if shutil.which("rg") is None:
|
||||
print("ripgrep (rg) not found", file=sys.stderr)
|
||||
return 1
|
||||
sources = [s for s in config.sources if not source_id or s.source_id == source_id]
|
||||
if not sources:
|
||||
print(f"unknown source {source_id!r}", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
roots = [path for source in sources for path in grep_roots(source)]
|
||||
if not roots:
|
||||
print("(nothing on disk to grep — has wiki_sync.py run?)")
|
||||
return 0
|
||||
|
||||
command = [
|
||||
"rg",
|
||||
"--line-number",
|
||||
"--no-heading",
|
||||
"--color",
|
||||
"never",
|
||||
"--smart-case",
|
||||
"--max-count",
|
||||
str(GREP_MAX_PER_FILE),
|
||||
"--glob",
|
||||
"!.git",
|
||||
pattern,
|
||||
*[str(path) for path in roots],
|
||||
]
|
||||
try:
|
||||
result = subprocess.run(command, capture_output=True, text=True, timeout=GREP_TIMEOUT, check=False)
|
||||
except subprocess.TimeoutExpired:
|
||||
print(f"grep timed out after {GREP_TIMEOUT}s", file=sys.stderr)
|
||||
return 1
|
||||
if result.returncode not in (0, 1):
|
||||
print(result.stderr.strip(), file=sys.stderr)
|
||||
return 1
|
||||
output = result.stdout.strip()
|
||||
print(output if output else "(no matches)")
|
||||
return 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# cli
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Search the wiki index over your notes.")
|
||||
subparsers = parser.add_subparsers(dest="command", required=True)
|
||||
|
||||
search = subparsers.add_parser("search", help="hybrid BM25 + vector search over chunks")
|
||||
search.add_argument("query", help="free text; Czech inflection is covered by prefix matching")
|
||||
search.add_argument("--limit", type=int, default=DEFAULT_LIMIT, help="chunks to return")
|
||||
|
||||
toc = subparsers.add_parser("toc", help="directory -> file -> title + tags from the catalog")
|
||||
toc.add_argument("--source", help="limit to one source id")
|
||||
toc.add_argument("--tag", help="only files carrying this frontmatter tag")
|
||||
|
||||
grep = subparsers.add_parser("grep", help="live ripgrep over the files on disk, index-free")
|
||||
grep.add_argument("pattern", help="ripgrep regex")
|
||||
grep.add_argument("--source", help="limit to one source id")
|
||||
|
||||
return parser.parse_args(argv)
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
args = parse_args(argv)
|
||||
try:
|
||||
config = load_config()
|
||||
except ConfigError as exc:
|
||||
print(f"config error: {exc}", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
if args.command == "search":
|
||||
return run_search(config, args.query, args.limit)
|
||||
if args.command == "toc":
|
||||
return run_toc(args.source, args.tag)
|
||||
return run_grep(config, args.pattern, args.source)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
Reference in New Issue
Block a user