diff --git a/agent/browser_registry.py b/agent/browser_registry.py index 4348237af1..93a88dd011 100644 --- a/agent/browser_registry.py +++ b/agent/browser_registry.py @@ -37,117 +37,18 @@ job is purely selection, not capability routing. from __future__ import annotations import logging -import threading -from typing import Dict, List, Optional +from typing import Optional from agent.browser_provider import BrowserProvider -from hermes_constants import hermes_home_key +from agent.provider_registry import ProviderRegistry, is_available_safe logger = logging.getLogger(__name__) -_providers: Dict[str, BrowserProvider] = {} -_scoped_providers: Dict[str, Dict[str, BrowserProvider]] = {} -_generation = 0 -_scoped_generations: Dict[str, int] = {} -_lock = threading.Lock() - - -def register_provider(provider: BrowserProvider, *, scope: Optional[str] = None) -> None: - """Register a cloud browser provider. - - Re-registration (same ``name``) overwrites the previous entry and logs - a debug message — makes hot-reload scenarios (tests, dev loops) behave - predictably. - """ - if not isinstance(provider, BrowserProvider): - raise TypeError( - f"register_provider() expects a BrowserProvider instance, " - f"got {type(provider).__name__}" - ) - raw_name = provider.name - if not isinstance(raw_name, str) or not raw_name.strip(): - raise ValueError("Browser provider .name must be a non-empty string") - name = raw_name.strip() - global _generation - with _lock: - target = _providers if scope is None else _scoped_providers.setdefault(scope, {}) - existing = target.get(name) - target[name] = provider - if scope is None: - _generation += 1 - else: - _scoped_generations[scope] = _scoped_generations.get(scope, 0) + 1 - if existing is not None: - logger.debug( - "Browser provider '%s' re-registered (was %r)", - name, type(existing).__name__, - ) - else: - logger.debug( - "Registered browser provider '%s' (%s)", - name, type(provider).__name__, - ) - - -def list_providers(*, scope: Optional[str] = None) -> List[BrowserProvider]: - """Return all registered providers, sorted by name.""" - with _lock: - merged = dict(_providers) - merged.update(_scoped_providers.get(scope or hermes_home_key(), {})) - items = list(merged.values()) - return sorted(items, key=lambda p: p.name) - - -def get_provider(name: str, *, scope: Optional[str] = None) -> Optional[BrowserProvider]: - """Return the provider registered under *name*, or None.""" - if not isinstance(name, str): - return None - with _lock: - key = name.strip() - return _scoped_providers.get(scope or hermes_home_key(), {}).get(key) or _providers.get(key) - - -def snapshot_registration( - name: str, *, scope: Optional[str] = None -) -> Optional[BrowserProvider]: - with _lock: - target = _providers if scope is None else _scoped_providers.get(scope, {}) - return target.get(name.strip()) - - -def registry_generation(*, scope: Optional[str] = None) -> tuple[int, int]: - """Return a cache fingerprint for the global base and one profile.""" - active_scope = scope or hermes_home_key() - with _lock: - return _generation, _scoped_generations.get(active_scope, 0) - - -def restore_registration( - name: str, - current: BrowserProvider, - previous: Optional[BrowserProvider], - *, - scope: Optional[str] = None, -) -> bool: - """Restore a plugin registration only when *current* is still installed.""" - key = name.strip() - global _generation - with _lock: - target = _providers if scope is None else _scoped_providers.setdefault(scope, {}) - if target.get(key) is not current: - return False - if previous is None: - target.pop(key, None) - else: - target[key] = previous - if scope is None: - _generation += 1 - else: - _scoped_generations[scope] = _scoped_generations.get(scope, 0) + 1 - if not target: - _scoped_providers.pop(scope, None) - return True +_registry: ProviderRegistry[BrowserProvider] = ProviderRegistry( + label="Browser", provider_cls=BrowserProvider, logger=logger, +) +_registry.export(globals()) # --------------------------------------------------------------------------- @@ -155,11 +56,9 @@ def restore_registration( # --------------------------------------------------------------------------- -# Legacy auto-detect order — used when no ``browser.cloud_provider`` is set. -# Matches the pre-migration walk in :func:`tools.browser_tool._get_cloud_provider`. -# Firecrawl is intentionally absent so users with ``FIRECRAWL_API_KEY`` set -# for web-extract don't get silently routed to a paid cloud browser. See -# :func:`_resolve` for the full rationale. +# Auto-detect order when ``browser.cloud_provider`` is unset (pre-migration +# walk of :func:`tools.browser_tool._get_cloud_provider`); see :func:`_resolve` +# for why Firecrawl is absent. _LEGACY_PREFERENCE = ( "browser-use", "browserbase", @@ -167,60 +66,22 @@ _LEGACY_PREFERENCE = ( def _resolve(configured: Optional[str]) -> Optional[BrowserProvider]: - """Resolve the active browser provider. + """Resolve the active browser provider (rules in the module docstring). - Resolution rules (in order): - - 1. **Explicit "local".** Returns None — the dispatcher disables cloud - mode entirely. Mirrors legacy short-circuit in - :func:`tools.browser_tool._get_cloud_provider`. - 2. **Explicit config wins, ignoring availability.** If ``configured`` - names a registered provider, return it even if its - :meth:`is_available` returns False — the dispatcher will surface a - precise "X_API_KEY is not set" error instead of silently routing - somewhere else. - 3. **Legacy preference walk, filtered by availability.** Walk - :data:`_LEGACY_PREFERENCE` (``browser-use`` → ``browserbase``) looking - for a provider whose ``is_available()`` is True. - - There is intentionally NO "single-eligible shortcut" rule here (unlike - :func:`agent.web_search_registry._resolve`). Pre-migration, the - auto-detect branch in ``tools.browser_tool._get_cloud_provider`` only - considered Browser Use and Browserbase; Firecrawl was reachable only - via an explicit ``browser.cloud_provider: firecrawl`` config key. - Preserving that gate matters because Firecrawl shares its API key with - the *web* extract plugin (``plugins/web/firecrawl/``), so users who set - ``FIRECRAWL_API_KEY`` for web extract must NOT get silently routed to a - paid cloud browser on a fresh install. Third-party browser-provider - plugins added under ``~/.hermes/plugins/browser//`` are subject - to the same gate — they must be explicitly configured to take effect. - - Returns None when no provider is configured AND no available provider - matches the legacy preference; the dispatcher then falls back to local - browser mode. + There is intentionally NO "single-eligible shortcut" (unlike + :func:`agent.web_search_registry._resolve`): only ``_LEGACY_PREFERENCE`` + names are auto-eligible. Firecrawl shares its API key with the *web* + extract plugin, so a user with ``FIRECRAWL_API_KEY`` must never be routed + to a paid cloud browser without setting ``browser.cloud_provider``; the + same gate applies to third-party browser-provider plugins. """ - with _lock: - snapshot = dict(_providers) - snapshot.update(_scoped_providers.get(hermes_home_key(), {})) + snapshot = _registry.merged() - def _is_available_safe(p: BrowserProvider) -> bool: - """Wrap ``is_available()`` so a buggy provider doesn't kill resolution.""" - try: - return bool(p.is_available()) - except Exception as exc: # noqa: BLE001 - logger.warning( - "Browser provider %s.is_available() raised %s — treating as unavailable", - p.name, exc, exc_info=True, - ) - return False - - # 1. Explicit "local" short-circuit. if configured == "local": return None - # 2. Explicit config wins — return regardless of is_available() so the - # user gets a precise downstream error message rather than a silent - # backend switch. Matches _get_cloud_provider() in browser_tool.py. + # Explicit config wins regardless of is_available(): the dispatcher then + # surfaces a precise "X_API_KEY is not set" error instead of a silent switch. if configured: provider = snapshot.get(configured) if provider is not None: @@ -231,23 +92,14 @@ def _resolve(configured: Optional[str]) -> Optional[BrowserProvider]: configured, ) - # 3. Legacy preference walk — only providers in _LEGACY_PREFERENCE are - # auto-eligible. Filtered by availability so we don't surface a - # provider the user has no credentials for. See docstring for why - # we do NOT fall back to "any single-eligible registered provider". for legacy in _LEGACY_PREFERENCE: provider = snapshot.get(legacy) - if provider is not None and _is_available_safe(provider): + if provider is not None and is_available_safe( + provider, logger, + "Browser provider %s.is_available() raised %s — treating as unavailable", + level=logging.WARNING, exc_info=True, + ): return provider return None - -def _reset_for_tests() -> None: - """Clear the registry. **Test-only.**""" - global _generation - with _lock: - _providers.clear() - _scoped_providers.clear() - _scoped_generations.clear() - _generation += 1 diff --git a/agent/image_gen_registry.py b/agent/image_gen_registry.py index 6239bbe891..087144f8be 100644 --- a/agent/image_gen_registry.py +++ b/agent/image_gen_registry.py @@ -11,9 +11,9 @@ Active selection The active provider is chosen by ``image_gen.provider`` in ``config.yaml``. If unset, :func:`get_active_provider` applies fallback logic: -1. If exactly one provider is registered, use it. -2. Otherwise if a provider named ``fal`` is registered, use it (legacy - default — matches pre-plugin behavior). +1. If exactly one *available* provider is registered, use it. +2. Otherwise if a provider named ``fal`` is registered and available, use it + (legacy default — matches pre-plugin behavior). 3. Otherwise return ``None`` (the tool surfaces a helpful error pointing the user at ``hermes tools``). """ @@ -21,151 +21,35 @@ If unset, :func:`get_active_provider` applies fallback logic: from __future__ import annotations import logging -import threading -from typing import Dict, List, Optional +from typing import Optional from agent.image_gen_provider import ImageGenProvider -from hermes_constants import hermes_home_key +from agent.provider_registry import ProviderRegistry, configured_provider_name, is_available_safe logger = logging.getLogger(__name__) -_providers: Dict[str, ImageGenProvider] = {} -_scoped_providers: Dict[str, Dict[str, ImageGenProvider]] = {} -_lock = threading.Lock() - - -def register_provider(provider: ImageGenProvider, *, scope: Optional[str] = None) -> None: - """Register an image generation provider. - - Re-registration (same ``name``) overwrites the previous entry and logs - a debug message — this makes hot-reload scenarios (tests, dev loops) - behave predictably. - """ - if not isinstance(provider, ImageGenProvider): - raise TypeError( - f"register_provider() expects an ImageGenProvider instance, " - f"got {type(provider).__name__}" - ) - raw_name = provider.name - if not isinstance(raw_name, str) or not raw_name.strip(): - raise ValueError("Image gen provider .name must be a non-empty string") - name = raw_name.strip() - with _lock: - target = _providers if scope is None else _scoped_providers.setdefault(scope, {}) - existing = target.get(name) - target[name] = provider - if existing is not None: - logger.debug("Image gen provider '%s' re-registered (was %r)", name, type(existing).__name__) - else: - logger.debug("Registered image gen provider '%s' (%s)", name, type(provider).__name__) - - -def list_providers(*, scope: Optional[str] = None) -> List[ImageGenProvider]: - """Return all registered providers, sorted by name.""" - with _lock: - merged = dict(_providers) - merged.update(_scoped_providers.get(scope or hermes_home_key(), {})) - items = list(merged.values()) - return sorted(items, key=lambda p: p.name) - - -def get_provider(name: str, *, scope: Optional[str] = None) -> Optional[ImageGenProvider]: - """Return the provider registered under *name*, or None.""" - if not isinstance(name, str): - return None - with _lock: - key = name.strip() - return _scoped_providers.get(scope or hermes_home_key(), {}).get(key) or _providers.get(key) - - -def snapshot_registration( - name: str, *, scope: Optional[str] = None -) -> Optional[ImageGenProvider]: - with _lock: - target = _providers if scope is None else _scoped_providers.get(scope, {}) - return target.get(name.strip()) - - -def restore_registration( - name: str, - current: ImageGenProvider, - previous: Optional[ImageGenProvider], - *, - scope: Optional[str] = None, -) -> bool: - """Restore a plugin registration only when *current* is still installed.""" - key = name.strip() - with _lock: - target = _providers if scope is None else _scoped_providers.setdefault(scope, {}) - if target.get(key) is not current: - return False - if previous is None: - target.pop(key, None) - else: - target[key] = previous - if scope is not None and not target: - _scoped_providers.pop(scope, None) - return True +_registry: ProviderRegistry[ImageGenProvider] = ProviderRegistry( + label="Image gen", provider_cls=ImageGenProvider, logger=logger, +) +_registry.export(globals()) def get_active_provider() -> Optional[ImageGenProvider]: """Resolve the currently-active provider. - Reads ``image_gen.provider`` from config.yaml; falls back per the - module docstring. - **Availability semantics** (mirrors :mod:`agent.web_search_registry`): - - - When ``image_gen.provider`` is explicitly set, the configured - provider is returned even if :meth:`ImageGenProvider.is_available` - reports False — the dispatcher surfaces a precise "X_API_KEY is not - set" error rather than silently switching backends. - - When ``image_gen.provider`` is unset, the fallback path (single- - provider shortcut and the FAL legacy preference) is filtered by - ``is_available()`` so we don't pick a provider the user has no - credentials for. + an explicitly configured provider is returned even if ``is_available()`` + is False, so the dispatcher surfaces a precise "X_API_KEY is not set" + error instead of silently switching backends. Only the unconfigured + fallback path is filtered by availability. """ - configured: Optional[str] = None - try: - from hermes_cli.config import load_config_readonly + configured = configured_provider_name("image_gen", logger) + snapshot = _registry.merged() - cfg = load_config_readonly() - section = cfg.get("image_gen") if isinstance(cfg, dict) else None - if isinstance(section, dict): - raw = section.get("provider") - if isinstance(raw, str) and raw.strip(): - configured = raw.strip() - except Exception as exc: - logger.debug("Could not read image_gen.provider from config: %s", exc) + def _available(p: ImageGenProvider) -> bool: + return is_available_safe(p, logger, "image_gen provider %s.is_available() raised %s") - # The managed "Nous Subscription" selection is serviced by the FAL - # plugin through the managed fal-queue gateway (the legacy FAL pipeline - # routes managed when the stored selection is "nous"). - if configured: - try: - from tools.tool_backend_helpers import NOUS_MANAGED_PROVIDER - - if configured.lower() == NOUS_MANAGED_PROVIDER: - configured = "fal" - except Exception: # pragma: no cover — helpers are in-repo - pass - - with _lock: - snapshot = dict(_providers) - snapshot.update(_scoped_providers.get(hermes_home_key(), {})) - - def _is_available_safe(p: ImageGenProvider) -> bool: - """Wrap ``is_available()`` so a buggy provider doesn't kill resolution.""" - try: - return bool(p.is_available()) - except Exception as exc: # noqa: BLE001 - logger.debug("image_gen provider %s.is_available() raised %s", p.name, exc) - return False - - # 1. Explicit config wins — return regardless of is_available() so the - # user gets a precise downstream error message rather than a silent - # backend switch. if configured: provider = snapshot.get(configured) if provider is not None: @@ -175,22 +59,12 @@ def get_active_provider() -> Optional[ImageGenProvider]: configured, ) - # 2. Fallback: single registered provider — but only if it's actually - # available (no credentials = don't surface it as "active"). - available = [p for p in snapshot.values() if _is_available_safe(p)] + available = [p for p in snapshot.values() if _available(p)] if len(available) == 1: return available[0] - # 3. Fallback: prefer legacy FAL for backward compat, when available. fal = snapshot.get("fal") - if fal is not None and _is_available_safe(fal): + if fal is not None and _available(fal): return fal return None - - -def _reset_for_tests() -> None: - """Clear the registry. **Test-only.**""" - with _lock: - _providers.clear() - _scoped_providers.clear() diff --git a/agent/provider_registry.py b/agent/provider_registry.py new file mode 100644 index 0000000000..f8c00ce0a4 --- /dev/null +++ b/agent/provider_registry.py @@ -0,0 +1,244 @@ +"""Shared engine behind the ``agent.*_registry`` provider registries. + +Every pluggable-backend registry (browser, TTS, image/video gen, transcription, +web search, terminal env) has the same shape: a global name->provider map plus +per-profile *scoped* maps (multiplexed gateways), a lock, registration with +re-registration logging, and the snapshot/restore pair that +:mod:`hermes_cli.plugins` uses to unwind a plugin's registrations. Each +``*_registry`` module instantiates one :class:`ProviderRegistry` and re-exports +its bound methods under the historical module-level names via +:meth:`ProviderRegistry.export`, so call sites, ``patch("agent.x_registry.get_provider")`` +targets, and the ``_providers`` / ``_scoped_providers`` / ``_lock`` test hooks are unchanged. +""" + +from __future__ import annotations + +import logging +import threading +from typing import Any, Callable, Dict, FrozenSet, Generic, List, Optional, TypeVar + +from hermes_constants import hermes_home_key + +P = TypeVar("P") + + +def strip_key(name: str) -> str: + return name.strip() + + +def lower_key(name: str) -> str: + return name.strip().lower() + + +class ProviderRegistry(Generic[P]): + """Global + per-scope provider map with plugin snapshot/restore support. + + Args: + label: Human label used in log/error strings (``"Browser"``, ``"TTS"``). + provider_cls: ABC every registered instance must satisfy (TypeError otherwise). + logger: The owning module's logger, so record names stay per-registry. + normalize: Key normalizer — ``strip_key`` or ``lower_key`` (case-insensitive + registries mirror how their dispatcher normalizes the configured name). + builtin_names: Reserved names owned by in-tree implementations; a collision + calls ``on_builtin_collision(key)`` and, if that returns, skips registration. + """ + + def __init__( + self, + *, + label: str, + provider_cls: type, + logger: logging.Logger, + normalize: Callable[[str], str] = strip_key, + builtin_names: FrozenSet[str] = frozenset(), + on_builtin_collision: Optional[Callable[[str], None]] = None, + ) -> None: + self.label = label + self.provider_cls = provider_cls + self.logger = logger + self.normalize = normalize + self.builtin_names = builtin_names + self._on_builtin_collision = on_builtin_collision + self._providers: Dict[str, P] = {} + self._scoped_providers: Dict[str, Dict[str, P]] = {} + self._generation = 0 + self._scoped_generations: Dict[str, int] = {} + self._lock = threading.Lock() + # "TTS provider" but "Registered browser provider": acronyms keep their case. + self._log_label = label if label.isupper() else label[0].lower() + label[1:] + + # -- internal helpers (caller holds the lock) --------------------------- + + def _target(self, scope: Optional[str], *, create: bool) -> Dict[str, P]: + if scope is None: + return self._providers + if create: + return self._scoped_providers.setdefault(scope, {}) + return self._scoped_providers.get(scope, {}) + + def _bump(self, scope: Optional[str]) -> None: + if scope is None: + self._generation += 1 + else: + self._scoped_generations[scope] = self._scoped_generations.get(scope, 0) + 1 + + # -- registration ------------------------------------------------------- + + def register(self, provider: P, *, scope: Optional[str] = None) -> None: + """Register a provider; same-name re-registration overwrites (hot reload).""" + if not isinstance(provider, self.provider_cls): + article = "an" if self.provider_cls.__name__[0] in "AEIOU" else "a" + raise TypeError( + f"register_provider() expects {article} {self.provider_cls.__name__} " + f"instance, got {type(provider).__name__}" + ) + raw_name = getattr(provider, "name") + if not isinstance(raw_name, str) or not raw_name.strip(): + raise ValueError(f"{self.label} provider .name must be a non-empty string") + key = self.normalize(raw_name) + if key in self.builtin_names: + if self._on_builtin_collision is not None: + self._on_builtin_collision(key) + return + with self._lock: + target = self._target(scope, create=True) + existing = target.get(key) + target[key] = provider + self._bump(scope) + if existing is not None: + self.logger.debug( + f"{self.label} provider '%s' re-registered (was %r)", + key, type(existing).__name__, + ) + else: + self.logger.debug( + f"Registered {self._log_label} provider '%s' (%s)", + key, type(provider).__name__, + ) + + # -- lookup --------------------------------------------------------------- + + def merged(self, scope: Optional[str] = None) -> Dict[str, P]: + """Global map overlaid with the active profile's scoped map (a copy).""" + with self._lock: + merged = dict(self._providers) + merged.update(self._scoped_providers.get(scope or hermes_home_key(), {})) + return merged + + def list_providers(self, *, scope: Optional[str] = None) -> List[P]: + """Return all registered providers, sorted by name.""" + return sorted(self.merged(scope).values(), key=lambda p: p.name) + + def get_provider(self, name: str, *, scope: Optional[str] = None) -> Optional[P]: + """Return the provider registered under *name* (scoped first), or None.""" + if not isinstance(name, str): + return None + key = self.normalize(name) + with self._lock: + return ( + self._scoped_providers.get(scope or hermes_home_key(), {}).get(key) + or self._providers.get(key) + ) + + def registry_generation(self, *, scope: Optional[str] = None) -> tuple: + """Cache fingerprint ``(global_generation, scoped_generation)``.""" + active_scope = scope or hermes_home_key() + with self._lock: + return self._generation, self._scoped_generations.get(active_scope, 0) + + # -- plugin unload support (hermes_cli.plugins) ----------------------------- + + def snapshot_registration(self, name: str, *, scope: Optional[str] = None) -> Optional[P]: + """Exact-slot lookup (no global fallback) used to detect plugin ownership.""" + with self._lock: + return self._target(scope, create=False).get(self.normalize(name)) + + def restore_registration( + self, name: str, current: P, previous: Optional[P], *, scope: Optional[str] = None + ) -> bool: + """Restore *previous* only when *current* is still installed under *name*.""" + key = self.normalize(name) + with self._lock: + target = self._target(scope, create=True) + if target.get(key) is not current: + return False + if previous is None: + target.pop(key, None) + else: + target[key] = previous + self._bump(scope) + if scope is not None and not target: + self._scoped_providers.pop(scope, None) + return True + + def reset_for_tests(self) -> None: + """Clear every registration. **Test-only.**""" + with self._lock: + self._providers.clear() + self._scoped_providers.clear() + self._scoped_generations.clear() + self._generation += 1 + + def export(self, namespace: Dict[str, Any]) -> None: + """Bind the historical module-level API into a ``*_registry`` module. + + Installs ``register_provider``/``list_providers``/``get_provider``/ + ``snapshot_registration``/``restore_registration``/``registry_generation``/ + ``_reset_for_tests`` plus the ``_providers``/``_scoped_providers``/``_lock`` + test hooks, so ``patch("agent.x_registry.get_provider")`` and direct + ``_providers`` manipulation in tests keep working unchanged. + """ + namespace.update( + _providers=self._providers, + _scoped_providers=self._scoped_providers, + _lock=self._lock, + register_provider=self.register, + list_providers=self.list_providers, + get_provider=self.get_provider, + snapshot_registration=self.snapshot_registration, + restore_registration=self.restore_registration, + registry_generation=self.registry_generation, + _reset_for_tests=self.reset_for_tests, + ) + + +def is_available_safe( + provider: Any, + logger: logging.Logger, + fmt: str, + *, + level: int = logging.DEBUG, + exc_info: bool = False, +) -> bool: + """``bool(provider.is_available())`` that treats a raising provider as unavailable.""" + try: + return bool(provider.is_available()) + except Exception as exc: # noqa: BLE001 + logger.log(level, fmt, provider.name, exc, exc_info=exc_info) + return False + + +def configured_provider_name(section: str, logger: logging.Logger) -> Optional[str]: + """Read ``
.provider`` from config.yaml, mapping the managed Nous + selection to ``fal`` (the FAL plugin services it via the managed gateway).""" + configured: Optional[str] = None + try: + from hermes_cli.config import load_config_readonly + + cfg = load_config_readonly() + block = cfg.get(section) if isinstance(cfg, dict) else None + if isinstance(block, dict): + raw = block.get("provider") + if isinstance(raw, str) and raw.strip(): + configured = raw.strip() + except Exception as exc: + logger.debug("Could not read %s.provider from config: %s", section, exc) + if configured: + try: + from tools.tool_backend_helpers import NOUS_MANAGED_PROVIDER + + if configured.lower() == NOUS_MANAGED_PROVIDER: + configured = "fal" + except Exception: # pragma: no cover — helpers are in-repo + pass + return configured diff --git a/agent/terminal_env_registry.py b/agent/terminal_env_registry.py index 3c65cd2242..251f84f28c 100644 --- a/agent/terminal_env_registry.py +++ b/agent/terminal_env_registry.py @@ -26,11 +26,10 @@ into a per-profile scope (multiplexed gateways) or the global base map. from __future__ import annotations import logging -import threading -from typing import Dict, List, Optional +from typing import List, Optional +from agent.provider_registry import ProviderRegistry, lower_key from agent.terminal_env_provider import TerminalEnvironmentProvider -from hermes_constants import hermes_home_key logger = logging.getLogger(__name__) @@ -43,86 +42,27 @@ BUILTIN_BACKEND_NAMES = frozenset({ }) -_providers: Dict[str, TerminalEnvironmentProvider] = {} -_scoped_providers: Dict[str, Dict[str, TerminalEnvironmentProvider]] = {} -_generation = 0 -_scoped_generations: Dict[str, int] = {} -_lock = threading.Lock() +def _reject_builtin_collision(name: str) -> None: + raise ValueError( + f"Terminal backend name '{name}' is reserved for the built-in " + f"{name} backend and cannot be registered by a plugin" + ) -def register_provider( - provider: TerminalEnvironmentProvider, *, scope: Optional[str] = None -) -> None: - """Register a terminal environment provider. - - Re-registration (same ``name``) overwrites the previous entry — makes - hot-reload scenarios (tests, dev loops) behave predictably. - - Raises: - TypeError: not a TerminalEnvironmentProvider instance. - ValueError: empty name or collision with a built-in backend name. - """ - if not isinstance(provider, TerminalEnvironmentProvider): - raise TypeError( - f"register_provider() expects a TerminalEnvironmentProvider " - f"instance, got {type(provider).__name__}" - ) - raw_name = provider.name - if not isinstance(raw_name, str) or not raw_name.strip(): - raise ValueError("Terminal environment provider .name must be a non-empty string") - name = raw_name.strip().lower() - if name in BUILTIN_BACKEND_NAMES: - raise ValueError( - f"Terminal backend name '{name}' is reserved for the built-in " - f"{name} backend and cannot be registered by a plugin" - ) - global _generation - with _lock: - target = _providers if scope is None else _scoped_providers.setdefault(scope, {}) - existing = target.get(name) - target[name] = provider - if scope is None: - _generation += 1 - else: - _scoped_generations[scope] = _scoped_generations.get(scope, 0) + 1 - if existing is not None: - logger.debug( - "Terminal environment provider '%s' re-registered (was %r)", - name, type(existing).__name__, - ) - else: - logger.debug( - "Registered terminal environment provider '%s' (%s)", - name, type(provider).__name__, - ) - - -def list_providers(*, scope: Optional[str] = None) -> List[TerminalEnvironmentProvider]: - """Return all registered providers, sorted by name.""" - with _lock: - merged = dict(_providers) - merged.update(_scoped_providers.get(scope or hermes_home_key(), {})) - items = list(merged.values()) - return sorted(items, key=lambda p: p.name) - - -def get_provider( - name: str, *, scope: Optional[str] = None -) -> Optional[TerminalEnvironmentProvider]: - """Return the provider registered under *name*, or None.""" - if not isinstance(name, str): - return None - key = name.strip().lower() - with _lock: - return ( - _scoped_providers.get(scope or hermes_home_key(), {}).get(key) - or _providers.get(key) - ) +_registry: ProviderRegistry[TerminalEnvironmentProvider] = ProviderRegistry( + label="Terminal environment", + provider_cls=TerminalEnvironmentProvider, + logger=logger, + normalize=lower_key, + builtin_names=BUILTIN_BACKEND_NAMES, + on_builtin_collision=_reject_builtin_collision, +) +_registry.export(globals()) def plugin_backend_names(*, scope: Optional[str] = None) -> List[str]: """Names of all registered plugin backends (sorted).""" - return [p.name.strip().lower() for p in list_providers(scope=scope)] + return [p.name.strip().lower() for p in _registry.list_providers(scope=scope)] def provider_flag(name: str, attr: str, default=False): @@ -132,7 +72,7 @@ def provider_flag(name: str, attr: str, default=False): misbehaving plugin degrades to built-in-equivalent behavior instead of taking the terminal tool down. """ - provider = get_provider(name) + provider = _registry.get_provider(name) if provider is None: return default try: @@ -154,9 +94,9 @@ def plugin_strip_env_keys() -> frozenset: the static tier-1 set unconditionally). """ keys: set = set() - with _lock: - all_providers = list(_providers.values()) - for scoped in _scoped_providers.values(): + with _registry._lock: + all_providers = list(_registry._providers.values()) + for scoped in _registry._scoped_providers.values(): all_providers.extend(scoped.values()) for provider in all_providers: try: @@ -167,55 +107,3 @@ def plugin_strip_env_keys() -> frozenset: exc_info=True, ) return frozenset(keys) - - -def snapshot_registration( - name: str, *, scope: Optional[str] = None -) -> Optional[TerminalEnvironmentProvider]: - with _lock: - target = _providers if scope is None else _scoped_providers.get(scope, {}) - return target.get(name.strip().lower()) - - -def registry_generation(*, scope: Optional[str] = None) -> tuple: - """Return a cache fingerprint for the global base and one profile.""" - active_scope = scope or hermes_home_key() - with _lock: - return _generation, _scoped_generations.get(active_scope, 0) - - -def restore_registration( - name: str, - current: TerminalEnvironmentProvider, - previous: Optional[TerminalEnvironmentProvider], - *, - scope: Optional[str] = None, -) -> bool: - """Restore a plugin registration only when *current* is still installed.""" - key = name.strip().lower() - global _generation - with _lock: - target = _providers if scope is None else _scoped_providers.setdefault(scope, {}) - if target.get(key) is not current: - return False - if previous is None: - target.pop(key, None) - else: - target[key] = previous - if scope is None: - _generation += 1 - else: - _scoped_generations[scope] = _scoped_generations.get(scope, 0) + 1 - if not target: - _scoped_providers.pop(scope, None) - return True - - -def _reset_for_tests() -> None: - """Clear all registrations. Test hook — mirrors sibling registries.""" - global _generation - with _lock: - _providers.clear() - _scoped_providers.clear() - _scoped_generations.clear() - _generation = 0 diff --git a/agent/transcription_registry.py b/agent/transcription_registry.py index a167e710b5..99e6b4a2b8 100644 --- a/agent/transcription_registry.py +++ b/agent/transcription_registry.py @@ -2,42 +2,31 @@ Transcription Provider Registry ================================ -Central map of registered STT providers. Populated by plugins at -import-time via :meth:`PluginContext.register_transcription_provider`; -consumed by :mod:`tools.transcription_tools` to dispatch -:func:`transcribe_audio` calls to the active plugin backend **when** -the configured ``stt.provider`` name is not a built-in. +Central map of registered STT providers. Populated by plugins at import-time +via :meth:`PluginContext.register_transcription_provider`; consumed by +:mod:`tools.transcription_tools` to dispatch :func:`transcribe_audio` calls +to the active plugin backend **when** the configured ``stt.provider`` name is +not a built-in. -Built-ins-always-win --------------------- -Plugin names that collide with a built-in STT provider (``local``, -``local_command``, ``groq``, ``openai``, ``mistral``, ``xai``) are -rejected at registration with a warning. This invariant is also -re-checked at dispatch time in -:func:`tools.transcription_tools._dispatch_to_plugin_provider`. +Built-ins-always-win: a plugin name colliding with a built-in STT provider is +rejected at registration with a warning (re-checked at dispatch time in +:func:`tools.transcription_tools._dispatch_to_plugin_provider`). """ from __future__ import annotations import logging -import threading -from typing import Dict, List, Optional +from agent.provider_registry import ProviderRegistry, lower_key from agent.transcription_provider import TranscriptionProvider -from hermes_constants import hermes_home_key logger = logging.getLogger(__name__) -# Names reserved for native built-in STT handlers. Plugins cannot -# register a name in this set — the registration call is rejected with -# a warning. **Kept in sync with ``BUILTIN_STT_PROVIDERS`` in -# :mod:`tools.transcription_tools`** — a regression test in -# ``tests/agent/test_transcription_registry.py::TestBuiltinSync`` -# fails if the two lists drift. Importing from -# ``tools.transcription_tools`` directly would create a circular -# dependency (``tools.transcription_tools`` imports -# ``agent.transcription_registry`` for dispatch). +# Names reserved for native built-in STT handlers. **Kept in sync with +# ``BUILTIN_STT_PROVIDERS`` in :mod:`tools.transcription_tools`** (a regression +# test in ``tests/agent/test_transcription_registry.py::TestBuiltinSync`` fails +# on drift); importing it directly would be a circular import. _BUILTIN_NAMES = frozenset({ "local", "local_command", @@ -50,114 +39,23 @@ _BUILTIN_NAMES = frozenset({ }) -_providers: Dict[str, TranscriptionProvider] = {} -_scoped_providers: Dict[str, Dict[str, TranscriptionProvider]] = {} -_lock = threading.Lock() +def _warn_builtin_collision(key: str) -> None: + logger.warning( + "Transcription provider '%s' shadows a built-in name; registration " + "ignored. Built-in STT providers (%s) always win — pick a different " + "name.", + key, ", ".join(sorted(_BUILTIN_NAMES)), + ) -def register_provider(provider: TranscriptionProvider, *, scope: Optional[str] = None) -> None: - """Register a transcription provider. - - Rejects: - - - Non-:class:`TranscriptionProvider` instances (raises :class:`TypeError`). - - Empty/whitespace ``.name`` (raises :class:`ValueError`). - - Names colliding with a built-in (logs a warning, silently - ignores — built-ins-always-win invariant). - - Re-registration (same ``name``) overwrites the previous entry and - logs a debug message — makes hot-reload scenarios (tests, dev - loops) behave predictably. - """ - if not isinstance(provider, TranscriptionProvider): - raise TypeError( - f"register_provider() expects a TranscriptionProvider instance, " - f"got {type(provider).__name__}" - ) - name = provider.name - if not isinstance(name, str) or not name.strip(): - raise ValueError("Transcription provider .name must be a non-empty string") - key = name.strip().lower() - if key in _BUILTIN_NAMES: - logger.warning( - "Transcription provider '%s' shadows a built-in name; registration " - "ignored. Built-in STT providers (%s) always win — pick a different " - "name.", - key, ", ".join(sorted(_BUILTIN_NAMES)), - ) - return - with _lock: - target = _providers if scope is None else _scoped_providers.setdefault(scope, {}) - existing = target.get(key) - target[key] = provider - if existing is not None: - logger.debug( - "Transcription provider '%s' re-registered (was %r)", - key, type(existing).__name__, - ) - else: - logger.debug( - "Registered transcription provider '%s' (%s)", - key, type(provider).__name__, - ) - - -def list_providers(*, scope: Optional[str] = None) -> List[TranscriptionProvider]: - """Return all registered providers, sorted by name.""" - with _lock: - merged = dict(_providers) - merged.update(_scoped_providers.get(scope or hermes_home_key(), {})) - items = list(merged.values()) - return sorted(items, key=lambda p: p.name) - - -def get_provider(name: str, *, scope: Optional[str] = None) -> Optional[TranscriptionProvider]: - """Return the provider registered under *name*, or None. - - Name matching is case-insensitive and whitespace-tolerant — mirrors - how ``tools.transcription_tools._get_provider`` normalizes the - configured ``stt.provider`` value. - """ - if not isinstance(name, str): - return None - key = name.strip().lower() - with _lock: - return _scoped_providers.get(scope or hermes_home_key(), {}).get(key) or _providers.get(key) - - -def snapshot_registration( - name: str, *, scope: Optional[str] = None -) -> Optional[TranscriptionProvider]: - key = name.strip().lower() - with _lock: - target = _providers if scope is None else _scoped_providers.get(scope, {}) - return target.get(key) - - -def restore_registration( - name: str, - current: TranscriptionProvider, - previous: Optional[TranscriptionProvider], - *, - scope: Optional[str] = None, -) -> bool: - """Restore a plugin registration only when *current* is still installed.""" - key = name.strip().lower() - with _lock: - target = _providers if scope is None else _scoped_providers.setdefault(scope, {}) - if target.get(key) is not current: - return False - if previous is None: - target.pop(key, None) - else: - target[key] = previous - if scope is not None and not target: - _scoped_providers.pop(scope, None) - return True - - -def _reset_for_tests() -> None: - """Clear the registry. **Test-only.**""" - with _lock: - _providers.clear() - _scoped_providers.clear() +# Case-insensitive, whitespace-tolerant keys mirror how +# ``tools.transcription_tools`` normalizes the configured ``stt.provider``. +_registry: ProviderRegistry[TranscriptionProvider] = ProviderRegistry( + label="Transcription", + provider_cls=TranscriptionProvider, + logger=logger, + normalize=lower_key, + builtin_names=_BUILTIN_NAMES, + on_builtin_collision=_warn_builtin_collision, +) +_registry.export(globals()) diff --git a/agent/tts_registry.py b/agent/tts_registry.py index 5e8f94b459..a94e1c095e 100644 --- a/agent/tts_registry.py +++ b/agent/tts_registry.py @@ -2,50 +2,36 @@ TTS Provider Registry ===================== -Central map of registered TTS providers. Populated by plugins at -import-time via :meth:`PluginContext.register_tts_provider`; consumed -by :mod:`tools.tts_tool` to dispatch ``text_to_speech`` tool calls to -the active plugin backend **when** the configured ``tts.provider`` -name is neither a built-in nor a command-type provider. +Central map of registered TTS providers. Populated by plugins at import-time +via :meth:`PluginContext.register_tts_provider`; consumed by +:mod:`tools.tts_tool` to dispatch ``text_to_speech`` calls to the active +plugin backend **when** the configured ``tts.provider`` name is neither a +built-in nor a command-type provider. -Built-ins-always-win --------------------- -Plugin names that collide with a built-in TTS provider (``edge``, -``openai``, ``elevenlabs``, ``minimax``, ``gemini``, ``mistral``, -``xai``, ``piper``, ``kittentts``, ``neutts``) are rejected at -registration with a warning. This invariant is also re-checked at -dispatch time in :func:`tools.tts_tool._dispatch_to_plugin_provider`. +Built-ins-always-win: a plugin name colliding with a built-in TTS provider is +rejected at registration with a warning (re-checked at dispatch time in +:func:`tools.tts_tool._dispatch_to_plugin_provider`). -Command-providers-win-over-plugins ----------------------------------- -This registry doesn't enforce the command-vs-plugin precedence — that -lives in the dispatcher, which checks for a same-name -``tts.providers.: type: command`` entry before consulting the -registry. The rationale is locality: a name declared in the user's -``config.yaml`` is more specific to their setup than a plugin that -happens to be installed. +Command-providers-win-over-plugins is enforced by the dispatcher, not here: +it checks for a same-name ``tts.providers.: type: command`` entry before +consulting the registry (a name declared in the user's config.yaml is more +specific to their setup than an installed plugin). """ from __future__ import annotations import logging -import threading -from typing import Dict, List, Optional +from agent.provider_registry import ProviderRegistry, lower_key from agent.tts_provider import TTSProvider -from hermes_constants import hermes_home_key logger = logging.getLogger(__name__) -# Names reserved for native built-in TTS handlers. Plugins cannot -# register a name in this set — the registration call is rejected with -# a warning. **Kept in sync with ``BUILTIN_TTS_PROVIDERS`` in -# :mod:`tools.tts_tool`** — a regression test in -# ``tests/agent/test_tts_registry.py::TestBuiltinSync`` fails if the -# two lists drift. Importing from ``tools.tts_tool`` directly would -# create a circular dependency (``tools.tts_tool`` imports -# ``agent.tts_registry`` for dispatch). +# Names reserved for native built-in TTS handlers. **Kept in sync with +# ``BUILTIN_TTS_PROVIDERS`` in :mod:`tools.tts_tool`** (a regression test in +# ``tests/agent/test_tts_registry.py::TestBuiltinSync`` fails on drift); +# importing it directly would be a circular import. _BUILTIN_NAMES = frozenset({ "edge", "elevenlabs", @@ -61,113 +47,22 @@ _BUILTIN_NAMES = frozenset({ }) -_providers: Dict[str, TTSProvider] = {} -_scoped_providers: Dict[str, Dict[str, TTSProvider]] = {} -_lock = threading.Lock() +def _warn_builtin_collision(key: str) -> None: + logger.warning( + "TTS provider '%s' shadows a built-in name; registration ignored. " + "Built-in TTS providers (%s) always win — pick a different name.", + key, ", ".join(sorted(_BUILTIN_NAMES)), + ) -def register_provider(provider: TTSProvider, *, scope: Optional[str] = None) -> None: - """Register a TTS provider. - - Rejects: - - - Non-:class:`TTSProvider` instances (raises :class:`TypeError`). - - Empty/whitespace ``.name`` (raises :class:`ValueError`). - - Names colliding with a built-in (logs a warning, silently - ignores — built-ins-always-win invariant). - - Re-registration (same ``name``) overwrites the previous entry and - logs a debug message — makes hot-reload scenarios (tests, dev - loops) behave predictably. - """ - if not isinstance(provider, TTSProvider): - raise TypeError( - f"register_provider() expects a TTSProvider instance, " - f"got {type(provider).__name__}" - ) - name = provider.name - if not isinstance(name, str) or not name.strip(): - raise ValueError("TTS provider .name must be a non-empty string") - key = name.strip().lower() - if key in _BUILTIN_NAMES: - logger.warning( - "TTS provider '%s' shadows a built-in name; registration ignored. " - "Built-in TTS providers (%s) always win — pick a different name.", - key, ", ".join(sorted(_BUILTIN_NAMES)), - ) - return - with _lock: - target = _providers if scope is None else _scoped_providers.setdefault(scope, {}) - existing = target.get(key) - target[key] = provider - if existing is not None: - logger.debug( - "TTS provider '%s' re-registered (was %r)", - key, type(existing).__name__, - ) - else: - logger.debug( - "Registered TTS provider '%s' (%s)", - key, type(provider).__name__, - ) - - -def list_providers(*, scope: Optional[str] = None) -> List[TTSProvider]: - """Return all registered providers, sorted by name.""" - with _lock: - merged = dict(_providers) - merged.update(_scoped_providers.get(scope or hermes_home_key(), {})) - items = list(merged.values()) - return sorted(items, key=lambda p: p.name) - - -def get_provider(name: str, *, scope: Optional[str] = None) -> Optional[TTSProvider]: - """Return the provider registered under *name*, or None. - - Name matching is case-insensitive and whitespace-tolerant — mirrors - how ``tools.tts_tool._get_provider`` normalizes the configured - ``tts.provider`` value. - """ - if not isinstance(name, str): - return None - key = name.strip().lower() - with _lock: - return _scoped_providers.get(scope or hermes_home_key(), {}).get(key) or _providers.get(key) - - -def snapshot_registration( - name: str, *, scope: Optional[str] = None -) -> Optional[TTSProvider]: - key = name.strip().lower() - with _lock: - target = _providers if scope is None else _scoped_providers.get(scope, {}) - return target.get(key) - - -def restore_registration( - name: str, - current: TTSProvider, - previous: Optional[TTSProvider], - *, - scope: Optional[str] = None, -) -> bool: - """Restore a plugin registration only when *current* is still installed.""" - key = name.strip().lower() - with _lock: - target = _providers if scope is None else _scoped_providers.setdefault(scope, {}) - if target.get(key) is not current: - return False - if previous is None: - target.pop(key, None) - else: - target[key] = previous - if scope is not None and not target: - _scoped_providers.pop(scope, None) - return True - - -def _reset_for_tests() -> None: - """Clear the registry. **Test-only.**""" - with _lock: - _providers.clear() - _scoped_providers.clear() +# Case-insensitive, whitespace-tolerant keys mirror how +# ``tools.tts_tool._get_provider`` normalizes the configured ``tts.provider``. +_registry: ProviderRegistry[TTSProvider] = ProviderRegistry( + label="TTS", + provider_cls=TTSProvider, + logger=logger, + normalize=lower_key, + builtin_names=_BUILTIN_NAMES, + on_builtin_collision=_warn_builtin_collision, +) +_registry.export(globals()) diff --git a/agent/video_gen_registry.py b/agent/video_gen_registry.py index 5b4c9cea54..733f32ca0f 100644 --- a/agent/video_gen_registry.py +++ b/agent/video_gen_registry.py @@ -15,138 +15,34 @@ If unset, :func:`get_active_provider` applies fallback logic: 2. Otherwise return ``None`` (the tool surfaces a helpful error pointing the user at ``hermes tools``). -Mirrors ``agent/image_gen_registry.py`` so the two surfaces behave the -same: the unconfigured fallback is filtered by ``is_available()`` so a box -that has credentials for only one backend (e.g. DeepInfra, while the -``fal``/``xai`` plugins also register unconditionally) auto-selects it -instead of returning ``None``. +Mirrors ``agent/image_gen_registry.py``: the unconfigured fallback is +filtered by ``is_available()`` so a box with credentials for only one backend +(e.g. DeepInfra, while ``fal``/``xai`` register unconditionally) auto-selects +it instead of returning ``None``. Unlike image gen there is no legacy ``fal`` +preference, and a configured-but-unregistered name fails closed. """ from __future__ import annotations import logging -import threading -from typing import Dict, List, Optional +from typing import Optional +from agent.provider_registry import ProviderRegistry, configured_provider_name, is_available_safe from agent.video_gen_provider import VideoGenProvider -from hermes_constants import hermes_home_key logger = logging.getLogger(__name__) -_providers: Dict[str, VideoGenProvider] = {} -_scoped_providers: Dict[str, Dict[str, VideoGenProvider]] = {} -_lock = threading.Lock() - - -def register_provider(provider: VideoGenProvider, *, scope: Optional[str] = None) -> None: - """Register a video generation provider. - - Re-registration (same ``name``) overwrites the previous entry and logs - a debug message — this makes hot-reload scenarios (tests, dev loops) - behave predictably. - """ - if not isinstance(provider, VideoGenProvider): - raise TypeError( - f"register_provider() expects a VideoGenProvider instance, " - f"got {type(provider).__name__}" - ) - raw_name = provider.name - if not isinstance(raw_name, str) or not raw_name.strip(): - raise ValueError("Video gen provider .name must be a non-empty string") - name = raw_name.strip() - with _lock: - target = _providers if scope is None else _scoped_providers.setdefault(scope, {}) - existing = target.get(name) - target[name] = provider - if existing is not None: - logger.debug("Video gen provider '%s' re-registered (was %r)", name, type(existing).__name__) - else: - logger.debug("Registered video gen provider '%s' (%s)", name, type(provider).__name__) - - -def list_providers(*, scope: Optional[str] = None) -> List[VideoGenProvider]: - """Return all registered providers, sorted by name.""" - with _lock: - merged = dict(_providers) - merged.update(_scoped_providers.get(scope or hermes_home_key(), {})) - items = list(merged.values()) - return sorted(items, key=lambda p: p.name) - - -def get_provider(name: str, *, scope: Optional[str] = None) -> Optional[VideoGenProvider]: - """Return the provider registered under *name*, or None.""" - if not isinstance(name, str): - return None - with _lock: - key = name.strip() - return _scoped_providers.get(scope or hermes_home_key(), {}).get(key) or _providers.get(key) - - -def snapshot_registration( - name: str, *, scope: Optional[str] = None -) -> Optional[VideoGenProvider]: - with _lock: - target = _providers if scope is None else _scoped_providers.get(scope, {}) - return target.get(name.strip()) - - -def restore_registration( - name: str, - current: VideoGenProvider, - previous: Optional[VideoGenProvider], - *, - scope: Optional[str] = None, -) -> bool: - """Restore a plugin registration only when *current* is still installed.""" - key = name.strip() - with _lock: - target = _providers if scope is None else _scoped_providers.setdefault(scope, {}) - if target.get(key) is not current: - return False - if previous is None: - target.pop(key, None) - else: - target[key] = previous - if scope is not None and not target: - _scoped_providers.pop(scope, None) - return True +_registry: ProviderRegistry[VideoGenProvider] = ProviderRegistry( + label="Video gen", provider_cls=VideoGenProvider, logger=logger, +) +_registry.export(globals()) def get_active_provider() -> Optional[VideoGenProvider]: - """Resolve the currently-active provider. - - Reads ``video_gen.provider`` from config.yaml; falls back per the - module docstring. - """ - configured: Optional[str] = None - try: - from hermes_cli.config import load_config_readonly - - cfg = load_config_readonly() - section = cfg.get("video_gen") if isinstance(cfg, dict) else None - if isinstance(section, dict): - raw = section.get("provider") - if isinstance(raw, str) and raw.strip(): - configured = raw.strip() - except Exception as exc: - logger.debug("Could not read video_gen.provider from config: %s", exc) - - # The managed "Nous Subscription" selection is serviced by the FAL - # plugin through the managed fal-queue gateway (the plugin's resolver - # routes managed when the stored selection is "nous"). - if configured: - try: - from tools.tool_backend_helpers import NOUS_MANAGED_PROVIDER - - if configured.lower() == NOUS_MANAGED_PROVIDER: - configured = "fal" - except Exception: # pragma: no cover — helpers are in-repo - pass - - with _lock: - snapshot = dict(_providers) - snapshot.update(_scoped_providers.get(hermes_home_key(), {})) + """Resolve the currently-active provider (see module docstring).""" + configured = configured_provider_name("video_gen", logger) + snapshot = _registry.merged() if configured: provider = snapshot.get(configured) @@ -158,27 +54,11 @@ def get_active_provider() -> Optional[VideoGenProvider]: ) return None - def _is_available_safe(p: VideoGenProvider) -> bool: - """Wrap ``is_available()`` so a buggy provider doesn't kill resolution.""" - try: - return bool(p.is_available()) - except Exception as exc: # noqa: BLE001 - logger.debug("video_gen provider %s.is_available() raised %s", p.name, exc) - return False - - # Fallback: single *available* provider — filter by is_available() so a - # box with credentials for only one backend auto-selects it even when - # other providers (fal/xai) register unconditionally without keys. - # Mirrors agent/image_gen_registry.get_active_provider(). - available = [p for p in snapshot.values() if _is_available_safe(p)] + available = [ + p for p in snapshot.values() + if is_available_safe(p, logger, "video_gen provider %s.is_available() raised %s") + ] if len(available) == 1: return available[0] return None - - -def _reset_for_tests() -> None: - """Clear the registry. **Test-only.**""" - with _lock: - _providers.clear() - _scoped_providers.clear() diff --git a/agent/web_search_registry.py b/agent/web_search_registry.py index 7f60ea838e..7f18be7e35 100644 --- a/agent/web_search_registry.py +++ b/agent/web_search_registry.py @@ -33,98 +33,18 @@ extract-capable backend. from __future__ import annotations import logging -import threading -from typing import Dict, List, Optional +from typing import Optional +from agent.provider_registry import ProviderRegistry, is_available_safe from agent.web_search_provider import WebSearchProvider -from hermes_constants import hermes_home_key logger = logging.getLogger(__name__) -_providers: Dict[str, WebSearchProvider] = {} -_scoped_providers: Dict[str, Dict[str, WebSearchProvider]] = {} -_lock = threading.Lock() - - -def register_provider(provider: WebSearchProvider, *, scope: Optional[str] = None) -> None: - """Register a web search/extract provider. - - Re-registration (same ``name``) overwrites the previous entry and logs - a debug message — makes hot-reload scenarios (tests, dev loops) behave - predictably. - """ - if not isinstance(provider, WebSearchProvider): - raise TypeError( - f"register_provider() expects a WebSearchProvider instance, " - f"got {type(provider).__name__}" - ) - raw_name = provider.name - if not isinstance(raw_name, str) or not raw_name.strip(): - raise ValueError("Web provider .name must be a non-empty string") - name = raw_name.strip() - with _lock: - target = _providers if scope is None else _scoped_providers.setdefault(scope, {}) - existing = target.get(name) - target[name] = provider - if existing is not None: - logger.debug( - "Web provider '%s' re-registered (was %r)", - name, type(existing).__name__, - ) - else: - logger.debug( - "Registered web provider '%s' (%s)", - name, type(provider).__name__, - ) - - -def list_providers(*, scope: Optional[str] = None) -> List[WebSearchProvider]: - """Return all registered providers, sorted by name.""" - with _lock: - merged = dict(_providers) - merged.update(_scoped_providers.get(scope or hermes_home_key(), {})) - items = list(merged.values()) - return sorted(items, key=lambda p: p.name) - - -def get_provider(name: str, *, scope: Optional[str] = None) -> Optional[WebSearchProvider]: - """Return the provider registered under *name*, or None.""" - if not isinstance(name, str): - return None - with _lock: - key = name.strip() - return _scoped_providers.get(scope or hermes_home_key(), {}).get(key) or _providers.get(key) - - -def snapshot_registration( - name: str, *, scope: Optional[str] = None -) -> Optional[WebSearchProvider]: - with _lock: - target = _providers if scope is None else _scoped_providers.get(scope, {}) - return target.get(name.strip()) - - -def restore_registration( - name: str, - current: WebSearchProvider, - previous: Optional[WebSearchProvider], - *, - scope: Optional[str] = None, -) -> bool: - """Restore a plugin registration only when *current* is still installed.""" - key = name.strip() - with _lock: - target = _providers if scope is None else _scoped_providers.setdefault(scope, {}) - if target.get(key) is not current: - return False - if previous is None: - target.pop(key, None) - else: - target[key] = previous - if scope is not None and not target: - _scoped_providers.pop(scope, None) - return True +_registry: ProviderRegistry[WebSearchProvider] = ProviderRegistry( + label="Web", provider_cls=WebSearchProvider, logger=logger, +) +_registry.export(globals()) # --------------------------------------------------------------------------- @@ -150,6 +70,11 @@ def _read_config_key(*path: str) -> Optional[str]: return None +def _configured_backend(capability: str) -> Optional[str]: + """``web._backend`` (preferred) or ``web.backend`` (shared fallback).""" + return _read_config_key("web", f"{capability}_backend") or _read_config_key("web", "backend") + + # Legacy preference order — preserves behaviour for users who set no # ``web.backend`` / ``web._backend`` config key at all. Matches # the historic candidate order in :func:`tools.web_tools._get_backend` @@ -207,36 +132,13 @@ def _keyless_preference() -> tuple: def _resolve(configured: Optional[str], *, capability: str) -> Optional[WebSearchProvider]: """Resolve the active provider for a capability ("search" | "extract"). - Resolution rules (in order): - - 1. **Explicit config wins, ignoring availability.** If - ``web.{capability}_backend`` or ``web.backend`` names a registered - provider that supports *capability*, return it even if its - :meth:`is_available` returns False — the dispatcher will surface a - precise "X_API_KEY is not set" error to the user instead of silently - routing somewhere else. Matches legacy - :func:`tools.web_tools._get_backend` behavior for configured names. - - 2. **Single-provider shortcut.** When only one registered provider - supports *capability* AND ``is_available()`` reports True, return it. - - 3. **Legacy preference walk, filtered by availability.** Walk the - :data:`_LEGACY_PREFERENCE` order (firecrawl → parallel → tavily → - exa → searxng → brave-free → ddgs) looking for a provider whose - ``supports_()`` is True AND whose ``is_available()`` is - True. Matches the historic ``tools.web_tools._get_backend()`` - candidate order so users with credentials but no explicit config - key keep landing on the same provider as pre-migration. This is - the path that fires when no config key is set — pick the - highest-priority backend the user actually has credentials for. - - Returns None when no provider is configured AND no available provider - matches the legacy preference; the dispatcher then returns a "set up a - provider" error to the user. + Rules, in order (see module docstring): explicit config wins even when + ``is_available()`` is False (the dispatcher surfaces a precise + "X_API_KEY is not set" error instead of a silent switch); then the single + available capable provider; then the availability-filtered legacy walk; + then the keyless free-tier walk; else None. """ - with _lock: - snapshot = dict(_providers) - snapshot.update(_scoped_providers.get(hermes_home_key(), {})) + snapshot = _registry.merged() def _capable(p: WebSearchProvider) -> bool: if capability == "search": @@ -245,17 +147,9 @@ def _resolve(configured: Optional[str], *, capability: str) -> Optional[WebSearc return bool(p.supports_extract()) return False - def _is_available_safe(p: WebSearchProvider) -> bool: - """Wrap ``is_available()`` so a buggy provider doesn't kill resolution.""" - try: - return bool(p.is_available()) - except Exception as exc: # noqa: BLE001 - logger.debug("provider %s.is_available() raised %s", p.name, exc) - return False + def _available(p: WebSearchProvider) -> bool: + return is_available_safe(p, logger, "provider %s.is_available() raised %s") - # 1. Explicit config wins — return regardless of is_available() so the - # user gets a precise downstream error message rather than a silent - # backend switch. Matches _get_backend() in web_tools.py. if configured: provider = snapshot.get(configured) if provider is not None and _capable(provider): @@ -271,31 +165,20 @@ def _resolve(configured: Optional[str], *, capability: str) -> Optional[WebSearc configured, capability, ) - # 2. + 3. Fallback path — filter by availability so we don't surface - # a provider the user has no credentials for. Without this filter, - # a registered-but-unconfigured provider could end up "active" on - # a fresh install with no API keys at all. - eligible = [ - p for p in snapshot.values() - if _capable(p) and _is_available_safe(p) - ] + # Fallbacks are availability-filtered so a registered-but-keyless provider + # never becomes "active" on a fresh install. + eligible = [p for p in snapshot.values() if _capable(p) and _available(p)] if len(eligible) == 1: return eligible[0] for legacy in _LEGACY_PREFERENCE: provider = snapshot.get(legacy) - if ( - provider is not None - and _capable(provider) - and _is_available_safe(provider) - ): + if provider is not None and _capable(provider) and _available(provider): return provider - # 4. Keyless free-tier walk — the user has NO credentialed/importable - # backend at all. Fall back to providers that can serve anonymously - # (public MCP free tiers), unless disabled via - # ``web.keyless_fallback: false``. This tier never pre-empts a keyed - # setup: it is only reachable when the legacy walk found nothing. + # Keyless free tier (anonymous public MCP tiers) is last-resort only: it is + # reachable solely when the legacy walk found nothing, never pre-empting a + # keyed setup. Disabled via ``web.keyless_fallback: false``. if _keyless_tier_enabled(): for name in _keyless_preference(): provider = snapshot.get(name) @@ -325,41 +208,23 @@ def _keyless_tier_enabled() -> bool: def _disabled_web_plugin_for(configured: Optional[str] = None, *, capability: Optional[str] = None) -> Optional[str]: - """Return the plugin key of a *disabled* bundled web plugin that would - have provided the configured backend, or None. + """Plugin key of a *disabled* bundled web plugin that would have provided + the configured backend (``web._backend`` → ``web.backend``), + or None. - When a user sets ``web.extract_backend: firecrawl`` (or the search - equivalent) but also lists ``web-firecrawl`` in ``plugins.disabled``, - the provider never registers and the dispatcher would otherwise emit a - misleading "No web extract provider configured. Set web.extract_backend - to ..." error — even though the backend IS configured correctly. The - real fix is to re-enable the plugin. This helper detects that case so - the dispatcher can point the user at the actual cause (issue #40190 - follow-up: pi314's disabled-plugin symptom). - - Pass ``capability`` ("search" | "extract") to resolve the configured - name straight from ``config.yaml`` (``web._backend`` → - ``web.backend``). This is more reliable than the resolved backend the - dispatcher fell back to, since a disabled provider fails the - ``_is_backend_available`` gate and the dispatcher silently drops to - the shared default. An explicit ``configured`` name still wins when - given. - - Matching is by convention: bundled web plugins live under the - ``web/`` key with the provider ``name`` differing only in - hyphen/underscore (``brave-free`` provider ⇄ ``web/brave_free`` key, - ``firecrawl`` ⇄ ``web/firecrawl``). We normalize both sides before - comparing so every bundled provider is covered without hardcoding a - per-vendor table. + Lets the dispatcher say "re-enable web-firecrawl" instead of a misleading + "No web extract provider configured" when the backend IS configured but + listed in ``plugins.disabled``. Resolving from config.yaml (rather than + the resolved backend) matters because a disabled provider fails the + availability gate and the dispatcher silently drops to the default. + Bundled web plugins live under ``web/`` with the provider name + differing only by hyphen/underscore, so both sides are normalized. """ def _norm(s: str) -> str: return s.strip().lower().replace("-", "_") if not configured and capability in ("search", "extract"): - configured = ( - _read_config_key("web", f"{capability}_backend") - or _read_config_key("web", "backend") - ) + configured = _configured_backend(capability) if not configured: return None @@ -384,27 +249,11 @@ def _disabled_web_plugin_for(configured: Optional[str] = None, *, capability: Op def get_active_search_provider() -> Optional[WebSearchProvider]: - """Resolve the currently-active web search provider. - - Reads ``web.search_backend`` (preferred) or ``web.backend`` (shared - fallback) from config.yaml; falls back per the module docstring. - """ - explicit = _read_config_key("web", "search_backend") or _read_config_key("web", "backend") - return _resolve(explicit, capability="search") + """Resolve the currently-active web search provider.""" + return _resolve(_configured_backend("search"), capability="search") def get_active_extract_provider() -> Optional[WebSearchProvider]: - """Resolve the currently-active web extract provider. + """Resolve the currently-active web extract provider.""" + return _resolve(_configured_backend("extract"), capability="extract") - Reads ``web.extract_backend`` (preferred) or ``web.backend`` (shared - fallback) from config.yaml; falls back per the module docstring. - """ - explicit = _read_config_key("web", "extract_backend") or _read_config_key("web", "backend") - return _resolve(explicit, capability="extract") - - -def _reset_for_tests() -> None: - """Clear the registry. **Test-only.**""" - with _lock: - _providers.clear() - _scoped_providers.clear()