fix(plugins): isolate ownership by profile

This commit is contained in:
doncazper
2026-08-01 20:03:12 -07:00
committed by Teknium
parent 4e1b2e436c
commit 85020f2238
29 changed files with 3101 additions and 479 deletions
+10
View File
@@ -1441,6 +1441,16 @@ def init_agent(
print(f"🔄 Fallback chain ({len(agent._fallback_chain)} providers): " +
" → ".join(f"{f['model']} ({f['provider']})" for f in agent._fallback_chain))
# A multiplexed gateway may enter a different HERMES_HOME after
# ``model_tools`` was first imported. Ensure that profile's keyed plugin
# manager has discovered its registrations before taking the tool snapshot.
try:
from hermes_cli.plugins import discover_plugins
discover_plugins()
except Exception:
logger.warning("Plugin discovery failed during agent setup", exc_info=True)
# Get available tools with filtering. Capture the registry generation this
# snapshot is derived from FIRST, so a later concurrent refresh can tell
# whether it holds a newer or staler view (see refresh_agent_mcp_tools).
+57 -12
View File
@@ -41,15 +41,19 @@ import threading
from typing import Dict, List, Optional
from agent.browser_provider import BrowserProvider
from hermes_constants import hermes_home_key
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) -> None:
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
@@ -61,12 +65,19 @@ def register_provider(provider: BrowserProvider) -> None:
f"register_provider() expects a BrowserProvider instance, "
f"got {type(provider).__name__}"
)
name = provider.name
if not isinstance(name, str) or not name.strip():
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:
existing = _providers.get(name)
_providers[name] = provider
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)",
@@ -79,34 +90,63 @@ def register_provider(provider: BrowserProvider) -> None:
)
def list_providers() -> List[BrowserProvider]:
def list_providers(*, scope: Optional[str] = None) -> List[BrowserProvider]:
"""Return all registered providers, sorted by name."""
with _lock:
items = list(_providers.values())
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) -> Optional[BrowserProvider]:
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:
return _providers.get(name.strip())
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:
if _providers.get(name) is not current:
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
if target.get(key) is not current:
return False
if previous is None:
_providers.pop(name, None)
target.pop(key, None)
else:
_providers[name] = previous
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
@@ -161,6 +201,7 @@ def _resolve(configured: Optional[str]) -> Optional[BrowserProvider]:
"""
with _lock:
snapshot = dict(_providers)
snapshot.update(_scoped_providers.get(hermes_home_key(), {}))
def _is_available_safe(p: BrowserProvider) -> bool:
"""Wrap ``is_available()`` so a buggy provider doesn't kill resolution."""
@@ -204,5 +245,9 @@ def _resolve(configured: Optional[str]) -> Optional[BrowserProvider]:
def _reset_for_tests() -> None:
"""Clear the registry. **Test-only.**"""
global _generation
with _lock:
_providers.clear()
_scoped_providers.clear()
_scoped_generations.clear()
_generation += 1
+35 -12
View File
@@ -25,15 +25,17 @@ import threading
from typing import Dict, List, Optional
from agent.image_gen_provider import ImageGenProvider
from hermes_constants import hermes_home_key
logger = logging.getLogger(__name__)
_providers: Dict[str, ImageGenProvider] = {}
_scoped_providers: Dict[str, Dict[str, ImageGenProvider]] = {}
_lock = threading.Lock()
def register_provider(provider: ImageGenProvider) -> None:
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
@@ -45,46 +47,65 @@ def register_provider(provider: ImageGenProvider) -> None:
f"register_provider() expects an ImageGenProvider instance, "
f"got {type(provider).__name__}"
)
name = provider.name
if not isinstance(name, str) or not name.strip():
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:
existing = _providers.get(name)
_providers[name] = provider
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() -> List[ImageGenProvider]:
def list_providers(*, scope: Optional[str] = None) -> List[ImageGenProvider]:
"""Return all registered providers, sorted by name."""
with _lock:
items = list(_providers.values())
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) -> Optional[ImageGenProvider]:
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:
return _providers.get(name.strip())
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:
if _providers.get(name) is not current:
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
if target.get(key) is not current:
return False
if previous is None:
_providers.pop(name, None)
target.pop(key, None)
else:
_providers[name] = previous
target[key] = previous
if scope is not None and not target:
_scoped_providers.pop(scope, None)
return True
@@ -120,6 +141,7 @@ def get_active_provider() -> Optional[ImageGenProvider]:
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."""
@@ -159,3 +181,4 @@ def _reset_for_tests() -> None:
"""Clear the registry. **Test-only.**"""
with _lock:
_providers.clear()
_scoped_providers.clear()
+107 -57
View File
@@ -30,6 +30,7 @@ from __future__ import annotations
import concurrent.futures
import logging
import os
import threading
from dataclasses import dataclass, field
from pathlib import Path
from typing import Dict, List, MutableMapping, Optional
@@ -43,6 +44,7 @@ from agent.secret_sources.base import (
reset_source_environment,
set_source_environment,
)
from hermes_constants import hermes_home_key
logger = logging.getLogger(__name__)
@@ -51,7 +53,9 @@ logger = logging.getLogger(__name__)
# recorded beside each source so consumers never infer ownership from names.
_SOURCES: Dict[str, SecretSource] = {}
_SOURCE_ORIGINS: Dict[str, str] = {}
_SCOPED_SOURCES: Dict[str, Dict[str, SecretSource]] = {}
_BUILTINS_LOADED = False
_REGISTRY_LOCK = threading.RLock()
@dataclass
@@ -101,6 +105,7 @@ def register_source(
*,
replace: bool = False,
builtin: bool = False,
scope: Optional[str] = None,
) -> bool:
"""Register a secret source. Returns True on success.
@@ -133,48 +138,80 @@ def register_source(
name, getattr(source, "shape", None),
)
return False
if name in _SOURCES and not replace:
logger.warning("Secret source '%s' already registered; ignoring duplicate", name)
return False
scheme = getattr(source, "scheme", None)
if scheme:
for other_name, other in _SOURCES.items():
if other_name != name and getattr(other, "scheme", None) == scheme:
logger.warning(
"Ignoring secret source '%s': scheme '%s://' is already "
"owned by source '%s'",
name, scheme, other_name,
)
return False
_SOURCES[name] = source
_SOURCE_ORIGINS[name] = "builtin" if builtin else "plugin"
with _REGISTRY_LOCK:
effective = dict(_SOURCES)
if scope is not None:
effective.update(_SCOPED_SOURCES.get(scope, {}))
if name in effective and not replace:
logger.warning(
"Secret source '%s' already registered; ignoring duplicate", name
)
return False
scheme = getattr(source, "scheme", None)
if scheme:
for other_name, other in effective.items():
if other_name != name and getattr(other, "scheme", None) == scheme:
logger.warning(
"Ignoring secret source '%s': scheme '%s://' is already "
"owned by source '%s'",
name,
scheme,
other_name,
)
return False
target = _SOURCES if scope is None else _SCOPED_SOURCES.setdefault(scope, {})
target[name] = source
if scope is None:
_SOURCE_ORIGINS[name] = "builtin" if builtin else "plugin"
return True
def get_source(name: str) -> Optional[SecretSource]:
def get_source(name: str, *, scope: Optional[str] = None) -> Optional[SecretSource]:
_ensure_builtin_sources()
return _SOURCES.get(name)
with _REGISTRY_LOCK:
return _SCOPED_SOURCES.get(scope or hermes_home_key(), {}).get(
name
) or _SOURCES.get(name)
def snapshot_registration(
name: str, *, scope: Optional[str] = None
) -> Optional[SecretSource]:
"""Return the registration owned by exactly one registry layer."""
_ensure_builtin_sources()
with _REGISTRY_LOCK:
target = _SOURCES if scope is None else _SCOPED_SOURCES.get(scope, {})
return target.get(name)
def restore_registration(
name: str,
current: SecretSource,
previous: Optional[SecretSource],
*,
scope: Optional[str] = None,
) -> bool:
"""Restore a host-owned source registration if it is still current."""
_ensure_builtin_sources()
if _SOURCES.get(name) is not current:
return False
if previous is None:
_SOURCES.pop(name, None)
else:
_SOURCES[name] = previous
with _REGISTRY_LOCK:
target = _SOURCES if scope is None else _SCOPED_SOURCES.setdefault(scope, {})
if target.get(name) is not current:
return False
if previous is None:
target.pop(name, None)
else:
target[name] = previous
if scope is not None and not target:
_SCOPED_SOURCES.pop(scope, None)
return True
def list_sources() -> List[SecretSource]:
def list_sources(*, scope: Optional[str] = None) -> List[SecretSource]:
_ensure_builtin_sources()
return list(_SOURCES.values())
with _REGISTRY_LOCK:
merged = dict(_SOURCES)
merged.update(_SCOPED_SOURCES.get(scope or hermes_home_key(), {}))
return list(merged.values())
def list_plugin_sources() -> List[SecretSource]:
@@ -194,37 +231,46 @@ def _ensure_builtin_sources() -> None:
source can never break registration of the others.
"""
global _BUILTINS_LOADED
if _BUILTINS_LOADED:
return
_BUILTINS_LOADED = True
try:
from agent.secret_sources.bitwarden import BitwardenSource
with _REGISTRY_LOCK:
if _BUILTINS_LOADED:
return
_BUILTINS_LOADED = True
try:
from agent.secret_sources.bitwarden import BitwardenSource
register_source(BitwardenSource(), builtin=True)
except Exception: # noqa: BLE001 — never block startup
logger.warning("Failed to register bundled Bitwarden secret source",
exc_info=True)
try:
from agent.secret_sources.onepassword import OnePasswordSource
register_source(BitwardenSource(), builtin=True)
except Exception: # noqa: BLE001 — never block startup
logger.warning(
"Failed to register bundled Bitwarden secret source",
exc_info=True,
)
try:
from agent.secret_sources.onepassword import OnePasswordSource
register_source(OnePasswordSource(), builtin=True)
except Exception: # noqa: BLE001 — never block startup
logger.warning("Failed to register bundled 1Password secret source",
exc_info=True)
try:
from agent.secret_sources.command import CommandSource
register_source(OnePasswordSource(), builtin=True)
except Exception: # noqa: BLE001 — never block startup
logger.warning(
"Failed to register bundled 1Password secret source",
exc_info=True,
)
try:
from agent.secret_sources.command import CommandSource
register_source(CommandSource(), builtin=True)
except Exception: # noqa: BLE001 — never block startup
logger.warning("Failed to register bundled command secret source",
exc_info=True)
register_source(CommandSource(), builtin=True)
except Exception: # noqa: BLE001 — never block startup
logger.warning(
"Failed to register bundled command secret source",
exc_info=True,
)
def _reset_registry_for_tests() -> None:
global _BUILTINS_LOADED
_SOURCES.clear()
_SOURCE_ORIGINS.clear()
_BUILTINS_LOADED = False
with _REGISTRY_LOCK:
_SOURCES.clear()
_SOURCE_ORIGINS.clear()
_SCOPED_SOURCES.clear()
_BUILTINS_LOADED = False
# ---------------------------------------------------------------------------
@@ -287,7 +333,9 @@ def _fetch_with_timeout(
return result
def _ordered_enabled_sources(secrets_cfg: dict) -> List[SecretSource]:
def _ordered_enabled_sources(
secrets_cfg: dict, *, scope: Optional[str] = None
) -> List[SecretSource]:
"""Resolve which sources run, in which order.
Order: the optional ``secrets.sources`` list wins; sources not named
@@ -295,28 +343,28 @@ def _ordered_enabled_sources(secrets_cfg: dict) -> List[SecretSource]:
``is_enabled`` says so for its config section. Mapped-vs-bulk
precedence is applied on top of this order by :func:`apply_all`.
"""
_ensure_builtin_sources()
sources = {source.name: source for source in list_sources(scope=scope)}
explicit = secrets_cfg.get("sources")
order: List[str] = []
if isinstance(explicit, list):
for entry in explicit:
if isinstance(entry, str) and entry in _SOURCES and entry not in order:
if isinstance(entry, str) and entry in sources and entry not in order:
order.append(entry)
unknown = [e for e in explicit
if isinstance(e, str) and e not in _SOURCES]
if isinstance(e, str) and e not in sources]
if unknown:
logger.warning(
"secrets.sources names unknown source(s): %s (known: %s)",
", ".join(unknown), ", ".join(_SOURCES) or "none",
", ".join(unknown), ", ".join(sources) or "none",
)
for name in _SOURCES:
for name in sources:
if name not in order:
order.append(name)
enabled: List[SecretSource] = []
for name in order:
source = _SOURCES[name]
source = sources[name]
cfg = secrets_cfg.get(name)
cfg = cfg if isinstance(cfg, dict) else {}
try:
@@ -398,7 +446,9 @@ def apply_all(secrets_cfg: dict, home_path: Path,
report = ApplyReport()
secrets_cfg = secrets_cfg if isinstance(secrets_cfg, dict) else {}
enabled = _ordered_enabled_sources(secrets_cfg)
enabled = _ordered_enabled_sources(
secrets_cfg, scope=hermes_home_key(home_path)
)
if not enabled:
return report
+1
View File
@@ -355,6 +355,7 @@ def _tool_search_scoped_names(agent) -> frozenset:
enabled = getattr(agent, "enabled_toolsets", None)
disabled = getattr(agent, "disabled_toolsets", None)
cache_key = (
_registry.current_scope_key(),
getattr(_registry, "_generation", 0),
frozenset(enabled) if enabled is not None else None,
frozenset(disabled) if disabled is not None else None,
+32 -10
View File
@@ -24,6 +24,7 @@ import threading
from typing import Dict, List, Optional
from agent.transcription_provider import TranscriptionProvider
from hermes_constants import hermes_home_key
logger = logging.getLogger(__name__)
@@ -50,10 +51,11 @@ _BUILTIN_NAMES = frozenset({
_providers: Dict[str, TranscriptionProvider] = {}
_scoped_providers: Dict[str, Dict[str, TranscriptionProvider]] = {}
_lock = threading.Lock()
def register_provider(provider: TranscriptionProvider) -> None:
def register_provider(provider: TranscriptionProvider, *, scope: Optional[str] = None) -> None:
"""Register a transcription provider.
Rejects:
@@ -85,8 +87,9 @@ def register_provider(provider: TranscriptionProvider) -> None:
)
return
with _lock:
existing = _providers.get(key)
_providers[key] = provider
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)",
@@ -99,14 +102,16 @@ def register_provider(provider: TranscriptionProvider) -> None:
)
def list_providers() -> List[TranscriptionProvider]:
def list_providers(*, scope: Optional[str] = None) -> List[TranscriptionProvider]:
"""Return all registered providers, sorted by name."""
with _lock:
items = list(_providers.values())
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) -> Optional[TranscriptionProvider]:
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
@@ -115,23 +120,39 @@ def get_provider(name: str) -> Optional[TranscriptionProvider]:
"""
if not isinstance(name, str):
return None
return _providers.get(name.strip().lower())
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:
if _providers.get(key) is not current:
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
if target.get(key) is not current:
return False
if previous is None:
_providers.pop(key, None)
target.pop(key, None)
else:
_providers[key] = previous
target[key] = previous
if scope is not None and not target:
_scoped_providers.pop(scope, None)
return True
@@ -139,3 +160,4 @@ def _reset_for_tests() -> None:
"""Clear the registry. **Test-only.**"""
with _lock:
_providers.clear()
_scoped_providers.clear()
+32 -10
View File
@@ -33,6 +33,7 @@ import threading
from typing import Dict, List, Optional
from agent.tts_provider import TTSProvider
from hermes_constants import hermes_home_key
logger = logging.getLogger(__name__)
@@ -61,10 +62,11 @@ _BUILTIN_NAMES = frozenset({
_providers: Dict[str, TTSProvider] = {}
_scoped_providers: Dict[str, Dict[str, TTSProvider]] = {}
_lock = threading.Lock()
def register_provider(provider: TTSProvider) -> None:
def register_provider(provider: TTSProvider, *, scope: Optional[str] = None) -> None:
"""Register a TTS provider.
Rejects:
@@ -95,8 +97,9 @@ def register_provider(provider: TTSProvider) -> None:
)
return
with _lock:
existing = _providers.get(key)
_providers[key] = provider
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)",
@@ -109,14 +112,16 @@ def register_provider(provider: TTSProvider) -> None:
)
def list_providers() -> List[TTSProvider]:
def list_providers(*, scope: Optional[str] = None) -> List[TTSProvider]:
"""Return all registered providers, sorted by name."""
with _lock:
items = list(_providers.values())
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) -> Optional[TTSProvider]:
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
@@ -125,23 +130,39 @@ def get_provider(name: str) -> Optional[TTSProvider]:
"""
if not isinstance(name, str):
return None
return _providers.get(name.strip().lower())
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:
if _providers.get(key) is not current:
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
if target.get(key) is not current:
return False
if previous is None:
_providers.pop(key, None)
target.pop(key, None)
else:
_providers[key] = previous
target[key] = previous
if scope is not None and not target:
_scoped_providers.pop(scope, None)
return True
@@ -149,3 +170,4 @@ def _reset_for_tests() -> None:
"""Clear the registry. **Test-only.**"""
with _lock:
_providers.clear()
_scoped_providers.clear()
+35 -12
View File
@@ -29,15 +29,17 @@ import threading
from typing import Dict, List, Optional
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) -> None:
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
@@ -49,46 +51,65 @@ def register_provider(provider: VideoGenProvider) -> None:
f"register_provider() expects a VideoGenProvider instance, "
f"got {type(provider).__name__}"
)
name = provider.name
if not isinstance(name, str) or not name.strip():
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:
existing = _providers.get(name)
_providers[name] = provider
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() -> List[VideoGenProvider]:
def list_providers(*, scope: Optional[str] = None) -> List[VideoGenProvider]:
"""Return all registered providers, sorted by name."""
with _lock:
items = list(_providers.values())
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) -> Optional[VideoGenProvider]:
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:
return _providers.get(name.strip())
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:
if _providers.get(name) is not current:
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
if target.get(key) is not current:
return False
if previous is None:
_providers.pop(name, None)
target.pop(key, None)
else:
_providers[name] = previous
target[key] = previous
if scope is not None and not target:
_scoped_providers.pop(scope, None)
return True
@@ -113,6 +134,7 @@ def get_active_provider() -> Optional[VideoGenProvider]:
with _lock:
snapshot = dict(_providers)
snapshot.update(_scoped_providers.get(hermes_home_key(), {}))
if configured:
provider = snapshot.get(configured)
@@ -147,3 +169,4 @@ def _reset_for_tests() -> None:
"""Clear the registry. **Test-only.**"""
with _lock:
_providers.clear()
_scoped_providers.clear()
+35 -12
View File
@@ -37,15 +37,17 @@ import threading
from typing import Dict, List, Optional
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) -> None:
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
@@ -57,12 +59,14 @@ def register_provider(provider: WebSearchProvider) -> None:
f"register_provider() expects a WebSearchProvider instance, "
f"got {type(provider).__name__}"
)
name = provider.name
if not isinstance(name, str) or not name.strip():
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:
existing = _providers.get(name)
_providers[name] = provider
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)",
@@ -75,34 +79,51 @@ def register_provider(provider: WebSearchProvider) -> None:
)
def list_providers() -> List[WebSearchProvider]:
def list_providers(*, scope: Optional[str] = None) -> List[WebSearchProvider]:
"""Return all registered providers, sorted by name."""
with _lock:
items = list(_providers.values())
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) -> Optional[WebSearchProvider]:
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:
return _providers.get(name.strip())
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:
if _providers.get(name) is not current:
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
if target.get(key) is not current:
return False
if previous is None:
_providers.pop(name, None)
target.pop(key, None)
else:
_providers[name] = previous
target[key] = previous
if scope is not None and not target:
_scoped_providers.pop(scope, None)
return True
@@ -178,6 +199,7 @@ def _resolve(configured: Optional[str], *, capability: str) -> Optional[WebSearc
"""
with _lock:
snapshot = dict(_providers)
snapshot.update(_scoped_providers.get(hermes_home_key(), {}))
def _capable(p: WebSearchProvider) -> bool:
if capability == "search":
@@ -318,3 +340,4 @@ def _reset_for_tests() -> None:
"""Clear the registry. **Test-only.**"""
with _lock:
_providers.clear()
_scoped_providers.clear()
+1
View File
@@ -564,6 +564,7 @@ class GatewayAuthorizationMixin:
if source.platform not in platform_env_map:
try:
from gateway.platform_registry import platform_registry
entry = platform_registry.get(source.platform.value)
if entry:
if entry.allowed_users_env:
+287 -58
View File
@@ -29,12 +29,36 @@ Usage (gateway side):
"""
import logging
import sys
import threading
from dataclasses import dataclass, field
from typing import Any, Awaitable, Callable, Optional
from hermes_constants import hermes_home_key
logger = logging.getLogger(__name__)
def _plugin_scope_from_callable(callback: Callable) -> Optional[str]:
"""Infer a plugin profile from code registered outside PluginContext."""
try:
from tools.registry import registry as tool_registry
return tool_registry.plugin_scope_for_callable(callback)
except (ImportError, AttributeError):
return None
def _caller_plugin_scope() -> Optional[str]:
try:
module_name = sys._getframe(2).f_globals.get("__name__", "") or ""
except Exception:
return None
return _plugin_scope_from_callable(
type("_Caller", (), {"__module__": module_name})
)
@dataclass
class PlatformEntry:
"""Metadata and factory for a single platform adapter."""
@@ -208,12 +232,17 @@ class PlatformEntry:
class PlatformRegistry:
"""Central registry of platform adapters.
Thread-safe for reads (dict lookups are atomic under GIL).
Writes happen at startup during sequential discovery.
Registrations are serialized, and concurrent lazy lookups share an
in-flight event while the loader runs outside the registry lock.
"""
def __init__(self) -> None:
self._lock = threading.RLock()
# Process-global registrations (for example the built-in relay).
self._entries: dict[str, PlatformEntry] = {}
# Plugin adapters are isolated per resolved HERMES_HOME and overlay the
# process-global entries for lookups in that profile's runtime scope.
self._scoped_entries: dict[str, dict[str, PlatformEntry]] = {}
# Deferred platform loaders: name -> zero-arg callable that imports the
# owning plugin module (which calls register() and populates _entries).
#
@@ -227,10 +256,50 @@ class PlatformRegistry:
# actually asks for that platform (gateway start, cron delivery,
# `hermes setup`/`gateway status`, send_message).
self._deferred: dict[str, Callable[[], None]] = {}
self._scoped_deferred: dict[str, dict[str, Callable[[], None]]] = {}
self._inflight: dict[tuple[Optional[str], str], threading.Event] = {}
self._inflight_loaders: dict[
tuple[Optional[str], str], Callable[[], None]
] = {}
self._inflight_owners: dict[tuple[Optional[str], str], int] = {}
self._cancelled_inflight: set[tuple[Optional[str], str]] = set()
# A failed loader is no longer discoverable, but its identity remains
# until ownership teardown can CAS-restore the displaced predecessor.
self._consumed_loaders: dict[
tuple[Optional[str], str], Callable[[], None]
] = {}
@staticmethod
def current_scope_key() -> str:
return hermes_home_key()
def _scope_maps(
self,
scope: Optional[str],
*,
create: bool = False,
) -> tuple[dict[str, PlatformEntry], dict[str, Callable[[], None]]]:
if scope is None:
return self._entries, self._deferred
if create:
return (
self._scoped_entries.setdefault(scope, {}),
self._scoped_deferred.setdefault(scope, {}),
)
return (
self._scoped_entries.get(scope, {}),
self._scoped_deferred.get(scope, {}),
)
# -- deferred loading ----------------------------------------------------
def register_deferred(self, name: str, loader: Callable[[], None]) -> None:
def register_deferred(
self,
name: str,
loader: Callable[[], None],
*,
scope: Optional[str] = None,
) -> None:
"""Register a lazy loader for a platform that hasn't been imported yet.
*loader* is a zero-arg callable that imports the owning plugin module,
@@ -240,14 +309,19 @@ class PlatformRegistry:
registered directly (e.g. a built-in) takes precedence -- the deferred
loader is then dropped.
"""
if name in self._entries:
# Already concretely registered; no need to defer.
return
self._deferred[name] = loader
with self._lock:
entries, deferred = self._scope_maps(scope, create=True)
self._consumed_loaders.pop((scope, name), None)
if name in entries:
# Already concretely registered; no need to defer.
return
deferred[name] = loader
def snapshot_registration(
self,
name: str,
*,
scope: Optional[str] = None,
) -> tuple[Optional[PlatformEntry], Optional[Callable[[], None]]]:
"""Return the concrete and deferred state for *name* without resolving it.
@@ -256,13 +330,22 @@ class PlatformRegistry:
importing the displaced adapter as a side effect of taking the
snapshot.
"""
return self._entries.get(name), self._deferred.get(name)
with self._lock:
entries, deferred = self._scope_maps(scope)
loader = deferred.get(name)
if entries.get(name) is None and loader is None:
loader = self._inflight_loaders.get((scope, name))
if entries.get(name) is None and loader is None:
loader = self._consumed_loaders.get((scope, name))
return entries.get(name), loader
def restore_registration(
self,
name: str,
current: tuple[Optional[PlatformEntry], Optional[Callable[[], None]]],
previous: tuple[Optional[PlatformEntry], Optional[Callable[[], None]]],
*,
scope: Optional[str] = None,
) -> bool:
"""Restore a platform registration if its full state is still current.
@@ -271,29 +354,89 @@ class PlatformRegistry:
it displaced. Both concrete entries and deferred loaders are part of
the state because bundled platform plugins load lazily.
"""
current_state = self.snapshot_registration(name)
if (
current_state[0] is not current[0]
or current_state[1] is not current[1]
):
return False
with self._lock:
entries, deferred = self._scope_maps(scope, create=True)
entry = entries.get(name)
loader = deferred.get(name)
load_key = (scope, name)
if entry is None and loader is None:
loader = self._inflight_loaders.get(load_key)
if entry is None and loader is None:
loader = self._consumed_loaders.get(load_key)
current_state = (entry, loader)
is_current = not (
current_state[0] is not current[0]
or current_state[1] is not current[1]
)
if not is_current:
return False
entry, loader = previous
if entry is None:
self._entries.pop(name, None)
else:
self._entries[name] = entry
if loader is None:
self._deferred.pop(name, None)
else:
self._deferred[name] = loader
return True
previous_entry, previous_loader = previous
if previous_entry is None:
entries.pop(name, None)
else:
entries[name] = previous_entry
if previous_loader is None:
deferred.pop(name, None)
else:
deferred[name] = previous_loader
if load_key in self._inflight:
self._cancelled_inflight.add(load_key)
self._consumed_loaders.pop(load_key, None)
if scope is not None:
if not entries:
self._scoped_entries.pop(scope, None)
if not deferred:
self._scoped_deferred.pop(scope, None)
return True
def _resolve(self, name: str) -> None:
def _resolve(self, name: str, scope: Optional[str] = None) -> None:
"""Run the deferred loader for *name* if one is pending."""
loader = self._deferred.pop(name, None)
if loader is None:
loader: Optional[Callable[[], None]] = None
event: Optional[threading.Event] = None
load_key: tuple[Optional[str], str]
is_loader = False
with self._lock:
active_scope = scope or self.current_scope_key()
entries, deferred = self._scope_maps(active_scope)
scoped_key = (active_scope, name)
global_key = (None, name)
event = self._inflight.get(scoped_key)
load_key = scoped_key
if event is None and name not in entries:
loader = deferred.pop(name, None)
if event is None and loader is None and name not in entries:
event = self._inflight.get(global_key)
load_key = global_key
if event is None and loader is None and name not in entries:
loader = self._deferred.pop(name, None)
load_key = global_key
if event is None and loader is not None:
event = threading.Event()
self._inflight[load_key] = event
self._inflight_loaders[load_key] = loader
self._inflight_owners[load_key] = threading.get_ident()
is_loader = True
if event is None:
return
if (
not is_loader
and self._inflight_owners.get(load_key) == threading.get_ident()
):
logger.warning(
"Deferred platform '%s' recursively requested while loading",
name,
)
return
if not is_loader:
event.wait()
# Teardown may have restored an older deferred generation while
# cancelling the one we waited for. Resolve that predecessor in
# the same lookup instead of returning a one-shot false negative.
self._resolve(name, active_scope)
return
try:
loader()
except Exception as e:
@@ -303,6 +446,34 @@ class PlatformRegistry:
e,
exc_info=True,
)
finally:
with self._lock:
was_cancelled = load_key in self._cancelled_inflight
load_scope, _load_name = load_key
entries, deferred = self._scope_maps(load_scope)
if (
not was_cancelled
and name not in entries
and name not in deferred
):
self._consumed_loaders[load_key] = loader
self._inflight.pop(load_key, None)
self._inflight_loaders.pop(load_key, None)
self._inflight_owners.pop(load_key, None)
self._cancelled_inflight.discard(load_key)
event.set()
if was_cancelled:
self._resolve(name, active_scope)
def is_deferred_load_cancelled(
self,
name: str,
*,
scope: Optional[str] = None,
) -> bool:
"""Return whether ownership teardown cancelled an in-flight loader."""
with self._lock:
return (scope, name) in self._cancelled_inflight
def _resolve_all(self) -> None:
"""Run every pending deferred loader.
@@ -312,59 +483,119 @@ class PlatformRegistry:
gateway startup, ``hermes setup``/``gateway status``, channel
directory. CLI chat never iterates the full set.
"""
if not self._deferred:
return
# Snapshot keys -- loaders mutate _deferred as they resolve.
for name in list(self._deferred):
self._resolve(name)
active_scope = self.current_scope_key()
with self._lock:
_entries, scoped_deferred = self._scope_maps(active_scope)
scoped_names = set(scoped_deferred)
global_names = set(self._deferred)
for inflight_scope, name in self._inflight:
if inflight_scope == active_scope:
scoped_names.add(name)
elif inflight_scope is None:
global_names.add(name)
# Load outside the registry lock; each name has an in-flight event so
# concurrent readers wait for the same materialization.
for name in sorted(scoped_names):
self._resolve(name, active_scope)
for name in sorted(global_names):
self._resolve(name, active_scope)
def register(self, entry: PlatformEntry) -> None:
def register(
self,
entry: PlatformEntry,
*,
scope: Optional[str] = None,
) -> None:
"""Register a platform adapter entry.
If an entry with the same name exists, it is replaced (last writer
wins -- this lets plugins override built-in adapters if desired).
"""
# A concrete registration supersedes any pending deferred loader.
self._deferred.pop(entry.name, None)
if entry.name in self._entries:
prev = self._entries[entry.name]
logger.info(
"Platform '%s' re-registered (was %s, now %s)",
entry.name,
prev.source,
entry.source,
)
self._entries[entry.name] = entry
logger.debug("Registered platform adapter: %s (%s)", entry.name, entry.source)
with self._lock:
if scope is None and entry.source == "plugin":
scope = _caller_plugin_scope()
if scope is None:
scope = _plugin_scope_from_callable(entry.adapter_factory)
if scope is None:
scope = _plugin_scope_from_callable(entry.check_fn)
# A concrete registration supersedes any pending deferred loader.
entries, deferred = self._scope_maps(scope, create=True)
self._consumed_loaders.pop((scope, entry.name), None)
deferred.pop(entry.name, None)
if entry.name in entries:
prev = entries[entry.name]
logger.info(
"Platform '%s' re-registered (was %s, now %s)",
entry.name,
prev.source,
entry.source,
)
entries[entry.name] = entry
logger.debug("Registered platform adapter: %s (%s)", entry.name, entry.source)
def unregister(self, name: str) -> bool:
def unregister(self, name: str, *, scope: Optional[str] = None) -> bool:
"""Remove a platform entry. Returns True if it existed."""
self._deferred.pop(name, None)
removed = self._entries.pop(name, None) is not None
return removed
with self._lock:
inferred_scope = scope if scope is not None else _caller_plugin_scope()
active_scope = inferred_scope or self.current_scope_key()
entries, deferred = self._scope_maps(active_scope)
if inferred_scope is not None or name in entries or name in deferred:
deferred.pop(name, None)
removed = entries.pop(name, None) is not None
if not entries:
self._scoped_entries.pop(active_scope, None)
if not deferred:
self._scoped_deferred.pop(active_scope, None)
return removed
self._deferred.pop(name, None)
return self._entries.pop(name, None) is not None
def get(self, name: str) -> Optional[PlatformEntry]:
"""Look up a platform entry by name."""
if name not in self._entries:
self._resolve(name)
return self._entries.get(name)
scope = self.current_scope_key()
with self._lock:
entries, deferred = self._scope_maps(scope)
needs_resolve = name not in entries and (
name in deferred
or (name not in self._entries and name in self._deferred)
or (scope, name) in self._inflight
or (None, name) in self._inflight
)
if needs_resolve:
self._resolve(name, scope)
with self._lock:
entries, _deferred = self._scope_maps(scope)
return entries.get(name) or self._entries.get(name)
def all_entries(self) -> list[PlatformEntry]:
"""Return all registered platform entries."""
self._resolve_all()
return list(self._entries.values())
with self._lock:
entries = dict(self._entries)
entries.update(self._scoped_entries.get(self.current_scope_key(), {}))
return list(entries.values())
def plugin_entries(self) -> list[PlatformEntry]:
"""Return only plugin-registered platform entries."""
self._resolve_all()
return [e for e in self._entries.values() if e.source == "plugin"]
return [e for e in self.all_entries() if e.source == "plugin"]
def is_registered(self, name: str) -> bool:
# A deferred (not-yet-imported) platform still counts as registered --
# the loader will materialize it on first real use. This keeps cheap
# membership checks (toolset resolution, webhook deliver-target checks)
# from triggering a heavy import.
return name in self._entries or name in self._deferred
with self._lock:
scope = self.current_scope_key()
entries, deferred = self._scope_maps(scope)
return (
name in entries
or name in deferred
or name in self._entries
or name in self._deferred
or (scope, name) in self._inflight
or (None, name) in self._inflight
)
def create_adapter(self, name: str, config: Any) -> Optional[Any]:
"""Create an adapter instance for the given platform name.
@@ -376,9 +607,7 @@ class PlatformRegistry:
- validate_config() returns False (misconfigured)
- The factory raises an exception
"""
if name not in self._entries:
self._resolve(name)
entry = self._entries.get(name)
entry = self.get(name)
if entry is None:
return None
+3
View File
@@ -13778,6 +13778,9 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
with _profile_runtime_scope(profile_home):
profile_runtime_cfg = _load_gateway_runtime_config()
from hermes_cli.plugins import discover_plugins
discover_plugins()
profile_cfg = load_gateway_config()
violation = _own_policy_open_startup_violation(profile_cfg)
self._snapshot_profile_busy_modes(profile_name, profile_runtime_cfg)
+46 -14
View File
@@ -10,6 +10,7 @@ import logging
import threading
from typing import List, Optional
from hermes_constants import hermes_home_key
from hermes_cli.dashboard_auth.base import (
DashboardAuthProvider,
assert_protocol_compliance,
@@ -18,9 +19,20 @@ from hermes_cli.dashboard_auth.base import (
_log = logging.getLogger(__name__)
_lock = threading.Lock()
_providers: dict[str, DashboardAuthProvider] = {}
_scoped_providers: dict[str, dict[str, DashboardAuthProvider]] = {}
def register_provider(provider: DashboardAuthProvider) -> None:
def _merged(scope: Optional[str] = None) -> dict[str, DashboardAuthProvider]:
providers = dict(_providers)
providers.update(_scoped_providers.get(scope or hermes_home_key(), {}))
return providers
def register_provider(
provider: DashboardAuthProvider,
*,
scope: Optional[str] = None,
) -> None:
"""Register a provider.
Raises:
@@ -29,43 +41,64 @@ def register_provider(provider: DashboardAuthProvider) -> None:
"""
assert_protocol_compliance(type(provider))
with _lock:
if provider.name in _providers:
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
effective = target if scope is None else _merged(scope)
if provider.name in effective:
raise ValueError(
f"dashboard-auth provider already registered: {provider.name!r}"
)
_providers[provider.name] = provider
target[provider.name] = provider
_log.info(
"dashboard-auth: registered provider %r (%s)",
provider.name, provider.display_name,
)
def get_provider(name: str) -> Optional[DashboardAuthProvider]:
def get_provider(
name: str,
*,
scope: Optional[str] = None,
) -> Optional[DashboardAuthProvider]:
"""Return the registered provider for ``name``, or None if unknown."""
with _lock:
return _providers.get(name)
return _merged(scope).get(name)
def snapshot_registration(
name: str,
*,
scope: Optional[str] = None,
) -> Optional[DashboardAuthProvider]:
with _lock:
target = _providers if scope is None else _scoped_providers.get(scope, {})
return target.get(name)
def restore_registration(
name: str,
current: DashboardAuthProvider,
previous: Optional[DashboardAuthProvider],
*,
scope: Optional[str] = None,
) -> bool:
"""Restore a host-owned provider registration if it is still current."""
with _lock:
if _providers.get(name) is not current:
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
if target.get(name) is not current:
return False
if previous is None:
_providers.pop(name, None)
target.pop(name, None)
else:
_providers[name] = previous
target[name] = previous
if scope is not None and not target:
_scoped_providers.pop(scope, None)
return True
def list_providers() -> List[DashboardAuthProvider]:
def list_providers(*, scope: Optional[str] = None) -> List[DashboardAuthProvider]:
"""All registered providers, in registration order."""
with _lock:
return list(_providers.values())
return list(_merged(scope).values())
def list_token_providers() -> List[DashboardAuthProvider]:
@@ -78,8 +111,7 @@ def list_token_providers() -> List[DashboardAuthProvider]:
no token provider is registered — a token-authable route then fails
closed (401), never open.
"""
with _lock:
return [p for p in _providers.values() if getattr(p, "supports_token", False)]
return [p for p in list_providers() if getattr(p, "supports_token", False)]
def list_session_providers() -> List[DashboardAuthProvider]:
@@ -87,11 +119,11 @@ def list_session_providers() -> List[DashboardAuthProvider]:
sessions). The login page, /auth/login, and the gate's verify/refresh loops
consult only these. Mirror of list_token_providers.
"""
with _lock:
return [p for p in _providers.values() if getattr(p, "supports_session", True)]
return [p for p in list_providers() if getattr(p, "supports_session", True)]
def clear_providers() -> None:
"""Test-only: drop all registrations."""
with _lock:
_providers.clear()
_scoped_providers.clear()
+11 -3
View File
@@ -685,7 +685,9 @@ def _apply_external_secret_sources(home_path: Path) -> None:
)
if src.result.error:
print(f" {src.label}: {src.result.error}", file=sys.stderr)
hint = _remediation_hint(src.name, src.result.error_kind, cfg)
hint = _remediation_hint(
src.name, src.result.error_kind, cfg, scope=home_key
)
if hint:
print(f" {src.label}: → {hint}", file=sys.stderr)
for warn in src.result.warnings:
@@ -694,7 +696,13 @@ def _apply_external_secret_sources(home_path: Path) -> None:
print(f" Secret sources: {conflict}", file=sys.stderr)
def _remediation_hint(source_name: str, error_kind, secrets_cfg: dict) -> str:
def _remediation_hint(
source_name: str,
error_kind,
secrets_cfg: dict,
*,
scope: str | None = None,
) -> str:
"""Ask the failed source for its one-line fix-it hint.
Defensive wrapper: remediation() is a pure mapping and shouldn't
@@ -704,7 +712,7 @@ def _remediation_hint(source_name: str, error_kind, secrets_cfg: dict) -> str:
try:
from agent.secret_sources.registry import get_source
source = get_source(source_name)
source = get_source(source_name, scope=scope)
if source is None:
return ""
src_cfg = secrets_cfg.get(source_name)
+488 -188
View File
File diff suppressed because it is too large Load Diff
+26 -17
View File
@@ -2732,8 +2732,8 @@ def _prompt_choice(question: str, choices: list, default: int = 0) -> int:
# ─── Token Estimation ────────────────────────────────────────────────────────
# Module-level cache so discovery + tokenization runs at most once per process.
_tool_token_cache: Optional[Dict[str, int]] = None
# Profile-keyed cache so one process can serve distinct plugin tool catalogs.
_tool_token_cache: Optional[Dict[tuple[str, int], Dict[str, int]]] = None
def _estimate_tool_tokens() -> Dict[str, int]:
@@ -2746,25 +2746,33 @@ def _estimate_tool_tokens() -> Dict[str, int]:
Returns an empty dict when tiktoken or the registry is unavailable.
"""
global _tool_token_cache
if _tool_token_cache is not None:
return _tool_token_cache
from hermes_constants import hermes_home_key
scope = hermes_home_key()
try:
# Trigger full tool discovery (imports all tool modules).
import model_tools # noqa: F401
from tools.registry import registry
cache_key = (scope, registry._generation)
except Exception:
logger.debug("Tool registry unavailable; skipping token estimation")
cache_key = (scope, -1)
_tool_token_cache = _tool_token_cache or {}
_tool_token_cache[cache_key] = {}
return _tool_token_cache[cache_key]
if _tool_token_cache is not None and cache_key in _tool_token_cache:
return _tool_token_cache[cache_key]
try:
import tiktoken
enc = tiktoken.get_encoding("cl100k_base")
except Exception:
logger.debug("tiktoken unavailable; skipping tool token estimation")
_tool_token_cache = {}
return _tool_token_cache
try:
# Trigger full tool discovery (imports all tool modules).
import model_tools # noqa: F401
from tools.registry import registry
except Exception:
logger.debug("Tool registry unavailable; skipping token estimation")
_tool_token_cache = {}
return _tool_token_cache
_tool_token_cache = _tool_token_cache or {}
_tool_token_cache[cache_key] = {}
return _tool_token_cache[cache_key]
counts: Dict[str, int] = {}
for name in registry.get_all_tool_names():
@@ -2774,8 +2782,9 @@ def _estimate_tool_tokens() -> Dict[str, int]:
# {"type": "function", "function": <schema>}
text = _json.dumps({"type": "function", "function": schema})
counts[name] = len(enc.encode(text))
_tool_token_cache = counts
return _tool_token_cache
_tool_token_cache = _tool_token_cache or {}
_tool_token_cache[cache_key] = counts
return counts
def _prompt_toolset_checklist(
+12
View File
@@ -139,6 +139,18 @@ def get_hermes_home() -> Path:
return _hermes_home_from_env()
def hermes_home_key(path: str | Path | None = None) -> str:
"""Return a stable key for a Hermes home/profile directory.
Runtime registries use this key to isolate plugin-owned entries while
keeping built-in registrations process-global. ``strict=False`` preserves
useful behavior for profiles whose directories have not been created yet.
"""
candidate = Path(path) if path is not None else get_hermes_home()
resolved = candidate.expanduser().resolve(strict=False)
return os.path.normcase(str(resolved))
def get_process_hermes_home() -> Path:
"""Return the Hermes home for the running process, ignoring task overrides.
+17 -7
View File
@@ -274,7 +274,7 @@ _LEGACY_TOOLSET_MAP = {
# =============================================================================
# Module-level memoization for get_tool_definitions(). Keyed on
# (frozenset(enabled_toolsets), frozenset(disabled_toolsets), registry._generation).
# (profile scope, enabled/disabled toolsets, registry generation).
# Hot callers (gateway runner, AIAgent.__init__) invoke this on every turn
# with quiet_mode=True; caching avoids ~7 ms of registry walking + schema
# filtering + check_fn probing per call. Only active when quiet_mode=True
@@ -285,6 +285,7 @@ _LEGACY_TOOLSET_MAP = {
# inner check_fn TTL cache in registry.py handles environment drift (Docker
# daemon start/stop, env var changes, etc.) on a 30 s horizon.
_tool_defs_cache: Dict[tuple, List[Dict[str, Any]]] = {}
_tool_defs_cache_lock = threading.Lock()
# Hard cap on memoized get_tool_definitions() results. A long-lived Gateway
# process sees many distinct toolset/config fingerprints over its lifetime
@@ -299,7 +300,8 @@ def _clear_tool_defs_cache() -> None:
"""Drop memoized get_tool_definitions() results. Called when dynamic
schema dependencies change (e.g. discord capability cache reset,
execute_code sandbox reconfigured)."""
_tool_defs_cache.clear()
with _tool_defs_cache_lock:
_tool_defs_cache.clear()
def get_tool_definitions(
@@ -346,6 +348,7 @@ def get_tool_definitions(
profile_scope = check_fn_cache_scope()
if profile_scope != CHECK_FN_CACHE_BYPASS:
cache_key = (
registry.current_scope_key(),
frozenset(enabled_toolsets) if enabled_toolsets is not None else None,
frozenset(disabled_toolsets) if disabled_toolsets else None,
registry._generation,
@@ -356,7 +359,8 @@ def get_tool_definitions(
_is_dispatcher_owned_worker(),
profile_scope,
)
cached = _tool_defs_cache.get(cache_key) if cache_key is not None else None
with _tool_defs_cache_lock:
cached = _tool_defs_cache.get(cache_key) if cache_key is not None else None
if cached is not None:
# Update _last_resolved_tool_names so downstream callers see
# consistent state even on a cache hit.
@@ -379,10 +383,16 @@ def get_tool_definitions(
# Bound the cache with LRU eviction so a long-lived Gateway process
# doesn't accumulate entries unboundedly across the many distinct
# toolset/config fingerprints it sees over its lifetime (#19251).
if len(_tool_defs_cache) >= _TOOL_DEFS_CACHE_MAX:
_tool_defs_cache.pop(next(iter(_tool_defs_cache))) # evict oldest
_tool_defs_cache[cache_key] = result
return list(result)
with _tool_defs_cache_lock:
# Another thread may have populated this exact key while this
# thread computed. Reuse it and serialize capacity eviction.
cached = _tool_defs_cache.get(cache_key)
if cached is None:
if len(_tool_defs_cache) >= _TOOL_DEFS_CACHE_MAX:
_tool_defs_cache.pop(next(iter(_tool_defs_cache)))
_tool_defs_cache[cache_key] = result
cached = result
return list(cached)
if quiet_mode:
return list(result)
return result
+128
View File
@@ -0,0 +1,128 @@
"""Ownership leases for replaceable runtime registrations.
The coordinator models registration *generations*, not just value identity.
That distinction matters when the same provider singleton is registered again
after an older ownership generation was unloaded.
"""
from __future__ import annotations
import threading
from collections.abc import Callable, Hashable
from contextlib import contextmanager
from dataclasses import dataclass, field
from typing import Any
def same_registration(left: Any, right: Any) -> bool:
"""Compare opaque registry snapshots using identity only."""
if isinstance(left, tuple) and isinstance(right, tuple):
return len(left) == len(right) and all(
same_registration(a, b) for a, b in zip(left, right)
)
return left is right
@dataclass
class ReplacementLease:
"""One ownership generation in a replaceable registry slot."""
coordinator: "ReplacementCoordinator"
slot: Hashable
current: Any
previous: Any
restore: Callable[[Any], bool]
finalize: Callable[[], None] | None = None
predecessor: "ReplacementLease | None" = None
active: bool = field(default=True, init=False)
def dispose(self) -> None:
self.coordinator.dispose(self)
class ReplacementCoordinator:
"""Link and remove registration generations in arbitrary unload order."""
def __init__(self) -> None:
self._active: dict[Hashable, list[ReplacementLease]] = {}
self._lock = threading.RLock()
@contextmanager
def transaction(self):
"""Serialize a registry snapshot/write/acquire with lease disposal."""
with self._lock:
yield
def acquire(
self,
slot: Hashable,
*,
current: Any,
previous: Any,
restore: Callable[[Any], bool],
finalize: Callable[[], None] | None = None,
) -> ReplacementLease:
"""Attach a new live generation to the matching active predecessor."""
with self._lock:
leases = self._active.setdefault(slot, [])
predecessor = next(
(
lease
for lease in reversed(leases)
if lease.active and same_registration(lease.current, previous)
),
None,
)
lease = ReplacementLease(
coordinator=self,
slot=slot,
current=current,
previous=previous,
restore=restore,
finalize=finalize,
predecessor=predecessor,
)
leases.append(lease)
return lease
def dispose(self, lease: ReplacementLease) -> None:
"""Remove *lease*, restoring the nearest still-live predecessor."""
with self._lock:
if not lease.active:
return
leases = self._active.get(lease.slot, [])
latest = next(
(candidate for candidate in reversed(leases) if candidate.active),
None,
)
lease.active = False
# An older generation can share the exact same object identity as
# a newer one. Registry-level CAS cannot distinguish those leases,
# so only the latest live generation is allowed to mutate the slot.
try:
try:
if latest is lease:
replacement = lease.previous
predecessor = lease.predecessor
while predecessor is not None:
if predecessor.active:
replacement = predecessor.current
break
replacement = predecessor.previous
predecessor = predecessor.predecessor
lease.restore(replacement)
finally:
if lease.finalize is not None:
lease.finalize()
finally:
if leases:
self._active[lease.slot] = [
item for item in leases if item.active
]
if not self._active[lease.slot]:
self._active.pop(lease.slot, None)
replacement_coordinator = ReplacementCoordinator()
+5 -1
View File
@@ -1567,8 +1567,13 @@ class TestAuthorizationEmailMatch:
from gateway.config import GatewayConfig
from gateway.run import GatewayRunner
from gateway.session import SessionSource
from hermes_cli.plugins import discover_plugins
monkeypatch.setenv("GOOGLE_CHAT_ALLOWED_USERS", "alice@example.com")
# Plugin platforms become available during the normal gateway startup
# discovery pass. This unit test constructs GatewayRunner directly,
# so perform that lifecycle step explicitly before testing auth.
discover_plugins()
cfg = GatewayConfig()
runner = GatewayRunner(cfg)
runner.pairing_store = MagicMock()
@@ -1739,4 +1744,3 @@ class TestGoogleChatStandaloneSend:
assert kwargs["headers"]["Authorization"] == "Bearer the-token"
assert kwargs["json"] == {"text": "hello cron"}
File diff suppressed because it is too large Load Diff
+13 -5
View File
@@ -1322,22 +1322,30 @@ class TestPluginCommands:
try:
manager_a = plugins_mod.get_plugin_manager()
manager_a.discover_and_load()
module_a = manager_a._plugins["stateful-plugin"].module
finally:
reset_hermes_home_override(token_a)
assert "hermes_plugins.stateful_plugin.state" in sys.modules
assert sys.modules["hermes_plugins.stateful_plugin.state"].MARKER == "marker-a"
assert module_a is not None
module_a_state = f"{module_a.__name__}.state"
assert module_a_state in sys.modules
assert sys.modules[module_a_state].MARKER == "marker-a"
token_b = set_hermes_home_override(str(home_b))
try:
manager_b = plugins_mod.get_plugin_manager()
manager_b.discover_and_load()
module_b = manager_b._plugins["stateful-plugin"].module
finally:
reset_hermes_home_override(token_b)
# The submodule cached under sys.modules must now reflect profile
# b's code, not a leftover from profile a.
assert sys.modules["hermes_plugins.stateful_plugin.state"].MARKER == "marker-b"
# Each profile keeps a stable namespace, so concurrent/runtime relative
# imports cannot resolve another profile's package or submodules.
assert module_b is not None
module_b_state = f"{module_b.__name__}.state"
assert module_a.__name__ != module_b.__name__
assert sys.modules[module_a_state].MARKER == "marker-a"
assert sys.modules[module_b_state].MARKER == "marker-b"
assert (
manager_b._plugin_skills["stateful-plugin::marker"]["marker"] == "marker-b"
)
+34 -1
View File
@@ -13,7 +13,7 @@ import pytest
from agent.secret_sources import bitwarden as bw
from agent.secret_sources import onepassword as op
from agent.secret_sources.base import ErrorKind, SecretSource
from agent.secret_sources.base import ErrorKind, FetchResult, SecretSource
from agent.secret_sources.bitwarden import (
BitwardenSource,
_classify_bws_error,
@@ -149,3 +149,36 @@ def test_env_loader_prints_remediation_hint(tmp_path, monkeypatch, capsys):
assert "hermes secrets bitwarden token" in err
def test_remediation_hint_uses_explicit_profile_scope(tmp_path, monkeypatch):
from agent.secret_sources import registry
from hermes_cli import env_loader
class ScopedSource(SecretSource):
name = "scoped_hint"
label = "Scoped hint"
shape = "mapped"
def __init__(self, marker):
self.marker = marker
def fetch(self, cfg, home_path):
return FetchResult()
def remediation(self, kind, cfg):
return self.marker
monkeypatch.setattr(registry, "_ensure_builtin_sources", lambda: None)
registry._reset_registry_for_tests()
home_a = str((tmp_path / "hint-a").resolve())
home_b = str((tmp_path / "hint-b").resolve())
source_a = ScopedSource("profile-a")
source_b = ScopedSource("profile-b")
assert registry.register_source(source_a, scope=home_a)
assert registry.register_source(source_b, scope=home_b)
try:
assert env_loader._remediation_hint(
"scoped_hint", ErrorKind.AUTH_FAILED, {}, scope=home_b
) == "profile-b"
finally:
registry._reset_registry_for_tests()
@@ -95,6 +95,38 @@ class TestRegistration:
def test_rejects_non_secretsource_instance(self):
assert reg.register_source(object()) is False
def test_same_name_is_isolated_by_profile(self, tmp_path):
from hermes_constants import (
reset_hermes_home_override,
set_hermes_home_override,
)
home_a = str((tmp_path / "secrets-a").resolve())
home_b = str((tmp_path / "secrets-b").resolve())
source_a = _make_source(name="profile_secret", secrets={"A": "a"})
source_b = _make_source(name="profile_secret", secrets={"B": "b"})
assert reg.register_source(source_a, scope=home_a)
assert reg.register_source(source_b, scope=home_b)
token = set_hermes_home_override(home_a)
try:
assert reg.get_source("profile_secret") is source_a
explicit_b_env = {}
report = reg.apply_all(
{"profile_secret": {"enabled": True}},
Path(home_b),
environ=explicit_b_env,
)
assert report.sources[0].result.secrets == {"B": "b"}
assert explicit_b_env == {"B": "b"}
finally:
reset_hermes_home_override(token)
token = set_hermes_home_override(home_b)
try:
assert reg.get_source("profile_secret") is source_b
finally:
reset_hermes_home_override(token)
@@ -16,6 +16,9 @@ These tests pin:
"""
from __future__ import annotations
from concurrent.futures import ThreadPoolExecutor
from threading import Barrier
import pytest
import model_tools
@@ -80,3 +83,28 @@ class TestQuietModeCacheIsolation:
explains why the bug only hit Gateway."""
model_tools.get_tool_definitions(quiet_mode=False)
assert len(model_tools._tool_defs_cache) == 0
def test_concurrent_capacity_misses_evict_atomically(self, monkeypatch):
"""Two profile/toolset misses at capacity cannot race on eviction."""
barrier = Barrier(2)
def compute(*args, **kwargs):
barrier.wait(timeout=2)
return []
monkeypatch.setattr(model_tools, "_compute_tool_definitions", compute)
for index in range(model_tools._TOOL_DEFS_CACHE_MAX):
model_tools._tool_defs_cache[("old", index)] = []
with ThreadPoolExecutor(max_workers=2) as pool:
futures = [
pool.submit(
model_tools.get_tool_definitions,
enabled_toolsets=[f"concurrent_{index}"],
quiet_mode=True,
)
for index in range(2)
]
assert [future.result(timeout=2) for future in futures] == [[], []]
assert len(model_tools._tool_defs_cache) == model_tools._TOOL_DEFS_CACHE_MAX
@@ -25,6 +25,179 @@ def _reset_resolver_state(monkeypatch):
class TestCloudProviderCachePolicy:
def test_cache_is_isolated_by_hermes_home(self, tmp_path, monkeypatch):
from hermes_constants import (
get_hermes_home,
reset_hermes_home_override,
set_hermes_home_override,
)
monkeypatch.setattr(
"hermes_cli.config.read_raw_config",
lambda: {"browser": {"cloud_provider": "profile-provider"}},
)
providers = {}
resolutions = []
def resolve(_name):
home = str(get_hermes_home())
resolutions.append(home)
return providers[home]
monkeypatch.setattr(browser_tool, "_ensure_browser_plugins_loaded", lambda: None)
monkeypatch.setattr(browser_tool, "_registry_get_browser_provider", resolve)
home_a = tmp_path / "browser-a"
home_b = tmp_path / "browser-b"
providers[str(home_a)] = Mock(name="provider-a")
providers[str(home_b)] = Mock(name="provider-b")
def resolve_for(home):
token = set_hermes_home_override(home)
try:
return browser_tool._get_cloud_provider()
finally:
reset_hermes_home_override(token)
assert resolve_for(home_a) is providers[str(home_a)]
assert resolve_for(home_b) is providers[str(home_b)]
assert resolve_for(home_a) is providers[str(home_a)]
assert resolutions == [str(home_a), str(home_b)]
def test_same_profile_registry_replacement_invalidates_cache(
self, tmp_path, monkeypatch
):
from agent.browser_provider import BrowserProvider
import agent.browser_registry as browser_registry
from hermes_constants import (
reset_hermes_home_override,
set_hermes_home_override,
)
class Provider(BrowserProvider):
def __init__(self, marker):
self.marker = marker
@property
def name(self):
return "cache-replacement"
def is_available(self):
return True
def create_session(self, task_id):
return {"marker": self.marker}
def close_session(self, session_id):
return True
def emergency_cleanup(self, session_id):
return None
home = str((tmp_path / "same-profile").resolve())
first = Provider("first")
second = Provider("second")
monkeypatch.setattr(
"hermes_cli.config.read_raw_config",
lambda: {"browser": {"cloud_provider": "cache-replacement"}},
)
monkeypatch.setattr(browser_tool, "_ensure_browser_plugins_loaded", lambda: None)
token = set_hermes_home_override(home)
try:
browser_registry.register_provider(first, scope=home)
assert browser_tool._get_cloud_provider() is first
browser_registry.register_provider(second, scope=home)
assert browser_tool._get_cloud_provider() is second
finally:
current = browser_registry.snapshot_registration(
"cache-replacement", scope=home
)
if current is not None:
browser_registry.restore_registration(
"cache-replacement", current, None, scope=home
)
reset_hermes_home_override(token)
def test_concurrent_registry_replacement_discards_stale_resolution(
self, tmp_path, monkeypatch
):
from concurrent.futures import ThreadPoolExecutor
from threading import Event
from agent.browser_provider import BrowserProvider
import agent.browser_registry as browser_registry
from hermes_constants import (
reset_hermes_home_override,
set_hermes_home_override,
)
class Provider(BrowserProvider):
def __init__(self, marker):
self.marker = marker
@property
def name(self):
return "cache-race"
def is_available(self):
return True
def create_session(self, task_id):
return {"marker": self.marker}
def close_session(self, session_id):
return True
def emergency_cleanup(self, session_id):
return None
home = str((tmp_path / "race-profile").resolve())
first = Provider("first")
second = Provider("second")
paused = Event()
release = Event()
calls = 0
original_get = browser_registry.get_provider
def racing_get(name):
nonlocal calls
calls += 1
resolved = original_get(name, scope=home)
if calls == 1:
paused.set()
assert release.wait(timeout=2)
return resolved
monkeypatch.setattr(
"hermes_cli.config.read_raw_config",
lambda: {"browser": {"cloud_provider": "cache-race"}},
)
monkeypatch.setattr(browser_tool, "_ensure_browser_plugins_loaded", lambda: None)
monkeypatch.setattr(browser_tool, "_registry_get_browser_provider", racing_get)
browser_registry.register_provider(first, scope=home)
def resolve():
token = set_hermes_home_override(home)
try:
return browser_tool._get_cloud_provider()
finally:
reset_hermes_home_override(token)
try:
with ThreadPoolExecutor(max_workers=1) as pool:
future = pool.submit(resolve)
assert paused.wait(timeout=1)
browser_registry.register_provider(second, scope=home)
release.set()
assert future.result(timeout=2) is second
assert calls == 2
finally:
release.set()
current = browser_registry.snapshot_registration("cache-race", scope=home)
if current is not None:
browser_registry.restore_registration(
"cache-race", current, None, scope=home
)
def test_explicit_local_caches_permanently(self, monkeypatch):
"""`cloud_provider: local` is a positive choice and must stick."""
monkeypatch.setattr(
+56 -2
View File
@@ -69,6 +69,7 @@ from hermes_constants import (
agent_browser_runnable,
get_hermes_home,
get_hermes_home_override,
hermes_home_key,
)
from utils import env_int, is_truthy_value
from hermes_cli.config import DEFAULT_CONFIG, cfg_get
@@ -165,6 +166,16 @@ from agent.browser_provider import BrowserProvider as CloudBrowserProvider # no
from agent.browser_registry import ( # noqa: F401 (test-patchable surface)
get_provider as _registry_get_browser_provider,
)
try:
from agent.browser_registry import (
registry_generation as _browser_registry_generation,
)
except ImportError:
# A few isolated compatibility tests intentionally install a minimal
# ``agent.browser_registry`` stub exposing only ``get_provider``. Those
# harnesses have no mutable registry, so a constant generation is exact.
def _browser_registry_generation(*, scope=None):
return (0, 0)
from plugins.browser.browserbase.provider import ( # noqa: F401 (legacy import surface)
BrowserbaseBrowserProvider as BrowserbaseProvider,
)
@@ -683,6 +694,11 @@ _DEFAULT_PROVIDER_REGISTRY: Dict[str, type] = dict(_PROVIDER_REGISTRY)
_cached_cloud_provider: Optional[CloudBrowserProvider] = None
_cloud_provider_resolved = False
_cached_cloud_provider_scope: Optional[str] = None
_cached_cloud_providers: Dict[
tuple[str, tuple[int, int]], Optional[CloudBrowserProvider]
] = {}
_cloud_provider_cache_lock = threading.RLock()
_allow_private_urls_resolved = False
_cached_allow_private_urls: Optional[bool] = None
_cached_agent_browser: Optional[str] = None
@@ -738,6 +754,46 @@ def _ensure_browser_plugins_loaded() -> None:
def _get_cloud_provider() -> Optional[CloudBrowserProvider]:
"""Return the provider cached for the active Hermes profile."""
global _cached_cloud_provider, _cloud_provider_resolved
global _cached_cloud_provider_scope
scope = hermes_home_key()
with _cloud_provider_cache_lock:
# Tests and legacy reset paths clear the boolean. Treat that as a full
# reset even if a previous scoped resolution remains mirrored here.
if not _cloud_provider_resolved:
_cached_cloud_provider_scope = None
_cached_cloud_providers.clear()
while True:
before_generation = _browser_registry_generation(scope=scope)
cache_key = (scope, before_generation)
if cache_key in _cached_cloud_providers:
_cached_cloud_provider = _cached_cloud_providers[cache_key]
_cloud_provider_resolved = True
_cached_cloud_provider_scope = scope
return _cached_cloud_provider
_cached_cloud_provider = None
_cloud_provider_resolved = False
resolved = _resolve_cloud_provider_uncached()
after_generation = _browser_registry_generation(scope=scope)
if before_generation != after_generation:
# A force reload replaced/unloaded this profile's provider
# while resolution was in progress. Discard the stale result
# and resolve against the new registry generation.
continue
if _cloud_provider_resolved:
_cached_cloud_provider_scope = scope
for stale_key in [
key for key in _cached_cloud_providers if key[0] == scope
]:
_cached_cloud_providers.pop(stale_key, None)
_cached_cloud_providers[cache_key] = resolved
return resolved
def _resolve_cloud_provider_uncached() -> Optional[CloudBrowserProvider]:
"""Return the configured cloud browser provider, or None for local mode.
Reads ``config["browser"]["cloud_provider"]`` once and caches the result
@@ -755,8 +811,6 @@ def _get_cloud_provider() -> Optional[CloudBrowserProvider]:
``_is_legacy_provider_registry_overridden``.
"""
global _cached_cloud_provider, _cloud_provider_resolved
if _cloud_provider_resolved:
return _cached_cloud_provider
resolved: Optional[CloudBrowserProvider] = None
try:
+306 -51
View File
@@ -15,6 +15,7 @@ Import chain (circular-import safe):
"""
import ast
import functools
import importlib
import json
import logging
@@ -24,6 +25,8 @@ import time
from pathlib import Path
from typing import Callable, Dict, List, Optional, Set
from hermes_constants import hermes_home_key
logger = logging.getLogger(__name__)
# Cap on a tool error body; only trims runaway interpolated exceptions (static msgs are ~115 chars).
@@ -230,6 +233,15 @@ class ToolEntry:
self.dynamic_schema_overrides = dynamic_schema_overrides
class _PluginOverridePolicy:
"""Identity-bearing authorization record for one plugin generation."""
__slots__ = ("allowed",)
def __init__(self, allowed: bool) -> None:
self.allowed = bool(allowed)
# ---------------------------------------------------------------------------
# check_fn TTL cache
#
@@ -415,13 +427,20 @@ class ToolRegistry:
"""Singleton registry that collects tool schemas + handlers from tool files."""
def __init__(self):
# Built-in and other process-global registrations.
self._tools: Dict[str, ToolEntry] = {}
# Durable map: plugin module namespace (handler.__globals__["__name__"])
# -> operator opt-in for built-in override. Populated at plugin load and
# never cleared, so a plugin's override authorization is bound to the
# code that defined the handler, independent of WHEN the register() call
# happens (sync during load, or a delayed/threaded callback afterwards).
self._plugin_override_policy: Dict[str, bool] = {}
# Plugin registrations are overlays keyed by resolved HERMES_HOME. A
# profile sees its own overlay first and then the global built-ins.
self._scoped_tools: Dict[str, Dict[str, ToolEntry]] = {}
# Plugin module namespace -> operator opt-in for built-in override.
# Authorization records are lifecycle-managed; the separate scope map
# remains durable so delayed callbacks stay profile-confined.
self._plugin_override_policy: Dict[
tuple[Optional[str], str], _PluginOverridePolicy
] = {}
# Scope attribution stays durable after policy removal so delayed code
# remains confined to the profile where its module was loaded.
self._plugin_module_scopes: Dict[str, Set[Optional[str]]] = {}
self._toolset_checks: Dict[str, Callable] = {}
self._toolset_aliases: Dict[str, str] = {}
# MCP dynamic refresh can mutate the registry while other threads are
@@ -435,10 +454,30 @@ class ToolRegistry:
# long as the generation hasn't changed.
self._generation: int = 0
def _snapshot_state(self) -> tuple[List[ToolEntry], Dict[str, Callable]]:
@staticmethod
def current_scope_key() -> str:
"""Return the active profile's canonical registry scope."""
return hermes_home_key()
def _merged_tools(self, scope: Optional[str] = None) -> Dict[str, ToolEntry]:
"""Return global tools overlaid with one profile's plugin tools."""
active_scope = scope or self.current_scope_key()
merged = dict(self._tools)
merged.update(self._scoped_tools.get(active_scope, {}))
return merged
def _snapshot_state(
self,
scope: Optional[str] = None,
) -> tuple[List[ToolEntry], Dict[str, Callable]]:
"""Return a coherent snapshot of registry entries and toolset checks."""
with self._lock:
return list(self._tools.values()), dict(self._toolset_checks)
entries = list(self._merged_tools(scope).values())
checks = dict(self._toolset_checks)
for entry in entries:
if entry.check_fn is not None:
checks[entry.toolset] = entry.check_fn
return entries, checks
def _snapshot_entries(self) -> List[ToolEntry]:
"""Return a stable snapshot of registered tool entries."""
@@ -468,15 +507,35 @@ class ToolRegistry:
return True
return False
def get_entry(self, name: str) -> Optional[ToolEntry]:
"""Return a registered tool entry by name, or None."""
def get_entry(
self,
name: str,
*,
scope: Optional[str] = None,
) -> Optional[ToolEntry]:
"""Return the active profile's entry by name, falling back to global."""
with self._lock:
return self._tools.get(name)
return self._merged_tools(scope).get(name)
def snapshot_registration(
self,
name: str,
*,
scope: Optional[str] = None,
) -> Optional[ToolEntry]:
"""Return the local slot state without following global fallback."""
with self._lock:
target = self._tools if scope is None else self._scoped_tools.get(scope, {})
return target.get(name)
def get_registered_toolset_names(self) -> List[str]:
"""Return sorted unique toolset names present in the registry."""
return sorted({entry.toolset for entry in self._snapshot_entries()})
def get_all_entries(self) -> List[ToolEntry]:
"""Return the active profile's merged tool entries."""
return self._snapshot_entries()
def get_tool_names_for_toolset(self, toolset: str) -> List[str]:
"""Return sorted tool names registered under a given toolset."""
return sorted(
@@ -510,14 +569,62 @@ class ToolRegistry:
# Registration
# ------------------------------------------------------------------
def register_plugin_override_policy(self, module_namespace: str, allowed: bool) -> None:
"""Bind a plugin module namespace to its operator opt-in for built-in
override. Called once per plugin at load time. Durable: never cleared,
so later (even threaded/delayed) register() calls from that module are
still gated by the same policy.
def register_plugin_override_policy(
self,
module_namespace: str,
allowed: bool,
*,
scope: Optional[str] = None,
) -> _PluginOverridePolicy:
"""Bind a plugin module namespace to its current operator opt-in.
The identity-bearing result lets plugin unload/reload revoke a stale
authorization without losing durable module-to-profile attribution.
"""
with self._lock:
self._plugin_override_policy[module_namespace] = bool(allowed)
policy = _PluginOverridePolicy(allowed)
self._plugin_override_policy[(scope, module_namespace)] = policy
self._plugin_module_scopes.setdefault(module_namespace, set()).add(scope)
return policy
def snapshot_plugin_override_policy(
self,
module_namespace: str,
*,
scope: Optional[str] = None,
) -> Optional[_PluginOverridePolicy]:
"""Return one local authorization generation without fallback."""
with self._lock:
return self._plugin_override_policy.get((scope, module_namespace))
def restore_plugin_override_policy(
self,
module_namespace: str,
current: _PluginOverridePolicy,
previous: Optional[_PluginOverridePolicy],
*,
scope: Optional[str] = None,
) -> bool:
"""CAS-restore policy state while retaining durable scope attribution."""
with self._lock:
key = (scope, module_namespace)
if self._plugin_override_policy.get(key) is not current:
return False
if previous is None:
self._plugin_override_policy.pop(key, None)
else:
self._plugin_override_policy[key] = previous
return True
def _plugin_override_allowed(
self,
scope: Optional[str],
module_namespace: str,
) -> bool:
policy = self._plugin_override_policy.get((scope, module_namespace))
if policy is None and scope is not None:
policy = self._plugin_override_policy.get((None, module_namespace))
return bool(policy and policy.allowed)
def _plugin_owner_of(self, handler: Callable) -> Optional[str]:
"""Return the plugin module namespace that defined *handler*, or None
@@ -531,18 +638,86 @@ class ToolRegistry:
handlers live outside the plugin namespace and return None (unchanged
behavior).
"""
try:
mod = handler.__globals__.get("__name__", "") # type: ignore[attr-defined]
except AttributeError:
mod = self._callable_module(handler)
if not mod:
return None
if mod in self._plugin_override_policy:
return mod
return self._plugin_namespace_of_module(mod)
@staticmethod
def _callable_module(handler: Callable) -> str:
"""Resolve defining module through wrappers, partials, and objects."""
current = handler
seen: Set[int] = set()
while id(current) not in seen:
seen.add(id(current))
if isinstance(current, functools.partial):
current = current.func
continue
func = getattr(current, "__func__", None)
if func is not None:
current = func
continue
globals_dict = getattr(current, "__globals__", None)
if isinstance(globals_dict, dict):
module_name = globals_dict.get("__name__", "")
if module_name:
return str(module_name)
wrapped = getattr(current, "__wrapped__", None)
if wrapped is not None:
current = wrapped
continue
break
module_name = getattr(current, "__module__", "")
if module_name:
return str(module_name)
return str(getattr(type(current), "__module__", "") or "")
def _plugin_namespace_of_module(
self,
module_namespace: str,
) -> Optional[str]:
"""Resolve a module/submodule to its durable plugin namespace."""
with self._lock:
matches = [
namespace
for namespace in self._plugin_module_scopes
if module_namespace == namespace
or module_namespace.startswith(f"{namespace}.")
]
if matches:
return max(matches, key=len)
# Also gate plugin modules currently loading but not yet policy-recorded
# (defensive: a handler defined in the plugin namespace is plugin code).
if isinstance(mod, str) and mod.startswith("hermes_plugins."):
return mod
if module_namespace.startswith("hermes_plugins."):
return ".".join(module_namespace.split(".")[:2])
return None
def _plugin_scope_of(self, module_namespace: str) -> Optional[str]:
"""Return the profile scope bound to a loaded plugin module."""
with self._lock:
scopes = self._plugin_module_scopes.get(module_namespace)
if not scopes:
return None
active_scope = self.current_scope_key()
if active_scope in scopes:
return active_scope
if len(scopes) == 1:
return next(iter(scopes))
raise PermissionError(
f"Plugin module {module_namespace!r} is active in multiple "
"profiles and cannot register outside one of those scopes."
)
def plugin_scope_for_module(self, module_namespace: str) -> Optional[str]:
"""Public host lookup for a loaded plugin module's immutable scope."""
owner = self._plugin_namespace_of_module(module_namespace)
return self._plugin_scope_of(owner or module_namespace)
def plugin_scope_for_callable(self, callback: Callable) -> Optional[str]:
"""Return the durable plugin scope for any supported callable shape."""
module_name = self._callable_module(callback)
return self.plugin_scope_for_module(module_name) if module_name else None
@staticmethod
def _caller_module() -> str:
"""Best-effort module name of whoever called the registry method that
@@ -573,6 +748,7 @@ class ToolRegistry:
max_result_size_chars: int | float | None = None,
dynamic_schema_overrides: Callable = None,
override: bool = False,
scope: Optional[str] = None,
):
"""Register a tool. Called at module-import time by each tool file.
@@ -582,22 +758,58 @@ class ToolRegistry:
registrations that would shadow an existing tool from a different
toolset are rejected to prevent accidental overwrites.
"""
handler_owner = self._plugin_owner_of(handler)
caller_owner = self._plugin_namespace_of_module(self._caller_module())
owner = caller_owner or handler_owner
if scope is None and owner is not None:
scope = self._plugin_scope_of(owner)
with self._lock:
existing = self._tools.get(name)
target = (
self._tools
if scope is None
else self._scoped_tools.setdefault(scope, {})
)
existing = (
self._tools.get(name)
if scope is None
else self._merged_tools(scope).get(name)
)
shadows_global = (
owner is not None
and scope is not None
and name not in target
and name in self._tools
)
if shadows_global:
if not override:
logger.error(
"Tool registration REJECTED: plugin %r attempted to "
"shadow global tool %r without override=True",
owner,
name,
)
return
if not self._plugin_override_allowed(scope, owner):
raise PermissionError(
f"Plugin module {owner!r} cannot override built-in "
f"tool {name!r} without operator opt-in "
f"(allow_tool_override)."
)
if existing and existing.toolset != toolset:
if override:
_owner = self._plugin_owner_of(handler)
if _owner is not None and not self._plugin_override_policy.get(_owner, False):
if owner is not None and not self._plugin_override_allowed(
scope, owner
):
logger.error(
"Tool registration REJECTED: plugin %r attempted to "
"override built-in tool %r (existing toolset %r) without "
"operator opt-in. Set "
"plugins.entries.<plugin_id>.allow_tool_override: true "
"in config.yaml to allow it.",
_owner, name, existing.toolset,
owner, name, existing.toolset,
)
raise PermissionError(
f"Plugin module {_owner!r} cannot override built-in "
f"Plugin module {owner!r} cannot override built-in "
f"tool {name!r} without operator opt-in "
f"(allow_tool_override)."
)
@@ -620,7 +832,7 @@ class ToolRegistry:
name, toolset, existing.toolset,
)
return
self._tools[name] = ToolEntry(
target[name] = ToolEntry(
name=name,
toolset=toolset,
schema=schema,
@@ -639,7 +851,7 @@ class ToolRegistry:
# banner.py reads (presence only, never called) to classify an
# already-unavailable toolset as lazy-init vs disabled. Keep the
# write path for that classification.
if check_fn and toolset not in self._toolset_checks:
if scope is None and check_fn and toolset not in self._toolset_checks:
self._toolset_checks[toolset] = check_fn
self._generation += 1
@@ -660,11 +872,30 @@ class ToolRegistry:
every refresh and has no plugin-override concept.
"""
with self._lock:
entry = self._tools.get(name)
caller_mod = self._caller_module()
caller_owner = self._plugin_namespace_of_module(caller_mod)
caller_scope = (
self._plugin_scope_of(caller_owner)
if caller_owner is not None
else None
)
target = (
self._scoped_tools.get(caller_scope, {})
if caller_scope is not None
else self._tools
)
entry = target.get(name)
if entry is None and caller_scope is not None:
if name in self._tools:
raise PermissionError(
f"Scoped plugin module {caller_mod!r} cannot deregister "
f"process-global tool {name!r}; register a scoped "
"override instead."
)
return
if entry is None:
return
if not entry.toolset.startswith("mcp-"):
caller_mod = self._caller_module()
owner = self._plugin_owner_of(entry.handler)
# Ownership check: bind to the plugin package root
# (``hermes_plugins.{name}``), not the exact module string.
@@ -673,13 +904,13 @@ class ToolRegistry:
# string equality would wrongly block root-module cleanup code
# from removing tools registered by a submodule of the same
# plugin (egilewski review on #55840).
caller_root = ".".join(caller_mod.split(".")[:2])
owner_root = ".".join(owner.split(".")[:2]) if owner else ""
same_plugin = bool(owner and caller_root == owner_root)
same_plugin = bool(owner and caller_owner == owner)
if (
caller_mod.startswith("hermes_plugins.")
caller_owner is not None
and not same_plugin
and not self._plugin_override_policy.get(caller_root, False)
and not self._plugin_override_allowed(
caller_scope, caller_owner
)
):
logger.error(
"Tool deregistration REJECTED: plugin %r attempted to "
@@ -694,11 +925,14 @@ class ToolRegistry:
f"{name!r} (toolset {entry.toolset!r}) without operator "
f"opt-in (allow_tool_override)."
)
del self._tools[name]
del target[name]
if caller_scope is not None and not target:
self._scoped_tools.pop(caller_scope, None)
# Drop the toolset check and aliases if this was the last tool in
# that toolset.
toolset_still_exists = any(
e.toolset == entry.toolset for e in self._tools.values()
e.toolset == entry.toolset
for e in self._merged_tools(caller_scope).values()
)
if not toolset_still_exists:
self._toolset_checks.pop(entry.toolset, None)
@@ -715,6 +949,8 @@ class ToolRegistry:
name: str,
current: ToolEntry,
previous: Optional[ToolEntry],
*,
scope: Optional[str] = None,
) -> bool:
"""Restore a host-owned registration if it is still current.
@@ -725,13 +961,20 @@ class ToolRegistry:
must leave the newer entry untouched.
"""
with self._lock:
if self._tools.get(name) is not current:
target = (
self._tools
if scope is None
else self._scoped_tools.setdefault(scope, {})
)
if target.get(name) is not current:
return False
if previous is None:
self._tools.pop(name, None)
target.pop(name, None)
else:
self._tools[name] = previous
target[name] = previous
if scope is not None and not target:
self._scoped_tools.pop(scope, None)
# Rebuild the affected toolset checks from the surviving entries.
# A plugin may have replaced an entry in the same toolset, so
@@ -742,18 +985,23 @@ class ToolRegistry:
affected_toolsets.add(previous.toolset)
for toolset in affected_toolsets:
surviving = [
entry for entry in self._tools.values()
entry for entry in self._merged_tools(scope).values()
if entry.toolset == toolset
]
check_fn = next(
(entry.check_fn for entry in surviving if entry.check_fn),
None,
)
if check_fn is None:
self._toolset_checks.pop(toolset, None)
else:
self._toolset_checks[toolset] = check_fn
if not surviving:
if scope is None:
if check_fn is None:
self._toolset_checks.pop(toolset, None)
else:
self._toolset_checks[toolset] = check_fn
if not surviving and not any(
entry.toolset == toolset
for entries in self._scoped_tools.values()
for entry in entries.values()
):
self._toolset_aliases = {
alias: target
for alias, target in self._toolset_aliases.items()
@@ -851,7 +1099,14 @@ class ToolRegistry:
result_type=result_type,
)
def dispatch(self, name: str, args: dict, **kwargs) -> str | dict:
def dispatch(
self,
name: str,
args: dict,
*,
scope: Optional[str] = None,
**kwargs,
) -> str | dict:
"""Execute a tool handler by name.
* Async handlers are bridged automatically via ``_run_async()``.
@@ -860,7 +1115,7 @@ class ToolRegistry:
* All exceptions are caught and returned as ``{"error": "..."}``
for consistent error format.
"""
entry = self.get_entry(name)
entry = self.get_entry(name, scope=scope)
if not entry:
return tool_error(f"Unknown tool: {name}")
try:
+1 -1
View File
@@ -810,7 +810,7 @@ def resolve_toolset(name: str, visited: Set[str] = None, *, include_registry: bo
try:
from tools.registry import registry
plugin_tools.update(
e.name for e in registry._tools.values()
e.name for e in registry.get_all_entries()
if e.toolset == platform_name
)
except Exception: