Files
hermes-agent/agent/provider_registry.py
Teknium 96e952a4f8 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)
2026-09-02 13:53:28 -07:00

245 lines
9.9 KiB
Python

"""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