refactor(tools): group H — BM25 loop fold, catalog/selection/limits micro-collapses, sync-manager walrus

This commit is contained in:
Teknium
2026-09-03 01:13:51 -07:00
parent 31017623e9
commit e0df9656bb
7 changed files with 33 additions and 48 deletions

View File

@@ -22,8 +22,8 @@ def tip_tool(text: str, selector: str, title: str = "", side: str = "") -> str:
"what's on screen and prefer a target reporting stable: true.")
if side and side not in SIDES:
return tool_error(f"side must be one of: {', '.join(SIDES)}.")
payload = {"selector": selector, "text": text}
payload.update({k: v for k, v in (("title", title), ("side", side)) if v})
payload = {"selector": selector, "text": text,
**{k: v for k, v in (("title", title), ("side", side)) if v}}
try:
ok = desktop_ui.emit("tip.show", payload)
except Exception as exc:

View File

@@ -123,8 +123,7 @@ def resolve_provider_secret(env_var: str, provider_id: str, config_value: str =
return ""
except Exception: # pragma: no cover — secret_scope is in-repo
pass
key = (str(env_getter(env_var) or "").strip() if env_getter is not None
else _dotenv_value(env_var))
key = str(env_getter(env_var) or "").strip() if env_getter else _dotenv_value(env_var)
if key or not provider_id:
return key
try:
@@ -195,9 +194,9 @@ def read_selection(section: str) -> str | None:
if is_truthy_value(raw.get("use_gateway")):
return NOUS_MANAGED_PROVIDER
for key in _SELECTION_NAME_KEYS.get(section, _DEFAULT_NAME_KEYS):
value = raw.get(key)
if value is not None and str(value).strip():
return str(value).strip().lower()
text = str(raw.get(key)).strip().lower() if raw.get(key) is not None else ""
if text:
return text
# use_gateway: false with no name key is not a usable selection shape;
# per-capability web keys still count as configured via selection_exists().
return None

View File

@@ -24,8 +24,7 @@ def _coerce_int(value: Any, default: int, minimum: int) -> int:
def _coerce_positive_int(value: Any, default: int) -> int:
"""Return ``value`` as a positive int, or ``default`` on any issue."""
return _coerce_int(value, default, 1)
return _coerce_int(value, default, 1) # positive int, or ``default`` on any issue
def get_tool_output_limits() -> Dict[str, int]:

View File

@@ -105,12 +105,11 @@ def _sandbox_visible_spillover_path(host_path: str, env) -> str | None:
except Exception as exc:
logger.debug("Spillover path translation failed: %s", exc)
return None
sync_manager = getattr(env, "_sync_manager", None)
if sync_manager is not None:
try:
try:
if (sync_manager := getattr(env, "_sync_manager", None)) is not None:
sync_manager.sync(force=True)
except Exception as exc:
logger.debug("Spillover sync failed: %s", exc)
except Exception as exc:
logger.debug("Spillover sync failed: %s", exc)
try:
if env.execute(f"test -r {shlex.quote(visible)}", timeout=15).get("returncode", 1) == 0:
return visible

View File

@@ -172,13 +172,12 @@ def _deferrable_in(tool_defs: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
def estimate_tokens_from_schemas(tool_defs: Iterable[Dict[str, Any]]) -> int:
"""Token cost via the chars/4 rule (order-of-magnitude precision suffices)."""
total_chars = 0
for td in tool_defs:
def _chars(td: Dict[str, Any]) -> int:
try:
total_chars += len(json.dumps(td, ensure_ascii=False, separators=(",", ":")))
return len(json.dumps(td, ensure_ascii=False, separators=(",", ":")))
except (TypeError, ValueError):
total_chars += len(str(td))
return int(math.ceil(total_chars / CHARS_PER_TOKEN))
return len(str(td))
return int(math.ceil(sum(map(_chars, tool_defs)) / CHARS_PER_TOKEN))
def should_activate(config: ToolSearchConfig, deferrable_tokens: int,
@@ -366,8 +365,7 @@ def _string_list_arg(args: Dict[str, Any], key: str, *, dedupe: bool, max_items:
"""Read a list-of-strings bridge argument -> ``(items, error_json)``. A bare string (a
common model slip) is a one-item list; rejects non-lists, all-blank lists, > ``max_items``."""
raw = args.get(key)
if isinstance(raw, str):
raw = [raw]
raw = [raw] if isinstance(raw, str) else raw
if not isinstance(raw, list):
return None, tool_error(f"{key} is required and must be an array of strings")
out: List[str] = []

View File

@@ -44,10 +44,9 @@ def _stem(token: str) -> str:
"""Stem one token, memoized across stateless catalog rebuilds. Snowball stemmers carry
mutable parsing state and bridge dispatch runs on parallel tool-call threads, so the
stemmer is one-per-thread, created lazily."""
st = getattr(_thread_local, "stemmer", None)
if st is None:
st = _thread_local.stemmer = snowballstemmer.stemmer("english")
return st.stemWord(token)
if getattr(_thread_local, "stemmer", None) is None:
_thread_local.stemmer = snowballstemmer.stemmer("english")
return _thread_local.stemmer.stemWord(token)
def _tokenize(text: str) -> List[str]:
@@ -122,22 +121,18 @@ def _bm25_score(query_tokens: List[str], doc_tokens: List[str], doc_lengths: Lis
b: float = 0.75) -> float:
"""Standard BM25 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)
tf = doc_tf.get(q, 0)
if df == 0 or tf == 0:
continue
idf = math.log(1 + (n_docs - df + 0.5) / (df + 0.5))
score += idf * tf * (k1 + 1) / (tf + k1 * (1 - b + b * dl / max(avg_dl, 1.0)))
df, tf = doc_freq.get(q, 0), doc_tf.get(q, 0)
if df and tf:
idf = math.log(1 + (n_docs - df + 0.5) / (df + 0.5))
score += idf * tf * (k1 + 1) / (tf + k1 * (1 - b + b * dl / max(avg_dl, 1.0)))
return score
_CorpusStats = Tuple[List[int], float, Dict[str, int], int]
_CorpusStats = Tuple[List[int], float, Dict[str, int], int] # doc_lengths, avg_dl, df, n_docs
def _corpus_stats(catalog: List[CatalogEntry]) -> _CorpusStats:
@@ -157,8 +152,7 @@ def search_catalog(catalog: List[CatalogEntry], query: str, limit: int = 5, *,
query_tokens = _tokenize(query) if catalog and limit > 0 else []
if not query_tokens:
return []
if corpus_stats is None:
corpus_stats = _corpus_stats(catalog)
corpus_stats = corpus_stats or _corpus_stats(catalog)
scored: List[Tuple[float, CatalogEntry]] = []
exact_name = query.strip().lower()
for entry in catalog:
@@ -182,13 +176,11 @@ def _short_desc(description: str, max_chars: int = 60) -> str:
search stay linear-time on hostile input."""
text = " ".join((description or "").split())
m = _SENTENCE_END_RE.search(text)
if m:
text = text[:m.end()]
text = text[:m.end()] if m else text
if len(text) <= max_chars:
return text
clipped = text[:max_chars]
if " " in clipped:
clipped = clipped.rsplit(" ", 1)[0]
clipped = clipped.rsplit(" ", 1)[0] if " " in clipped else clipped
return clipped.rstrip(",;: ") + "…"

View File

@@ -23,10 +23,9 @@ def _schema_for_local_validation(node: Any) -> Any:
if not isinstance(node, dict):
return node
# Literal keywords hold instance data, not schemas: copy byte-for-byte.
normalized = {
key: (copy.deepcopy(value) if key in _SCHEMA_LITERAL_KEYS
else _schema_for_local_validation(value))
for key, value in node.items() if key != "nullable"}
normalized = {key: (copy.deepcopy(value) if key in _SCHEMA_LITERAL_KEYS
else _schema_for_local_validation(value))
for key, value in node.items() if key != "nullable"}
if node.get("nullable") is not True:
return normalized
schema_type = normalized.get("type")
@@ -49,10 +48,9 @@ def _schema_has_external_ref(node: Any) -> bool:
if not isinstance(node, dict):
return False
ref = node.get("$ref")
if isinstance(ref, str) and not ref.startswith("#"):
return True
return any(_schema_has_external_ref(value) for key, value in node.items()
if key not in _SCHEMA_LITERAL_KEYS)
return (isinstance(ref, str) and not ref.startswith("#")) or any(
_schema_has_external_ref(value) for key, value in node.items()
if key not in _SCHEMA_LITERAL_KEYS)
def _validation_path(error: Any) -> str: