From 6d6b244a90ff1891e78938c901fbf01738dc3e7d Mon Sep 17 00:00:00 2001 From: Eugeniusz Gilewski Date: Thu, 23 Jul 2026 11:26:42 +0200 Subject: [PATCH] fix(security): enforce disabled toolsets after dynamic injection (#49386) Static registry filtering happened before memory-provider tools were injected or refreshed, so a disabled memory toolset could reappear in published schemas, dispatch ownership, and prompt guidance. Several entrypoints also failed to carry the disabled policy into late tool discovery. Normalize toolset policy at each entrypoint and apply final subtraction to provider schemas, ownership, prompts, API discovery, and reload publication. Publish policy and tool snapshots atomically so concurrent or failed refreshes cannot restore stale provider capabilities, while preserving enabled provider tools, passive provider context, and generic dispatch for name collisions. Focused owner-file validation passes 392 tests. The final full-suite run reached 36,123 passes before host /tmp exhaustion made later file-fixture results nondiagnostic; two earlier unrelated readiness assertions also failed. Ruff, byte-compilation, and diff checks pass. CodeRabbit's prior cycle was clean; the requested final cycle remained rate-limited after its retry. Fixes #49386 Related #46171 --- acp_adapter/server.py | 8 +- acp_adapter/session.py | 9 + agent/agent_init.py | 4 + agent/agent_runtime_helpers.py | 13 +- agent/memory_manager.py | 131 +++- agent/system_prompt.py | 25 +- agent/tool_executor.py | 3 +- cli.py | 5 +- cron/scheduler.py | 5 +- gateway/platforms/api_server.py | 38 +- gateway/run.py | 9 +- hermes_cli/cli_agent_setup_mixin.py | 6 +- hermes_cli/cli_commands_mixin.py | 4 + hermes_cli/oneshot.py | 5 + hermes_cli/tools_config.py | 5 +- hermes_cli/toolset_validation.py | 18 +- tests/acp/test_server.py | 20 + tests/acp/test_session.py | 73 ++ tests/agent/test_memory_provider.py | 184 ++++- tests/agent/test_system_prompt.py | 128 ++++ tests/cli/test_cli_approval_ui.py | 116 ++- tests/cron/test_scheduler.py | 21 + tests/gateway/test_api_server.py | 71 ++ tests/gateway/test_api_server_toolset.py | 45 ++ .../test_session_model_override_routing.py | 137 +++- .../hermes_cli/test_oneshot_toolset_policy.py | 53 ++ tests/hermes_cli/test_tools_config.py | 13 + tests/run_agent/test_run_agent.py | 12 +- tests/tools/test_refresh_agent_mcp_tools.py | 690 +++++++++++++++++- tests/tui_gateway/test_make_agent_provider.py | 69 ++ tests/tui_gateway/test_mcp_reload_rev.py | 165 +++++ tools/mcp_tool.py | 316 ++++++-- tui_gateway/server.py | 88 ++- 33 files changed, 2365 insertions(+), 124 deletions(-) create mode 100644 tests/hermes_cli/test_oneshot_toolset_policy.py diff --git a/acp_adapter/server.py b/acp_adapter/server.py index 7b10f2c23a247..df81f099414f6 100644 --- a/acp_adapter/server.py +++ b/acp_adapter/server.py @@ -2092,7 +2092,12 @@ def _cmd_tools(self, args: str, state: SessionState) -> str: toolsets = _expand_acp_enabled_toolsets( getattr(state.agent, "enabled_toolsets", None) or ["hermes-acp"] ) - tools = get_tool_definitions(enabled_toolsets=toolsets, quiet_mode=True) + disabled_toolsets = getattr(state.agent, "disabled_toolsets", None) + tools = get_tool_definitions( + enabled_toolsets=toolsets, + disabled_toolsets=disabled_toolsets, + quiet_mode=True, + ) tool_view = SimpleNamespace( tools=list(tools or []), valid_tool_names={ @@ -2101,6 +2106,7 @@ def _cmd_tools(self, args: str, state: SessionState) -> str: if isinstance(tool, dict) }, enabled_toolsets=toolsets, + disabled_toolsets=disabled_toolsets, _memory_manager=getattr(state.agent, "_memory_manager", None), ) inject_memory_provider_tools(tool_view) diff --git a/acp_adapter/session.py b/acp_adapter/session.py index 6f1e17a07f57a..a1e53b6ea39a3 100644 --- a/acp_adapter/session.py +++ b/acp_adapter/session.py @@ -9,6 +9,7 @@ from __future__ import annotations from hermes_constants import get_hermes_home +from hermes_cli.toolset_validation import normalize_toolset_names import copy import json @@ -620,12 +621,20 @@ def _make_agent( if not isinstance(cfg, dict) or cfg.get("enabled", True) is not False ] + agent_cfg = config.get("agent") + disabled_toolsets = normalize_toolset_names( + agent_cfg.get("disabled_toolsets") + if isinstance(agent_cfg, dict) + else None + ) + kwargs = { "platform": "acp", "enabled_toolsets": _expand_acp_enabled_toolsets( ["hermes-acp"], mcp_server_names=configured_mcp_servers, ), + "disabled_toolsets": disabled_toolsets, "quiet_mode": True, "session_id": session_id, "session_db": self._get_db(), diff --git a/agent/agent_init.py b/agent/agent_init.py index 955ca11cff8d3..e4c378c7a8967 100644 --- a/agent/agent_init.py +++ b/agent/agent_init.py @@ -870,6 +870,10 @@ def init_agent( # late/concurrent refresh reject a stale (older-generation) rebuild instead # of clobbering a newer one. Set adjacent to the tool snapshot below. agent._tool_snapshot_generation = 0 + # Policy publication has its own ordering because registry generation does + # not change when an explicit reload alters only enabled/disabled toolsets. + agent._tool_policy_epoch = 0 + agent._tool_published_policy_epoch = 0 # Rate limit tracking — updated from x-ratelimit-* response headers # after each API call. Accessed by /usage slash command. agent._rate_limit_state: Optional["RateLimitState"] = None diff --git a/agent/agent_runtime_helpers.py b/agent/agent_runtime_helpers.py index 6458ad88fbebc..58f15889ff844 100644 --- a/agent/agent_runtime_helpers.py +++ b/agent/agent_runtime_helpers.py @@ -63,14 +63,21 @@ def _ra(): ) +def memory_provider_owns_tool(agent: Any, function_name: str) -> bool: + """Return whether the published agent surface assigns provider routing.""" + owned_names = getattr(agent, "_memory_provider_tool_names", None) + return isinstance(owned_names, set) and function_name in owned_names + + def agent_runtime_owns_post_tool_hook(agent: Any, function_name: str) -> bool: """Return True when an agent-level tool path emits its own post hook.""" if function_name in AGENT_RUNTIME_POST_HOOK_TOOL_NAMES: return True if getattr(agent, "_context_engine_tool_names", None) and function_name in agent._context_engine_tool_names: return True - memory_manager = getattr(agent, "_memory_manager", None) - return bool(memory_manager and memory_manager.has_tool(function_name)) + return bool(getattr(agent, "_memory_manager", None)) and memory_provider_owns_tool( + agent, function_name + ) def convert_to_trajectory_format(agent, messages: List[Dict[str, Any]], user_query: str, completed: bool) -> List[Dict[str, Any]]: @@ -2604,7 +2611,7 @@ def _execute(next_args: dict) -> Any: ), ) return _finish_agent_tool(result, next_args) - elif agent._memory_manager and agent._memory_manager.has_tool(function_name): + elif agent._memory_manager and memory_provider_owns_tool(agent, function_name): def _execute(next_args: dict) -> Any: return _finish_agent_tool(agent._memory_manager.handle_tool_call(function_name, next_args), next_args) elif function_name == "clarify": diff --git a/agent/memory_manager.py b/agent/memory_manager.py index 6f3bbadd6f059..1aab1f98022d6 100644 --- a/agent/memory_manager.py +++ b/agent/memory_manager.py @@ -107,25 +107,91 @@ def memory_provider_tools_enabled( return False +def memory_provider_denied_tool_names( + disabled_toolsets: Optional[List[str]], +) -> Optional[set[str]]: + """Resolve provider-tool subtraction. + + Returns the exact denied tool names, or ``None`` when all provider tools + must be denied. Unknown toolsets are ignored, matching the registry tool + filtering path. Resolver failures fail closed. + """ + if not disabled_toolsets: + return set() + if any(name in {"memory", "all", "*"} for name in disabled_toolsets): + return None + + try: + from toolsets import bundle_non_core_tools, resolve_toolset, validate_toolset + + denied: set[str] = set() + for name in disabled_toolsets: + if not validate_toolset(name): + continue + resolved = set( + bundle_non_core_tools(name) + if name.startswith("hermes-") + else resolve_toolset(name) + ) + if "memory" in resolved: + return None + denied.update(resolved) + except Exception: + logger.debug("Failed to resolve disabled toolsets for memory-provider tools", exc_info=True) + return None + return denied + + +def memory_provider_tools_disabled(disabled_toolsets: Optional[List[str]]) -> bool: + """Return whether subtraction denies the complete provider tool family.""" + return memory_provider_denied_tool_names(disabled_toolsets) is None + + +def effective_memory_provider_tool_schemas( + raw_schemas, + *, + enabled_toolsets: Optional[List[str]], + disabled_toolsets: Optional[List[str]], + memory_selected: bool = False, +) -> List[Dict[str, Any]]: + """Return normalized provider schemas surviving selection and subtraction.""" + if not memory_selected and not memory_provider_tools_enabled(enabled_toolsets): + return [] + + denied_names = memory_provider_denied_tool_names(disabled_toolsets) + if denied_names is None: + return [] + + effective = [] + for raw_schema in raw_schemas: + schema = normalize_tool_schema(raw_schema) + if schema is None: + logger.warning( + "Memory provider returned a tool schema with no resolvable " + "name; skipping to avoid poisoning the request (%r)", + raw_schema, + ) + continue + if schema["name"] not in denied_names: + effective.append(schema) + return effective + + def inject_memory_provider_tools(agent: Any) -> int: """Append external memory-provider tool schemas to an agent tool surface.""" + # This is routing authorization, not merely diagnostic metadata. Only + # schemas this injector actually appends are owned by the provider; a + # registry/MCP/plugin collision keeps generic dispatch precedence. + agent._memory_provider_tool_names = set() memory_manager = getattr(agent, "_memory_manager", None) tools = getattr(agent, "tools", None) if not memory_manager or tools is None: return 0 - existing_tool_names = { tool.get("function", {}).get("name") for tool in tools if isinstance(tool, dict) } - if not memory_provider_tools_enabled( - getattr(agent, "enabled_toolsets", None), - getattr(agent, "disabled_toolsets", None), - memory_tool_present="memory" in existing_tool_names, - ): - return 0 - get_schemas = getattr(memory_manager, "get_all_tool_schemas", None) if not callable(get_schemas): return 0 @@ -136,21 +202,20 @@ def inject_memory_provider_tools(agent: Any) -> int: agent.valid_tool_names = valid_tool_names added = 0 - for raw_schema in get_schemas(): - schema = normalize_tool_schema(raw_schema) - if schema is None: - logger.warning( - "Memory provider returned a tool schema with no resolvable " - "name; skipping to avoid poisoning the request (%r)", - raw_schema, - ) - continue + schemas = effective_memory_provider_tool_schemas( + get_schemas(), + enabled_toolsets=getattr(agent, "enabled_toolsets", None), + disabled_toolsets=getattr(agent, "disabled_toolsets", None), + memory_selected="memory" in existing_tool_names, + ) + for schema in schemas: tool_name = schema["name"] if tool_name in existing_tool_names: continue tools.append({"type": "function", "function": schema}) valid_tool_names.add(tool_name) existing_tool_names.add(tool_name) + agent._memory_provider_tool_names.add(tool_name) added += 1 return added @@ -483,15 +548,33 @@ def get_provider(self, name: str) -> Optional[MemoryProvider]: # -- System prompt ------------------------------------------------------- - def build_system_prompt(self) -> str: + def build_system_prompt( + self, *, available_tool_names: Optional[set[str]] = None + ) -> str: """Collect system prompt blocks from all providers. Returns combined text, or empty string if no providers contribute. Each non-empty block is labeled with the provider name. + + When ``available_tool_names`` is provided, providers whose schemas are + all absent are omitted so their tool-use instructions cannot advertise + unavailable capabilities. Providers without schemas remain visible + because their prompt describes passive context injection. """ blocks = [] for provider in self._providers: try: + provider_schemas = provider.get_tool_schemas() + if available_tool_names is not None and provider_schemas: + schema_names = { + schema["name"] + for raw_schema in provider_schemas + if (schema := normalize_tool_schema(raw_schema)) is not None + } + # Provider blocks are opaque: retaining one when even one + # callable is unavailable can advertise the denied tool. + if not schema_names.issubset(available_tool_names): + continue block = provider.system_prompt_block() if block and block.strip(): blocks.append(block) @@ -781,7 +864,7 @@ def flush_pending(self, timeout: Optional[float] = None) -> bool: # -- Tools --------------------------------------------------------------- - def get_all_tool_schemas(self) -> List[Dict[str, Any]]: + def _collect_all_tool_schemas(self, *, strict: bool) -> List[Dict[str, Any]]: """Collect tool schemas from all providers. Reserved core tool names (``clarify``, ``delegate_task``, etc.) are @@ -812,12 +895,22 @@ def get_all_tool_schemas(self) -> List[Dict[str, Any]]: schemas.append(schema) seen.add(name) except Exception as e: + if strict: + raise logger.warning( "Memory provider '%s' get_tool_schemas() failed: %s", provider.name, e, ) return schemas + def get_all_tool_schemas(self) -> List[Dict[str, Any]]: + """Collect provider schemas fail-soft for optional runtime callers.""" + return self._collect_all_tool_schemas(strict=False) + + def get_all_tool_schemas_strict(self) -> List[Dict[str, Any]]: + """Collect every provider schema or fail without a partial result.""" + return self._collect_all_tool_schemas(strict=True) + def get_all_tool_names(self) -> set: """Return set of all tool names across all providers.""" return set(self._tool_to_provider.keys()) diff --git a/agent/system_prompt.py b/agent/system_prompt.py index 6cc7554b98661..f64017d7c6f59 100644 --- a/agent/system_prompt.py +++ b/agent/system_prompt.py @@ -511,7 +511,30 @@ def build_system_prompt_parts(agent: Any, system_message: Optional[str] = None) # External memory provider system prompt block (additive to built-in) if agent._memory_manager: try: - _ext_mem_block = agent._memory_manager.build_system_prompt() + from agent.memory_manager import ( + effective_memory_provider_tool_schemas, + ) + + _provider_schemas = agent._memory_manager.get_all_tool_schemas() + _effective_provider_schemas = effective_memory_provider_tool_schemas( + _provider_schemas, + enabled_toolsets=getattr(agent, "enabled_toolsets", None), + disabled_toolsets=getattr(agent, "disabled_toolsets", None), + memory_selected="memory" in getattr(agent, "valid_tool_names", set()), + ) + _effective_provider_names = { + schema["name"] for schema in _effective_provider_schemas + } + _published_provider_names = getattr( + agent, "_memory_provider_tool_names", None + ) + if isinstance(_published_provider_names, set): + _effective_provider_names.intersection_update( + _published_provider_names + ) + _ext_mem_block = agent._memory_manager.build_system_prompt( + available_tool_names=_effective_provider_names + ) if _ext_mem_block: volatile_parts.append(_ext_mem_block) except Exception: diff --git a/agent/tool_executor.py b/agent/tool_executor.py index a5729a2e8f503..00303432a700b 100644 --- a/agent/tool_executor.py +++ b/agent/tool_executor.py @@ -32,6 +32,7 @@ _detect_tool_failure, ) from agent.tool_guardrails import ToolGuardrailDecision +from agent.agent_runtime_helpers import memory_provider_owns_tool from agent.tool_dispatch_helpers import ( _is_destructive_command, _is_multimodal_tool_result, @@ -1463,7 +1464,7 @@ def _execute(next_args: dict) -> Any: spinner.stop(cute_msg) elif agent._should_emit_quiet_tool_messages(): agent._vprint(f" {cute_msg}") - elif agent._memory_manager and agent._memory_manager.has_tool(function_name): + elif agent._memory_manager and memory_provider_owns_tool(agent, function_name): # Memory provider tools (hindsight_retain, honcho_search, etc.) # These are not in the tool registry — route through MemoryManager. spinner = None diff --git a/cli.py b/cli.py index 7df68530ae293..276212079b468 100644 --- a/cli.py +++ b/cli.py @@ -53,6 +53,7 @@ import yaml from hermes_cli.fallback_config import get_fallback_chain +from hermes_cli.toolset_validation import normalize_toolset_names from hermes_cli.cli_agent_setup_mixin import CLIAgentSetupMixin from hermes_cli.cli_commands_mixin import CLICommandsMixin from hermes_cli.cli_billing_mixin import CLIBillingMixin @@ -4118,7 +4119,9 @@ def __init__( # Parse and validate toolsets self.enabled_toolsets = toolsets - self.disabled_toolsets = CLI_CONFIG["agent"].get("disabled_toolsets") or [] + self.disabled_toolsets = normalize_toolset_names( + CLI_CONFIG["agent"].get("disabled_toolsets") + ) or [] if toolsets and "all" not in toolsets and "*" not in toolsets: # Validate each toolset — MCP server names are resolved via diff --git a/cron/scheduler.py b/cron/scheduler.py index 5bffb00e14612..82f064d6fa70d 100644 --- a/cron/scheduler.py +++ b/cron/scheduler.py @@ -43,6 +43,7 @@ from hermes_cli._subprocess_compat import windows_hide_flags from hermes_cli.config import load_config, _expand_env_vars from hermes_cli.fallback_config import get_fallback_chain +from hermes_cli.toolset_validation import normalize_toolset_names from hermes_time import now as _hermes_now logger = logging.getLogger(__name__) @@ -168,7 +169,9 @@ def _resolve_cron_disabled_toolsets(cfg: dict) -> list[str]: """ disabled = ["cronjob", "messaging", "clarify"] agent_cfg = (cfg or {}).get("agent") or {} - user_disabled = agent_cfg.get("disabled_toolsets") or [] + user_disabled = normalize_toolset_names( + agent_cfg.get("disabled_toolsets") + ) or [] for name in user_disabled: name = str(name).strip() if name and name not in disabled: diff --git a/gateway/platforms/api_server.py b/gateway/platforms/api_server.py index a1b066a85b4a5..867fee7180c36 100644 --- a/gateway/platforms/api_server.py +++ b/gateway/platforms/api_server.py @@ -2352,6 +2352,7 @@ def _create_agent( GatewayRunner, ) from hermes_cli.tools_config import _get_platform_tools + from hermes_cli.toolset_validation import normalize_toolset_names # Catch RuntimeError ONLY around this call, not the wider # _create_agent()+run_conversation() span -- @@ -2571,6 +2572,10 @@ def _resolve_provider_runtime( user_config = _load_gateway_config() enabled_toolsets = sorted(_get_platform_tools(user_config, "api_server")) + agent_cfg = user_config.get("agent") or {} + disabled_toolsets = normalize_toolset_names( + agent_cfg.get("disabled_toolsets") + ) max_iterations = _current_max_iterations() @@ -2591,6 +2596,7 @@ def _resolve_provider_runtime( "verbose_logging": False, "ephemeral_system_prompt": ephemeral_system_prompt or None, "enabled_toolsets": enabled_toolsets, + "disabled_toolsets": disabled_toolsets, "session_id": session_id, "platform": "api_server", "stream_delta_callback": stream_delta_callback, @@ -2905,12 +2911,18 @@ async def _handle_toolsets(self, request: "web.Request") -> "web.Response": try: from hermes_cli.config import load_config + from hermes_cli.toolset_validation import normalize_toolset_names from hermes_cli.tools_config import ( _get_effective_configurable_toolsets, _get_platform_tools, _toolset_has_keys, ) - from toolsets import resolve_toolset + from toolsets import ( + bundle_non_core_tools, + get_toolset, + resolve_toolset, + validate_toolset, + ) config = load_config() enabled_toolsets = _get_platform_tools( @@ -2918,13 +2930,35 @@ async def _handle_toolsets(self, request: "web.Request") -> "web.Response": "api_server", include_default_mcp_servers=False, ) + agent_cfg = config.get("agent") or {} + disabled_toolsets = [ + str(name) + for name in normalize_toolset_names( + agent_cfg.get("disabled_toolsets") + ) or [] + ] + disable_all = any(name in {"all", "*"} for name in disabled_toolsets) + disabled_tools: set[str] = set() + for disabled_name in disabled_toolsets: + if not validate_toolset(disabled_name): + continue + disabled_def = get_toolset(disabled_name) or {} + if disabled_name.startswith("hermes-") or disabled_def.get("posture"): + disabled_tools.update(bundle_non_core_tools(disabled_name)) + else: + disabled_tools.update(resolve_toolset(disabled_name)) + data: List[Dict[str, Any]] = [] for name, label, desc in _get_effective_configurable_toolsets(): try: tools = sorted(set(resolve_toolset(name))) except Exception: tools = [] - is_enabled = name in enabled_toolsets + is_enabled = ( + name in enabled_toolsets + and not disable_all + and (not tools or bool(set(tools) - disabled_tools)) + ) data.append({ "name": name, "label": label, diff --git a/gateway/run.py b/gateway/run.py index 22d7c58e9e113..68b3daf9c1390 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -69,6 +69,7 @@ from agent.i18n import t from hermes_cli.config import cfg_get from hermes_cli.fallback_config import get_fallback_chain +from hermes_cli.toolset_validation import normalize_toolset_names # --- Agent cache tuning --------------------------------------------------- # Bounds the per-session AIAgent cache to prevent unbounded growth in @@ -15731,7 +15732,9 @@ async def _run_background_task_inner( from hermes_cli.tools_config import _get_platform_tools enabled_toolsets = sorted(_get_platform_tools(user_config, platform_key)) agent_cfg = user_config.get("agent") or {} - disabled_toolsets = agent_cfg.get("disabled_toolsets") or None + disabled_toolsets = normalize_toolset_names( + agent_cfg.get("disabled_toolsets") + ) pr = self._provider_routing max_iterations = _current_max_iterations() @@ -20151,7 +20154,9 @@ def _run_still_current() -> bool: from hermes_cli.tools_config import _get_platform_tools enabled_toolsets = sorted(_get_platform_tools(user_config, platform_key)) agent_cfg_local = user_config.get("agent") or {} - disabled_toolsets = agent_cfg_local.get("disabled_toolsets") or None + disabled_toolsets = normalize_toolset_names( + agent_cfg_local.get("disabled_toolsets") + ) display_config = user_config.get("display", {}) if not isinstance(display_config, dict): diff --git a/hermes_cli/cli_agent_setup_mixin.py b/hermes_cli/cli_agent_setup_mixin.py index e7ef885240bd8..406236ed9fd80 100644 --- a/hermes_cli/cli_agent_setup_mixin.py +++ b/hermes_cli/cli_agent_setup_mixin.py @@ -18,6 +18,8 @@ from rich.markup import escape as _escape +from hermes_cli.toolset_validation import normalize_toolset_names + class CLIAgentSetupMixin: """Agent construction + session-resume display methods for ``HermesCLI``.""" @@ -370,7 +372,9 @@ def _init_agent(self, *, model_override: str = None, runtime_override: dict = No max_tokens=self.max_tokens, max_iterations=self.max_turns, enabled_toolsets=self.enabled_toolsets, - disabled_toolsets=self.disabled_toolsets, + disabled_toolsets=( + normalize_toolset_names(self.disabled_toolsets) or [] + ), verbose_logging=self.verbose, quiet_mode=not self.verbose, tool_progress_mode=getattr(self, "tool_progress_mode", "all"), diff --git a/hermes_cli/cli_commands_mixin.py b/hermes_cli/cli_commands_mixin.py index e447c2cac6be3..3bb600ee0297c 100644 --- a/hermes_cli/cli_commands_mixin.py +++ b/hermes_cli/cli_commands_mixin.py @@ -38,6 +38,7 @@ local_port_in_use, manual_chrome_debug_command, ) +from hermes_cli.toolset_validation import normalize_toolset_names class CLICommandsMixin: @@ -1696,6 +1697,9 @@ def run_background(): max_tokens=turn_route["runtime"].get("max_tokens"), max_iterations=self.max_turns, enabled_toolsets=self.enabled_toolsets, + disabled_toolsets=( + normalize_toolset_names(self.disabled_toolsets) or [] + ), quiet_mode=True, verbose_logging=False, session_id=task_id, diff --git a/hermes_cli/oneshot.py b/hermes_cli/oneshot.py index 320c61c8cc12e..43a53650af982 100644 --- a/hermes_cli/oneshot.py +++ b/hermes_cli/oneshot.py @@ -30,6 +30,7 @@ from gateway.session_context import declare_stateless_channel from hermes_cli.fallback_config import get_fallback_chain +from hermes_cli.toolset_validation import normalize_toolset_names def _normalize_toolsets(toolsets: object = None) -> list[str] | None: @@ -328,6 +329,7 @@ def _run_agent( from run_agent import AIAgent cfg = load_config() + agent_cfg = cfg.get("agent") or {} # Resolve effective model: explicit arg → env var → config. model_cfg = cfg.get("model") or {} @@ -415,6 +417,9 @@ def _run_agent( api_mode=runtime.get("api_mode"), model=effective_model, enabled_toolsets=toolsets_list, + disabled_toolsets=normalize_toolset_names( + agent_cfg.get("disabled_toolsets") + ), quiet_mode=True, platform="cli", session_db=session_db, diff --git a/hermes_cli/tools_config.py b/hermes_cli/tools_config.py index 2614ead045ecb..4139448f4e973 100644 --- a/hermes_cli/tools_config.py +++ b/hermes_cli/tools_config.py @@ -30,6 +30,7 @@ get_nous_subscription_features, ) from hermes_cli.nous_account import format_nous_portal_entitlement_message +from hermes_cli.toolset_validation import normalize_toolset_names from tools.tool_backend_helpers import fal_key_is_configured from utils import base_url_hostname, is_truthy_value @@ -1990,7 +1991,9 @@ def _get_platform_tools( # platforms without per-platform toolset configuration. This runs # last so it overrides everything above. agent_cfg = config.get("agent") or {} - disabled_toolsets = agent_cfg.get("disabled_toolsets") or [] + disabled_toolsets = normalize_toolset_names( + agent_cfg.get("disabled_toolsets") + ) or [] if disabled_toolsets: disabled_set = {str(ts) for ts in disabled_toolsets} enabled_toolsets -= disabled_set diff --git a/hermes_cli/toolset_validation.py b/hermes_cli/toolset_validation.py index cce8140791232..2121fcee7266c 100644 --- a/hermes_cli/toolset_validation.py +++ b/hermes_cli/toolset_validation.py @@ -12,7 +12,23 @@ zero-tools end state) loudly turns that silent failure into an actionable one. """ -from typing import Callable, Dict, List +from typing import Callable, Dict, List, Optional + + +def normalize_toolset_names( + value: Optional[List[str] | str], +) -> Optional[List[str]]: + """Normalize a scalar config toolset name without changing list semantics. + + YAML accepts both ``disabled_toolsets: memory`` and the documented list + form. Treating the scalar as an iterable later turns it into character + names, so normalize only that ambiguous shape at config-consumption + boundaries. Existing lists pass through unchanged; falsey values retain + the callers' established ``None`` semantics. + """ + if isinstance(value, str): + return [value] if value else None + return value or None def validate_platform_toolsets( diff --git a/tests/acp/test_server.py b/tests/acp/test_server.py index 96101b93ef93d..80268d6c24e8c 100644 --- a/tests/acp/test_server.py +++ b/tests/acp/test_server.py @@ -1867,6 +1867,26 @@ def test_context_says_compression_due_when_past_threshold(self, agent, mock_mana assert "Context usage: ~82,000 / 100,000 tokens (82.0%)" in result assert "Compression: due now (threshold ~80,000, 80%). Run /compress." in result + def test_tools_respects_disabled_memory_toolset(self, agent, mock_manager): + state = self._make_state(mock_manager) + state.agent.enabled_toolsets = ["hermes-acp"] + state.agent.disabled_toolsets = ["memory"] + state.agent._memory_manager = SimpleNamespace( + get_all_tool_schemas=lambda: [ + {"name": "fact_store", "description": "Store", "parameters": {}} + ] + ) + + with patch("model_tools.get_tool_definitions", return_value=[]) as mock_defs: + result = agent._cmd_tools("", state) + + mock_defs.assert_called_once_with( + enabled_toolsets=["hermes-acp"], + disabled_toolsets=["memory"], + quiet_mode=True, + ) + assert result == "No tools available." + def test_reset_clears_history(self, agent, mock_manager): state = self._make_state(mock_manager) state.history = [{"role": "user", "content": "hello"}] diff --git a/tests/acp/test_session.py b/tests/acp/test_session.py index 199454b39dbab..b5051ad3be9a6 100644 --- a/tests/acp/test_session.py +++ b/tests/acp/test_session.py @@ -440,6 +440,79 @@ def fake_agent(**kwargs): assert captured["enabled_toolsets"] == ["hermes-acp", "mcp-olympus", "mcp-exa"] + def test_create_session_propagates_global_disabled_toolsets(self, tmp_path, monkeypatch): + captured = {} + + monkeypatch.setattr("hermes_cli.config.load_config", lambda: { + "model": {"provider": "openrouter", "default": "test-model"}, + "agent": {"disabled_toolsets": ["memory", "terminal"]}, + "mcp_servers": {}, + }) + monkeypatch.setattr( + "hermes_cli.runtime_provider.resolve_runtime_provider", + lambda requested=None: {"provider": requested}, + ) + + with patch( + "run_agent.AIAgent", + side_effect=lambda **kwargs: ( + captured.update(kwargs) + or SimpleNamespace(model=kwargs.get("model")) + ), + ): + SessionManager(db=SessionDB(tmp_path / "state.db")).create_session(cwd="/work") + + assert captured["platform"] == "acp" + assert captured["enabled_toolsets"] == ["hermes-acp"] + assert captured["disabled_toolsets"] == ["memory", "terminal"] + + def test_create_session_normalizes_bare_disabled_toolset(self, tmp_path, monkeypatch): + captured = {} + + monkeypatch.setattr("hermes_cli.config.load_config", lambda: { + "model": {"provider": "openrouter", "default": "test-model"}, + "agent": {"disabled_toolsets": "memory"}, + "mcp_servers": {}, + }) + monkeypatch.setattr( + "hermes_cli.runtime_provider.resolve_runtime_provider", + lambda requested=None: {"provider": requested}, + ) + + with patch( + "run_agent.AIAgent", + side_effect=lambda **kwargs: ( + captured.update(kwargs) + or SimpleNamespace(model=kwargs.get("model")) + ), + ): + SessionManager(db=SessionDB(tmp_path / "state.db")).create_session(cwd="/work") + + assert captured["disabled_toolsets"] == ["memory"] + + def test_create_session_defaults_missing_disabled_toolsets_to_none(self, tmp_path, monkeypatch): + captured = {} + + monkeypatch.setattr("hermes_cli.config.load_config", lambda: { + "model": {"provider": "openrouter", "default": "test-model"}, + "mcp_servers": {}, + }) + monkeypatch.setattr( + "hermes_cli.runtime_provider.resolve_runtime_provider", + lambda requested=None: {"provider": requested}, + ) + + with patch( + "run_agent.AIAgent", + side_effect=lambda **kwargs: ( + captured.update(kwargs) + or SimpleNamespace(model=kwargs.get("model")) + ), + ): + SessionManager(db=SessionDB(tmp_path / "state.db")).create_session(cwd="/work") + + assert captured["disabled_toolsets"] is None + def test_create_session_writes_to_db(self, manager): state = manager.create_session(cwd="/project") db = manager._get_db() diff --git a/tests/agent/test_memory_provider.py b/tests/agent/test_memory_provider.py index a3cf671ab6614..f259c7f0bd485 100644 --- a/tests/agent/test_memory_provider.py +++ b/tests/agent/test_memory_provider.py @@ -5,7 +5,7 @@ import time import pytest from types import SimpleNamespace -from unittest.mock import MagicMock +from unittest.mock import MagicMock, patch from agent.memory_provider import MemoryProvider from agent.memory_manager import MemoryManager, inject_memory_provider_tools @@ -316,6 +316,20 @@ def test_tool_schemas_collected(self): names = {s["name"] for s in schemas} assert names == {"recall_builtin", "recall_ext"} + def test_strict_tool_schema_collection_propagates_provider_failure(self): + mgr = MemoryManager() + mgr.add_provider(FakeMemoryProvider("builtin", tools=[{"name": "good"}])) + broken = FakeMemoryProvider("broken") + mgr.add_provider(broken) + broken.get_tool_schemas = lambda: (_ for _ in ()).throw( + RuntimeError("schema callback failed") + ) + + with pytest.raises(RuntimeError, match="schema callback failed"): + mgr.get_all_tool_schemas_strict() + + assert {schema["name"] for schema in mgr.get_all_tool_schemas()} == {"good"} + def test_tool_name_conflict_first_wins(self): mgr = MemoryManager() p1 = FakeMemoryProvider("builtin", tools=[ @@ -1379,14 +1393,16 @@ class TestMemoryToolToolsetGate: These tests exercise the shared gate used by agent init and ACP refreshes. The gate condition is: - disabled_toolsets includes memory → skip injection + disabled toolsets include memory → skip injection (deny wins) enabled_toolsets is None → no filter, inject (backward compat) selected toolsets include memory → user opted in, inject otherwise (incl. []) → skip injection """ @staticmethod - def _run_memory_injection(enabled_toolsets, memory_manager, disabled_toolsets=None): + def _run_memory_injection( + enabled_toolsets, memory_manager, disabled_toolsets=None + ): """Run the shared memory-tool injection helper against a fake agent.""" fake_agent = SimpleNamespace( _memory_manager=memory_manager, @@ -1421,6 +1437,168 @@ def test_memory_in_toolsets_injects(self): tools, names = self._run_memory_injection(["terminal", "memory", "web"], mgr) assert "fact_store" in names + def test_registry_collision_is_not_provider_owned_or_dispatched(self): + manager = self._mgr_with_tools("shared_tool") + agent = SimpleNamespace( + _memory_manager=manager, + enabled_toolsets=None, + disabled_toolsets=None, + tools=[{ + "type": "function", + "function": { + "name": "shared_tool", + "description": "registry", + "parameters": {}, + }, + }], + valid_tool_names={"shared_tool"}, + session_id="session", + _current_turn_id="turn", + _current_api_request_id="request", + ) + inject_memory_provider_tools(agent) + assert agent._memory_provider_tool_names == set() + + from agent.agent_runtime_helpers import invoke_tool + with ( + patch( + "hermes_cli.plugins.get_pre_tool_call_block_message", + return_value=None, + ), + patch("hermes_cli.middleware.apply_tool_request_middleware") as request_mw, + patch( + "hermes_cli.middleware.run_tool_execution_middleware", + side_effect=lambda _name, args, execute, **_kw: execute(args), + ), + patch("run_agent.handle_function_call", return_value="registry") as generic, + ): + request_mw.return_value = SimpleNamespace(payload={}, trace=[]) + assert invoke_tool(agent, "shared_tool", {}, "task") == "registry" + generic.assert_called_once() + + def test_injected_provider_tool_keeps_provider_dispatch(self): + manager = self._mgr_with_tools("provider_tool") + agent = SimpleNamespace( + _memory_manager=manager, + enabled_toolsets=None, + disabled_toolsets=None, + tools=[], + valid_tool_names=set(), + session_id="session", + _current_turn_id="turn", + _current_api_request_id="request", + ) + inject_memory_provider_tools(agent) + assert agent._memory_provider_tool_names == {"provider_tool"} + + from agent.agent_runtime_helpers import invoke_tool + with ( + patch( + "hermes_cli.plugins.get_pre_tool_call_block_message", + return_value=None, + ), + patch("hermes_cli.middleware.apply_tool_request_middleware") as request_mw, + patch( + "hermes_cli.middleware.run_tool_execution_middleware", + side_effect=lambda _name, args, execute, **_kw: execute(args), + ), + patch("run_agent.handle_function_call") as generic, + ): + request_mw.return_value = SimpleNamespace(payload={}, trace=[]) + result = json.loads(invoke_tool(agent, "provider_tool", {}, "task")) + assert result["handled"] == "provider_tool" + generic.assert_not_called() + + def test_disabled_memory_toolset_blocks_default_injection(self): + """An explicit memory denial overrides the default-open tool surface.""" + mgr = self._mgr_with_tools("fact_store", "fact_feedback") + tools, names = self._run_memory_injection( + None, mgr, disabled_toolsets=["memory"] + ) + assert tools == [] + assert names == set() + + def test_disabled_memory_toolset_overrides_explicit_enable(self): + """The global disabled-toolset policy takes precedence over enablement.""" + mgr = self._mgr_with_tools("fact_store") + tools, names = self._run_memory_injection( + ["memory"], mgr, disabled_toolsets=["memory"] + ) + assert tools == [] + assert names == set() + + def test_disabled_all_toolsets_blocks_injection(self): + """The wildcard global denial also suppresses provider memory tools.""" + mgr = self._mgr_with_tools("fact_store") + tools, names = self._run_memory_injection( + None, mgr, disabled_toolsets=["*"] + ) + assert tools == [] + assert names == set() + + def test_disabled_composite_with_memory_blocks_injection(self): + """Composite subtraction also denies its provider-owned equivalent.""" + mgr = self._mgr_with_tools("fact_store") + tools, names = self._run_memory_injection( + None, mgr, disabled_toolsets=["coding"] + ) + assert tools == [] + assert names == set() + + def test_disabled_toolset_resolution_failure_blocks_injection(self, monkeypatch): + """Resolver failures must not reopen provider-owned memory tools.""" + import toolsets + + monkeypatch.setattr( + toolsets, + "resolve_toolset", + lambda _name: (_ for _ in ()).throw(RuntimeError("resolution failed")), + ) + mgr = self._mgr_with_tools("fact_store") + + tools, names = self._run_memory_injection( + None, mgr, disabled_toolsets=["coding"] + ) + + assert tools == [] + assert names == set() + + def test_unrelated_disabled_toolset_preserves_injection(self): + """A denial outside the memory toolset must not affect provider tools.""" + mgr = self._mgr_with_tools("fact_store") + tools, names = self._run_memory_injection( + None, mgr, disabled_toolsets=["terminal"] + ) + assert "fact_store" in names + assert any(t["function"]["name"] == "fact_store" for t in tools) + + def test_custom_toolset_subtracts_exact_provider_schema(self, monkeypatch): + """Provider schemas participate in ordinary resolved subtraction.""" + import toolsets + + monkeypatch.setitem( + toolsets.TOOLSETS, + "deny-fact-store", + {"description": "test", "tools": ["fact_store"], "includes": []}, + ) + mgr = self._mgr_with_tools("fact_store", "fact_search") + + tools, names = self._run_memory_injection( + None, mgr, disabled_toolsets=["deny-fact-store"] + ) + + assert names == {"fact_search"} + assert [tool["function"]["name"] for tool in tools] == ["fact_search"] + + def test_disabled_platform_bundle_preserves_shared_memory_tool(self): + """Bundle subtraction preserves core memory just like static tools.""" + mgr = self._mgr_with_tools("fact_store") + tools, names = self._run_memory_injection( + None, mgr, disabled_toolsets=["hermes-acp"] + ) + assert "fact_store" in names + assert any(t["function"]["name"] == "fact_store" for t in tools) + def test_composite_toolset_with_memory_injects(self): """Composite toolsets that include memory should inject provider tools.""" mgr = self._mgr_with_tools("hindsight_recall") diff --git a/tests/agent/test_system_prompt.py b/tests/agent/test_system_prompt.py index f37716a72e256..db38a439e5221 100644 --- a/tests/agent/test_system_prompt.py +++ b/tests/agent/test_system_prompt.py @@ -5,6 +5,7 @@ from types import SimpleNamespace from unittest.mock import patch +from agent.memory_manager import MemoryManager from agent.system_prompt import build_system_prompt, build_system_prompt_parts @@ -82,6 +83,133 @@ def _prompt_parts(agent): return build_system_prompt_parts(agent) +def _volatile_prompt(agent): + return _prompt_parts(agent)["volatile"] + + +def _memory_manager(*, prompt, tools): + manager = MemoryManager() + provider = SimpleNamespace( + name="prompt-test-provider", + get_tool_schemas=lambda: tools, + system_prompt_block=lambda: prompt, + ) + manager.add_provider(provider) + return manager + + +class TestMemoryProviderPromptPolicy: + def test_denied_provider_tools_do_not_leave_tool_instructions(self): + manager = _memory_manager( + prompt="Use provider_search to retrieve memory.", + tools=[{"name": "provider_search", "parameters": {}}], + ) + agent = _make_agent( + _memory_manager=manager, + enabled_toolsets=["memory"], + disabled_toolsets=["memory"], + ) + + assert "provider_search" not in _volatile_prompt(agent) + + def test_unselected_provider_tools_do_not_leave_tool_instructions(self): + manager = _memory_manager( + prompt="Use provider_search to retrieve memory.", + tools=[{"name": "provider_search", "parameters": {}}], + ) + agent = _make_agent( + _memory_manager=manager, + enabled_toolsets=["terminal"], + disabled_toolsets=None, + ) + + assert "provider_search" not in _volatile_prompt(agent) + + def test_allowed_provider_tools_keep_provider_prompt(self): + manager = _memory_manager( + prompt="Use provider_search to retrieve memory.", + tools=[{"name": "provider_search", "parameters": {}}], + ) + agent = _make_agent( + _memory_manager=manager, + enabled_toolsets=["memory"], + disabled_toolsets=None, + ) + + assert "provider_search" in _volatile_prompt(agent) + + def test_denied_tools_keep_passive_context_provider_prompt(self): + manager = _memory_manager( + prompt="Relevant memory context is injected automatically.", + tools=[], + ) + agent = _make_agent( + _memory_manager=manager, + enabled_toolsets=["memory"], + disabled_toolsets=["memory"], + ) + + assert "injected automatically" in _volatile_prompt(agent) + + def test_exact_denial_hides_only_affected_provider_prompt(self, monkeypatch): + import toolsets + + monkeypatch.setitem( + toolsets.TOOLSETS, + "deny-fact-store", + {"description": "test", "tools": ["fact_store"], "includes": []}, + ) + manager = MemoryManager() + manager.add_provider( + SimpleNamespace( + name="builtin", + get_tool_schemas=lambda: [{"name": "fact_store", "parameters": {}}], + system_prompt_block=lambda: "Use fact_store.", + ) + ) + manager.add_provider( + SimpleNamespace( + name="search-provider", + get_tool_schemas=lambda: [{"name": "fact_search", "parameters": {}}], + system_prompt_block=lambda: "Use fact_search.", + ) + ) + agent = _make_agent( + _memory_manager=manager, + enabled_toolsets=["memory"], + disabled_toolsets=["deny-fact-store"], + ) + + prompt = _volatile_prompt(agent) + + assert "fact_store" not in prompt + assert "fact_search" in prompt + + def test_partial_denial_hides_indivisible_provider_prompt(self): + manager = _memory_manager( + prompt="Use provider_store and provider_search.", + tools=[ + {"name": "provider_store", "parameters": {}}, + {"name": "provider_search", "parameters": {}}, + ], + ) + + assert manager.build_system_prompt(available_tool_names={"provider_search"}) == "" + + def test_all_provider_schemas_allowed_keeps_indivisible_prompt(self): + manager = _memory_manager( + prompt="Use provider_store and provider_search.", + tools=[ + {"name": "provider_store", "parameters": {}}, + {"name": "provider_search", "parameters": {}}, + ], + ) + + assert "provider_store" in manager.build_system_prompt( + available_tool_names={"provider_store", "provider_search"} + ) + + def _init_code_repo(path): """A git repo that actually holds code — the coding posture requires a source file (or manifest), not a bare ``.git`` (a prose/notes repo stays general).""" diff --git a/tests/cli/test_cli_approval_ui.py b/tests/cli/test_cli_approval_ui.py index ebca5bd85e927..b5bd8c553e067 100644 --- a/tests/cli/test_cli_approval_ui.py +++ b/tests/cli/test_cli_approval_ui.py @@ -48,6 +48,7 @@ def _make_background_cli_stub(): }) cli.max_turns = 90 cli.enabled_toolsets = [] + cli.disabled_toolsets = [] cli._session_db = None cli.reasoning_config = {} cli.service_tier = None @@ -66,7 +67,75 @@ def _make_background_cli_stub(): return cli +class _MemoryProviderSurfaceAgent: + """Build the real registry plus external-provider final tool surface.""" + + last_disabled_toolsets = None + last_valid_tool_names = None + + def __init__(self, **kwargs): + import model_tools + from agent.memory_manager import inject_memory_provider_tools + + self.enabled_toolsets = kwargs.get("enabled_toolsets") + self.disabled_toolsets = kwargs.get("disabled_toolsets") + self.tools = model_tools.get_tool_definitions( + enabled_toolsets=self.enabled_toolsets, + disabled_toolsets=self.disabled_toolsets, + quiet_mode=True, + ) + self.valid_tool_names = { + definition["function"]["name"] for definition in self.tools + } + self._memory_manager = SimpleNamespace( + get_all_tool_schemas=lambda: [ + { + "name": "fact_store", + "description": "store a fact", + "parameters": {}, + } + ] + ) + inject_memory_provider_tools(self) + self._print_fn = None + self.thinking_callback = None + type(self).last_disabled_toolsets = self.disabled_toolsets + type(self).last_valid_tool_names = set(self.valid_tool_names) + + def run_conversation(self, **_kwargs): + return {"final_response": "done", "messages": []} + + class TestCliApprovalUi: + def test_foreground_scalar_memory_absent_from_final_agent_surface(self): + with patch.dict( + cli_module.CLI_CONFIG["agent"], + {"disabled_toolsets": "memory"}, + ): + cli = HermesCLI(toolsets=["memory"], compact=True, max_turns=1) + + cli._session_db = object() + cli._resumed = False + cli.conversation_history = [] + cli._install_tool_callbacks = MagicMock() + cli._ensure_tirith_security = MagicMock() + cli._ensure_runtime_credentials = MagicMock(return_value=True) + + _MemoryProviderSurfaceAgent.last_disabled_toolsets = None + _MemoryProviderSurfaceAgent.last_valid_tool_names = None + with patch.object(cli_module, "AIAgent", _MemoryProviderSurfaceAgent), \ + patch.object(cli_module, "_prepare_deferred_agent_startup"), \ + patch("hermes_cli.mcp_startup.wait_for_mcp_discovery"): + assert cli._init_agent() is True + + assert cli.disabled_toolsets == ["memory"] + assert _MemoryProviderSurfaceAgent.last_disabled_toolsets == ["memory"] + assert _MemoryProviderSurfaceAgent.last_valid_tool_names is not None + assert _MemoryProviderSurfaceAgent.last_valid_tool_names.isdisjoint( + {"memory", "fact_store"} + ) + cli_module._active_agent_ref = None + def test_smart_denied_callback_offers_only_once_and_deny(self): cli = _make_cli_stub() result = {} @@ -398,6 +467,52 @@ def run_conversation(self, **kwargs): assert seen["sudo"].__func__ is HermesCLI._sudo_password_callback assert not cli._background_tasks + def test_background_task_inherits_disabled_toolsets(self): + """Classic CLI background agents retain global final subtraction.""" + cli = _make_background_cli_stub() + cli.disabled_toolsets = ["memory", "deny-provider-store"] + seen = {} + + class FakeAgent: + def __init__(self, **kwargs): + seen.update(kwargs) + self._print_fn = None + self.thinking_callback = None + + def run_conversation(self, **_kwargs): + return {"final_response": "done", "messages": []} + + with patch.object(cli_module, "AIAgent", FakeAgent), \ + patch.object(cli_module, "_cprint"), \ + patch.object(cli_module, "ChatConsole") as chat_console: + chat_console.return_value.print = MagicMock() + cli._handle_background_command("/btw check policy") + for thread in list(cli._background_tasks.values()): + thread.join(timeout=10) + + assert seen["disabled_toolsets"] == cli.disabled_toolsets + + def test_background_scalar_memory_absent_from_final_agent_surface(self): + cli = _make_background_cli_stub() + cli.enabled_toolsets = ["memory"] + cli.disabled_toolsets = "memory" + + _MemoryProviderSurfaceAgent.last_disabled_toolsets = None + _MemoryProviderSurfaceAgent.last_valid_tool_names = None + with patch.object(cli_module, "AIAgent", _MemoryProviderSurfaceAgent), \ + patch.object(cli_module, "_cprint"), \ + patch.object(cli_module, "ChatConsole") as chat_console: + chat_console.return_value.print = MagicMock() + cli._handle_background_command("/btw check policy") + for thread in list(cli._background_tasks.values()): + thread.join(timeout=10) + + assert _MemoryProviderSurfaceAgent.last_disabled_toolsets == ["memory"] + assert _MemoryProviderSurfaceAgent.last_valid_tool_names is not None + assert _MemoryProviderSurfaceAgent.last_valid_tool_names.isdisjoint( + {"memory", "fact_store"} + ) + def _make_real_paint_cli_stub(): """A stub whose modal repaint path runs the REAL _paint_now / _invalidate. @@ -768,4 +883,3 @@ def _worker(): assert not t.is_alive(), "worker thread never unblocked" assert result["value"] == "deny" - diff --git a/tests/cron/test_scheduler.py b/tests/cron/test_scheduler.py index ec8be40631318..d56dabee512eb 100644 --- a/tests/cron/test_scheduler.py +++ b/tests/cron/test_scheduler.py @@ -1619,6 +1619,27 @@ def test_run_job_disabled_toolsets_layer_user_config_on_baseline(self, tmp_path) "cronjob", "messaging", "clarify", "terminal", "file", } + def test_scalar_terminal_policy_is_absent_from_final_cron_tool_surface(self): + """A scalar terminal denial reaches final cron schema subtraction.""" + import model_tools + from cron.scheduler import _resolve_cron_disabled_toolsets + from toolsets import resolve_toolset + + disabled = _resolve_cron_disabled_toolsets( + {"agent": {"disabled_toolsets": "terminal"}} + ) + definitions = model_tools.get_tool_definitions( + enabled_toolsets=["terminal"], + disabled_toolsets=disabled, + quiet_mode=True, + ) + final_names = { + definition["function"]["name"] for definition in definitions + } + + assert "terminal" in disabled + assert final_names.isdisjoint(resolve_toolset("terminal")) + def test_run_job_enabled_toolsets_resolves_from_platform_config_when_not_set(self, tmp_path): """When a job has no explicit enabled_toolsets, the scheduler now resolves them from ``hermes tools`` platform config for ``cron`` diff --git a/tests/gateway/test_api_server.py b/tests/gateway/test_api_server.py index 347dc423a86d3..827bf4fa714f2 100644 --- a/tests/gateway/test_api_server.py +++ b/tests/gateway/test_api_server.py @@ -1371,6 +1371,77 @@ async def test_toolsets_returns_resolved_tools(self, adapter): assert by_name["web"]["tools"] == ["web_search"] assert by_name["default"]["configured"] is True + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("platform_toolsets", "disabled_toolsets", "expected"), + [ + pytest.param( + ["web", "terminal"], + ["web"], + {"web": False, "terminal": True}, + id="direct", + ), + pytest.param( + ["web", "terminal"], + "web", + {"web": False, "terminal": True}, + id="bare-string", + ), + pytest.param( + ["web", "terminal"], + None, + {"web": True, "terminal": True}, + id="none", + ), + pytest.param( + ["web", "terminal"], + ["*"], + {"web": False, "terminal": False}, + id="wildcard", + ), + pytest.param( + ["web", "yuanbao"], + ["hermes-yuanbao"], + {"web": True, "yuanbao": False}, + id="composite", + ), + pytest.param( + ["web", "terminal"], + ["web-extra"], + {"web": True, "terminal": True}, + id="nonmatching-exact-name", + ), + ], + ) + async def test_toolsets_enabled_applies_global_disabled_subtraction( + self, + adapter, + platform_toolsets, + disabled_toolsets, + expected, + ): + """Discovery must describe the same surface the agent receives.""" + config = { + "platform_toolsets": {"api_server": platform_toolsets}, + "agent": {"disabled_toolsets": disabled_toolsets}, + } + fake_toolsets = [ + (name, name.title(), f"{name} tools") for name in platform_toolsets + ] + with patch("hermes_cli.config.load_config", return_value=config), patch( + "hermes_cli.tools_config._get_effective_configurable_toolsets", + return_value=fake_toolsets, + ), patch( + "hermes_cli.tools_config._toolset_has_keys", + return_value=False, + ): + app = _create_app(adapter) + async with TestClient(TestServer(app)) as cli: + resp = await cli.get("/v1/toolsets") + assert resp.status == 200 + data = await resp.json() + assert {row["name"]: row["enabled"] for row in data["data"]} == expected + @pytest.mark.asyncio async def test_toolsets_handles_resolution_failure_per_toolset(self, adapter): """If one toolset fails to resolve, others still appear with empty tools.""" diff --git a/tests/gateway/test_api_server_toolset.py b/tests/gateway/test_api_server_toolset.py index 5940ee8c2f392..2e8939c9ff480 100644 --- a/tests/gateway/test_api_server_toolset.py +++ b/tests/gateway/test_api_server_toolset.py @@ -1,6 +1,8 @@ """Tests for hermes-api-server toolset and API server tool availability.""" from unittest.mock import patch, MagicMock +import pytest + from toolsets import resolve_toolset, get_toolset, validate_toolset @@ -176,3 +178,46 @@ def test_create_agent_respects_config_override(self): call_kwargs = mock_agent_cls.call_args toolsets = call_kwargs.kwargs.get("enabled_toolsets") assert sorted(toolsets) == ["terminal", "web"] + + @patch("gateway.platforms.api_server.AIOHTTP_AVAILABLE", True) + @pytest.mark.parametrize( + ("configured", "expected"), + [ + pytest.param("memory", ["memory"], id="bare-string"), + pytest.param( + ["*", "deny-provider-store"], + ["*", "deny-provider-store"], + id="list", + ), + pytest.param(None, None, id="none"), + ], + ) + def test_create_agent_forwards_global_disabled_toolsets( + self, + configured, + expected, + ): + """Enabled-name inference cannot replace final global subtraction.""" + from gateway.platforms.api_server import APIServerAdapter + from gateway.config import PlatformConfig + + adapter = APIServerAdapter(PlatformConfig()) + with patch("gateway.run._resolve_runtime_agent_kwargs") as mock_kwargs, \ + patch("gateway.run._resolve_gateway_model", return_value="test/model"), \ + patch("gateway.run._load_gateway_config") as mock_config, \ + patch("run_agent.AIAgent") as mock_agent_cls: + mock_kwargs.return_value = { + "api_key": "test-key", + "base_url": None, + "provider": None, + "api_mode": None, + "command": None, + "args": [], + } + mock_config.return_value = { + "agent": {"disabled_toolsets": configured} + } + + adapter._create_agent() + + assert mock_agent_cls.call_args.kwargs["disabled_toolsets"] == expected diff --git a/tests/gateway/test_session_model_override_routing.py b/tests/gateway/test_session_model_override_routing.py index b1e50c07bf31e..aa8442ef49f70 100644 --- a/tests/gateway/test_session_model_override_routing.py +++ b/tests/gateway/test_session_model_override_routing.py @@ -36,6 +36,40 @@ def run_conversation(self, user_message: str, conversation_history=None, task_id } +class _MemoryProviderSurfaceAgent(_CapturingAgent): + """Capture the real final registry plus external-provider tool surface.""" + + last_valid_tool_names = None + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + import model_tools + from agent.memory_manager import inject_memory_provider_tools + + self.enabled_toolsets = kwargs.get("enabled_toolsets") + self.disabled_toolsets = kwargs.get("disabled_toolsets") + self.tools = model_tools.get_tool_definitions( + enabled_toolsets=self.enabled_toolsets, + disabled_toolsets=self.disabled_toolsets, + quiet_mode=True, + ) + self.valid_tool_names = { + definition["function"]["name"] for definition in self.tools + } + self._memory_manager = types.SimpleNamespace( + get_all_tool_schemas=lambda: [ + { + "name": "fact_store", + "description": "store a fact", + "parameters": {}, + } + ] + ) + inject_memory_provider_tools(self) + type(self).last_valid_tool_names = set(self.valid_tool_names) + + def _make_runner(): runner = object.__new__(gateway_run.GatewayRunner) runner.adapters = {} @@ -126,6 +160,58 @@ def test_run_agent_prefers_session_override_over_global_runtime(monkeypatch): assert _CapturingAgent.last_init["reasoning_config"] == {"enabled": True, "effort": "high"} +def test_primary_messaging_scalar_memory_absent_from_final_agent_surface(monkeypatch): + monkeypatch.setattr( + gateway_run, + "_load_gateway_config", + lambda: {"agent": {"disabled_toolsets": "memory"}}, + ) + monkeypatch.setattr(gateway_run, "load_dotenv", lambda *args, **kwargs: None) + monkeypatch.setattr(gateway_run, "_resolve_runtime_agent_kwargs", _explode_runtime_resolution) + + import hermes_cli.tools_config as tools_config + + monkeypatch.setattr( + tools_config, + "_get_platform_tools", + lambda *_args, **_kwargs: {"memory"}, + ) + fake_run_agent = types.ModuleType("run_agent") + fake_run_agent.AIAgent = _MemoryProviderSurfaceAgent + monkeypatch.setitem(sys.modules, "run_agent", fake_run_agent) + + _MemoryProviderSurfaceAgent.last_init = None + _MemoryProviderSurfaceAgent.last_valid_tool_names = None + runner = _make_runner() + source = SessionSource( + platform=Platform.LOCAL, + chat_id="cli", + chat_name="CLI", + chat_type="dm", + user_id="user-1", + ) + session_key = "agent:main:local:dm" + runner._session_model_overrides[session_key] = _codex_override() + + result = asyncio.run( + runner._run_agent( + message="ping", + context_prompt="", + history=[], + source=source, + session_id="session-1", + session_key=session_key, + ) + ) + + assert result["final_response"] == "ok" + assert _MemoryProviderSurfaceAgent.last_init["disabled_toolsets"] == ["memory"] + assert _MemoryProviderSurfaceAgent.last_valid_tool_names is not None + assert _MemoryProviderSurfaceAgent.last_valid_tool_names.isdisjoint( + {"memory", "fact_store"} + ) + + @pytest.mark.asyncio async def test_background_task_prefers_session_override_over_global_runtime(monkeypatch): monkeypatch.setattr(gateway_run, "_load_gateway_config", lambda: {}) @@ -164,6 +250,56 @@ async def test_background_task_prefers_session_override_over_global_runtime(monk assert _CapturingAgent.last_init["api_key"] == "***" assert _CapturingAgent.last_init["reasoning_config"] == {"enabled": True, "effort": "high"} + +@pytest.mark.asyncio +async def test_gateway_background_scalar_memory_absent_from_final_agent_surface( + monkeypatch, +): + monkeypatch.setattr( + gateway_run, + "_load_gateway_config", + lambda: {"agent": {"disabled_toolsets": "memory"}}, + ) + monkeypatch.setattr(gateway_run, "_resolve_runtime_agent_kwargs", _explode_runtime_resolution) + + import hermes_cli.tools_config as tools_config + + monkeypatch.setattr( + tools_config, + "_get_platform_tools", + lambda *_args, **_kwargs: {"memory"}, + ) + fake_run_agent = types.ModuleType("run_agent") + fake_run_agent.AIAgent = _MemoryProviderSurfaceAgent + monkeypatch.setitem(sys.modules, "run_agent", fake_run_agent) + + _MemoryProviderSurfaceAgent.last_init = None + _MemoryProviderSurfaceAgent.last_valid_tool_names = None + runner = _make_runner() + adapter = AsyncMock() + adapter.send = AsyncMock() + adapter.extract_media = MagicMock(return_value=([], "ok")) + adapter.extract_images = MagicMock(return_value=([], "ok")) + runner.adapters[Platform.TELEGRAM] = adapter + + source = SessionSource( + platform=Platform.TELEGRAM, + user_id="12345", + chat_id="67890", + user_name="testuser", + ) + session_key = runner._session_key_for_source(source) + runner._session_model_overrides[session_key] = _codex_override() + + await runner._run_background_task("say hello", source, "bg_test") + + assert _MemoryProviderSurfaceAgent.last_init["disabled_toolsets"] == ["memory"] + assert _MemoryProviderSurfaceAgent.last_valid_tool_names is not None + assert _MemoryProviderSurfaceAgent.last_valid_tool_names.isdisjoint( + {"memory", "fact_store"} + ) + + def test_gateway_auth_fallback_uses_fallback_model_from_config(tmp_path, monkeypatch): """Regression: fallback provider must not inherit the primary model. @@ -260,4 +396,3 @@ def fake_resolve_runtime_provider(*, requested=None, explicit_base_url=None, exp assert runtime_kwargs["api_key"] == "env-secret" assert runtime_kwargs["base_url"] == "https://fallback.example/v1" assert runtime_kwargs["model"] == "fallback-model" - diff --git a/tests/hermes_cli/test_oneshot_toolset_policy.py b/tests/hermes_cli/test_oneshot_toolset_policy.py new file mode 100644 index 0000000000000..be51da7068b71 --- /dev/null +++ b/tests/hermes_cli/test_oneshot_toolset_policy.py @@ -0,0 +1,53 @@ +"""Security regressions for one-shot AIAgent toolset policy propagation.""" + +from unittest.mock import patch + +import pytest + + +@pytest.mark.parametrize( + ("configured", "expected"), + [ + pytest.param("memory", ["memory"], id="bare-string"), + pytest.param( + ["*", "deny-provider-store"], + ["*", "deny-provider-store"], + id="list", + ), + pytest.param(None, None, id="none"), + ], +) +def test_oneshot_forwards_global_disabled_toolsets(configured, expected): + from hermes_cli import oneshot + + captured = {} + + class FakeAgent: + def __init__(self, **kwargs): + captured.update(kwargs) + + def run_conversation(self, *_args, **_kwargs): + return {"final_response": "ok"} + + config = { + "model": {"default": "test/model"}, + "agent": {"disabled_toolsets": configured}, + } + runtime = { + "api_key": "test-key", + "base_url": "https://example.test/v1", + "provider": "test", + "api_mode": "chat_completions", + "credential_pool": None, + } + with ( + patch("hermes_cli.config.load_config", return_value=config), + patch("hermes_cli.runtime_provider.resolve_runtime_provider", return_value=runtime), + patch("run_agent.AIAgent", FakeAgent), + patch.object(oneshot, "_create_session_db_for_oneshot", return_value=None), + ): + text, result = oneshot._run_agent("check policy") + + assert text == "ok" + assert result["final_response"] == "ok" + assert captured["disabled_toolsets"] == expected diff --git a/tests/hermes_cli/test_tools_config.py b/tests/hermes_cli/test_tools_config.py index 0b3e842511854..857fa06089395 100644 --- a/tests/hermes_cli/test_tools_config.py +++ b/tests/hermes_cli/test_tools_config.py @@ -44,6 +44,19 @@ def test_agent_disabled_toolsets_suppresses_across_platforms(): assert "memory" not in discord_enabled +def test_agent_scalar_disabled_toolset_suppresses_exact_shared_selection(): + """A YAML scalar is one exact global toolset name, not characters.""" + config = { + "agent": {"disabled_toolsets": "memory"}, + "platform_toolsets": {"cli": ["memory", "terminal"]}, + } + + enabled = _get_platform_tools(config, "cli") + + assert "memory" not in enabled + assert "terminal" in enabled + + def test_agent_disabled_toolsets_with_explicit_platform_config(): """agent.disabled_toolsets should still suppress even when the platform has an explicit toolset list that includes the disabled toolset. diff --git a/tests/run_agent/test_run_agent.py b/tests/run_agent/test_run_agent.py index 3dfa2aff0d4d4..b48360b0e6abf 100644 --- a/tests/run_agent/test_run_agent.py +++ b/tests/run_agent/test_run_agent.py @@ -3404,10 +3404,18 @@ def test_agent_runtime_post_hook_ownership_predicate_covers_agent_tools(self, ag agent._context_engine_tool_names = {"context_query"} assert agent_runtime_owns_post_tool_hook(agent, "context_query") is True - agent._memory_manager = SimpleNamespace(has_tool=lambda name: name == "memory_extra") - assert agent_runtime_owns_post_tool_hook(agent, "memory_extra") is True assert agent_runtime_owns_post_tool_hook(agent, "web_search") is False + def test_memory_provider_post_hook_ownership_requires_live_manager(self, agent): + from agent.agent_runtime_helpers import agent_runtime_owns_post_tool_hook + + agent._memory_provider_tool_names = {"memory_extra"} + agent._memory_manager = SimpleNamespace() + assert agent_runtime_owns_post_tool_hook(agent, "memory_extra") is True + + agent._memory_manager = None + assert agent_runtime_owns_post_tool_hook(agent, "memory_extra") is False + def test_blocked_memory_tool_does_not_reset_counter(self, agent, monkeypatch): """Blocked memory tool should not reset the nudge counter.""" agent._turns_since_memory = 5 diff --git a/tests/tools/test_refresh_agent_mcp_tools.py b/tests/tools/test_refresh_agent_mcp_tools.py index 30f9791045075..96c788ad38719 100644 --- a/tests/tools/test_refresh_agent_mcp_tools.py +++ b/tests/tools/test_refresh_agent_mcp_tools.py @@ -98,6 +98,114 @@ def _capture(**kw): assert seen["disabled_toolsets"] == ["messaging"] +def test_failed_tightening_remains_pending_and_automatic_refresh_recovers( + monkeypatch, +): + """Assembly failure must not consume a policy epoch or lose tightening.""" + agent = _agent( + ["read_file", "fact_store"], + ) + agent._memory_provider_tool_names = {"fact_store"} + agent._memory_manager = types.SimpleNamespace( + get_all_tool_schemas=lambda: [ + {"name": "fact_store", "description": "", "parameters": {}} + ] + ) + original_tools = agent.tools + original_names = agent.valid_tool_names + definitions_available = False + + import model_tools + + def _definitions(*, enabled_toolsets, **_kw): + if not definitions_available: + raise RuntimeError("definition failure") + return [_tool(enabled_toolsets[0])] + + monkeypatch.setattr(model_tools, "get_tool_definitions", _definitions) + + import pytest + + with pytest.raises(RuntimeError, match="definition failure"): + mcp_tool.refresh_agent_mcp_tools( + agent, + enabled_override=["terminal"], + ) + + assert agent.enabled_toolsets is None + assert agent.tools is original_tools + assert agent.valid_tool_names is original_names + assert agent._memory_provider_tool_names == {"fact_store"} + assert agent._tool_policy_epoch == 1 + assert getattr(agent, "_tool_published_policy_epoch", 0) == 0 + + definitions_available = True + mcp_tool.refresh_agent_mcp_tools(agent) + + assert agent.enabled_toolsets == ["terminal"] + assert agent.valid_tool_names == {"terminal"} + assert agent._memory_provider_tool_names == set() + assert agent._tool_published_policy_epoch == agent._tool_policy_epoch == 1 + + +def test_pending_policy_recovery_cannot_overwrite_newer_policy(monkeypatch): + """An older recovery finishing last cannot replace the latest policy.""" + agent = _agent(["old-tool"], enabled=["old-policy"]) + fail_first = True + recovery_entered = threading.Event() + release_recovery = threading.Event() + + import model_tools + + def _definitions(*, enabled_toolsets, **_kw): + nonlocal fail_first + policy = enabled_toolsets[0] + if policy == "pending-policy": + if fail_first: + fail_first = False + raise RuntimeError("definition failure") + recovery_entered.set() + assert release_recovery.wait(5) + return [_tool(policy)] + + monkeypatch.setattr(model_tools, "get_tool_definitions", _definitions) + + import pytest + + with pytest.raises(RuntimeError, match="definition failure"): + mcp_tool.refresh_agent_mcp_tools( + agent, + enabled_override=["pending-policy"], + ) + + recovery_errors = [] + + def _recover(): + try: + mcp_tool.refresh_agent_mcp_tools(agent) + except Exception as exc: # pragma: no cover - failure diagnostic + recovery_errors.append(exc) + + recovery = threading.Thread( + target=_recover, + ) + recovery.start() + assert recovery_entered.wait(5) + + mcp_tool.refresh_agent_mcp_tools( + agent, + enabled_override=["latest-policy"], + ) + release_recovery.set() + recovery.join(5) + + assert not recovery.is_alive() + assert recovery_errors == [] + assert agent.enabled_toolsets == ["latest-policy"] + assert agent.valid_tool_names == {"latest-policy"} + assert agent._tool_published_policy_epoch == agent._tool_policy_epoch == 2 + + def test_refresh_preserves_memory_provider_and_context_engine_tools(monkeypatch): """B1 regression: a rebuild must NOT drop post-build-injected tools. @@ -116,6 +224,7 @@ def test_refresh_preserves_memory_provider_and_context_engine_tools(monkeypatch) {"name": "memory_search", "description": "", "parameters": {}} ] ) + agent.context_compressor = types.SimpleNamespace( get_tool_schemas=lambda: [ {"name": "lcm_grep", "description": "", "parameters": {}} @@ -141,7 +250,7 @@ def test_refresh_preserves_memory_provider_and_context_engine_tools(monkeypatch) def test_refresh_does_not_reinject_disabled_memory_provider_tools(monkeypatch): - """A refresh removes stale provider tools when memory becomes disabled.""" + """An MCP rebuild must preserve the session's final memory denial.""" agent = _agent( ["read_file", "memory_search"], enabled=["all"], @@ -157,13 +266,563 @@ def test_refresh_does_not_reinject_disabled_memory_provider_tools(monkeypatch): monkeypatch.setattr( model_tools, "get_tool_definitions", - lambda **kw: [_tool("read_file")], + lambda **kw: [_tool("read_file"), _tool("mcp_new_server_tool")], + ) + + mcp_tool.refresh_agent_mcp_tools(agent) + + assert "mcp_new_server_tool" in agent.valid_tool_names + assert "memory_search" not in agent.valid_tool_names + assert all( + tool["function"]["name"] != "memory_search" for tool in agent.tools + ) + + +def test_refresh_subtracts_only_exact_provider_tool_name(monkeypatch): + """A custom denial removes its provider schema without hiding siblings.""" + agent = _agent(["read_file", "fact_store", "fact_search"]) + agent._memory_manager = types.SimpleNamespace( + get_all_tool_schemas=lambda: [ + {"name": "fact_store", "description": "", "parameters": {}}, + {"name": "fact_search", "description": "", "parameters": {}}, + ] + ) + + import model_tools + import toolsets + + monkeypatch.setitem( + toolsets.TOOLSETS, + "deny-fact-store", + {"description": "test", "tools": ["fact_store"], "includes": []}, + ) + monkeypatch.setattr( + model_tools, + "get_tool_definitions", + lambda **_kw: [_tool("read_file"), _tool("mcp_new_tool")], + ) + + mcp_tool.refresh_agent_mcp_tools( + agent, + disabled_override=["deny-fact-store"], + ) + + assert "fact_store" not in agent.valid_tool_names + assert "fact_search" in agent.valid_tool_names + assert "mcp_new_tool" in agent.valid_tool_names + + +def test_refresh_provider_change_invalidates_prompt_cache(monkeypatch): + agent = _agent(["read_file", "fact_store"]) + agent._memory_provider_tool_names = {"fact_store"} + agent._memory_manager = types.SimpleNamespace( + get_all_tool_schemas=lambda: [ + {"name": "fact_store", "description": "", "parameters": {}} + ] + ) + agent._cached_system_prompt = "Use fact_store." + + import model_tools + import toolsets + + monkeypatch.setitem( + toolsets.TOOLSETS, + "deny-fact-store", + {"description": "test", "tools": ["fact_store"], "includes": []}, + ) + monkeypatch.setattr( + model_tools, + "get_tool_definitions", + lambda **_kw: [_tool("read_file")], + ) + + mcp_tool.refresh_agent_mcp_tools( + agent, + disabled_override=["deny-fact-store"], + ) + + assert agent._cached_system_prompt is None + + +def test_refresh_provider_collision_invalidates_prompt_cache(monkeypatch): + """A registry collision transfers ownership away from the provider.""" + agent = _agent(["read_file", "fact_store"]) + provider_calls = 0 + provider_lock_was_free = [] + + def _provider_schemas(): + nonlocal provider_calls + provider_calls += 1 + lock_was_free = mcp_tool._agent_tools_lock.acquire(blocking=False) + provider_lock_was_free.append(lock_was_free) + if lock_was_free: + mcp_tool._agent_tools_lock.release() + return [{"name": "fact_store", "description": "", "parameters": {}}] + + agent._memory_manager = types.SimpleNamespace( + get_all_tool_schemas=_provider_schemas, + ) + agent._memory_provider_tool_names = {"fact_store"} + agent._cached_system_prompt = "Use fact_store as provider memory." + original_tools = agent.tools + + import model_tools + + monkeypatch.setattr( + model_tools, + "get_tool_definitions", + lambda **_kw: [_tool("read_file"), _tool("fact_store")], + ) + + mcp_tool.refresh_agent_mcp_tools(agent) + + assert agent.tools is not original_tools + assert agent.tools == original_tools + assert agent._cached_system_prompt is None + assert agent._memory_provider_tool_names == set() + from agent.agent_runtime_helpers import memory_provider_owns_tool + assert not memory_provider_owns_tool(agent, "fact_store") + assert provider_calls == 1 + assert provider_lock_was_free == [True] + + +def test_refresh_equal_names_publishes_registry_schema_on_provider_transfer( + monkeypatch, +): + """Provider -> registry transfer must atomically replace the contract.""" + provider_tool = { + "type": "function", + "function": { + "name": "shared_tool", + "description": "Provider-owned lookup by memory query.", + "parameters": { + "type": "object", + "properties": {"query": {"type": "string"}}, + "required": ["query"], + }, + }, + } + registry_tool = { + "type": "function", + "function": { + "name": "shared_tool", + "description": "Registry-owned lookup by numeric record id.", + "parameters": { + "type": "object", + "properties": {"record_id": {"type": "integer"}}, + "required": ["record_id"], + }, + }, + } + agent = _agent([]) + agent.tools = [provider_tool] + agent.valid_tool_names = {"shared_tool"} + agent._memory_provider_tool_names = {"shared_tool"} + agent._memory_manager = types.SimpleNamespace( + get_all_tool_schemas=lambda: [provider_tool["function"]], + handle_tool_call=lambda _name, _args: "provider-dispatch", + ) + agent.session_id = "session" + + import model_tools + import run_agent + from agent.agent_runtime_helpers import invoke_tool + from hermes_cli import middleware + + monkeypatch.setattr( + model_tools, "get_tool_definitions", lambda **_kw: [registry_tool] + ) + monkeypatch.setattr( + middleware, + "run_tool_execution_middleware", + lambda _name, args, execute, **_kw: execute(args), + ) + registry_calls = [] + monkeypatch.setattr( + run_agent, + "handle_function_call", + lambda name, args, *_a, **_kw: ( + registry_calls.append((name, args)) or "registry-dispatch" + ), + ) + + mcp_tool.refresh_agent_mcp_tools(agent) + + assert agent.tools == [registry_tool] + assert agent._memory_provider_tool_names == set() + assert ( + invoke_tool( + agent, + "shared_tool", + {"record_id": 7}, + "task", + pre_tool_block_checked=True, + skip_tool_request_middleware=True, + ) + == "registry-dispatch" + ) + assert registry_calls == [("shared_tool", {"record_id": 7})] + + +def test_refresh_equal_names_publishes_provider_schema_on_registry_transfer( + monkeypatch, +): + """Registry -> provider transfer must atomically replace the contract.""" + registry_tool = { + "type": "function", + "function": { + "name": "shared_tool", + "description": "Registry-owned lookup by numeric record id.", + "parameters": { + "type": "object", + "properties": {"record_id": {"type": "integer"}}, + "required": ["record_id"], + }, + }, + } + provider_schema = { + "name": "shared_tool", + "description": "Provider-owned lookup by memory query.", + "parameters": { + "type": "object", + "properties": {"query": {"type": "string"}}, + "required": ["query"], + }, + } + agent = _agent([]) + agent.tools = [registry_tool] + agent.valid_tool_names = {"shared_tool"} + agent._memory_provider_tool_names = set() + provider_calls = [] + agent._memory_manager = types.SimpleNamespace( + get_all_tool_schemas=lambda: [provider_schema], + handle_tool_call=lambda name, args: ( + provider_calls.append((name, args)) or "provider-dispatch" + ), + ) + agent.session_id = "session" + + import model_tools + import run_agent + from agent.agent_runtime_helpers import invoke_tool + from hermes_cli import middleware + + monkeypatch.setattr(model_tools, "get_tool_definitions", lambda **_kw: []) + monkeypatch.setattr( + middleware, + "run_tool_execution_middleware", + lambda _name, args, execute, **_kw: execute(args), + ) + registry_calls = [] + monkeypatch.setattr( + run_agent, + "handle_function_call", + lambda name, args, *_a, **_kw: ( + registry_calls.append((name, args)) or "registry-dispatch" + ), + ) + + mcp_tool.refresh_agent_mcp_tools(agent) + + assert agent.tools == [{"type": "function", "function": provider_schema}] + assert agent._memory_provider_tool_names == {"shared_tool"} + assert ( + invoke_tool( + agent, + "shared_tool", + {"query": "needle"}, + "task", + pre_tool_block_checked=True, + skip_tool_request_middleware=True, + ) + == "provider-dispatch" + ) + assert provider_calls == [("shared_tool", {"query": "needle"})] + assert registry_calls == [] + + +def test_refresh_records_genuinely_injected_provider_ownership(monkeypatch): + agent = _agent(["read_file"]) + agent._memory_manager = types.SimpleNamespace( + get_all_tool_schemas=lambda: [ + {"name": "fact_store", "description": "", "parameters": {}} + ] + ) + import model_tools + monkeypatch.setattr( + model_tools, + "get_tool_definitions", + lambda **_kw: [_tool("read_file")], + ) + + mcp_tool.refresh_agent_mcp_tools(agent) + + assert agent._memory_provider_tool_names == {"fact_store"} + assert "fact_store" in agent.valid_tool_names + from agent.agent_runtime_helpers import memory_provider_owns_tool + assert memory_provider_owns_tool(agent, "fact_store") + + +def test_schema_callback_failure_preserves_all_refresh_state(monkeypatch): + agent = _agent(["read_file", "fact_store"], enabled=["memory"]) + agent._memory_provider_tool_names = {"fact_store"} + agent._context_engine_tool_names = {"lcm_grep"} + agent._cached_system_prompt = "provider prompt" + agent._memory_manager = types.SimpleNamespace( + get_all_tool_schemas=lambda: (_ for _ in ()).throw( + RuntimeError("schema callback failed") + ) + ) + original_tools = agent.tools + original_names = agent.valid_tool_names + import model_tools + monkeypatch.setattr( + model_tools, + "get_tool_definitions", + lambda **_kw: [_tool("read_file")], + ) + import pytest + + with pytest.raises(RuntimeError, match="schema callback failed"): + mcp_tool.refresh_agent_mcp_tools( + agent, + enabled_override=["coding"], + ) + + assert agent.tools is original_tools + assert agent.valid_tool_names is original_names + assert agent.enabled_toolsets == ["memory"] + assert agent._memory_provider_tool_names == {"fact_store"} + assert agent._context_engine_tool_names == {"lcm_grep"} + assert agent._cached_system_prompt == "provider prompt" + + +def test_schema_failure_cannot_veto_full_provider_family_revocation(monkeypatch): + """A denied provider family is removable without enumerating schemas.""" + agent = _agent(["read_file", "fact_store"], enabled=["memory"]) + agent._memory_provider_tool_names = {"fact_store"} + agent._cached_system_prompt = "Use fact_store." + provider_calls = [] + agent._memory_manager = types.SimpleNamespace( + get_all_tool_schemas=lambda: (_ for _ in ()).throw( + RuntimeError("schema callback failed") + ), + handle_tool_call=lambda name, args: provider_calls.append((name, args)), + ) + + import model_tools + monkeypatch.setattr( + model_tools, + "get_tool_definitions", + lambda **_kw: [_tool("read_file")], + ) + + mcp_tool.refresh_agent_mcp_tools(agent, disabled_override=["memory"]) + + assert agent.disabled_toolsets == ["memory"] + assert agent.valid_tool_names == {"read_file"} + assert agent._memory_provider_tool_names == set() + assert agent._cached_system_prompt is None + from agent.agent_runtime_helpers import memory_provider_owns_tool + assert not memory_provider_owns_tool(agent, "fact_store") + assert provider_calls == [] + + +def test_schema_failure_still_applies_exact_provider_name_revocation( + monkeypatch, +): + """Tightening may retain allowed published schemas without adding any.""" + agent = _agent( + ["read_file", "fact_store", "fact_search"], + enabled=["memory"], + ) + agent._memory_provider_tool_names = {"fact_store", "fact_search"} + agent._memory_manager = types.SimpleNamespace( + get_all_tool_schemas=lambda: (_ for _ in ()).throw( + RuntimeError("schema callback failed") + ) + ) + + import model_tools + import toolsets + monkeypatch.setitem( + toolsets.TOOLSETS, + "deny-fact-store", + {"description": "test", "tools": ["fact_store"], "includes": []}, + ) + monkeypatch.setattr( + model_tools, + "get_tool_definitions", + lambda **_kw: [_tool("read_file")], + ) + + mcp_tool.refresh_agent_mcp_tools( + agent, + disabled_override=["deny-fact-store"], + ) + + assert agent.valid_tool_names == {"read_file", "fact_search"} + assert agent._memory_provider_tool_names == {"fact_search"} + + +def test_stale_generation_explicit_policy_retries_before_publishing_epoch( + monkeypatch, +): + """A winning registry generation cannot consume a tightening policy.""" + from tools import registry as registry_module + + agent = _agent(["read_file", "fact_store"], enabled=["memory"]) + agent._memory_provider_tool_names = {"fact_store"} + agent._memory_manager = types.SimpleNamespace( + get_all_tool_schemas=lambda: [ + {"name": "fact_store", "description": "", "parameters": {}} + ] + ) + explicit_entered = threading.Event() + release_explicit = threading.Event() + explicit_calls = 0 + + import model_tools + + def _definitions(*, disabled_toolsets, **_kw): + nonlocal explicit_calls + if disabled_toolsets == ["memory"]: + explicit_calls += 1 + if explicit_calls == 1: + explicit_entered.set() + assert release_explicit.wait(5) + return [_tool("read_file")] + + monkeypatch.setattr(model_tools, "get_tool_definitions", _definitions) + refresh_error = [] + + def _refresh(): + try: + mcp_tool.refresh_agent_mcp_tools( + agent, + disabled_override=["memory"], + ) + except Exception as exc: # pragma: no cover - failure diagnostic + refresh_error.append(exc) + + refresh = threading.Thread(target=_refresh) + refresh.start() + assert explicit_entered.wait(5) + + sentinel = "test_policy_generation_sentinel" + registry_module.registry.register( + name=sentinel, + toolset="test", + schema={"name": sentinel, "description": "", "parameters": {}}, + handler=lambda _args, **_kw: "{}", + ) + try: + # Model a newer registry snapshot publishing while the explicit policy + # rebuild is still staging its older generation. + agent._tool_snapshot_generation = registry_module.registry._generation + release_explicit.set() + refresh.join(5) + + assert not refresh.is_alive() + assert refresh_error == [] + assert explicit_calls == 2 + assert agent.disabled_toolsets == ["memory"] + assert agent.valid_tool_names == {"read_file"} + assert agent._memory_provider_tool_names == set() + assert agent._tool_published_policy_epoch == agent._tool_policy_epoch + finally: + registry_module.registry.deregister(sentinel) + + +def test_same_generation_automatic_refresh_cannot_beat_newer_policy(monkeypatch): + agent = _agent(["old_tool"], enabled=["old-policy"]) + automatic_entered = threading.Event() + release_automatic = threading.Event() + + import model_tools + + def _definitions(*, enabled_toolsets, **_kw): + if enabled_toolsets == ["old-policy"]: + automatic_entered.set() + assert release_automatic.wait(5) + return [_tool("old_tool")] + return [_tool("new_tool")] + + monkeypatch.setattr(model_tools, "get_tool_definitions", _definitions) + automatic = threading.Thread( + target=mcp_tool.refresh_agent_mcp_tools, + args=(agent,), + ) + automatic.start() + assert automatic_entered.wait(5) + + mcp_tool.refresh_agent_mcp_tools( + agent, + enabled_override=["new-policy"], + ) + release_automatic.set() + automatic.join(5) + + assert not automatic.is_alive() + assert agent.enabled_toolsets == ["new-policy"] + assert agent.valid_tool_names == {"new_tool"} + + +def test_failed_refresh_preserves_prompt_cache(monkeypatch): + agent = _agent(["read_file", "fact_store"]) + agent._cached_system_prompt = "Use fact_store." + + import model_tools + + monkeypatch.setattr( + model_tools, + "get_tool_definitions", + lambda **_kw: (_ for _ in ()).throw(RuntimeError("definition failure")), + ) + + import pytest + + with pytest.raises(RuntimeError, match="definition failure"): + mcp_tool.refresh_agent_mcp_tools( + agent, + disabled_override=["memory"], + ) + + assert agent._cached_system_prompt == "Use fact_store." + + +def test_refresh_resolution_failure_does_not_reinject_memory_provider_tools( + monkeypatch, +): + """A disabled-toolset resolver failure must keep provider tools denied.""" + agent = _agent(["read_file", "memory_search"], disabled=["coding"]) + agent._memory_manager = types.SimpleNamespace( + get_all_tool_schemas=lambda: [ + {"name": "memory_search", "description": "", "parameters": {}} + ] + ) + + import model_tools + import toolsets + + monkeypatch.setattr( + model_tools, + "get_tool_definitions", + lambda **kw: [_tool("read_file"), _tool("mcp_new_server_tool")], + ) + monkeypatch.setattr( + toolsets, + "resolve_toolset", + lambda _name: (_ for _ in ()).throw(RuntimeError("resolution failed")), ) mcp_tool.refresh_agent_mcp_tools(agent) + assert "mcp_new_server_tool" in agent.valid_tool_names assert "memory_search" not in agent.valid_tool_names - assert all(t["function"]["name"] != "memory_search" for t in agent.tools) + assert all( + tool["function"]["name"] != "memory_search" for tool in agent.tools + ) def test_refresh_respects_context_engine_toolset_gate(monkeypatch): @@ -188,6 +847,31 @@ def test_refresh_respects_context_engine_toolset_gate(monkeypatch): assert "lcm_grep" not in agent.valid_tool_names # gated out (#5544) +def test_refresh_equal_names_transfers_context_engine_ownership(monkeypatch): + """A registry replacement with the same name must stop engine routing.""" + agent = _agent(["read_file", "lcm_grep"]) + agent.context_compressor = types.SimpleNamespace( + get_tool_schemas=lambda: [ + {"name": "lcm_grep", "description": "", "parameters": {}} + ] + ) + agent._context_engine_tool_names = {"lcm_grep"} + original_tools = agent.tools + + import model_tools + + monkeypatch.setattr( + model_tools, + "get_tool_definitions", + lambda **_kw: [_tool("read_file"), _tool("lcm_grep")], + ) + + mcp_tool.refresh_agent_mcp_tools(agent) + + assert agent.tools is original_tools + assert agent._context_engine_tool_names == set() + + def test_refreshed_tool_is_callable_through_valid_tool_names_guard(monkeypatch): """The whole point: a late tool, once refreshed, passes the name guard the run loop uses to accept/reject tool calls (agent.valid_tool_names).""" diff --git a/tests/tui_gateway/test_make_agent_provider.py b/tests/tui_gateway/test_make_agent_provider.py index 94b606dbd385f..76b7a796a489c 100644 --- a/tests/tui_gateway/test_make_agent_provider.py +++ b/tests/tui_gateway/test_make_agent_provider.py @@ -8,6 +8,8 @@ import os from unittest.mock import MagicMock, patch +import pytest + def test_make_agent_passes_resolved_provider(): """_make_agent forwards provider/base_url/api_key/api_mode from @@ -60,6 +62,73 @@ def test_make_agent_passes_resolved_provider(): assert call_kwargs.kwargs["api_mode"] == "anthropic_messages" +@pytest.mark.parametrize( + ("configured", "expected"), + [ + pytest.param("memory", ["memory"], id="bare-string"), + pytest.param( + ["memory", "deny-provider-store"], + ["memory", "deny-provider-store"], + id="list", + ), + pytest.param(None, None, id="none"), + ], +) +def test_make_agent_forwards_global_disabled_toolsets(configured, expected): + """Desktop/TUI construction must retain final global subtraction.""" + fake_runtime = { + "provider": "openrouter", + "base_url": "https://openrouter.ai/api/v1", + "api_key": "sk-test", + "api_mode": "chat_completions", + "command": None, + "args": None, + "credential_pool": None, + } + fake_cfg = { + "agent": { + "system_prompt": "", + "disabled_toolsets": configured, + }, + "model": {"default": "test/model"}, + } + + with ( + patch("tui_gateway.server._load_cfg", return_value=fake_cfg), + patch("tui_gateway.server._get_db", return_value=MagicMock()), + patch("tui_gateway.server._load_enabled_toolsets", return_value=["coding"]), + patch( + "hermes_cli.runtime_provider.resolve_runtime_provider", + return_value=fake_runtime, + ), + patch("run_agent.AIAgent") as mock_agent, + ): + from tui_gateway.server import _make_agent + + _make_agent("sid-policy", "key-policy") + + assert mock_agent.call_args.kwargs["disabled_toolsets"] == expected + + +def test_tui_background_agent_inherits_resolved_disabled_toolsets(monkeypatch): + """Child agents inherit the parent's resolved subtraction verbatim.""" + from types import SimpleNamespace + from tui_gateway import server + + parent = SimpleNamespace( + enabled_toolsets=["coding"], + disabled_toolsets=["memory", "deny-provider-store"], + model="test/model", + ) + monkeypatch.setattr(server, "_load_cfg", lambda: {}) + monkeypatch.setattr(server, "_get_db", lambda: None) + monkeypatch.setattr(server, "_load_fallback_model", lambda: None) + + kwargs = server._background_agent_kwargs(parent, "bg-policy") + + assert kwargs["disabled_toolsets"] == parent.disabled_toolsets + + def test_make_agent_forwards_provider_routing(): """Parity with the messaging gateway + CLI: ``provider_routing`` in config.yaml must reach AIAgent so OpenRouter honors the user's sort / diff --git a/tests/tui_gateway/test_mcp_reload_rev.py b/tests/tui_gateway/test_mcp_reload_rev.py index 866473d1a42e2..4c7c043bd7d4d 100644 --- a/tests/tui_gateway/test_mcp_reload_rev.py +++ b/tests/tui_gateway/test_mcp_reload_rev.py @@ -81,6 +81,171 @@ def _boom(): assert srv._mcp_reload_loaded_rev == "" +def test_failed_session_refresh_is_an_error_and_no_generation_advance( + reload_env, + monkeypatch, +): + """A rebuilt registry is not complete until the session snapshot refreshes.""" + agent = object() + monkeypatch.setitem(srv._sessions, "session-a", {"agent": agent}) + monkeypatch.setattr(srv, "_session_uses_compute_host", lambda _session: False) + monkeypatch.setattr(srv, "_load_enabled_toolsets", lambda: ["terminal"]) + monkeypatch.setattr( + mcp_tool, + "refresh_agent_mcp_tools", + lambda *_args, **_kwargs: (_ for _ in ()).throw( + RuntimeError("definition failure") + ), + ) + + envelope = srv._methods["reload.mcp"]( + 1, + {"session_id": "session-a", "confirm": True, "rev": "rev-a"}, + ) + + assert envelope["error"]["message"] == "definition failure" + assert srv._mcp_reload_gen == 0 + assert srv._mcp_reload_loaded_rev == "" + + +@pytest.mark.parametrize( + "host_ack", + [ + pytest.param( + { + "type": "control.error", + "message": "definition failure", + }, + id="control-error", + ), + pytest.param( + { + "type": "control.ack", + "response": { + "jsonrpc": "2.0", + "id": "reload-mcp-1", + "result": {"status": "reloaded", "loaded_rev": "rev-a"}, + }, + }, + id="unexpected-ack-type", + ), + pytest.param( + {"type": "reload_mcp.ack"}, + id="missing-response", + ), + pytest.param( + {"type": "reload_mcp.ack", "response": []}, + id="malformed-response", + ), + pytest.param( + { + "type": "reload_mcp.ack", + "response": { + "jsonrpc": "2.0", + "id": "reload-mcp-1", + "error": {"code": 5015, "message": "definition failure"}, + }, + }, + id="nested-json-rpc-error", + ), + pytest.param( + { + "type": "reload_mcp.ack", + "response": { + "jsonrpc": "2.0", + "id": "reload-mcp-1", + "result": {"status": "pending", "loaded_rev": "rev-a"}, + }, + }, + id="unexpected-result-status", + ), + ], +) +def test_compute_host_reload_rejects_failed_child_ack_without_state_advance( + reload_env, + monkeypatch, + host_ack, +): + """A child reload is accepted only after its nested RPC reports success.""" + monkeypatch.setitem( + srv._sessions, + "session-a", + {"agent": object(), "_compute_host_active": True}, + ) + monkeypatch.setattr(srv, "_session_uses_compute_host", lambda _session: True) + + class _Supervisor: + def reload_mcp(self, _sid, *, request_id): + assert request_id == "reload-mcp-1" + return host_ack + + supervisor = _Supervisor() + monkeypatch.setattr(srv, "_get_compute_host_supervisor", lambda: supervisor) + + envelope = srv._methods["reload.mcp"]( + 1, + {"session_id": "session-a", "confirm": True, "rev": "rev-a"}, + ) + + assert "error" in envelope + assert srv._mcp_reload_gen == 0 + assert srv._mcp_reload_loaded_rev == "" + + +def test_compute_host_reload_recovers_and_propagates_child_loaded_rev( + reload_env, + monkeypatch, +): + """The same revision remains retryable and carries the recovered child rev.""" + monkeypatch.setitem( + srv._sessions, + "session-a", + {"agent": object(), "_compute_host_active": True}, + ) + monkeypatch.setattr(srv, "_session_uses_compute_host", lambda _session: True) + + class _Supervisor: + calls = 0 + + def reload_mcp(self, _sid, *, request_id): + self.calls += 1 + if self.calls == 1: + return { + "type": "reload_mcp.ack", + "response": { + "jsonrpc": "2.0", + "id": request_id, + "error": {"code": 5015, "message": "definition failure"}, + }, + } + return { + "type": "reload_mcp.ack", + "response": { + "jsonrpc": "2.0", + "id": request_id, + "result": {"status": "reloaded", "loaded_rev": "rev-recovered"}, + }, + } + + supervisor = _Supervisor() + monkeypatch.setattr(srv, "_get_compute_host_supervisor", lambda: supervisor) + + failed = srv._methods["reload.mcp"]( + 1, + {"session_id": "session-a", "confirm": True, "rev": "rev-a"}, + ) + recovered = srv._methods["reload.mcp"]( + 1, + {"session_id": "session-a", "confirm": True, "rev": "rev-a"}, + ) + + assert "error" in failed + assert recovered["result"]["status"] == "reloaded" + assert recovered["result"]["loaded_rev"] == "rev-recovered" + assert srv._mcp_reload_gen == 1 + assert srv._mcp_reload_loaded_rev == "rev-recovered" + + def test_leader_rehashes_until_stable_when_config_changes_mid_reload(reload_env, monkeypatch): """Revision A starts a reload; the config changes to revision B while discovery is connecting servers. The leader must not mark A complete — diff --git a/tools/mcp_tool.py b/tools/mcp_tool.py index 3a1964b0cf084..28b505b771724 100644 --- a/tools/mcp_tool.py +++ b/tools/mcp_tool.py @@ -5994,6 +5994,7 @@ def refresh_agent_mcp_tools( enabled_override=None, disabled_override=None, quiet_mode: bool = True, + _policy_epoch_override=None, ) -> set: """Re-derive an already-built agent's tool snapshot from the live registry. @@ -6036,14 +6037,78 @@ def refresh_agent_mcp_tools( # the user just ENABLED in config is picked up; the agent's stored selection # is then updated to match. The automatic paths (between-turns, late-binding) # pass nothing and reuse the agent's build-time selection unchanged. - if enabled_override is not None or disabled_override is not None: - enabled = enabled_override if enabled_override is not None else getattr(agent, "enabled_toolsets", None) - disabled = disabled_override if disabled_override is not None else getattr(agent, "disabled_toolsets", None) + explicit_policy = enabled_override is not None or disabled_override is not None + # Policy epochs order explicit reloads against automatic rebuilds even when + # both derive from the same registry generation. Policy is still staged: + # only a successful winning refresh publishes enabled/disabled toolsets. + with _agent_tools_lock: + published_policy_epoch = getattr(agent, "_tool_published_policy_epoch", 0) + if not isinstance(published_policy_epoch, int): + published_policy_epoch = 0 + latest_policy_epoch = getattr( + agent, "_tool_policy_epoch", published_policy_epoch + ) + if not isinstance(latest_policy_epoch, int): + latest_policy_epoch = published_policy_epoch + pending_policy = getattr(agent, "_tool_pending_policy", None) + if not ( + isinstance(pending_policy, tuple) + and len(pending_policy) == 3 + and isinstance(pending_policy[0], int) + and pending_policy[0] == latest_policy_epoch + and pending_policy[0] > published_policy_epoch + ): + pending_policy = None + if _policy_epoch_override is not None: + refresh_policy_epoch = _policy_epoch_override + enabled = enabled_override + disabled = disabled_override + policy_refresh = True + elif explicit_policy: + base_enabled = ( + pending_policy[1] + if pending_policy is not None + else getattr(agent, "enabled_toolsets", None) + ) + base_disabled = ( + pending_policy[2] + if pending_policy is not None + else getattr(agent, "disabled_toolsets", None) + ) + enabled = ( + enabled_override if enabled_override is not None else base_enabled + ) + disabled = ( + disabled_override if disabled_override is not None else base_disabled + ) + refresh_policy_epoch = latest_policy_epoch + 1 + agent._tool_policy_epoch = refresh_policy_epoch + agent._tool_pending_policy = ( + refresh_policy_epoch, + enabled, + disabled, + ) + policy_refresh = True + elif pending_policy is not None: + refresh_policy_epoch, enabled, disabled = pending_policy + policy_refresh = True + else: + refresh_policy_epoch = published_policy_epoch + enabled = getattr(agent, "enabled_toolsets", None) + disabled = getattr(agent, "disabled_toolsets", None) + policy_refresh = False + + def _publish_staged_policy() -> None: agent.enabled_toolsets = enabled agent.disabled_toolsets = disabled - else: - enabled = getattr(agent, "enabled_toolsets", None) - disabled = getattr(agent, "disabled_toolsets", None) + agent._tool_published_policy_epoch = refresh_policy_epoch + pending = getattr(agent, "_tool_pending_policy", None) + if ( + isinstance(pending, tuple) + and pending + and pending[0] == refresh_policy_epoch + ): + agent._tool_pending_policy = None # Capture the registry generation this rebuild is derived from BEFORE the # (potentially slow) get_tool_definitions call. Used at publish time to @@ -6057,15 +6122,16 @@ def refresh_agent_mcp_tools( # Computed OUTSIDE the lock (get_tool_definitions can be slow); the diff and # publish below happen together in ONE critical section so two concurrent # callers can't torn-publish or compute overlapping ``added`` sets. - new_defs = list( - get_tool_definitions( - enabled_toolsets=enabled, - disabled_toolsets=disabled, - quiet_mode=quiet_mode, + try: + new_defs = list( + get_tool_definitions( + enabled_toolsets=enabled, + disabled_toolsets=disabled, + quiet_mode=quiet_mode, + ) + or [] ) - or [] - ) - new_names = {t["function"]["name"] for t in new_defs} + new_names = {t["function"]["name"] for t in new_defs} # Re-append the post-build injected families that get_tool_definitions does # NOT reproduce, so a refresh never strips them (memory-provider + context- @@ -6075,42 +6141,106 @@ def refresh_agent_mcp_tools( # (``build_api_kwargs``) can't see a partial rebuild or a cross-attribute # half-swap. ``staged_engine_names`` are the context-engine routing names # this rebuild actually appended (matching agent_init's dedup-aware add). - staged_engine_names = _reinject_post_build_tools(agent, new_defs, new_names) + ( + staged_engine_names, + staged_provider_names, + ) = _reinject_post_build_tools( + agent, + new_defs, + new_names, + enabled_toolsets=enabled, + disabled_toolsets=disabled, + strict_memory_schemas=True, + ) + except Exception: + # A failed policy rebuild remains pending at its unpublished epoch. + # The next automatic refresh retries that exact policy; a newer + # explicit policy replaces it under the lock and keeps ordering. + raise # Single atomic read-diff-publish so the returned ``added`` is consistent # with what was actually published, even under concurrent callers, and a # stale (older-generation) rebuild can't overwrite a newer published one. - with _agent_tools_lock: - # Defensive: the published generation should be an int, but tolerate an - # agent that never set it (or set a non-int, e.g. a test mock) rather - # than throwing TypeError on the comparison and silently failing the - # whole refresh. - published_gen_raw = getattr(agent, "_tool_snapshot_generation", -1) - published_gen = published_gen_raw if isinstance(published_gen_raw, int) else -1 - if snapshot_generation < published_gen: - # A newer snapshot already won; our set is stale — drop it. - return set() - current = { - t["function"]["name"] - for t in (getattr(agent, "tools", None) or []) - } - if new_names == current: - # No change → leave the live snapshot untouched (no churn), but - # record the generation so an in-flight older caller can't clobber. + class _RetryWinningRegistryGeneration(Exception): + pass + + try: + with _agent_tools_lock: + # Defensive: the published generation should be an int, but tolerate an + # agent that never set it (or set a non-int, e.g. a test mock) rather + # than throwing TypeError on the comparison and silently failing the + # whole refresh. + published_gen_raw = getattr(agent, "_tool_snapshot_generation", -1) + published_gen = published_gen_raw if isinstance(published_gen_raw, int) else -1 + if snapshot_generation < published_gen: + # A newer snapshot already won. An automatic refresh can be + # dropped, but an explicit policy must be rebuilt against the + # winning generation before its epoch can be published. + if ( + policy_refresh + and getattr(agent, "_tool_policy_epoch", None) + == refresh_policy_epoch + ): + raise _RetryWinningRegistryGeneration + return set() + if refresh_policy_epoch < getattr(agent, "_tool_policy_epoch", 0): + # A newer explicit policy reload started after this snapshot. + return set() + current = { + t["function"]["name"] + for t in (getattr(agent, "tools", None) or []) + } + current_provider_names = set( + getattr(agent, "_memory_provider_tool_names", set()) or set() + ) + provider_availability_changed = current_provider_names != staged_provider_names + if new_names == current and not provider_availability_changed: + # No name or provider-ownership change → leave the live snapshot + # untouched (no churn). Equal names with changed provider ownership + # fall through so the newly owning tool's staged schema is published. + engine_names = getattr(agent, "_context_engine_tool_names", None) + if isinstance(engine_names, set): + engine_names.clear() + engine_names.update(staged_engine_names) + agent._memory_provider_tool_names = set(staged_provider_names) + # Record the generation so an in-flight older caller can't clobber. + agent._tool_snapshot_generation = max(published_gen, snapshot_generation) + if policy_refresh: + _publish_staged_policy() + return set() + agent.tools = new_defs + agent.valid_tool_names = new_names + if policy_refresh: + _publish_staged_policy() + # Publish context-engine routing names atomically with the snapshot. + engine_names = getattr(agent, "_context_engine_tool_names", None) + if isinstance(engine_names, set): + engine_names.clear() + engine_names.update(staged_engine_names) + agent._memory_provider_tool_names = set(staged_provider_names) agent._tool_snapshot_generation = max(published_gen, snapshot_generation) - return set() - agent.tools = new_defs - agent.valid_tool_names = new_names - # Publish context-engine routing names atomically with the snapshot. - engine_names = getattr(agent, "_context_engine_tool_names", None) - if isinstance(engine_names, set): - engine_names.clear() - engine_names.update(staged_engine_names) - agent._tool_snapshot_generation = max(published_gen, snapshot_generation) - return new_names - current - - -def _reinject_post_build_tools(agent, tools_list: list, name_set: set) -> set: + if provider_availability_changed: + agent._cached_system_prompt = None + return new_names - current + except _RetryWinningRegistryGeneration: + return refresh_agent_mcp_tools( + agent, + enabled_override=enabled, + disabled_override=disabled, + quiet_mode=quiet_mode, + _policy_epoch_override=refresh_policy_epoch, + ) + + +def _reinject_post_build_tools( + agent, + tools_list: list, + name_set: set, + *, + enabled_toolsets=None, + disabled_toolsets=None, + strict_memory_schemas: bool = False, +) -> tuple[set, set]: """Append memory-provider and context-engine tools onto staged locals. Mirrors the post-``get_tool_definitions`` injection in ``agent_init`` so a @@ -6119,11 +6249,10 @@ def _reinject_post_build_tools(agent, tools_list: list, name_set: set) -> set: / ``name_set`` (never the live agent attributes) so the rebuild stays atomic. Idempotent (skips names already present) and fail-soft. - Returns the set of context-engine routing names actually appended by THIS - rebuild — matching ``agent_init``'s dedup behavior (a name already provided - by a registry/plugin tool is NOT claimed for context-engine routing). The - caller publishes this into ``agent._context_engine_tool_names`` atomically - with the snapshot. + Returns the context-engine and memory-provider names actually appended by + THIS rebuild. The sets match ``agent_init``'s dedup behavior (a name + already provided by a registry/plugin tool is NOT claimed by an injected + family). """ def _add(schema: dict) -> bool: name = schema.get("name", "") @@ -6134,22 +6263,82 @@ def _add(schema: dict) -> bool: return True # Memory-provider tools (mem0/honcho/byterover/supermemory/…). + staged_provider_names: set = set() try: memory_manager = getattr(agent, "_memory_manager", None) get_mem_schemas = getattr(memory_manager, "get_all_tool_schemas", None) if memory_manager else None if callable(get_mem_schemas): - # Honor the same toolset gate inject_memory_provider_tools uses. - from agent.memory_manager import memory_provider_tools_enabled - if memory_provider_tools_enabled( - getattr(agent, "enabled_toolsets", None), - getattr(agent, "disabled_toolsets", None), - memory_tool_present="memory" in name_set, - ): - for schema in get_mem_schemas(): - if isinstance(schema, dict): - _add(schema) + # Honor the same final toolset policy inject_memory_provider_tools uses. + from agent.memory_manager import ( + effective_memory_provider_tool_schemas, + memory_provider_tools_disabled, + normalize_tool_schema, + ) + strict_get_schemas = getattr( + memory_manager, "get_all_tool_schemas_strict", None + ) + schema_callback = ( + strict_get_schemas + if strict_memory_schemas and callable(strict_get_schemas) + else get_mem_schemas + ) + memory_selected = "memory" in name_set + if memory_provider_tools_disabled(disabled_toolsets): + # Complete family revocation is independent of provider code; + # do not let a broken plugin veto its own removal. + effective_schemas = [] + else: + try: + raw_schemas = list(schema_callback()) + normalized_schemas = [ + schema + for raw_schema in raw_schemas + if (schema := normalize_tool_schema(raw_schema)) is not None + ] + effective_schemas = effective_memory_provider_tool_schemas( + normalized_schemas, + enabled_toolsets=enabled_toolsets, + disabled_toolsets=disabled_toolsets, + memory_selected=memory_selected, + ) + except Exception: + if not strict_memory_schemas: + raise + # Tightening can safely filter the already-published + # provider contracts when fresh enumeration is unavailable. + # It may remove names, never invent or add them. + current_provider_names = set( + getattr(agent, "_memory_provider_tool_names", set()) or set() + ) + current_provider_schemas = [ + schema + for tool in (getattr(agent, "tools", None) or []) + if isinstance(tool, dict) + and ( + schema := normalize_tool_schema(tool) + ) is not None + and schema["name"] in current_provider_names + ] + effective_schemas = effective_memory_provider_tool_schemas( + current_provider_schemas, + enabled_toolsets=enabled_toolsets, + disabled_toolsets=disabled_toolsets, + memory_selected=memory_selected, + ) + effective_names = { + schema["name"] for schema in effective_schemas + } + if not effective_names < current_provider_names: + raise + for schema in effective_schemas: + name = schema.get("name", "") + if _add(schema) and name: + staged_provider_names.add(name) except Exception: + if strict_memory_schemas: + raise logger.debug("Memory-provider tool re-injection skipped", exc_info=True) + staged_provider_names = set() # Context-engine tools (lcm_grep/lcm_describe/…) — the `context_engine` # toolset is intentionally empty, so these only exist via this append. @@ -6159,8 +6348,9 @@ def _add(schema: dict) -> bool: # model latency penalty. staged_engine_names: set = set() try: - enabled = getattr(agent, "enabled_toolsets", None) - context_engine_allowed = enabled is None or "context_engine" in enabled + context_engine_allowed = ( + enabled_toolsets is None or "context_engine" in enabled_toolsets + ) compressor = getattr(agent, "context_compressor", None) get_schemas = getattr(compressor, "get_tool_schemas", None) if compressor else None if context_engine_allowed and callable(get_schemas): @@ -6176,7 +6366,7 @@ def _add(schema: dict) -> bool: except Exception: logger.debug("Context-engine tool re-injection skipped", exc_info=True) - return staged_engine_names + return staged_engine_names, staged_provider_names def shutdown_mcp_servers(): diff --git a/tui_gateway/server.py b/tui_gateway/server.py index 06c15641d97ac..b492544719b30 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -25,6 +25,7 @@ set_hermes_home_override, ) from hermes_cli.env_loader import load_hermes_dotenv +from hermes_cli.toolset_validation import normalize_toolset_names from utils import is_truthy_value from tools.environments.local import hermes_subprocess_env from agent.replay_cleanup import sanitize_replay_history @@ -5181,6 +5182,7 @@ def _background_agent_kwargs(agent, task_id: str) -> dict: "max_iterations": _cfg_max_turns(cfg, 25), "enabled_toolsets": getattr(agent, "enabled_toolsets", None) or _load_enabled_toolsets(), + "disabled_toolsets": getattr(agent, "disabled_toolsets", None), "quiet_mode": True, "verbose_logging": False, "ephemeral_system_prompt": getattr(agent, "ephemeral_system_prompt", None) @@ -5673,6 +5675,9 @@ def _make_agent( else _load_service_tier() ), enabled_toolsets=_load_enabled_toolsets(), + disabled_toolsets=normalize_toolset_names( + agent_cfg.get("disabled_toolsets") + ), # OpenRouter provider-routing prefs (config.yaml `provider_routing`). # Mirrors the messaging gateway + CLI so the desktop/TUI honors the same # routing instead of letting OpenRouter pick providers at random. @@ -14370,6 +14375,8 @@ def _finish_reload(rid, params: dict, *, coalesced: bool) -> dict: @method("reload.mcp") def _(rid, params: dict) -> dict: + global _mcp_reload_gen, _mcp_reload_loaded_rev + session = _sessions.get(params.get("session_id", "")) try: # Gate: /reload-mcp invalidates the prompt cache for this session. @@ -14418,7 +14425,60 @@ def _(rid, params: dict) -> dict: ) except Exception as exc: return _err(rid, 5019, f"compute-host reload_mcp failed: {exc}") - return _ok(rid, {"status": "reloaded", "turn_isolation": True, "host_ack": ack}) + + if not isinstance(ack, dict): + return _err( + rid, + 5019, + "compute-host reload_mcp returned a malformed acknowledgement", + ) + ack_type = ack.get("type") + if ack_type == "control.error": + message = str(ack.get("message") or "compute-host reload_mcp failed") + return _err(rid, 5019, message) + if ack_type != "reload_mcp.ack": + return _err( + rid, + 5019, + f"compute-host reload_mcp returned unexpected acknowledgement type: {ack_type!r}", + ) + + child_response = ack.get("response") + if not isinstance(child_response, dict): + return _err(rid, 5019, "compute-host reload_mcp returned a malformed response") + if "error" in child_response: + child_error = child_response.get("error") + if isinstance(child_error, dict): + message = str(child_error.get("message") or child_error) + else: + message = str(child_error or "child reload returned a JSON-RPC error") + return _err(rid, 5019, f"compute-host reload_mcp failed: {message}") + + child_result = child_response.get("result") + if ( + not isinstance(child_result, dict) + or child_result.get("status") != "reloaded" + ): + return _err( + rid, + 5019, + "compute-host reload_mcp did not report a completed reload", + ) + + loaded_rev = str(child_result.get("loaded_rev") or "") + with _mcp_reload_lock: + _mcp_reload_loaded_rev = loaded_rev + _mcp_reload_gen += 1 + + return _ok( + rid, + { + "status": "reloaded", + "loaded_rev": loaded_rev, + "turn_isolation": True, + "host_ack": ack, + }, + ) from tools.mcp_tool import shutdown_mcp_servers, discover_mcp_tools @@ -14432,25 +14492,19 @@ def _refresh_session_agent() -> None: if not session: return agent = session["agent"] - try: - from tools.mcp_tool import refresh_agent_mcp_tools + from tools.mcp_tool import refresh_agent_mcp_tools - # Explicit reload: re-resolve enabled toolsets so a server the - # user just enabled in config this session is picked up. - refresh_agent_mcp_tools( - agent, - enabled_override=_load_enabled_toolsets(), - quiet_mode=True, - ) - except Exception as _exc: - logger.warning( - "Failed to refresh cached agent tools after /reload-mcp: %s", - _exc, - ) + # Explicit reload: re-resolve enabled toolsets so a server the + # user just enabled in config this session is picked up. A failed + # snapshot rebuild must propagate: the reload is not complete and + # its generation/revision must not be acknowledged. + refresh_agent_mcp_tools( + agent, + enabled_override=_load_enabled_toolsets(), + quiet_mode=True, + ) _emit("session.info", params.get("session_id", ""), _session_info(agent, session)) - global _mcp_reload_gen, _mcp_reload_loaded_rev - # The revision the CALLER is asking to load (the mcp_rev its poll # observed). Empty on legacy clients and manual /reload-mcp — those # coalesce on generation alone, as before.