- 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)
245 lines
9.9 KiB
Python
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
|