From cccc84a237d8ce27abb46e8a4d924d06e8db2ff2 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 23:14:03 -0700 Subject: [PATCH] refactor(tools): fold debug-error into _finish_debug, clamp-or-default helper, compact cache/policy control flow --- tools/web_result_cache.py | 24 +++++------------ tools/web_tools.py | 51 ++++++++++++++++--------------------- tools/web_tools_extract.py | 23 +++++++---------- tools/web_tools_rescue.py | 28 +++++++++++--------- tools/web_tools_truncate.py | 33 +++++++++++------------- tools/website_policy.py | 15 ++++------- 6 files changed, 74 insertions(+), 100 deletions(-) diff --git a/tools/web_result_cache.py b/tools/web_result_cache.py index e6a99c8567..3a8fa66f36 100644 --- a/tools/web_result_cache.py +++ b/tools/web_result_cache.py @@ -39,8 +39,7 @@ def _web_config() -> dict: def cache_enabled() -> bool: """Both caches honor ``web.cache_enabled`` (default: on).""" - val = _web_config().get("cache_enabled") - return True if val is None else bool(val) + return True if (val := _web_config().get("cache_enabled")) is None else bool(val) def ttl_seconds() -> float: @@ -78,9 +77,7 @@ def _deep_copy(response: dict) -> dict: return json.loads(json.dumps(response)) -# --------------------------------------------------------------------------- -# Search memo (in-memory, single-flight) -# --------------------------------------------------------------------------- +# ─── Search memo (in-memory, single-flight) ─────────────────────────────────── class SearchMemo: """TTL memo + single-flight coalescer for search responses. Thread-safe: the parallel tool-dispatch pool @@ -88,7 +85,7 @@ class SearchMemo: wait for (and share) the winner's response.""" def __init__(self) -> None: - self._store: Dict[tuple, Tuple[float, dict]] = {} + self._store: Dict[tuple, Tuple[float, dict]] = {} # key -> (expires_at, response) self._store_lock = threading.Lock() self._key_locks: Dict[tuple, threading.Lock] = {} @@ -156,9 +153,7 @@ def slice_search_response(response: dict, limit: int) -> dict: return response -# --------------------------------------------------------------------------- -# Extract cache (disk-backed, reuses cache/web) -# --------------------------------------------------------------------------- +# ─── Extract cache (disk-backed, reuses cache/web) ──────────────────────────── _index_lock = threading.Lock() @@ -208,8 +203,7 @@ def _url_digest(url: str, format: Optional[str], provider: str = "") -> str: def _entry_file_path(url: str, format: Optional[str], provider: str) -> Optional[Path]: """Dedicated cache file per (url, format, provider) — deliberately NOT the truncate-store file (keyed on URL alone), which html/markdown or two providers' copies of one URL would overwrite.""" - d = _cache_dir() - if d is None: + if (d := _cache_dir()) is None: return None try: slug = _host_slug(url) @@ -222,12 +216,8 @@ def _host_matches_pattern(host: str, pattern: str) -> bool: """Case-insensitive: exact, ``*.wildcard``, or bare-domain suffix (``mysite.dev`` also matches ``preview.mysite.dev``).""" host = host.lower().strip(".") - pattern = (pattern or "").lower().strip().strip(".") - if not pattern: - return False - if pattern.startswith("*."): - pattern = pattern[2:] - return host == pattern or host.endswith("." + pattern) + pattern = (pattern or "").lower().strip().strip(".").removeprefix("*.") + return bool(pattern) and (host == pattern or host.endswith("." + pattern)) def _is_cache_exempt_host(url: str) -> bool: diff --git a/tools/web_tools.py b/tools/web_tools.py index 14730963a4..0f76c42f44 100644 --- a/tools/web_tools.py +++ b/tools/web_tools.py @@ -127,12 +127,10 @@ def _get_backend() -> str: Autodetect runs ONLY when no web selection has ever been stored.""" configured = _configured_backend() if configured: - # "nous" (managed subscription) is serviced by the firecrawl provider, whose client - # resolver routes it through the managed Tool Gateway. + # "nous" (managed subscription) is serviced by firecrawl, routed through the managed Tool Gateway. return "firecrawl" if configured == NOUS_MANAGED_PROVIDER else configured if selection_exists("web"): - # Selection exists (use_gateway / per-capability keys) but no shared name: keep the - # firecrawl default rather than credential-laddering. + # Selection exists (use_gateway / per-capability keys) but no shared name: firecrawl, no ladder. return "firecrawl" # Never-configured install. Explicit user credentials beat the managed-gateway probe (a Nous OAuth @@ -226,10 +224,9 @@ def _is_backend_available(backend: str) -> bool: """True when *backend* is usable — the single availability chokepoint. Non-legacy names delegate to the registered provider's ``is_available()`` (unregistered names fall through); built-ins use cheap probes.""" backend = (backend or "").lower().strip() - if backend not in _LEGACY_WEB_BACKENDS: - provider = _registered_web_provider(backend) - if provider is not None: - return _probe(provider, "is_available") or False + provider = None if backend in _LEGACY_WEB_BACKENDS else _registered_web_provider(backend) + if provider is not None: + return _probe(provider, "is_available") or False probe = _BUILTIN_AVAILABILITY.get(backend) return probe() if probe else False @@ -262,17 +259,14 @@ def _ensure_web_plugins_loaded() -> None: logger.warning("Web plugin discovery failed (non-fatal): %s", exc) -def _finish_debug(call_name: str, debug_call_data: dict) -> None: +def _finish_debug(call_name: str, debug_call_data: dict, error_msg: Optional[str] = None) -> Optional[str]: + """Log the call into the debug session; with *error_msg*, record it and return its ``tool_error`` envelope.""" + if error_msg is not None: + logger.debug("%s", error_msg) + debug_call_data["error"] = error_msg _debug.log_call(call_name, debug_call_data) _debug.save() - - -def _debug_error(call_name: str, debug_call_data: dict, error_msg: str) -> str: - """Record *error_msg* in the debug session and return the ``tool_error`` envelope for it.""" - logger.debug("%s", error_msg) - debug_call_data["error"] = error_msg - _finish_debug(call_name, debug_call_data) - return tool_error(error_msg) + return None if error_msg is None else tool_error(error_msg) def web_search_tool(query: str, limit: int = 5) -> str: @@ -324,7 +318,7 @@ def web_search_tool(query: str, limit: int = 5) -> str: _finish_debug("web_search_tool", debug_call_data) return result_json except Exception as e: - return _debug_error("web_search_tool", debug_call_data, f"Error searching web: {str(e)}") + return _finish_debug("web_search_tool", debug_call_data, f"Error searching web: {str(e)}") def _memoized_search(provider, query: str, limit: int) -> dict: @@ -333,6 +327,7 @@ def _memoized_search(provider, query: str, limit: int) -> dict: the caller's count is sliced out. Only successful, non-rescued responses are cached — caching a rescue would make the one-shot ring fallback sticky for a whole TTL.""" from tools.web_result_cache import bucket_limit, search_memo, slice_search_response + def _paid_search() -> tuple[dict, bool]: fetch_limit = bucket_limit(limit) try: @@ -407,9 +402,10 @@ async def web_extract_tool(urls: List[Any], format: str = None, char_limit: Opti debug_call_data["processing_applied"].append("truncate_and_store") _truncate_results(results, _effective_char_limit(char_limit), debug_call_data) trimmed = _trim_results(results) - result_json = json.dumps({"results": trimmed}, indent=2, ensure_ascii=False) if trimmed else tool_error( - "Content was inaccessible or not found" - ) + if not trimmed: + result_json = tool_error("Content was inaccessible or not found") + else: + result_json = json.dumps({"results": trimmed}, indent=2, ensure_ascii=False) # Belt-and-suspenders sweep of the serialized JSON: a provider may tuck a base64 blob in metadata. cleaned_result = convert_base64_images_to_links(result_json) debug_call_data["final_response_size"] = len(cleaned_result) @@ -417,7 +413,7 @@ async def web_extract_tool(urls: List[Any], format: str = None, char_limit: Opti _finish_debug("web_extract_tool", debug_call_data) return cleaned_result except Exception as e: - return _debug_error("web_extract_tool", debug_call_data, f"Error extracting content: {str(e)}") + return _finish_debug("web_extract_tool", debug_call_data, f"Error extracting content: {str(e)}") def _provider_is_ready(provider) -> bool: @@ -429,13 +425,10 @@ def _provider_is_ready(provider) -> bool: """ if provider is None: return False - for method in ("is_available", "is_keyless_available"): - ready = _probe(provider, method, " during readiness check") - if ready is None: # broken provider == not ready; don't try the next probe - return False - if ready: - return True - return False + ready = _probe(provider, "is_available", " during readiness check") + if ready is None: # broken provider == not ready; don't try the keyless probe + return False + return bool(ready or _probe(provider, "is_keyless_available", " during readiness check")) def check_web_api_key() -> bool: diff --git a/tools/web_tools_extract.py b/tools/web_tools_extract.py index fecb4f1ffc..fa95babe7d 100644 --- a/tools/web_tools_extract.py +++ b/tools/web_tools_extract.py @@ -20,6 +20,9 @@ logger = logging.getLogger("tools.web_tools") _NO_RESULT_ERROR = "Extract backend returned no result for this URL" _EXTRACT_BACKENDS_HINT = "firecrawl, tavily, keenable, exa, or parallel." +_INVALID_ITEM_ERROR = ( + "Invalid URL item at index {}: expected a URL string or an object with a string 'url' or 'href' field" +) def _web_extract_url(value: Any) -> Optional[str]: @@ -30,9 +33,7 @@ def _web_extract_url(value: Any) -> Optional[str]: """ if isinstance(value, dict): value = value.get("url") or value.get("href") - if not isinstance(value, str): - return None - return value.strip() or None + return (value.strip() or None) if isinstance(value, str) else None def _disabled_plugin_error(capability: str, disabled_key: str) -> str: @@ -86,25 +87,19 @@ def _merge_in_order( def _validate_extract_urls(urls: List[Any]): - """Normalize model-supplied items and block URLs carrying secrets. - - Returns ``(normalized_urls, normalized_indices, invalid_urls, blocked_json)``; - ``blocked_json`` is a whole-call refusal (exfiltration prevention) or None. - Percent-encoded secrets are caught by checking the unquoted forms too. - """ + """Normalize model-supplied items and block URLs carrying secrets (percent-encoded forms are unquoted + and checked too). Returns ``(normalized_urls, normalized_indices, invalid_urls, blocked_json)``; + ``blocked_json`` is a whole-call refusal (exfiltration prevention) or None.""" from agent.redact import _PREFIX_RE from urllib.parse import unquote + normalized_urls: List[str] = [] normalized_indices: List[int] = [] invalid_urls: Dict[int, Dict[str, Any]] = {} for index, item in enumerate(urls): _url = _web_extract_url(item) if _url is None: - invalid_urls[index] = _result_entry( - "", - f"Invalid URL item at index {index}: expected a URL string " - "or an object with a string 'url' or 'href' field", - ) + invalid_urls[index] = _result_entry("", _INVALID_ITEM_ERROR.format(index)) continue normalized_url = normalize_url_for_request(_url) if any(_PREFIX_RE.search(c) for c in (_url, unquote(_url), normalized_url, unquote(normalized_url))): diff --git a/tools/web_tools_rescue.py b/tools/web_tools_rescue.py index 5f548ebc5b..bfa57c6360 100644 --- a/tools/web_tools_rescue.py +++ b/tools/web_tools_rescue.py @@ -60,12 +60,13 @@ def _rescue_search(provider_name: str, original_error: str, query: str, limit: i ) rescued = search_with_failover(provider_name, query, limit) if rescued.get("success"): - data = rescued.setdefault("data", {}) - data["rescued_from"] = provider_name - data["backend_error"] = ( - f"Configured backend '{provider_name}' failed this call " - f"({(original_error or 'unknown error')[:300]}); result served by the keyless free tier. " - f"The next call will use '{provider_name}' again." + rescued.setdefault("data", {}).update( + rescued_from=provider_name, + backend_error=( + f"Configured backend '{provider_name}' failed this call " + f"({(original_error or 'unknown error')[:300]}); result served by the keyless free tier. " + f"The next call will use '{provider_name}' again." + ), ) return rescued # Ring also failed: the ORIGINAL error names the user's setup, so lead with it. @@ -80,16 +81,19 @@ def _rescue_search(provider_name: str, original_error: str, query: str, limit: i def _policy_blocked_result(result: dict) -> bool: """True for a website-policy refusal — intentional, never rescued (it would fetch blocked content).""" - if result.get("blocked_by_policy"): - return True - return "blocked by website policy" in str(result.get("error") or "").lower() + error = str(result.get("error") or "").lower() + return bool(result.get("blocked_by_policy")) or "blocked by website policy" in error def _rescue_extract(provider_name: str, urls: list, results: list) -> list: - """Rescue a whole-batch extract failure via the ring. Only genuine failures are re-fetched; policy-blocked - entries are preserved verbatim. If the provider broke url/result order parity, every entry is treated as - rescueable and the ring's list replaces the batch wholesale.""" + """Rescue a whole-batch extract failure via the ring. + + Only genuine failures are re-fetched; policy-blocked entries are preserved verbatim. If the provider + broke url/result order parity, every entry is treated as rescueable and the ring's list replaces the + batch wholesale. + """ from plugins.web.keyless_mcp import extract_with_failover + parity = len(results) == len(urls) rescue_idx = [i for i, r in enumerate(results) if not parity or not _policy_blocked_result(r)] if not rescue_idx: diff --git a/tools/web_tools_truncate.py b/tools/web_tools_truncate.py index c930b89ec3..c5e935f9bd 100644 --- a/tools/web_tools_truncate.py +++ b/tools/web_tools_truncate.py @@ -12,11 +12,11 @@ from typing import Any, List, Optional logger = logging.getLogger("tools.web_tools") -# Per-page char budget sent to the model (override: web.extract_char_limit); larger pages are -# head+tail truncated and the full text stored on disk. +# Per-page char budget sent to the model (override: web.extract_char_limit); larger pages are head+tail +# truncated, full text stored on disk. DEFAULT_EXTRACT_CHAR_LIMIT = 15000 -# Ceiling on the full-text file written to cache/web so a multi-MB page can't write unbounded bytes -# on every extract; the model only ever sees char_limit. +# Ceiling on the full-text file written to cache/web so a multi-MB page can't write unbounded bytes on +# every extract; the model only ever sees char_limit. MAX_STORED_TEXT_CHARS = 2_000_000 _CHAR_LIMIT_FLOOR, _CHAR_LIMIT_CEILING = 2000, 500_000 @@ -28,16 +28,18 @@ def _clamp_char_limit(value: Any) -> int: return max(_CHAR_LIMIT_FLOOR, min(int(value), _CHAR_LIMIT_CEILING)) +def _clamp_or_default(value: Any) -> int: + """``_clamp_char_limit(value)``; ``None`` or non-numeric input falls back to the default.""" + try: + return DEFAULT_EXTRACT_CHAR_LIMIT if value is None else _clamp_char_limit(value) + except (TypeError, ValueError): + return DEFAULT_EXTRACT_CHAR_LIMIT + + def _get_extract_char_limit() -> int: """``web.extract_char_limit`` clamped to a sane range, else the default.""" from tools.web_tools import _load_web_config # lazy: tests patch tools.web_tools._load_web_config - try: - configured = _load_web_config().get("extract_char_limit") - if configured is not None: - return _clamp_char_limit(configured) - except (TypeError, ValueError): - pass - return DEFAULT_EXTRACT_CHAR_LIMIT + return _clamp_or_default(_load_web_config().get("extract_char_limit")) def convert_base64_images_to_links(text: str) -> str: @@ -45,8 +47,7 @@ def convert_base64_images_to_links(text: str) -> str: (alt kept), parenthesised blobs, and bare ``data:image/...;base64,`` payloads. Real http(s) markdown image links are left untouched so the agent can ``web_extract`` / ``vision_analyze`` them.""" def _md_repl(m: "re.Match[str]") -> str: - alt = (m.group("alt") or "").strip() - return f"[IMAGE: {alt}]" if alt else "[IMAGE]" + return f"[IMAGE: {alt}]" if (alt := (m.group("alt") or "").strip()) else "[IMAGE]" out = re.sub(r"!\[(?P[^\]]*)\]\(\s*data:image/[^;]+;base64,[A-Za-z0-9+/=\s]+\)", _md_repl, text) out = re.sub(r"\(\s*data:image/[^;]+;base64,[A-Za-z0-9+/=\s]+\)", "[IMAGE]", out) @@ -126,11 +127,7 @@ def _truncate_with_footer(content: str, url: str, char_limit: int) -> tuple[str, def _effective_char_limit(char_limit: Optional[int]) -> int: """Caller's ``char_limit`` (else config) clamped; non-numeric input falls back to the default.""" - value = char_limit if char_limit is not None else _get_extract_char_limit() - try: - return _clamp_char_limit(value) - except (TypeError, ValueError): - return DEFAULT_EXTRACT_CHAR_LIMIT + return _clamp_or_default(char_limit) if char_limit is not None else _get_extract_char_limit() def _truncate_results(results: List[dict], char_limit: int, debug_call_data: dict) -> None: diff --git a/tools/website_policy.py b/tools/website_policy.py index 7abbb3d042..868ef4e1b4 100644 --- a/tools/website_policy.py +++ b/tools/website_policy.py @@ -80,8 +80,7 @@ def _load_policy_config(config_path: Path) -> Dict[str, Any]: logger.debug("PyYAML not installed — website blocklist disabled") return dict(_DEFAULT_WEBSITE_BLOCKLIST) try: - with open(config_path, encoding="utf-8") as f: - config = yaml.safe_load(f) or {} + config = yaml.safe_load(config_path.read_text(encoding="utf-8")) or {} except yaml.YAMLError as exc: raise WebsitePolicyError(f"Invalid config YAML at {config_path}: {exc}") from exc except OSError as exc: @@ -119,12 +118,11 @@ def load_website_blocklist(config_path: Optional[Path] = None) -> Dict[str, Any] config_path = config_path or default_path policy = _load_policy_config(config_path) - raw_domains = _require_type(policy, "domains", list, []) - raw_shared_files = _require_type(policy, "shared_files", list, []) + domains = map(_normalize_rule, _require_type(policy, "domains", list, [])) + pairs: List[Tuple[str, str]] = [(p, "config") for p in domains if p] + shared_files = _require_type(policy, "shared_files", list, []) enabled = _require_type(policy, "enabled", bool, True) - - pairs: List[Tuple[str, str]] = [(p, "config") for p in map(_normalize_rule, raw_domains) if p] - for shared_file in raw_shared_files: + for shared_file in shared_files: if not isinstance(shared_file, str) or not shared_file.strip(): continue path = Path(shared_file).expanduser() @@ -169,11 +167,9 @@ def check_website_access(url: str, config_path: Optional[Path] = None) -> Option with _cache_lock: if _cached_policy is not None and not _cached_policy.get("enabled"): return None - host = _extract_host_from_urlish(url) if not host: return None - try: policy = load_website_blocklist(config_path) except WebsitePolicyError as exc: @@ -184,7 +180,6 @@ def check_website_access(url: str, config_path: Optional[Path] = None) -> Option except Exception as exc: logger.warning("Unexpected error loading website policy (failing open): %s", exc) return None - if not policy.get("enabled"): return None for rule in policy.get("rules", []):