fix(plugins): isolate ownership by profile
This commit is contained in:
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
+26
-17
@@ -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(
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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()
|
||||
@@ -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
@@ -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"
|
||||
)
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user