Files
hermes-agent/tools/tool_search_catalog.py
T

326 lines
12 KiB
Python

"""Deferred-tool catalog for tool search: BM25 retrieval over deferrable tool
defs plus the budgeted, byte-stable catalog listing embedded in the bridge."""
from __future__ import annotations
import functools
import math
import re
import threading
from collections import Counter
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Tuple
import snowballstemmer
from tools.tool_search_names import TOOL_CALL_NAME, TOOL_DESCRIBE_NAME, TOOL_SEARCH_NAME
# Chars-per-token rule of thumb for budget estimates; 4.0 slightly
# underestimates, which is the safer direction (fewer false activations).
CHARS_PER_TOKEN = 4.0
@dataclass
class CatalogEntry:
"""One deferrable tool, in a form the bridge tools can search and serve."""
name: str
description: str
schema: Dict[str, Any] # The full {"type":"function", "function": {...}} entry.
source: str # "mcp" | "plugin" | "other"
source_name: str # Toolset name, e.g. "mcp-github" or "kanban"
_tokens: List[str] = field(default_factory=list) # pre-tokenized for BM25
_TOKEN_RE = re.compile(r"[A-Za-z0-9]+")
# Snowball stemmers carry mutable parsing state and bridge dispatch runs on
# parallel tool-call threads, so: one stemmer per thread, created lazily.
_thread_local = threading.local()
def _stemmer() -> Any:
st = getattr(_thread_local, "stemmer", None)
if st is None:
st = snowballstemmer.stemmer("english")
_thread_local.stemmer = st
return st
@functools.lru_cache(maxsize=16384)
def _stem(token: str) -> str:
"""Stem one token, memoized across stateless catalog rebuilds."""
return _stemmer().stemWord(token)
def _tokenize(text: str) -> List[str]:
"""Lowercase alphanumeric tokens, Snowball-stemmed (English).
Shared by the index path and the query path so "issues" matches ``create_issue``.
"""
if not text:
return []
return [_stem(token.lower()) for token in _TOKEN_RE.findall(text)]
def _fn(td: Dict[str, Any]) -> Dict[str, Any]:
"""The ``function`` block of a tool-def (``{}`` when absent/None)."""
return td.get("function") or {}
def _registry_entry(name: str) -> Any:
"""Registry entry for ``name``; None when unregistered OR when the registry
is unavailable/raises (lookup failures must never fail a bridge call).
The import stays lazy: tests patch ``tools.registry.registry``."""
try:
from tools.registry import registry
return registry.get_entry(name)
except Exception:
return None
def _entry_search_text(td: Dict[str, Any], source_label: str = "") -> str:
"""Search-text blob: split name words + source label + description +
top-level parameter names. Schema bodies are excluded (noise, no recall
gain). The ``mcp__`` prefix is dropped — it is in every MCP document, so
its IDF is ~0. The source label lets a service-name query ("linear") reach
a tool whose own name omits the vendor.
"""
fn = _fn(td)
name = fn.get("name", "")
if name.startswith("mcp__"):
name = name[len("mcp__"):]
desc = fn.get("description", "") or ""
params = ((fn.get("parameters") or {}).get("properties") or {})
param_names = " ".join(params.keys())
name_words = re.sub(r"[_.:-]", " ", name)
extra = source_label if source_label and source_label not in name_words.split() else ""
return f"{name_words} {extra} {desc} {param_names}"
def _classify_source(name: str) -> Tuple[str, str]:
"""Return (source_kind, source_name) for a registered tool name."""
entry = _registry_entry(name)
if entry is None:
return ("other", "")
try:
return ("mcp" if entry.toolset.startswith("mcp-") else "plugin", entry.toolset)
except Exception: # malformed entry (no str toolset)
return ("other", "")
def build_catalog(tool_defs: List[Dict[str, Any]]) -> List[CatalogEntry]:
"""Build the deferred-tool catalog from the deferrable subset of tool-defs."""
catalog: List[CatalogEntry] = []
for td in tool_defs:
fn = _fn(td)
name = fn.get("name", "")
if not name:
continue
source, source_name = _classify_source(name)
# Index the human-facing label ("linear", not "mcp-linear").
source_label = _listing_group_label(source_name) if source_name else ""
catalog.append(CatalogEntry(
name=name,
description=fn.get("description", "") or "",
schema=td,
source=source,
source_name=source_name,
_tokens=_tokenize(_entry_search_text(td, source_label)),
))
return catalog
def _bm25_score(query_tokens: List[str], doc_tokens: List[str],
doc_lengths: List[int], avg_dl: float,
doc_freq: Dict[str, int], n_docs: int,
k1: float = 1.5, b: float = 0.75) -> float:
"""Standard BM25 score for one query against one document (inlined; the
catalog is bounded — typically < 500 tools — so a dependency is not worth it)."""
if not doc_tokens:
return 0.0
score = 0.0
dl = len(doc_tokens)
doc_tf = Counter(doc_tokens)
for q in query_tokens:
df = doc_freq.get(q, 0)
if df == 0:
continue
idf = math.log(1 + (n_docs - df + 0.5) / (df + 0.5))
tf = doc_tf.get(q, 0)
if tf == 0:
continue
norm = tf * (k1 + 1) / (tf + k1 * (1 - b + b * dl / max(avg_dl, 1.0)))
score += idf * norm
return score
_CorpusStats = Tuple[List[int], float, Dict[str, int], int]
def _corpus_stats(catalog: List[CatalogEntry]) -> _CorpusStats:
"""Compute the BM25 statistics shared by every query over a catalog."""
doc_lengths = [len(entry._tokens) for entry in catalog]
avg_dl = sum(doc_lengths) / max(len(doc_lengths), 1)
doc_freq: Dict[str, int] = Counter()
for entry in catalog:
doc_freq.update(set(entry._tokens))
return doc_lengths, avg_dl, dict(doc_freq), len(catalog)
def search_catalog(
catalog: List[CatalogEntry],
query: str,
limit: int = 5,
*,
corpus_stats: Optional[_CorpusStats] = None,
) -> List[CatalogEntry]:
"""Top-``limit`` catalog entries for ``query`` by BM25 (exact name match
ranks first). Falls back to a name-substring match only when NO query
token appears in any document (e.g. "hub" vs ``github_*``); the IDF
variant is strictly positive, so a hit anywhere suppresses the fallback.
"""
if not catalog or limit <= 0:
return []
query_tokens = _tokenize(query)
if not query_tokens:
return []
if corpus_stats is None:
corpus_stats = _corpus_stats(catalog)
doc_lengths, avg_dl, doc_freq, n_docs = corpus_stats
scored: List[Tuple[float, CatalogEntry]] = []
exact_name = query.strip().lower()
for entry in catalog:
if entry.name.lower() == exact_name:
scored.append((float("inf"), entry))
continue
s = _bm25_score(query_tokens, entry._tokens, doc_lengths, avg_dl,
doc_freq, n_docs)
if s > 0:
scored.append((s, entry))
if not scored:
ql = query.lower()
for entry in catalog:
if ql in entry.name.lower():
scored.append((0.1, entry))
scored.sort(key=lambda x: x[0], reverse=True)
return [e for _, e in scored[:limit]]
# Sentence end: ., !, ? followed by whitespace/EOS, not inside e.g./i.e./etc.
_SENTENCE_END_RE = re.compile(r"(?<!\be\.g)(?<!\bi\.e)(?<!\betc)[.!?](?=\s|$)")
def _short_desc(description: str, max_chars: int = 60) -> str:
"""First sentence of a tool description, clipped to ``max_chars`` on a
word boundary. ``e.g.``/``i.e.``/``etc.`` do not end a sentence; whitespace
normalization and the regex search stay linear-time on hostile input."""
text = " ".join((description or "").split())
if not text:
return ""
m = _SENTENCE_END_RE.search(text)
if m:
text = text[:m.end()]
if len(text) <= max_chars:
return text
clipped = text[:max_chars]
if " " in clipped:
clipped = clipped.rsplit(" ", 1)[0]
return clipped.rstrip(",;: ") + "…"
def _listing_group_label(source_name: str) -> str:
"""Human-facing group heading for a toolset, e.g. ``mcp-github`` -> ``github``."""
label = source_name or "other"
return label[4:] if label.startswith("mcp-") else label
def build_catalog_listing(
deferrable: List[Dict[str, Any]],
*,
max_tokens: int = 4000,
) -> Optional[str]:
"""Render the deferred-catalog manifest; text only (see
:func:`build_catalog_listing_with_form` for the degradation ladder)."""
return build_catalog_listing_with_form(deferrable, max_tokens=max_tokens)[0]
def build_catalog_listing_with_form(
deferrable: List[Dict[str, Any]],
*,
max_tokens: int = 4000,
) -> Tuple[Optional[str], str]:
"""Render the skills-style deferred-catalog manifest: ``- name: short desc``
lines grouped under a heading per source (MCP server / plugin toolset).
Returns ``(text, form)``; ``form`` is ``"full"``, ``"names"`` (names-only),
``"mixed"`` (oversized servers collapsed to a name + count summary line,
small ones keep per-tool lines), ``"groups"`` (every server summarized),
or ``"none"`` (over budget even summarized -> text is None).
Ordering is deterministic (sorted groups and tools) so the block is
byte-stable across assemblies — the request prefix stays cacheable.
Degradation is PER SERVER (largest first): one huge server must not cost
a small co-attached server its listing.
"""
if not deferrable:
return None, "none"
groups: Dict[str, List[Tuple[str, str]]] = {}
for td in deferrable:
fn = _fn(td)
name = fn.get("name", "")
if not name:
continue
# ``_classify_source`` returns ("other", "") for unregistered names and
# ``_listing_group_label("")`` is "other", so one call covers both.
label = _listing_group_label(_classify_source(name)[1])
groups.setdefault(label, []).append((name, _short_desc(fn.get("description", ""))))
if not groups:
return None, "none"
def render_group(label: str, mode: str) -> str:
"""Render one server's block. mode: 'full' | 'names' | 'summary'."""
tools = sorted(groups[label])
if mode == "summary":
return (f"{label} ({len(tools)} tools — names not listed; "
f"discover via `{TOOL_SEARCH_NAME}`)")
lines = [f"{label} tools ({len(tools)}):"]
if mode == "full":
lines.extend(f"- {name}: {desc}" if desc else f"- {name}" for name, desc in tools)
else:
lines.append(", ".join(name for name, _ in tools))
return "\n".join(lines)
header = ("Deferred tool catalog (call schemas via "
f"`{TOOL_DESCRIBE_NAME}`, invoke via `{TOOL_CALL_NAME}`):")
def assemble_if_fits(modes: Dict[str, str]) -> Optional[str]:
text = "\n".join([header] + [render_group(lbl, modes[lbl]) for lbl in sorted(groups)])
return text if math.ceil(len(text) / CHARS_PER_TOKEN) <= max_tokens else None
# 1. Everything full. 2. Everything names-only.
for mode in ("full", "names"):
modes = {lbl: mode for lbl in groups}
text = assemble_if_fits(modes)
if text is not None:
return text, mode
# 3. Per-server degradation: collapse the LARGEST rendered groups first
# (deterministic: size then label) so one oversized server does not
# cost a small co-attached server its per-tool names.
by_size = sorted(groups, key=lambda lbl: (-len(render_group(lbl, "names")), lbl))
for lbl in by_size:
modes[lbl] = "summary"
text = assemble_if_fits(modes)
if text is not None:
form = "groups" if all(m == "summary" for m in modes.values()) else "mixed"
return text, form
return None, "none"