refactor(tools): fold debug-error into _finish_debug, clamp-or-default helper, compact cache/policy control flow

This commit is contained in:
Teknium
2026-09-02 23:14:03 -07:00
parent 3106197c40
commit cccc84a237
6 changed files with 74 additions and 100 deletions

View File

@@ -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:

View File

@@ -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:

View File

@@ -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))):

View File

@@ -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:

View File

@@ -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<alt>[^\]]*)\]\(\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:

View File

@@ -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", []):