refactor(agent/providers): one ProviderRegistry engine behind every *_registry module
- provider_registry.py: ProviderRegistry (global + per-scope maps, lock, generation counters, register/list/get/snapshot/restore/reset) with export() binding the historical module-level names and _providers/ _scoped_providers/_lock test hooks into each *_registry module - is_available_safe / configured_provider_name replace the 4 nested _is_available_safe closures and 2 config-reading blocks - browser/image_gen/video_gen/web_search/terminal_env/tts/transcription registries keep their public API, log strings, error strings, builtin collision policy (warn vs raise) and key normalization (strip vs lower)
This commit is contained in:
+24
-172
@@ -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/<vendor>/`` 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
|
||||
|
||||
+19
-145
@@ -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()
|
||||
|
||||
@@ -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 ``<section>.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
|
||||
+21
-133
@@ -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
|
||||
|
||||
+31
-133
@@ -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())
|
||||
|
||||
+34
-139
@@ -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.<name>: 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.<name>: 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())
|
||||
|
||||
+18
-138
@@ -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()
|
||||
|
||||
+41
-192
@@ -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.<capability>_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.<capability>_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_<capability>()`` 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.<capability>_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.<capability>_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/<vendor>`` 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/<vendor>`` 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()
|
||||
|
||||
Reference in New Issue
Block a user