From e0df9656bb1cefb720e7ad363dd8377a15291950 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:13:51 -0700 Subject: [PATCH] =?UTF-8?q?refactor(tools):=20group=20H=20=E2=80=94=20BM25?= =?UTF-8?q?=20loop=20fold,=20catalog/selection/limits=20micro-collapses,?= =?UTF-8?q?=20sync-manager=20walrus?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/tip_tool.py | 4 ++-- tools/tool_backend_helpers.py | 9 ++++----- tools/tool_output_limits.py | 3 +-- tools/tool_result_storage.py | 9 ++++----- tools/tool_search.py | 12 +++++------- tools/tool_search_catalog.py | 30 +++++++++++------------------- tools/tool_search_validation.py | 14 ++++++-------- 7 files changed, 33 insertions(+), 48 deletions(-) diff --git a/tools/tip_tool.py b/tools/tip_tool.py index 37c3d44133..e4d33259a2 100644 --- a/tools/tip_tool.py +++ b/tools/tip_tool.py @@ -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: diff --git a/tools/tool_backend_helpers.py b/tools/tool_backend_helpers.py index 25ce805244..1087852ffd 100644 --- a/tools/tool_backend_helpers.py +++ b/tools/tool_backend_helpers.py @@ -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 diff --git a/tools/tool_output_limits.py b/tools/tool_output_limits.py index 550cc52a5a..38e539c9fb 100644 --- a/tools/tool_output_limits.py +++ b/tools/tool_output_limits.py @@ -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]: diff --git a/tools/tool_result_storage.py b/tools/tool_result_storage.py index 08f956b333..c96619e6f2 100644 --- a/tools/tool_result_storage.py +++ b/tools/tool_result_storage.py @@ -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 diff --git a/tools/tool_search.py b/tools/tool_search.py index 511c668d00..8387ce7772 100644 --- a/tools/tool_search.py +++ b/tools/tool_search.py @@ -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] = [] diff --git a/tools/tool_search_catalog.py b/tools/tool_search_catalog.py index 06b305d25e..52f80dc583 100644 --- a/tools/tool_search_catalog.py +++ b/tools/tool_search_catalog.py @@ -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(",;: ") + "…" diff --git a/tools/tool_search_validation.py b/tools/tool_search_validation.py index 1fbefc6337..511e8d5446 100644 --- a/tools/tool_search_validation.py +++ b/tools/tool_search_validation.py @@ -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: