From 0193de53474ada6836fdb21bd92f92a892ccb578 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=D0=9A=D0=B8=D1=80=D0=B8=D0=BB=D0=BB=20=D0=92=D0=B5=D1=87?= =?UTF-8?q?=D0=BA=D0=B0=D1=81=D0=BE=D0=B2?= Date: Fri, 26 Jun 2026 15:41:25 +0200 Subject: [PATCH 1/2] feat(gateway): per-topic profile isolation for forum topics MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Bind a chat/forum topic to a named profile (`/profile `) and run the entire turn scoped to that profile's `HERMES_HOME` — its own model, SOUL.md, memory, session history, skills, MCP servers, and credentials. Topics with no binding keep current behaviour exactly (zero regression). Mechanism: - `_profile_runtime_scope`: per-turn context that layers a home override plus a profile secret-scope (no os.environ mutation), restored on exit. - `RoutingSessionStoreProxy` / `RoutingSessionDBProxy`: resolve the active profile home on every call so session history persists to the profile's state.db even from background tasks. - `.hermes_profile.json` identity marker: validated when resolving a profile home; auto-migrated if missing and logs a loud warning instead of silently falling back to the global home. - MCP servers are fingerprinted per profile so same-named servers with different credentials get separate connections; file tools are hard-guarded to the profile home; subprocess/cron inherit the scoped home. - Per-topic `/model` and `/profile` bindings persist across `/new` and restart (`topic_models.json` / `topic_profiles.json`). Adds tests covering routing, session/DB isolation, MCP and skills isolation, and subprocess home isolation. --- agent/auxiliary_client.py | 507 +++++++-- agent/skill_commands.py | 29 +- cron/jobs.py | 92 +- gateway/run.py | 1000 +++++++++++++---- gateway/session.py | 13 +- gateway/session_context.py | 10 + gateway/slash_commands.py | 124 +- hermes_cli/auth.py | 144 ++- hermes_cli/env_loader.py | 9 +- hermes_cli/main.py | 30 +- hermes_cli/profiles.py | 444 +++++++- hermes_cli/runtime_provider.py | 182 ++- hermes_cli/subcommands/profile.py | 21 + hermes_state.py | 4 +- model_tools.py | 4 +- tests/agent/test_auxiliary_client.py | 14 +- .../test_model_command_flat_string_config.py | 16 +- tests/gateway/test_model_picker_persist.py | 2 +- .../test_profile_isolation_rework_suite.py | 440 ++++++++ tests/gateway/test_reasoning_command.py | 3 +- tests/gateway/test_telegram_topic_mode.py | 167 +++ tests/gateway/test_topic_profile_routing.py | 238 ++++ tests/gateway/test_update_command.py | 3 +- tests/hermes_cli/test_profiles.py | 248 ++++ .../test_runtime_provider_resolution.py | 6 +- tests/tools/test_mcp_profile_isolation.py | 109 ++ tests/tools/test_process_registry.py | 33 + tests/tools/test_skills_profile_isolation.py | 101 ++ tools/file_tools.py | 199 +++- tools/mcp_tool.py | 291 +++-- tools/memory_tool.py | 8 +- tools/process_registry.py | 174 ++- tools/registry.py | 25 +- tools/session_search_tool.py | 6 +- tools/skill_manager_tool.py | 27 +- tools/skills_tool.py | 38 +- tools/slash_confirm.py | 38 +- tools/terminal_tool.py | 8 + 38 files changed, 4198 insertions(+), 609 deletions(-) create mode 100644 tests/gateway/test_profile_isolation_rework_suite.py create mode 100644 tests/gateway/test_topic_profile_routing.py create mode 100644 tests/tools/test_mcp_profile_isolation.py create mode 100644 tests/tools/test_skills_profile_isolation.py diff --git a/agent/auxiliary_client.py b/agent/auxiliary_client.py index a4fe065b6d386..b195ff4609640 100644 --- a/agent/auxiliary_client.py +++ b/agent/auxiliary_client.py @@ -70,6 +70,13 @@ _OPENAI_CLS_CACHE: Optional[type] = None +def _env_lookup(env: Optional[Dict[str, Any]], key: str, default: str = "") -> str: + if env is None: + return os.getenv(key, default) + value = env.get(key, default) + return "" if value is None else str(value) + + def _load_openai_cls() -> type: """Import and cache ``openai.OpenAI``.""" global _OPENAI_CLS_CACHE @@ -522,8 +529,6 @@ def _nous_extra_body() -> dict: _NOUS_MODEL = "google/gemini-3-flash-preview" _NOUS_DEFAULT_BASE_URL = "https://inference-api.nousresearch.com/v1" _ANTHROPIC_DEFAULT_BASE_URL = "https://api.anthropic.com" -_AUTH_JSON_PATH = get_hermes_home() / "auth.json" - # Codex OAuth endpoint used when a caller explicitly requests # provider="openai-codex". There is deliberately no hardcoded default # model: the set of models OpenAI accepts on this endpoint for @@ -604,10 +609,67 @@ def _to_openai_base_url(base_url: str) -> str: return url -def _select_pool_entry(provider: str) -> Tuple[bool, Optional[Any]]: +def _scoped_auth_store_allowed(env: Optional[Dict[str, Any]]) -> bool: + if env is None: + return True + try: + from gateway.session_context import get_session_env + + if get_session_env("HERMES_SESSION_AGENT_HERMES_HOME", "").strip(): + return True + if get_session_env("HERMES_SESSION_AGENT_PROFILE", "").strip(): + return False + except Exception: + pass + try: + from hermes_constants import get_default_hermes_root + + profile_home = get_hermes_home().resolve(strict=False) + default_home = get_default_hermes_root().resolve(strict=False) + return profile_home != default_home + except Exception: + return False + + +def _load_scoped_auth_store(env: Optional[Dict[str, Any]]) -> Dict[str, Any]: + if not _scoped_auth_store_allowed(env): + return {} + auth_path = get_hermes_home() / "auth.json" + try: + if not auth_path.is_file(): + return {} + data = json.loads(auth_path.read_text()) + return data if isinstance(data, dict) else {} + except Exception as exc: + logger.debug("Auxiliary client: could not read scoped auth store %s: %s", auth_path, exc) + return {} + + +def _load_scoped_pool(provider: str, env: Optional[Dict[str, Any]]): + if env is None: + return load_pool(provider) + store = _load_scoped_auth_store(env) + pool_data = store.get("credential_pool") + if not isinstance(pool_data, dict): + raw_entries = [] + else: + raw_entries = pool_data.get(provider) + if not isinstance(raw_entries, list): + raw_entries = [] + try: + from agent.credential_pool import CredentialPool, PooledCredential + + entries = [PooledCredential.from_dict(provider, payload) for payload in raw_entries] + return CredentialPool(provider, entries) + except Exception as exc: + logger.debug("Auxiliary client: could not load scoped pool for %s: %s", provider, exc) + return None + + +def _select_pool_entry(provider: str, env: Optional[Dict[str, Any]] = None) -> Tuple[bool, Optional[Any]]: """Return (pool_exists_for_provider, selected_entry).""" try: - pool = load_pool(provider) + pool = _load_scoped_pool(provider, env) except Exception as exc: logger.debug("Auxiliary client: could not load pool for %s: %s", provider, exc) return False, None @@ -620,10 +682,10 @@ def _select_pool_entry(provider: str) -> Tuple[bool, Optional[Any]]: return True, None -def _peek_pool_entry(provider: str) -> Optional[Any]: +def _peek_pool_entry(provider: str, env: Optional[Dict[str, Any]] = None) -> Optional[Any]: """Best-effort current/next pool entry without mutating selection order.""" try: - pool = load_pool(provider) + pool = _load_scoped_pool(provider, env) except Exception as exc: logger.debug("Auxiliary client: could not load pool for %s (peek): %s", provider, exc) return None @@ -1305,13 +1367,13 @@ def _maybe_wrap_anthropic( ) -def _read_nous_auth() -> Optional[dict]: +def _read_nous_auth(env: Optional[Dict[str, Any]] = None) -> Optional[dict]: """Read and validate ~/.hermes/auth.json for an active Nous provider. Returns the provider state dict if Nous is active with tokens, otherwise None. """ - pool_present, entry = _select_pool_entry("nous") + pool_present, entry = _select_pool_entry("nous", env=env) if pool_present: if entry is None: return None @@ -1328,9 +1390,12 @@ def _read_nous_auth() -> Optional[dict]: } try: - if not _AUTH_JSON_PATH.is_file(): + if env is not None and not _scoped_auth_store_allowed(env): + return None + auth_json_path = get_hermes_home() / "auth.json" + if not auth_json_path.is_file(): return None - data = json.loads(_AUTH_JSON_PATH.read_text()) + data = json.loads(auth_json_path.read_text()) if data.get("active_provider") != "nous": return None provider = data.get("providers", {}).get("nous", {}) @@ -1363,9 +1428,9 @@ def _nous_api_key(provider: dict) -> str: return "" -def _nous_base_url() -> str: +def _nous_base_url(env: Optional[Dict[str, Any]] = None) -> str: """Resolve the Nous inference base URL from env or default.""" - return os.getenv("NOUS_INFERENCE_BASE_URL", _NOUS_DEFAULT_BASE_URL) + return _env_lookup(env, "NOUS_INFERENCE_BASE_URL", _NOUS_DEFAULT_BASE_URL) def _resolve_nous_pool_runtime_api(*, force_refresh: bool = False) -> Optional[tuple[str, str]]: @@ -1449,7 +1514,7 @@ def _resolve_nous_runtime_api(*, force_refresh: bool = False) -> Optional[tuple[ return api_key, base_url -def _resolve_xai_oauth_for_aux() -> Optional[Tuple[str, str]]: +def _resolve_xai_oauth_for_aux(env: Optional[Dict[str, Any]] = None) -> Optional[Tuple[str, str]]: """Resolve a fresh xAI OAuth (api_key, base_url) for auxiliary clients. Prefer the credential pool, matching the main runtime/provider status @@ -1468,7 +1533,7 @@ def _resolve_xai_oauth_for_aux() -> Optional[Tuple[str, str]]: _xai_validate_inference_base_url, ) - pool = load_pool("xai-oauth") + pool = _load_scoped_pool("xai-oauth", env) if pool and pool.has_credentials(): entry = pool.select() if entry is not None: @@ -1478,8 +1543,8 @@ def _resolve_xai_oauth_for_aux() -> Optional[Tuple[str, str]]: or "" ).strip() base_url = _xai_validate_inference_base_url( - os.getenv("HERMES_XAI_BASE_URL", "").strip().rstrip("/") - or os.getenv("XAI_BASE_URL", "").strip().rstrip("/") + _env_lookup(env, "HERMES_XAI_BASE_URL").strip().rstrip("/") + or _env_lookup(env, "XAI_BASE_URL").strip().rstrip("/") or str(getattr(entry, "runtime_base_url", None) or "").strip().rstrip("/") or str(getattr(entry, "base_url", None) or "").strip().rstrip("/"), fallback=DEFAULT_XAI_OAUTH_BASE_URL, @@ -1492,6 +1557,8 @@ def _resolve_xai_oauth_for_aux() -> Optional[Tuple[str, str]]: try: from hermes_cli.auth import resolve_xai_oauth_runtime_credentials + if env is not None and not _scoped_auth_store_allowed(env): + return None creds = resolve_xai_oauth_runtime_credentials() except Exception as exc: logger.debug("Auxiliary xAI OAuth runtime credential resolution failed: %s", exc) @@ -1504,7 +1571,7 @@ def _resolve_xai_oauth_for_aux() -> Optional[Tuple[str, str]]: return api_key, base_url -def _read_codex_access_token() -> Optional[str]: +def _read_codex_access_token(env: Optional[Dict[str, Any]] = None) -> Optional[str]: """Read a valid, non-expired Codex OAuth access token from Hermes auth store. If a credential pool exists but currently has no selectable runtime entry @@ -1513,13 +1580,15 @@ def _read_codex_access_token() -> Optional[str]: fallback-to-Codex working when the pool state is stale but the stored OAuth token is still valid. """ - pool_present, entry = _select_pool_entry("openai-codex") + pool_present, entry = _select_pool_entry("openai-codex", env=env) if pool_present: token = _pool_runtime_api_key(entry) if token: return token try: + if env is not None and not _scoped_auth_store_allowed(env): + return None from hermes_cli.auth import _read_codex_tokens data = _read_codex_tokens() tokens = data.get("tokens", {}) @@ -1547,7 +1616,7 @@ def _read_codex_access_token() -> Optional[str]: return None -def _resolve_api_key_provider() -> Tuple[Optional[OpenAI], Optional[str]]: +def _resolve_api_key_provider(env: Optional[Dict[str, Any]] = None) -> Tuple[Optional[OpenAI], Optional[str]]: """Try each API-key provider in PROVIDER_REGISTRY order. Returns (client, model) for the first provider with usable runtime @@ -1575,9 +1644,9 @@ def _resolve_api_key_provider() -> Tuple[Optional[OpenAI], Optional[str]]: continue except ImportError: pass - return _try_anthropic() + return _try_anthropic(env=env) - pool_present, entry = _select_pool_entry(provider_id) + pool_present, entry = _select_pool_entry(provider_id, env=env) if pool_present: api_key = _pool_runtime_api_key(entry) if not api_key: @@ -1618,7 +1687,7 @@ def _resolve_api_key_provider() -> Tuple[Optional[OpenAI], Optional[str]]: _client = _maybe_wrap_anthropic(_client, model, api_key, raw_base_url) return _client, model - creds = resolve_api_key_provider_credentials(provider_id) + creds = resolve_api_key_provider_credentials(provider_id, env=env) api_key = str(creds.get("api_key", "")).strip() if not api_key: continue @@ -1665,8 +1734,12 @@ def _resolve_api_key_provider() -> Tuple[Optional[OpenAI], Optional[str]]: -def _try_openrouter(explicit_api_key: str = None, model: str = None) -> Tuple[Optional[OpenAI], Optional[str]]: - pool_present, entry = _select_pool_entry("openrouter") +def _try_openrouter( + explicit_api_key: str = None, + model: str = None, + env: Optional[Dict[str, Any]] = None, +) -> Tuple[Optional[OpenAI], Optional[str]]: + pool_present, entry = _select_pool_entry("openrouter", env=env) if pool_present: or_key = explicit_api_key or _pool_runtime_api_key(entry) if not or_key: @@ -1677,7 +1750,7 @@ def _try_openrouter(explicit_api_key: str = None, model: str = None) -> Tuple[Op return OpenAI(api_key=or_key, base_url=base_url, default_headers=build_or_headers()), model or _OPENROUTER_MODEL - or_key = explicit_api_key or os.getenv("OPENROUTER_API_KEY") + or_key = explicit_api_key or _env_lookup(env, "OPENROUTER_API_KEY").strip() if not or_key: _mark_provider_unhealthy("openrouter", ttl=60) return None, None @@ -1686,20 +1759,23 @@ def _try_openrouter(explicit_api_key: str = None, model: str = None) -> Tuple[Op default_headers=build_or_headers()), model or _OPENROUTER_MODEL -def _describe_openrouter_unavailable() -> str: +def _describe_openrouter_unavailable(env: Optional[Dict[str, Any]] = None) -> str: """Return a more precise OpenRouter auth failure reason for logs.""" - pool_present, entry = _select_pool_entry("openrouter") + pool_present, entry = _select_pool_entry("openrouter", env=env) if pool_present: if entry is None: return "OpenRouter credential pool has no usable entries (credentials may be exhausted)" if not _pool_runtime_api_key(entry): return "OpenRouter credential pool entry is missing a runtime API key" - if not str(os.getenv("OPENROUTER_API_KEY") or "").strip(): + if not _env_lookup(env, "OPENROUTER_API_KEY").strip(): return "OPENROUTER_API_KEY not set" return "no usable OpenRouter credentials found" -def _try_nous(vision: bool = False) -> Tuple[Optional[OpenAI], Optional[str]]: +def _try_nous( + vision: bool = False, + env: Optional[Dict[str, Any]] = None, +) -> Tuple[Optional[OpenAI], Optional[str]]: # Check cross-session rate limit guard before attempting Nous — # if another session already recorded a 429, skip Nous entirely # to avoid piling more requests onto the tapped RPH bucket. @@ -1716,8 +1792,12 @@ def _try_nous(vision: bool = False) -> Tuple[Optional[OpenAI], Optional[str]]: except Exception: pass - nous = _read_nous_auth() - runtime = _resolve_nous_runtime_api(force_refresh=False) + nous = _read_nous_auth(env=env) + runtime = ( + _resolve_nous_runtime_api(force_refresh=False) + if env is None or nous + else None + ) if runtime is None and not nous: logger.warning( "Auxiliary Nous client unavailable: no Nous authentication found " @@ -1773,7 +1853,7 @@ def _try_nous(vision: bool = False) -> Tuple[Optional[OpenAI], Optional[str]]: ) _mark_provider_unhealthy("nous", ttl=60) return None, None - base_url = str((nous or {}).get("inference_base_url") or _nous_base_url()).rstrip("/") + base_url = str((nous or {}).get("inference_base_url") or _nous_base_url(env)).rstrip("/") return ( OpenAI( api_key=api_key, @@ -1927,7 +2007,7 @@ def clear_runtime_main() -> None: _RUNTIME_MAIN_API_MODE = "" -def _resolve_custom_runtime() -> Tuple[Optional[str], Optional[str], Optional[str]]: +def _resolve_custom_runtime(env: Optional[Dict[str, Any]] = None) -> Tuple[Optional[str], Optional[str], Optional[str]]: """Resolve the active custom/main endpoint the same way the main CLI does. This covers both env-driven OPENAI_BASE_URL setups and config-saved custom @@ -1937,14 +2017,14 @@ def _resolve_custom_runtime() -> Tuple[Optional[str], Optional[str], Optional[st try: from hermes_cli.runtime_provider import resolve_runtime_provider - runtime = resolve_runtime_provider(requested="custom") + runtime = resolve_runtime_provider(requested="custom", env=env) except Exception as exc: logger.debug("Auxiliary client: custom runtime resolution failed: %s", exc) runtime = None if not isinstance(runtime, dict): - openai_base = os.getenv("OPENAI_BASE_URL", "").strip().rstrip("/") - openai_key = os.getenv("OPENAI_API_KEY", "").strip() + openai_base = _env_lookup(env, "OPENAI_BASE_URL").strip().rstrip("/") + openai_key = _env_lookup(env, "OPENAI_API_KEY").strip() if not openai_base: return None, None, None runtime = { @@ -1977,8 +2057,8 @@ def _resolve_custom_runtime() -> Tuple[Optional[str], Optional[str], Optional[st return custom_base, custom_key.strip(), custom_mode -def _current_custom_base_url() -> str: - custom_base, _, _ = _resolve_custom_runtime() +def _current_custom_base_url(env: Optional[Dict[str, Any]] = None) -> str: + custom_base, _, _ = _resolve_custom_runtime(env=env) return custom_base or "" @@ -2029,8 +2109,8 @@ def _validate_base_url(base_url: str) -> None: ) from exc -def _try_custom_endpoint() -> Tuple[Optional[Any], Optional[str]]: - runtime = _resolve_custom_runtime() +def _try_custom_endpoint(env: Optional[Dict[str, Any]] = None) -> Tuple[Optional[Any], Optional[str]]: + runtime = _resolve_custom_runtime(env=env) if len(runtime) == 2: custom_base, custom_key = runtime custom_mode = None @@ -2080,7 +2160,10 @@ def _try_custom_endpoint() -> Tuple[Optional[Any], Optional[str]]: return _fallback_client, model -def _build_xai_oauth_aux_client(model: str) -> Tuple[Optional[Any], Optional[str]]: +def _build_xai_oauth_aux_client( + model: str, + env: Optional[Dict[str, Any]] = None, +) -> Tuple[Optional[Any], Optional[str]]: """Build a CodexAuxiliaryClient for an xAI Grok OAuth-authenticated session. xAI's ``/v1/responses`` endpoint speaks the OpenAI Responses API, so we @@ -2097,7 +2180,7 @@ def _build_xai_oauth_aux_client(model: str) -> Tuple[Optional[Any], Optional[str "pass model explicitly (auxiliary..model in config.yaml)." ) return None, None - resolved = _resolve_xai_oauth_for_aux() + resolved = _resolve_xai_oauth_for_aux(env=env) if resolved is None: return None, None api_key, base_url = resolved @@ -2106,7 +2189,10 @@ def _build_xai_oauth_aux_client(model: str) -> Tuple[Optional[Any], Optional[str return CodexAuxiliaryClient(real_client, model), model -def _build_codex_client(model: str) -> Tuple[Optional[Any], Optional[str]]: +def _build_codex_client( + model: str, + env: Optional[Dict[str, Any]] = None, +) -> Tuple[Optional[Any], Optional[str]]: """Build a CodexAuxiliaryClient for an explicitly-requested model. There is no auto-selection of the Codex model: the ChatGPT-account @@ -2123,18 +2209,18 @@ def _build_codex_client(model: str) -> Tuple[Optional[Any], Optional[str]]: "pass model explicitly (auxiliary..model in config.yaml)." ) return None, None - pool_present, entry = _select_pool_entry("openai-codex") + pool_present, entry = _select_pool_entry("openai-codex", env=env) if pool_present: codex_token = _pool_runtime_api_key(entry) if codex_token: base_url = _pool_runtime_base_url(entry, _CODEX_AUX_BASE_URL) or _CODEX_AUX_BASE_URL else: - codex_token = _read_codex_access_token() + codex_token = _read_codex_access_token(env=env) if not codex_token: return None, None base_url = _CODEX_AUX_BASE_URL else: - codex_token = _read_codex_access_token() + codex_token = _read_codex_access_token(env=env) if not codex_token: return None, None base_url = _CODEX_AUX_BASE_URL @@ -2153,6 +2239,7 @@ def _try_azure_foundry( explicit_api_key: Optional[str] = None, explicit_base_url: Optional[str] = None, api_mode: Optional[str] = None, + env: Optional[Dict[str, Any]] = None, ) -> Tuple[Optional[Any], Optional[str]]: """Resolve an Azure Foundry auxiliary client via the runtime resolver. @@ -2198,6 +2285,7 @@ def _try_azure_foundry( explicit_api_key=explicit_api_key, explicit_base_url=explicit_base_url, target_model=model, + env=env, ) except AuthError as exc: logger.debug("Auxiliary azure-foundry: %s", exc) @@ -2261,20 +2349,35 @@ def _try_azure_foundry( return client, final_model -def _try_anthropic(explicit_api_key: str = None) -> Tuple[Optional[Any], Optional[str]]: +def _try_anthropic( + explicit_api_key: str = None, + env: Optional[Dict[str, Any]] = None, +) -> Tuple[Optional[Any], Optional[str]]: try: from agent.anthropic_adapter import build_anthropic_client, resolve_anthropic_token except ImportError: return None, None - pool_present, entry = _select_pool_entry("anthropic") + pool_present, entry = _select_pool_entry("anthropic", env=env) if pool_present: if entry is None: return None, None token = explicit_api_key or _pool_runtime_api_key(entry) else: entry = None - token = explicit_api_key or resolve_anthropic_token() + token = explicit_api_key + if not token: + if env is None: + token = resolve_anthropic_token() + else: + try: + from hermes_cli.auth import PROVIDER_REGISTRY + for env_name in PROVIDER_REGISTRY["anthropic"].api_key_env_vars: + token = _env_lookup(env, env_name).strip() + if token: + break + except Exception: + token = "" if not token: return None, None @@ -2352,6 +2455,47 @@ def _normalize_main_runtime(main_runtime: Optional[Dict[str, Any]]) -> Dict[str, return normalized +def _runtime_env_cache_key(env: Optional[Dict[str, Any]]) -> tuple: + if env is None: + return () + try: + from hermes_constants import get_hermes_home + home = str(get_hermes_home()) + except Exception: + home = "" + items = sorted((str(k), "" if v is None else str(v)) for k, v in env.items()) + payload = json.dumps(items, separators=(",", ":"), ensure_ascii=True) + digest = hashlib.sha256(payload.encode("utf-8")).hexdigest() + return (home, digest) + + +def _effective_runtime_env(env: Optional[Dict[str, Any]]) -> Optional[Dict[str, Any]]: + if env is not None: + return env + try: + from gateway.session_context import get_runtime_env + runtime_env = get_runtime_env() + if runtime_env is not None: + return runtime_env + except Exception: + pass + try: + from gateway.session_context import get_session_env + profile_home = get_session_env("HERMES_SESSION_AGENT_HERMES_HOME", "").strip() + profile = get_session_env("HERMES_SESSION_AGENT_PROFILE", "").strip() + if profile_home: + from hermes_cli.env_loader import read_hermes_dotenv_values + runtime_env = read_hermes_dotenv_values(hermes_home=Path(profile_home)) + runtime_env["HERMES_PROFILE_STRICT_AUTH"] = "1" + runtime_env["HERMES_HOME"] = profile_home + return runtime_env + if profile: + return {"HERMES_PROFILE_STRICT_AUTH": "1"} + except Exception: + pass + return None + + def _get_provider_chain() -> List[tuple]: """Return the ordered provider detection chain. @@ -2373,6 +2517,15 @@ def _get_provider_chain() -> List[tuple]: ] +_ENV_AWARE_PROVIDER_CHAIN_LABELS = {"openrouter", "nous", "local/custom", "api-key"} + + +def _try_provider_chain_entry(label: str, try_fn, env: Optional[Dict[str, Any]] = None): + if env is None or label not in _ENV_AWARE_PROVIDER_CHAIN_LABELS: + return try_fn() + return try_fn(env=env) + + # ── Auxiliary "recently 402'd" unhealthy-provider cache ──────────────────── # # When an auxiliary provider returns HTTP 402 (Payment Required / credit @@ -2862,6 +3015,7 @@ def _pool_cache_hint( provider: str, *, main_runtime: Optional[Dict[str, Any]] = None, + env: Optional[Dict[str, Any]] = None, ) -> str: """Return a stable cache discriminator for pooled providers.""" normalized = _normalize_aux_provider(provider) @@ -2870,7 +3024,7 @@ def _pool_cache_hint( normalized = _normalize_aux_provider(runtime.get("provider") or _read_main_provider()) if normalized in {"", "auto", "custom"}: return "" - entry = _peek_pool_entry(normalized) + entry = _peek_pool_entry(normalized, env=env) if entry is None: return "" entry_id = str(getattr(entry, "id", "") or "").strip() @@ -2930,7 +3084,13 @@ def _recoverable_pool_provider( return None -def _recover_provider_pool(provider: str, exc: Exception, *, failed_api_key: str = "") -> bool: +def _recover_provider_pool( + provider: str, + exc: Exception, + *, + failed_api_key: str = "", + env: Optional[Dict[str, Any]] = None, +) -> bool: """Try same-provider credential-pool recovery for auxiliary calls. ``failed_api_key`` is the API key that was actually used for the failing @@ -2940,7 +3100,7 @@ def _recover_provider_pool(provider: str, exc: Exception, *, failed_api_key: str """ normalized = _normalize_aux_provider(provider) try: - pool = load_pool(normalized) + pool = _load_scoped_pool(normalized, env) except Exception as load_exc: logger.debug("Auxiliary client: could not load pool for %s recovery: %s", normalized, load_exc) return False @@ -2988,6 +3148,7 @@ def _retry_same_provider_sync( resolved_api_key: Optional[str], resolved_api_mode: Optional[str], main_runtime: Optional[Dict[str, Any]], + env: Optional[Dict[str, Any]], final_model: Optional[str], messages: list, temperature: Optional[float], @@ -3003,6 +3164,7 @@ def _retry_same_provider_sync( base_url=resolved_base_url, api_key=resolved_api_key, async_mode=False, + env=env, ) else: retry_client, retry_model = _get_cached_client( @@ -3012,6 +3174,7 @@ def _retry_same_provider_sync( api_key=resolved_api_key, api_mode=resolved_api_mode, main_runtime=main_runtime, + env=env, ) if retry_client is None: raise RuntimeError( @@ -3029,6 +3192,7 @@ def _retry_same_provider_sync( timeout=effective_timeout, extra_body=effective_extra_body, base_url=retry_base or resolved_base_url, + env=env, ) if _is_anthropic_compat_endpoint(resolved_provider, retry_base): retry_kwargs["messages"] = _convert_openai_images_to_anthropic(retry_kwargs["messages"]) @@ -3045,6 +3209,8 @@ async def _retry_same_provider_async( resolved_base_url: Optional[str], resolved_api_key: Optional[str], resolved_api_mode: Optional[str], + main_runtime: Optional[Dict[str, Any]], + env: Optional[Dict[str, Any]], final_model: Optional[str], messages: list, temperature: Optional[float], @@ -3060,6 +3226,7 @@ async def _retry_same_provider_async( base_url=resolved_base_url, api_key=resolved_api_key, async_mode=True, + env=env, ) else: retry_client, retry_model = _get_cached_client( @@ -3069,6 +3236,8 @@ async def _retry_same_provider_async( base_url=resolved_base_url, api_key=resolved_api_key, api_mode=resolved_api_mode, + main_runtime=main_runtime, + env=env, ) if retry_client is None: raise RuntimeError( @@ -3086,6 +3255,7 @@ async def _retry_same_provider_async( timeout=effective_timeout, extra_body=effective_extra_body, base_url=retry_base or resolved_base_url, + env=env, ) if _is_anthropic_compat_endpoint(resolved_provider, retry_base): retry_kwargs["messages"] = _convert_openai_images_to_anthropic(retry_kwargs["messages"]) @@ -3094,13 +3264,18 @@ async def _retry_same_provider_async( ) -def _refresh_provider_credentials(provider: str) -> bool: +def _refresh_provider_credentials( + provider: str, + env: Optional[Dict[str, Any]] = None, +) -> bool: """Refresh short-lived credentials for OAuth-backed auxiliary providers.""" normalized = _normalize_aux_provider(provider) try: if normalized == "openai-codex": from hermes_cli.auth import resolve_codex_runtime_credentials + if env is not None and not _scoped_auth_store_allowed(env): + return False creds = resolve_codex_runtime_credentials(force_refresh=True) if not str(creds.get("api_key", "") or "").strip(): return False @@ -3109,6 +3284,8 @@ def _refresh_provider_credentials(provider: str) -> bool: if normalized == "nous": from hermes_cli.auth import resolve_nous_runtime_credentials + if env is not None and _read_nous_auth(env=env) is None: + return False creds = resolve_nous_runtime_credentials( timeout_seconds=env_float("HERMES_NOUS_TIMEOUT_SECONDS", 15), force_refresh=True, @@ -3120,9 +3297,9 @@ def _refresh_provider_credentials(provider: str) -> bool: if normalized == "anthropic": from agent.anthropic_adapter import read_claude_code_credentials, _refresh_oauth_token, resolve_anthropic_token - creds = read_claude_code_credentials() + creds = read_claude_code_credentials() if env is None else None token = _refresh_oauth_token(creds) if isinstance(creds, dict) and creds.get("refreshToken") else None - if not str(token or "").strip(): + if not str(token or "").strip() and env is None: token = resolve_anthropic_token() if not str(token or "").strip(): return False @@ -3156,6 +3333,7 @@ def _try_payment_fallback( failed_provider: str, task: str = None, reason: str = "payment error", + env: Optional[Dict[str, Any]] = None, ) -> Tuple[Optional[Any], Optional[str], str]: """Try alternative providers after a payment/credit or connection error. @@ -3187,7 +3365,7 @@ def _try_payment_fallback( _log_skip_unhealthy(label, task) tried.append(f"{label} (unhealthy)") continue - client, model = try_fn() + client, model = _try_provider_chain_entry(label, try_fn, env=env) if client is not None: logger.info( "Auxiliary %s: %s on %s — falling back to %s (%s)", @@ -3207,6 +3385,7 @@ def _try_main_agent_model_fallback( failed_provider: str, task: str = None, reason: str = "error", + env: Optional[Dict[str, Any]] = None, ) -> Tuple[Optional[Any], Optional[str], str]: """Last-resort fallback to the user's main agent provider + model. @@ -3237,6 +3416,7 @@ def _try_main_agent_model_fallback( try: client, resolved_model = resolve_provider_client( provider=main_provider, model=main_model, + env=env, ) except Exception: client, resolved_model = None, None @@ -3338,6 +3518,7 @@ def _try_configured_fallback_chain( task: str, failed_provider: str, reason: str = "error", + env: Optional[Dict[str, Any]] = None, ) -> Tuple[Optional[Any], Optional[str], str]: """Try user-configured fallback_chain for a specific auxiliary task. @@ -3371,17 +3552,27 @@ def _try_configured_fallback_chain( label = f"fallback_chain[{i}]({fb_provider})" try: - fb_client, resolved_model = _resolve_fallback_entry(entry) + try: + fb_client, resolved_model = _resolve_fallback_entry(entry, env=env) + except TypeError as te: + if "unexpected keyword argument" in str(te) or "takes" in str(te): + fb_client, resolved_model = _resolve_fallback_entry(entry) + else: + raise except Exception: fb_client, resolved_model = None, None if fb_client is not None: if min_ctx is not None and resolved_model: + try: + fb_key = _fallback_entry_api_key(entry, env=env) + except TypeError: + fb_key = _fallback_entry_api_key(entry) fb_ctx = _candidate_context_window( fb_provider, resolved_model, base_url=str(entry.get("base_url") or ""), - api_key=_fallback_entry_api_key(entry) or "", + api_key=fb_key or "", ) if fb_ctx is not None and fb_ctx < min_ctx: logger.info( @@ -3408,6 +3599,7 @@ def _try_configured_fallback_chain( def _try_configured_fallback_for_unavailable_client( task: Optional[str], failed_provider: str, + env: Optional[Dict[str, Any]] = None, ) -> Tuple[Optional[Any], Optional[str], str]: """Try task fallback_chain when an explicit aux provider cannot build. @@ -3420,32 +3612,37 @@ def _try_configured_fallback_for_unavailable_client( explicit = (failed_provider or "").strip().lower() if not task or not explicit or explicit in {"auto"}: return None, None, "" + kwargs = {"reason": "provider unavailable"} + if env is not None: + kwargs["env"] = env return _try_configured_fallback_chain( task, explicit, - reason="provider unavailable", + **kwargs ) -def _fallback_entry_api_key(entry: Dict[str, Any]) -> Optional[str]: +def _fallback_entry_api_key(entry: Dict[str, Any], env: Optional[Dict[str, Any]] = None) -> Optional[str]: """Resolve inline or env-backed API key from a fallback-chain entry.""" explicit = str(entry.get("api_key") or "").strip() if explicit: return explicit key_env = str(entry.get("key_env") or entry.get("api_key_env") or "").strip() if key_env: + if env is not None: + return env.get(key_env, "").strip() or None return os.getenv(key_env, "").strip() or None return None -def _resolve_fallback_entry(entry: Dict[str, Any]) -> Tuple[Optional[Any], Optional[str]]: +def _resolve_fallback_entry(entry: Dict[str, Any], env: Optional[Dict[str, Any]] = None) -> Tuple[Optional[Any], Optional[str]]: """Resolve one fallback entry through the central provider router.""" provider = str(entry.get("provider") or "").strip() model = str(entry.get("model") or "").strip() or None if not provider or not model: return None, None base_url = str(entry.get("base_url") or "").strip() or None - api_key = _fallback_entry_api_key(entry) + api_key = _fallback_entry_api_key(entry, env=env) api_mode = str(entry.get("api_mode") or entry.get("transport") or "").strip() or None return resolve_provider_client( provider, @@ -3453,6 +3650,7 @@ def _resolve_fallback_entry(entry: Dict[str, Any]) -> Tuple[Optional[Any], Optio explicit_base_url=base_url, explicit_api_key=api_key, api_mode=api_mode, + env=env, ) @@ -3544,6 +3742,7 @@ def _resolve_single_provider( model: Optional[str] = None, base_url: Optional[str] = None, api_key: Optional[str] = None, + env: Optional[Dict[str, Any]] = None, ) -> Optional[Any]: """Resolve a single provider entry from fallback_chain to an OpenAI client. @@ -3555,12 +3754,14 @@ def _resolve_single_provider( model=model, explicit_base_url=base_url, explicit_api_key=api_key, + env=env, ) return client def _resolve_auto( main_runtime: Optional[Dict[str, Any]] = None, task: Optional[str] = None, + env: Optional[Dict[str, Any]] = None, ) -> Tuple[Optional[OpenAI], Optional[str]]: """Full auto-detection chain. @@ -3653,6 +3854,7 @@ def _resolve_auto( explicit_base_url=explicit_base_url, explicit_api_key=explicit_api_key, api_mode=runtime_api_mode or None, + env=env, ) if client is not None: logger.info("Auxiliary auto-detect: using main provider %s (%s)", @@ -3681,7 +3883,7 @@ def _resolve_auto( _log_skip_unhealthy(label) tried.append(f"{label} (unhealthy)") continue - client, model = try_fn() + client, model = _try_provider_chain_entry(label, try_fn, env=env) if client is not None: if tried: logger.info("Auxiliary auto-detect: using %s (%s) — skipped: %s", @@ -3794,6 +3996,7 @@ def resolve_provider_client( explicit_api_key: str = None, api_mode: str = None, main_runtime: Optional[Dict[str, Any]] = None, + env: Optional[Dict[str, Any]] = None, is_vision: bool = False, task: Optional[str] = None, ) -> Tuple[Optional[Any], Optional[str]]: @@ -3827,6 +4030,7 @@ def resolve_provider_client( Returns: (client, resolved_model) or (None, None) if auth is unavailable. """ + env = _effective_runtime_env(env) _validate_proxy_env_urls() # Preserve the original provider name before alias normalization so a # user-declared ``custom_providers`` entry whose name coincidentally @@ -3916,7 +4120,7 @@ def _wrap_if_needed(client_obj, final_model_str: str, base_url_str: str = "", # ── Auto: try all providers in priority order ──────────────────── if provider == "auto": - client, resolved = _resolve_auto(main_runtime=main_runtime, task=task) + client, resolved = _resolve_auto(main_runtime=main_runtime, task=task, env=env) if client is None: return None, None # When auto-detection lands on a non-OpenRouter provider (e.g. a @@ -3934,11 +4138,11 @@ def _wrap_if_needed(client_obj, final_model_str: str, base_url_str: str = "", # ── OpenRouter ─────────────────────────────────────────── if provider == "openrouter": - client, default = _try_openrouter(explicit_api_key=explicit_api_key) + client, default = _try_openrouter(explicit_api_key=explicit_api_key, env=env) if client is None: logger.warning( "resolve_provider_client: openrouter requested but %s", - _describe_openrouter_unavailable(), + _describe_openrouter_unavailable(env=env), ) return None, None final_model = _normalize_resolved_model(model or default, provider) @@ -3953,7 +4157,7 @@ def _wrap_if_needed(client_obj, final_model_str: str, base_url_str: str = "", model in _PROVIDER_VISION_MODELS.values() or (model or "").strip().lower() == "mimo-v2-omni" ) - client, default = _try_nous(vision=_is_vision) + client, default = _try_nous(vision=_is_vision, env=env) if client is None: logger.warning("resolve_provider_client: nous requested " "but Nous Portal not configured (run: hermes auth)") @@ -3974,7 +4178,7 @@ def _wrap_if_needed(client_obj, final_model_str: str, base_url_str: str = "", if raw_codex: # Return the raw OpenAI client for callers that need direct # access to responses.stream() (e.g., the main agent loop). - codex_token = _read_codex_access_token() + codex_token = _read_codex_access_token(env=env) if not codex_token: logger.warning("resolve_provider_client: openai-codex requested " "but no Codex OAuth token found (run: hermes model)") @@ -3987,7 +4191,7 @@ def _wrap_if_needed(client_obj, final_model_str: str, base_url_str: str = "", ) return (raw_client, final_model) # Standard path: wrap in CodexAuxiliaryClient adapter - client, default = _build_codex_client(model) + client, default = _build_codex_client(model, env=env) if client is None: logger.warning("resolve_provider_client: openai-codex requested " "but no Codex OAuth token found (run: hermes model)") @@ -4005,7 +4209,7 @@ def _wrap_if_needed(client_obj, final_model_str: str, base_url_str: str = "", # OpenRouter / Nous bills for side tasks they thought were running on # their xAI subscription. if provider == "xai-oauth": - client, default = _build_xai_oauth_aux_client(model) + client, default = _build_xai_oauth_aux_client(model, env=env) if client is None: logger.warning( "resolve_provider_client: xai-oauth requested but no xAI " @@ -4020,17 +4224,17 @@ def _wrap_if_needed(client_obj, final_model_str: str, base_url_str: str = "", if provider == "custom": if explicit_base_url: custom_base = _to_openai_base_url(explicit_base_url).strip() - custom_key = ( - (explicit_api_key or "").strip() - or os.getenv("OPENAI_API_KEY", "").strip() - or "no-key-required" # local servers don't need auth - ) + custom_key = (explicit_api_key or "").strip() if not custom_base: logger.warning( "resolve_provider_client: explicit custom endpoint requested " "but base_url is empty" ) return None, None + if not custom_key: + custom_key = _env_lookup(env, "OPENAI_API_KEY").strip() + if not custom_key: + custom_key = "no-key-required" # local servers don't need auth final_model = _normalize_resolved_model( model or (main_runtime.get("model") if main_runtime else None) or "gpt-4o-mini", provider, @@ -4068,7 +4272,7 @@ def _wrap_if_needed(client_obj, final_model_str: str, base_url_str: str = "", # Try custom first, then API-key providers (Codex excluded here: # falling through to Codex with no model is a stale-constant trap). for try_fn in (_try_custom_endpoint, _resolve_api_key_provider): - client, default = try_fn() + client, default = try_fn(env=env) if client is not None: final_model = _normalize_resolved_model(model or default, provider) _cbase = str(getattr(client, "base_url", "") or "") @@ -4096,15 +4300,23 @@ def _wrap_if_needed(client_obj, final_model_str: str, base_url_str: str = "", # still defer to the built-in per `_get_named_custom_provider`'s guard. custom_entry = None if original_provider and original_provider != provider: - custom_entry = _get_named_custom_provider(original_provider) + custom_entry = _get_named_custom_provider(original_provider, env=env) if custom_entry is None: - custom_entry = _get_named_custom_provider(provider) + custom_entry = _get_named_custom_provider(provider, env=env) if custom_entry: custom_base = custom_entry.get("base_url", "").strip() custom_key = custom_entry.get("api_key", "").strip() custom_key_env = (custom_entry.get("key_env") or custom_entry.get("api_key_env") or "").strip() if not custom_key and custom_key_env: - custom_key = os.getenv(custom_key_env, "").strip() + custom_key = _env_lookup(env, custom_key_env).strip() + if env is not None and custom_key_env and not custom_key: + logger.debug( + "resolve_provider_client: named custom provider %r skipped " + "because key_env %s is absent from scoped runtime env", + custom_entry.get("name") or provider, + custom_key_env, + ) + return None, None custom_key = custom_key or "no-key-required" if custom_key == "no-key-required": logger.warning( @@ -4216,6 +4428,7 @@ def _wrap_if_needed(client_obj, final_model_str: str, base_url_str: str = "", explicit_api_key=explicit_api_key, explicit_base_url=explicit_base_url, api_mode=api_mode, + env=env, ) if client is None: logger.warning( @@ -4246,14 +4459,14 @@ def _wrap_if_needed(client_obj, final_model_str: str, base_url_str: str = "", if pconfig.auth_type == "api_key": if provider == "anthropic": - client, default_model = _try_anthropic(explicit_api_key=explicit_api_key) + client, default_model = _try_anthropic(explicit_api_key=explicit_api_key, env=env) if client is None: logger.warning("resolve_provider_client: anthropic requested but no Anthropic credentials found") return None, None final_model = _normalize_resolved_model(model or default_model, provider) return (_to_async_client(client, final_model, is_vision=is_vision) if async_mode else (client, final_model)) - creds = resolve_api_key_provider_credentials(provider) + creds = resolve_api_key_provider_credentials(provider, env=env) api_key = str(creds.get("api_key", "")).strip() # Honour an explicit api_key override (e.g. from a fallback_model entry # or a custom_providers entry) so callers that pass an explicit @@ -4423,11 +4636,11 @@ def _wrap_if_needed(client_obj, final_model_str: str, base_url_str: str = "", elif pconfig.auth_type in {"oauth_device_code", "oauth_external"}: # OAuth providers — route through their specific try functions if provider == "nous": - return resolve_provider_client("nous", model, async_mode) + return resolve_provider_client("nous", model, async_mode, env=env) if provider == "openai-codex": - return resolve_provider_client("openai-codex", model, async_mode) + return resolve_provider_client("openai-codex", model, async_mode, env=env) if provider == "xai-oauth": - return resolve_provider_client("xai-oauth", model, async_mode) + return resolve_provider_client("xai-oauth", model, async_mode, env=env) # Other OAuth providers not directly supported logger.warning("resolve_provider_client: OAuth provider %s not " "directly supported, try 'auto'", provider) @@ -4444,6 +4657,7 @@ def get_text_auxiliary_client( task: str = "", *, main_runtime: Optional[Dict[str, Any]] = None, + env: Optional[Dict[str, Any]] = None, ) -> Tuple[Optional[OpenAI], Optional[str]]: """Return (client, default_model_slug) for text-only auxiliary tasks. @@ -4462,10 +4676,16 @@ def get_text_auxiliary_client( explicit_api_key=api_key, api_mode=api_mode, main_runtime=main_runtime, + env=env, ) -def get_async_text_auxiliary_client(task: str = "", *, main_runtime: Optional[Dict[str, Any]] = None): +def get_async_text_auxiliary_client( + task: str = "", + *, + main_runtime: Optional[Dict[str, Any]] = None, + env: Optional[Dict[str, Any]] = None, +): """Return (async_client, model_slug) for async consumers. For standard providers returns (AsyncOpenAI, model). For Codex returns @@ -4481,6 +4701,7 @@ def get_async_text_auxiliary_client(task: str = "", *, main_runtime: Optional[Di explicit_api_key=api_key, api_mode=api_mode, main_runtime=main_runtime, + env=env, ) @@ -4528,31 +4749,35 @@ def _normalize_vision_provider(provider: Optional[str]) -> str: def _resolve_strict_vision_backend( provider: str, model: Optional[str] = None, + env: Optional[Dict[str, Any]] = None, ) -> Tuple[Optional[Any], Optional[str]]: provider = _normalize_vision_provider(provider) if provider == "copilot": - return resolve_provider_client("copilot", model, is_vision=True) + return resolve_provider_client("copilot", model, env=env, is_vision=True) if provider == "openrouter": - return _try_openrouter(model=model) + return _try_openrouter(model=model, env=env) if provider == "nous": - return _try_nous(vision=True) + return _try_nous(vision=True, env=env) if provider == "openai-codex": # Route through resolve_provider_client so the caller's explicit # model is used. There is no safe default Codex model (shifting # allow-list); callers must specify via auxiliary..model. - return resolve_provider_client("openai-codex", model, is_vision=True) + return resolve_provider_client("openai-codex", model, env=env, is_vision=True) if provider == "anthropic": - return _try_anthropic() + return _try_anthropic(env=env) if provider == "custom": - return _try_custom_endpoint() + return _try_custom_endpoint(env=env) return None, None -def _strict_vision_backend_available(provider: str) -> bool: - return _resolve_strict_vision_backend(provider)[0] is not None +def _strict_vision_backend_available( + provider: str, + env: Optional[Dict[str, Any]] = None, +) -> bool: + return _resolve_strict_vision_backend(provider, env=env)[0] is not None -def get_available_vision_backends() -> List[str]: +def get_available_vision_backends(env: Optional[Dict[str, Any]] = None) -> List[str]: """Return the currently available vision backends in auto-selection order. Order: active provider → OpenRouter → Nous → stop. This is the single @@ -4564,15 +4789,15 @@ def get_available_vision_backends() -> List[str]: main_provider = _read_main_provider() if main_provider and main_provider not in {"auto", ""}: if main_provider in _VISION_AUTO_PROVIDER_ORDER: - if _strict_vision_backend_available(main_provider): + if _strict_vision_backend_available(main_provider, env=env): available.append(main_provider) else: - client, _ = resolve_provider_client(main_provider, _read_main_model()) + client, _ = resolve_provider_client(main_provider, _read_main_model(), env=env) if client is not None: available.append(main_provider) # 2. OpenRouter, 3. Nous — skip if already covered by main provider. for p in _VISION_AUTO_PROVIDER_ORDER: - if p not in available and _strict_vision_backend_available(p): + if p not in available and _strict_vision_backend_available(p, env=env): available.append(p) return available @@ -4584,6 +4809,7 @@ def resolve_vision_provider_client( base_url: Optional[str] = None, api_key: Optional[str] = None, async_mode: bool = False, + env: Optional[Dict[str, Any]] = None, ) -> Tuple[Optional[str], Optional[Any], Optional[str]]: """Resolve the client actually used for vision tasks. @@ -4617,6 +4843,7 @@ def _finalize(resolved_provider: str, sync_client: Any, default_model: Optional[ explicit_base_url=resolved_base_url, explicit_api_key=resolved_api_key, api_mode=resolved_api_mode, + env=env, ) if client is None: return provider_for_base_override, None, None @@ -4640,7 +4867,7 @@ def _finalize(resolved_provider: str, sync_client: Any, default_model: Optional[ vision_model = _PROVIDER_VISION_MODELS.get(main_provider, main_model) if main_provider == "nous": sync_client, default_model = _resolve_strict_vision_backend( - main_provider, vision_model + main_provider, vision_model, env=env ) if sync_client is not None: logger.info( @@ -4682,6 +4909,7 @@ def _finalize(resolved_provider: str, sync_client: Any, default_model: Optional[ rpc_client, rpc_model = resolve_provider_client( main_provider, vision_model, api_mode=resolved_api_mode, + env=env, is_vision=True) if rpc_client is not None: logger.info( @@ -4696,7 +4924,7 @@ def _finalize(resolved_provider: str, sync_client: Any, default_model: Optional[ for candidate in _VISION_AUTO_PROVIDER_ORDER: if candidate == main_provider: continue # already tried above - sync_client, default_model = _resolve_strict_vision_backend(candidate) + sync_client, default_model = _resolve_strict_vision_backend(candidate, env=env) if sync_client is not None: return _finalize(candidate, sync_client, default_model) @@ -4705,7 +4933,7 @@ def _finalize(resolved_provider: str, sync_client: Any, default_model: Optional[ if requested in _VISION_AUTO_PROVIDER_ORDER: sync_client, default_model = _resolve_strict_vision_backend( - requested, resolved_model + requested, resolved_model, env=env ) return _finalize(requested, sync_client, default_model) @@ -4724,6 +4952,7 @@ def _finalize(resolved_provider: str, sync_client: Any, default_model: Optional[ base_url=_zai_url, api_key=resolved_api_key or None, api_mode="chat_completions", + env=env, is_vision=True, ) if client is not None: @@ -4731,6 +4960,7 @@ def _finalize(resolved_provider: str, sync_client: Any, default_model: Optional[ # Fallback: try without explicit base_url (old behavior) client, final_model = _get_cached_client(requested, resolved_model, async_mode, api_mode=resolved_api_mode, + env=env, is_vision=True) if client is None: return requested, None, None @@ -4738,6 +4968,7 @@ def _finalize(resolved_provider: str, sync_client: Any, default_model: Optional[ client, final_model = _get_cached_client(requested, resolved_model, async_mode, api_mode=resolved_api_mode, + env=env, is_vision=True) if client is None: return requested, None, None @@ -4789,7 +5020,7 @@ def auxiliary_max_tokens_param(value: int, *, model: Optional[str] = None) -> di # Every auxiliary LLM consumer should use these instead of manually # constructing clients and calling .chat.completions.create(). -# Client cache: (provider, async_mode, base_url, api_key, api_mode, runtime_key) -> (client, default_model, loop) +# Client cache: (provider, async_mode, base_url, api_key, api_mode, runtime_key, env_key) -> (client, default_model, loop) # NOTE: loop identity is NOT part of the key. On async cache hits we check # whether the cached loop is the *current* loop; if not, the stale entry is # replaced in-place. This bounds cache growth to one entry per unique @@ -4808,17 +5039,19 @@ def _client_cache_key( api_key: Optional[str] = None, api_mode: Optional[str] = None, main_runtime: Optional[Dict[str, Any]] = None, + env: Optional[Dict[str, Any]] = None, is_vision: bool = False, task: Optional[str] = None, ) -> tuple: runtime = _normalize_main_runtime(main_runtime) - runtime_key = tuple(runtime.get(field, "") for field in _MAIN_RUNTIME_FIELDS) if provider == "auto" else () + runtime_key = tuple(runtime.get(field, "") for field in _MAIN_RUNTIME_FIELDS) if runtime else () # `auto` can now resolve through task-specific or main fallback policy, # so the task participates in the cache key. Non-auto providers keep the # old cache shape because the explicit provider/model tuple is sufficient. task_key = (task or "") if provider == "auto" else "" - pool_hint = _pool_cache_hint(provider, main_runtime=main_runtime) - return (provider, async_mode, base_url or "", api_key or "", api_mode or "", runtime_key, is_vision, task_key, pool_hint) + env_key = _runtime_env_cache_key(env) + pool_hint = _pool_cache_hint(provider, main_runtime=main_runtime, env=env) + return (provider, async_mode, base_url or "", api_key or "", api_mode or "", runtime_key, env_key, is_vision, task_key, pool_hint) def _store_cached_client(cache_key: tuple, client: Any, default_model: Optional[str], *, bound_loop: Any = None) -> None: @@ -4844,9 +5077,12 @@ def _refresh_nous_auxiliary_client( api_key: Optional[str] = None, api_mode: Optional[str] = None, main_runtime: Optional[Dict[str, Any]] = None, + env: Optional[Dict[str, Any]] = None, is_vision: bool = False, ) -> Tuple[Optional[Any], Optional[str]]: """Refresh Nous runtime creds, rebuild the client, and replace the cache entry.""" + if env is not None and _read_nous_auth(env=env) is None: + return None, model runtime = _resolve_nous_runtime_api(force_refresh=True) if runtime is None: return None, model @@ -4873,6 +5109,7 @@ def _refresh_nous_auxiliary_client( api_key=api_key, api_mode=api_mode, main_runtime=main_runtime, + env=env, is_vision=is_vision, ) _store_cached_client(cache_key, client, final_model, bound_loop=current_loop) @@ -5010,6 +5247,7 @@ def _get_cached_client( api_key: str = None, api_mode: str = None, main_runtime: Optional[Dict[str, Any]] = None, + env: Optional[Dict[str, Any]] = None, is_vision: bool = False, task: Optional[str] = None, ) -> Tuple[Optional[Any], Optional[str]]: @@ -5027,6 +5265,7 @@ def _get_cached_client( preventing the fd-exhaustion that previously occurred in long-running gateways where recycled worker threads created unbounded entries (#10200). """ + env = _effective_runtime_env(env) # Resolve the current event loop for async clients so we can validate # cached entries. Loop identity is NOT in the cache key — instead we # check at hit time whether the cached loop is still current and open. @@ -5048,6 +5287,7 @@ def _get_cached_client( api_key=api_key, api_mode=api_mode, main_runtime=main_runtime, + env=env, is_vision=is_vision, task=task, ) @@ -5395,6 +5635,7 @@ def _build_call_kwargs( timeout: float = 30.0, extra_body: Optional[dict] = None, base_url: Optional[str] = None, + env: Optional[Dict[str, Any]] = None, ) -> dict: """Build kwargs for .chat.completions.create() with model/provider adjustments.""" kwargs: Dict[str, Any] = { @@ -5436,7 +5677,7 @@ def _build_call_kwargs( # max_tokens is a MANDATORY field — omitting it is a hard 400. Keep it only # there. _effective_base = base_url or ( - _current_custom_base_url() if provider == "custom" else "" + _current_custom_base_url(env=env) if provider == "custom" else "" ) if _is_anthropic_compat_endpoint(provider, _effective_base): kwargs["max_tokens"] = max_tokens @@ -5573,6 +5814,7 @@ def call_llm( base_url: str = None, api_key: str = None, main_runtime: Optional[Dict[str, Any]] = None, + env: Optional[Dict[str, Any]] = None, messages: list, temperature: float = None, max_tokens: int = None, @@ -5604,6 +5846,7 @@ def call_llm( Raises: RuntimeError: If no provider is configured. """ + env = _effective_runtime_env(env) resolved_provider, resolved_model, resolved_base_url, resolved_api_key, resolved_api_mode = _resolve_task_provider_model( task, provider, model, base_url, api_key) effective_extra_body = _get_task_extra_body(task) @@ -5616,6 +5859,7 @@ def call_llm( base_url=resolved_base_url or base_url, api_key=resolved_api_key or api_key, async_mode=False, + env=env, ) if client is None and resolved_provider != "auto" and not resolved_base_url: logger.warning( @@ -5626,6 +5870,7 @@ def call_llm( provider="auto", model=resolved_model, async_mode=False, + env=env, ) if client is None: raise RuntimeError( @@ -5641,6 +5886,7 @@ def call_llm( api_key=resolved_api_key, api_mode=resolved_api_mode, main_runtime=main_runtime, + env=env, ) if client is None: # When the user explicitly chose a non-OpenRouter provider but no @@ -5651,7 +5897,7 @@ def call_llm( _explicit = (resolved_provider or "").strip().lower() if _explicit and _explicit not in {"auto", "openrouter", "custom"}: fb_client, fb_model, fb_label = _try_configured_fallback_for_unavailable_client( - task, _explicit, + task, _explicit, env=env, ) if fb_client is not None: client, final_model = fb_client, fb_model @@ -5670,7 +5916,7 @@ def call_llm( if client is None and not resolved_base_url: logger.info("Auxiliary %s: provider %s unavailable, trying auto-detection chain", task or "call", resolved_provider) - client, final_model = _get_cached_client("auto", main_runtime=main_runtime, task=task) + client, final_model = _get_cached_client("auto", main_runtime=main_runtime, env=env, task=task) if client is None: raise RuntimeError( f"No LLM provider configured for task={task} provider={resolved_provider}. " @@ -5692,7 +5938,9 @@ def call_llm( resolved_provider, final_model, messages, temperature=temperature, max_tokens=max_tokens, tools=tools, timeout=effective_timeout, extra_body=effective_extra_body, - base_url=_base_info or resolved_base_url) + base_url=_base_info or resolved_base_url, + env=env, + ) # Convert image blocks for Anthropic-compatible endpoints (e.g. MiniMax) _client_base = str(getattr(client, "base_url", "") or "") @@ -5824,6 +6072,7 @@ def call_llm( api_key=resolved_api_key, api_mode=resolved_api_mode, main_runtime=main_runtime, + env=env, is_vision=(task == "vision"), ) if refreshed_client is not None: @@ -5869,7 +6118,7 @@ def call_llm( if (_is_auth_error(first_err) and resolved_provider not in {"auto", "", None} and not client_is_nous): - if _refresh_provider_credentials(resolved_provider): + if _refresh_provider_credentials(resolved_provider, env=env): logger.info( "Auxiliary %s: refreshed %s credentials after auth error, retrying", task or "call", resolved_provider, @@ -5882,6 +6131,7 @@ def call_llm( resolved_api_key=resolved_api_key, resolved_api_mode=resolved_api_mode, main_runtime=main_runtime, + env=env, final_model=final_model, messages=messages, temperature=temperature, @@ -5910,7 +6160,7 @@ def call_llm( if not (_is_auth_error(retry_err) or _is_payment_error(retry_err) or _is_rate_limit_error(retry_err)): raise recovery_err = retry_err - if _recover_provider_pool(pool_provider, recovery_err, failed_api_key=_client_api_key): + if _recover_provider_pool(pool_provider, recovery_err, failed_api_key=_client_api_key, env=env): logger.info( "Auxiliary %s: recovered %s via credential-pool rotation after %s", task or "call", pool_provider, type(recovery_err).__name__, @@ -5924,6 +6174,7 @@ def call_llm( resolved_api_key=resolved_api_key, resolved_api_mode=resolved_api_mode, main_runtime=main_runtime, + env=env, final_model=final_model, messages=messages, temperature=temperature, @@ -5940,7 +6191,7 @@ def call_llm( # alternative providers can still serve the request. if (_is_payment_error(retry2_err) or _is_auth_error(retry2_err) or _is_rate_limit_error(retry2_err)): - _recover_provider_pool(pool_provider, retry2_err) + _recover_provider_pool(pool_provider, retry2_err, env=env) first_err = retry2_err else: raise @@ -6041,7 +6292,9 @@ def call_llm( temperature=temperature, max_tokens=max_tokens, tools=tools, timeout=effective_timeout, extra_body=effective_extra_body, - base_url=str(getattr(fb_client, "base_url", "") or "")) + base_url=str(getattr(fb_client, "base_url", "") or ""), + env=env, + ) return _validate_llm_response( fb_client.chat.completions.create(**fb_kwargs), task) # All fallback layers exhausted — emit a single user-visible @@ -6130,6 +6383,7 @@ async def async_call_llm( base_url: str = None, api_key: str = None, main_runtime: Optional[Dict[str, Any]] = None, + env: Optional[Dict[str, Any]] = None, messages: list, temperature: float = None, max_tokens: int = None, @@ -6141,6 +6395,7 @@ async def async_call_llm( Same as call_llm() but async. See call_llm() for full documentation. """ + env = _effective_runtime_env(env) resolved_provider, resolved_model, resolved_base_url, resolved_api_key, resolved_api_mode = _resolve_task_provider_model( task, provider, model, base_url, api_key) effective_extra_body = _get_task_extra_body(task) @@ -6153,6 +6408,7 @@ async def async_call_llm( base_url=resolved_base_url or base_url, api_key=resolved_api_key or api_key, async_mode=True, + env=env, ) if client is None and resolved_provider != "auto" and not resolved_base_url: logger.warning( @@ -6163,6 +6419,7 @@ async def async_call_llm( provider="auto", model=resolved_model, async_mode=True, + env=env, ) if client is None: raise RuntimeError( @@ -6178,12 +6435,14 @@ async def async_call_llm( base_url=resolved_base_url, api_key=resolved_api_key, api_mode=resolved_api_mode, + main_runtime=main_runtime, + env=env, ) if client is None: _explicit = (resolved_provider or "").strip().lower() if _explicit and _explicit not in {"auto", "openrouter", "custom"}: fb_client, fb_model, fb_label = _try_configured_fallback_for_unavailable_client( - task, _explicit, + task, _explicit, env=env, ) if fb_client is not None: client, final_model = _to_async_client( @@ -6199,7 +6458,7 @@ async def async_call_llm( if client is None and not resolved_base_url: logger.info("Auxiliary %s: provider %s unavailable, trying auto-detection chain", task or "call", resolved_provider) - client, final_model = _get_cached_client("auto", async_mode=True, main_runtime=main_runtime, task=task) + client, final_model = _get_cached_client("auto", async_mode=True, main_runtime=main_runtime, env=env, task=task) if client is None: raise RuntimeError( f"No LLM provider configured for task={task} provider={resolved_provider}. " @@ -6215,7 +6474,9 @@ async def async_call_llm( resolved_provider, final_model, messages, temperature=temperature, max_tokens=max_tokens, tools=tools, timeout=effective_timeout, extra_body=effective_extra_body, - base_url=_client_base or resolved_base_url) + base_url=_client_base or resolved_base_url, + env=env, + ) # Convert image blocks for Anthropic-compatible endpoints (e.g. MiniMax) if _is_anthropic_compat_endpoint(resolved_provider, _client_base): @@ -6332,6 +6593,8 @@ async def async_call_llm( base_url=resolved_base_url, api_key=resolved_api_key, api_mode=resolved_api_mode, + main_runtime=main_runtime, + env=env, is_vision=(task == "vision"), ) if refreshed_client is not None: @@ -6376,7 +6639,7 @@ async def async_call_llm( if (_is_auth_error(first_err) and resolved_provider not in {"auto", "", None} and not client_is_nous): - if _refresh_provider_credentials(resolved_provider): + if _refresh_provider_credentials(resolved_provider, env=env): logger.info( "Auxiliary %s (async): refreshed %s credentials after auth error, retrying", task or "call", resolved_provider, @@ -6388,6 +6651,8 @@ async def async_call_llm( resolved_base_url=resolved_base_url, resolved_api_key=resolved_api_key, resolved_api_mode=resolved_api_mode, + main_runtime=main_runtime, + env=env, final_model=final_model, messages=messages, temperature=temperature, @@ -6412,7 +6677,7 @@ async def async_call_llm( if not (_is_auth_error(retry_err) or _is_payment_error(retry_err) or _is_rate_limit_error(retry_err)): raise recovery_err = retry_err - if _recover_provider_pool(pool_provider, recovery_err, failed_api_key=_client_api_key): + if _recover_provider_pool(pool_provider, recovery_err, failed_api_key=_client_api_key, env=env): logger.info( "Auxiliary %s (async): recovered %s via credential-pool rotation after %s", task or "call", pool_provider, type(recovery_err).__name__, @@ -6425,6 +6690,8 @@ async def async_call_llm( resolved_base_url=resolved_base_url, resolved_api_key=resolved_api_key, resolved_api_mode=resolved_api_mode, + main_runtime=main_runtime, + env=env, final_model=final_model, messages=messages, temperature=temperature, @@ -6436,7 +6703,7 @@ async def async_call_llm( except Exception as retry2_err: if (_is_payment_error(retry2_err) or _is_auth_error(retry2_err) or _is_rate_limit_error(retry2_err)): - _recover_provider_pool(pool_provider, retry2_err) + _recover_provider_pool(pool_provider, retry2_err, env=env) first_err = retry2_err else: raise @@ -6509,7 +6776,9 @@ async def async_call_llm( temperature=temperature, max_tokens=max_tokens, tools=tools, timeout=effective_timeout, extra_body=effective_extra_body, - base_url=str(getattr(fb_client, "base_url", "") or "")) + base_url=str(getattr(fb_client, "base_url", "") or ""), + env=env, + ) # Convert sync fallback client to async async_fb, async_fb_model = _to_async_client( fb_client, fb_model or "", is_vision=(task == "vision") diff --git a/agent/skill_commands.py b/agent/skill_commands.py index 18264c44bd3bc..0107b372922da 100644 --- a/agent/skill_commands.py +++ b/agent/skill_commands.py @@ -18,10 +18,13 @@ substitute_template_vars as _substitute_template_vars, ) +from hermes_constants import get_hermes_home + logger = logging.getLogger(__name__) _skill_commands: Dict[str, Dict[str, Any]] = {} _skill_commands_platform: Optional[str] = None +_skill_commands_home: Optional[str] = None # Patterns for sanitizing skill names into clean hyphen-separated slugs. _SKILL_INVALID_CHARS = re.compile(r"[^a-z0-9-]") _SKILL_MULTI_HYPHEN = re.compile(r"-{2,}") @@ -142,13 +145,15 @@ def _load_skill_payload(skill_identifier: str, task_id: str | None = None) -> tu return None try: - from tools.skills_tool import SKILLS_DIR, skill_view + from tools.skills_tool import get_skills_dir, skill_view from agent.skill_utils import get_external_skills_dirs + skills_dir = get_skills_dir() + identifier_path = Path(raw_identifier).expanduser() if identifier_path.is_absolute(): normalized = None - trusted_roots = [SKILLS_DIR] + trusted_roots = [skills_dir] try: trusted_roots.extend(get_external_skills_dirs()) except Exception: @@ -169,7 +174,7 @@ def _load_skill_payload(skill_identifier: str, task_id: str | None = None) -> tu if normalized is None: try: - normalized = str(identifier_path.resolve().relative_to(SKILLS_DIR.resolve())) + normalized = str(identifier_path.resolve().relative_to(skills_dir.resolve())) except Exception: normalized = raw_identifier else: @@ -196,7 +201,7 @@ def _load_skill_payload(skill_identifier: str, task_id: str | None = None) -> tu skill_dir = Path(abs_skill_dir) elif skill_path: try: - skill_dir = SKILLS_DIR / Path(skill_path).parent + skill_dir = get_skills_dir() / Path(skill_path).parent except Exception: skill_dir = None @@ -251,9 +256,10 @@ def _build_skill_message( session_id: str | None = None, ) -> str: """Format a loaded skill into a user/system message payload.""" - from tools.skills_tool import SKILLS_DIR + from tools.skills_tool import get_skills_dir content = str(loaded_skill.get("content") or "") + skills_dir = get_skills_dir() # ── Template substitution and inline-shell expansion ── # Done before anything else so downstream blocks (setup notes, @@ -320,7 +326,7 @@ def _build_skill_message( if supporting and skill_dir: try: - skill_view_target = str(skill_dir.relative_to(SKILLS_DIR)) + skill_view_target = str(skill_dir.relative_to(skills_dir)) except ValueError: # Skill is from an external dir — use the skill name instead skill_view_target = skill_dir.name @@ -351,19 +357,21 @@ def scan_skill_commands() -> Dict[str, Dict[str, Any]]: Returns: Dict mapping "/skill-name" to {name, description, skill_md_path, skill_dir}. """ - global _skill_commands, _skill_commands_platform + global _skill_commands, _skill_commands_platform, _skill_commands_home _skill_commands_platform = _resolve_skill_commands_platform() + _skill_commands_home = str(get_hermes_home()) _skill_commands = {} try: - from tools.skills_tool import SKILLS_DIR, _parse_frontmatter, skill_matches_platform, skill_matches_environment, _get_disabled_skill_names + from tools.skills_tool import get_skills_dir, _parse_frontmatter, skill_matches_platform, skill_matches_environment, _get_disabled_skill_names from agent.skill_utils import get_external_skills_dirs, iter_skill_index_files disabled = _get_disabled_skill_names() seen_names: set = set() # Scan local dir first, then external dirs dirs_to_scan = [] - if SKILLS_DIR.exists(): - dirs_to_scan.append(SKILLS_DIR) + skills_dir = get_skills_dir() + if skills_dir.exists(): + dirs_to_scan.append(skills_dir) dirs_to_scan.extend(get_external_skills_dirs()) for scan_dir in dirs_to_scan: @@ -425,6 +433,7 @@ def get_skill_commands() -> Dict[str, Dict[str, Any]]: if ( not _skill_commands or _skill_commands_platform != _resolve_skill_commands_platform() + or _skill_commands_home != str(get_hermes_home()) ): scan_skill_commands() return _skill_commands diff --git a/cron/jobs.py b/cron/jobs.py index 3cab8488d2f40..dbf5cc832cf6d 100644 --- a/cron/jobs.py +++ b/cron/jobs.py @@ -49,18 +49,51 @@ # Configuration # ============================================================================= -HERMES_DIR = get_hermes_home().resolve() -CRON_DIR = HERMES_DIR / "cron" -JOBS_FILE = CRON_DIR / "jobs.json" -# Heartbeat file the in-process ticker touches on every loop iteration. The -# gateway process and the (separate) ``hermes cron status`` process share it -# so status can tell whether the ticker THREAD is alive, not just whether the -# gateway PROCESS exists — a ticker that dies silently inside a live gateway -# would otherwise report healthy (#32612, #32895). -TICKER_HEARTBEAT_FILE = CRON_DIR / "ticker_heartbeat" -# Last tick that completed WITHOUT raising. Distinguishing this from the plain -# heartbeat lets status detect a ticker that is alive but failing every tick. -TICKER_SUCCESS_FILE = CRON_DIR / "ticker_last_success" +# Module-level attributes default to None; they can be overridden by tests/web_server +HERMES_DIR = None +CRON_DIR = None +JOBS_FILE = None +TICKER_HEARTBEAT_FILE = None +TICKER_SUCCESS_FILE = None +OUTPUT_DIR = None + + +def _hermes_dir() -> Path: + if HERMES_DIR is not None: + return HERMES_DIR + return get_hermes_home().resolve() + + +def _cron_dir() -> Path: + if CRON_DIR is not None: + return CRON_DIR + return _hermes_dir() / "cron" + + +def _jobs_file() -> Path: + if JOBS_FILE is not None: + return JOBS_FILE + return _cron_dir() / "jobs.json" + + +def _ticker_heartbeat_file() -> Path: + if TICKER_HEARTBEAT_FILE is not None: + return TICKER_HEARTBEAT_FILE + return _cron_dir() / "ticker_heartbeat" + + +def _ticker_success_file() -> Path: + if TICKER_SUCCESS_FILE is not None: + return TICKER_SUCCESS_FILE + return _cron_dir() / "ticker_last_success" + + +def _output_dir() -> Path: + if OUTPUT_DIR is not None: + return OUTPUT_DIR + return _cron_dir() / "output" + + # Default ticker loop interval (seconds). The single source of truth shared by # the in-process ticker (cron/scheduler_provider.py) and the staleness # threshold in `hermes cron status` (hermes_cli/cron.py), so the two never @@ -72,13 +105,12 @@ # concurrent mark_job_run / advance_next_run calls can clobber each other. _jobs_file_lock = threading.RLock() _jobs_lock_state = threading.local() -OUTPUT_DIR = CRON_DIR / "output" ONESHOT_GRACE_SECONDS = 120 def _jobs_lock_file() -> Path: """Return the advisory lock path for the current cron directory.""" - return CRON_DIR / ".jobs.lock" + return _cron_dir() / ".jobs.lock" @contextlib.contextmanager @@ -163,7 +195,7 @@ def _job_output_dir(job_id: str) -> Path: raise ValueError(f"Invalid cron job id for output path: {job_id!r}") if Path(text).is_absolute() or Path(text).drive: raise ValueError(f"Invalid cron job id for output path: {job_id!r}") - return OUTPUT_DIR / text + return _output_dir() / text def _normalize_skill_list(skill: Optional[str] = None, skills: Optional[Any] = None) -> List[str]: @@ -270,10 +302,10 @@ def _secure_file(path: Path): def ensure_dirs(): """Ensure cron directories exist with secure permissions.""" - CRON_DIR.mkdir(parents=True, exist_ok=True) - OUTPUT_DIR.mkdir(parents=True, exist_ok=True) - _secure_dir(CRON_DIR) - _secure_dir(OUTPUT_DIR) + _cron_dir().mkdir(parents=True, exist_ok=True) + _output_dir().mkdir(parents=True, exist_ok=True) + _secure_dir(_cron_dir()) + _secure_dir(_output_dir()) # ============================================================================= @@ -562,7 +594,7 @@ def _atomic_write_epoch(path: Path) -> None: torn/truncated file. Best-effort: failures are swallowed by callers. """ ensure_dirs() - fd, tmp_path = tempfile.mkstemp(dir=str(CRON_DIR), suffix=".tmp", prefix=".hb_") + fd, tmp_path = tempfile.mkstemp(dir=str(_cron_dir()), suffix=".tmp", prefix=".hb_") try: with os.fdopen(fd, "w", encoding="utf-8") as f: f.write(str(time.time())) @@ -590,12 +622,12 @@ def record_ticker_heartbeat(success: bool = False) -> None: Best-effort: a write failure must never disrupt the tick loop. """ try: - _atomic_write_epoch(TICKER_HEARTBEAT_FILE) + _atomic_write_epoch(_ticker_heartbeat_file()) except Exception: pass if success: try: - _atomic_write_epoch(TICKER_SUCCESS_FILE) + _atomic_write_epoch(_ticker_success_file()) except Exception: pass @@ -614,12 +646,12 @@ def get_ticker_heartbeat_age() -> Optional[float]: None = heartbeat file missing/unreadable (older build, never ran, or a torn read). Callers treat None as "cannot determine", not "dead". """ - return _epoch_file_age(TICKER_HEARTBEAT_FILE) + return _epoch_file_age(_ticker_heartbeat_file()) def get_ticker_success_age() -> Optional[float]: """Seconds since the ticker last completed a tick WITHOUT raising, or None.""" - return _epoch_file_age(TICKER_SUCCESS_FILE) + return _epoch_file_age(_ticker_success_file()) # ============================================================================= @@ -629,19 +661,19 @@ def get_ticker_success_age() -> Optional[float]: def load_jobs() -> List[Dict[str, Any]]: """Load all jobs from storage.""" ensure_dirs() - if not JOBS_FILE.exists(): + if not _jobs_file().exists(): return [] _strict_retry = False # track whether we used the strict=False fallback try: - with open(JOBS_FILE, 'r', encoding='utf-8') as f: + with open(_jobs_file(), 'r', encoding='utf-8') as f: data = json.load(f) except json.JSONDecodeError: # Retry with strict=False to handle bare control chars in string values _strict_retry = True try: - with open(JOBS_FILE, 'r', encoding='utf-8') as f: + with open(_jobs_file(), 'r', encoding='utf-8') as f: data = json.loads(f.read(), strict=False) except Exception as e: logger.error("Failed to auto-repair jobs.json: %s", e) @@ -677,14 +709,14 @@ def load_jobs() -> List[Dict[str, Any]]: def _save_jobs_unlocked(jobs: List[Dict[str, Any]]): """Save all jobs to storage. Caller must hold _jobs_lock().""" ensure_dirs() - fd, tmp_path = tempfile.mkstemp(dir=str(JOBS_FILE.parent), suffix='.tmp', prefix='.jobs_') + fd, tmp_path = tempfile.mkstemp(dir=str(_jobs_file().parent), suffix='.tmp', prefix='.jobs_') try: with os.fdopen(fd, 'w', encoding='utf-8') as f: json.dump({"jobs": jobs, "updated_at": _hermes_now().isoformat()}, f, indent=2) f.flush() os.fsync(f.fileno()) - atomic_replace(tmp_path, JOBS_FILE) - _secure_file(JOBS_FILE) + atomic_replace(tmp_path, _jobs_file()) + _secure_file(_jobs_file()) except BaseException: try: os.unlink(tmp_path) diff --git a/gateway/run.py b/gateway/run.py index a84d3ca6cf714..1c584d147798e 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -1413,6 +1413,173 @@ class MultiplexConfigError(RuntimeError): """ +class RoutingSessionStoreProxy: + def __init__(self, runner): + self._runner = runner + self._db_proxy = RoutingSessionDBProxy(runner) + + @property + def _db(self): + return self._db_proxy + + @property + def _lock(self): + return self._resolve_store_for_args()._lock + + @property + def _entries(self): + return self._resolve_store_for_args()._entries + + def _resolve_store_for_args(self, *args, **kwargs): + # 1. Check if active home override is active + try: + from hermes_constants import get_hermes_home, get_default_hermes_root + cur = get_hermes_home().resolve() + root = get_default_hermes_root().resolve() + if cur != root: + return self._runner._get_or_create_store_for_home(cur) + except Exception: + pass + + # 2. Check for SessionSource in args/kwargs + from gateway.session import SessionSource + for arg in args: + if isinstance(arg, SessionSource): + return self._runner.session_store_for_source(arg) + for val in kwargs.values(): + if isinstance(val, SessionSource): + return self._runner.session_store_for_source(val) + + # 3. Check for session_key (string starting with agent:) in args/kwargs + for arg in args: + if isinstance(arg, str) and arg.startswith("agent:"): + parts = arg.split(":") + if len(parts) > 1: + profile_name = parts[1] + if profile_name != "main": + from hermes_cli.profiles import get_profile_dir + try: + ph = get_profile_dir(profile_name).resolve() + return self._runner._get_or_create_store_for_home(ph) + except Exception: + pass + break + for val in kwargs.values(): + if isinstance(val, str) and val.startswith("agent:"): + parts = val.split(":") + if len(parts) > 1: + profile_name = parts[1] + if profile_name != "main": + from hermes_cli.profiles import get_profile_dir + try: + ph = get_profile_dir(profile_name).resolve() + return self._runner._get_or_create_store_for_home(ph) + except Exception: + pass + break + + # 4. Check for session_id (which might be in args/kwargs) + session_id = kwargs.get("session_id") + if not session_id: + for arg in args: + if isinstance(arg, str) and len(arg) >= 32 and "-" in arg: + session_id = arg + break + if session_id: + if hasattr(self._runner, "_profile_session_stores"): + for store in self._runner._profile_session_stores.values(): + if hasattr(store, "_entries"): + for entry in store._entries.values(): + if entry.session_id == session_id: + return store + # Fallback scan of all profiles + for home in self._runner._all_profile_homes(): + store = self._runner._get_or_create_store_for_home(home) + if hasattr(store, "_entries"): + for entry in store._entries.values(): + if entry.session_id == session_id: + return store + + # Default fallback to the global session store + from hermes_constants import get_default_hermes_root + try: + root = get_default_hermes_root().resolve() + return self._runner._get_or_create_store_for_home(root) + except Exception: + return self._runner.session_store + + def __getattr__(self, name): + store = self._resolve_store_for_args() + return getattr(store, name) + + def _generate_session_key(self, source, *args, **kwargs): + return self._resolve_store_for_args(source, *args, **kwargs)._generate_session_key(source, *args, **kwargs) + + def get_or_create_session(self, source, *args, **kwargs): + return self._resolve_store_for_args(source, *args, **kwargs).get_or_create_session(source, *args, **kwargs) + + def switch_session(self, session_key, *args, **kwargs): + return self._resolve_store_for_args(session_key, *args, **kwargs).switch_session(session_key, *args, **kwargs) + + def append_to_transcript(self, session_id, *args, **kwargs): + return self._resolve_store_for_args(session_id, *args, **kwargs).append_to_transcript(session_id, *args, **kwargs) + + def load_transcript(self, session_id, *args, **kwargs): + return self._resolve_store_for_args(session_id, *args, **kwargs).load_transcript(session_id, *args, **kwargs) + + def rewrite_transcript(self, session_id, *args, **kwargs): + return self._resolve_store_for_args(session_id, *args, **kwargs).rewrite_transcript(session_id, *args, **kwargs) + + def rewind_session(self, session_id, *args, **kwargs): + return self._resolve_store_for_args(session_id, *args, **kwargs).rewind_session(session_id, *args, **kwargs) + + def lookup_by_session_id(self, session_id, *args, **kwargs): + return self._resolve_store_for_args(session_id, *args, **kwargs).lookup_by_session_id(session_id, *args, **kwargs) + + def has_platform_message_id(self, session_id, *args, **kwargs): + return self._resolve_store_for_args(session_id, *args, **kwargs).has_platform_message_id(session_id, *args, **kwargs) + + def reset_session(self, session_key, *args, **kwargs): + return self._resolve_store_for_args(session_key, *args, **kwargs).reset_session(session_key, *args, **kwargs) + + def suspend_session(self, session_key, *args, **kwargs): + return self._resolve_store_for_args(session_key, *args, **kwargs).suspend_session(session_key, *args, **kwargs) + + def mark_resume_pending(self, session_key, *args, **kwargs): + return self._resolve_store_for_args(session_key, *args, **kwargs).mark_resume_pending(session_key, *args, **kwargs) + + def clear_resume_pending(self, session_key, *args, **kwargs): + return self._resolve_store_for_args(session_key, *args, **kwargs).clear_resume_pending(session_key, *args, **kwargs) + + @property + def config(self): + return self._resolve_store_for_args().config + + +class RoutingSessionDBProxy: + def __init__(self, runner): + self._runner = runner + + def _resolve_db_for_args(self, *args, **kwargs): + store_proxy = RoutingSessionStoreProxy(self._runner) + store = store_proxy._resolve_store_for_args(*args, **kwargs) + if store: + return store._db + return None + + def __getattr__(self, name): + db = self._resolve_db_for_args() + return getattr(db, name) + + def delete_telegram_topic_binding(self, chat_id, thread_id, *args, **kwargs): + from gateway.session import SessionSource, Platform + source = SessionSource(platform=Platform.TELEGRAM, chat_id=chat_id, thread_id=thread_id) + store = self._runner.session_store_for_source(source) + if store and store._db: + return store._db.delete_telegram_topic_binding(chat_id, thread_id, *args, **kwargs) + return False + + @_contextmanager def _profile_runtime_scope(profile_home: "Path"): """Scope config/skills/memory AND credentials to a profile for one turn. @@ -2223,6 +2390,94 @@ def _resolve_gateway_model(config: dict | None = None) -> str: return "" +def _load_topic_models() -> dict: + """Load the persistent topic-specific model overrides from ~/.hermes/topic_models.json.""" + import json + path = _hermes_home / "topic_models.json" + if path.exists(): + try: + with open(path, "r", encoding="utf-8") as f: + return json.load(f) + except Exception: + return {} + return {} + + +def _save_topic_model(session_key: str, model_data: dict) -> None: + """Save a persistent topic-specific model override to ~/.hermes/topic_models.json.""" + path = _hermes_home / "topic_models.json" + data = _load_topic_models() + data[session_key] = model_data + try: + atomic_json_write(path, data) + except Exception as e: + logger.warning("Failed to save topic model: %s", e) + + +def _remove_topic_model(session_key: str) -> None: + """Remove a persistent topic-specific model override from ~/.hermes/topic_models.json.""" + path = _hermes_home / "topic_models.json" + data = _load_topic_models() + if session_key in data: + data.pop(session_key) + try: + atomic_json_write(path, data) + except Exception as e: + logger.warning("Failed to remove topic model: %s", e) + + +def _topic_profile_key(source: Any) -> str: + """Derive a profile lookup key from a session source. + + We use a representation of the source's platform, chat type, chat ID, and thread ID + which is completely independent of the profile namespace. + """ + if not source: + return "" + platform = getattr(source, "platform", None) + platform_str = platform.value if platform else "" + chat_type = getattr(source, "chat_type", "dm") or "dm" + chat_id = getattr(source, "chat_id", "") or "" + thread_id = getattr(source, "thread_id", "") or "" + return f"{platform_str}:{chat_type}:{chat_id}:{thread_id}" + + +def _load_topic_profiles() -> dict: + """Load the persistent topic-specific profile overrides from ~/.hermes/topic_profiles.json.""" + import json + path = _hermes_home / "topic_profiles.json" + if path.exists(): + try: + with open(path, "r", encoding="utf-8") as f: + return json.load(f) + except Exception: + return {} + return {} + + +def _save_topic_profile(topic_key: str, profile_name: str) -> None: + """Save a persistent topic-specific profile override to ~/.hermes/topic_profiles.json.""" + path = _hermes_home / "topic_profiles.json" + data = _load_topic_profiles() + data[topic_key] = profile_name + try: + atomic_json_write(path, data) + except Exception as e: + logger.warning("Failed to save topic profile: %s", e) + + +def _remove_topic_profile(topic_key: str) -> None: + """Remove a persistent topic-specific profile override from ~/.hermes/topic_profiles.json.""" + path = _hermes_home / "topic_profiles.json" + data = _load_topic_profiles() + if topic_key in data: + data.pop(topic_key) + try: + atomic_json_write(path, data) + except Exception as e: + logger.warning("Failed to remove topic profile: %s", e) + + def _resolve_hermes_bin() -> Optional[list[str]]: """Resolve the Hermes update command as argv parts. @@ -2544,6 +2799,102 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew _session_reasoning_overrides: Dict[str, Dict[str, Any]] = {} _startup_restore_in_progress: bool = False + def _get_or_create_store_for_home(self, home: "Path"): + home = home.resolve() + if not hasattr(self, "_profile_cache_lock"): + import threading + self._profile_cache_lock = threading.Lock() + if not hasattr(self, "_profile_session_stores"): + self._profile_session_stores = {} + with self._profile_cache_lock: + if home not in self._profile_session_stores: + from gateway.session import SessionStore + from tools.process_registry import process_registry + profile_sessions_dir = home / "sessions" + store = SessionStore( + profile_sessions_dir, + self.config, + has_active_processes_fn=lambda key: process_registry.has_active_for_session(key), + db_path=home / "state.db", + ) + self._profile_session_stores[home] = store + return self._profile_session_stores[home] + + @property + def session_store(self): + from hermes_constants import get_hermes_home + current_home = get_hermes_home().resolve() + return self._get_or_create_store_for_home(current_home) + + @session_store.setter + def session_store(self, value): + from hermes_constants import get_hermes_home + current_home = get_hermes_home().resolve() + if not hasattr(self, "_profile_cache_lock"): + import threading + self._profile_cache_lock = threading.Lock() + if not hasattr(self, "_profile_session_stores"): + self._profile_session_stores = {} + with self._profile_cache_lock: + self._profile_session_stores[current_home] = value + + def session_store_for_source(self, source: SessionSource): + profile_name = self._routed_profile_for_source(source) + if profile_name: + from hermes_cli.profiles import get_profile_dir + home = get_profile_dir(profile_name).resolve() + else: + from hermes_constants import get_default_hermes_root + home = get_default_hermes_root().resolve() + return self._get_or_create_store_for_home(home) + + def _all_profile_homes(self) -> List["Path"]: + from hermes_cli.profiles import list_profiles + try: + return [p.path.resolve() for p in list_profiles()] + except Exception: + from hermes_constants import get_default_hermes_root + root = get_default_hermes_root().resolve() + homes = [root] + profiles_dir = root / "profiles" + if profiles_dir.is_dir(): + for p in profiles_dir.iterdir(): + if p.is_dir(): + homes.append(p.resolve()) + return homes + + @property + def _session_db(self): + from hermes_constants import get_hermes_home + current_home = get_hermes_home().resolve() + if not hasattr(self, "_profile_cache_lock"): + import threading + self._profile_cache_lock = threading.Lock() + if not hasattr(self, "_profile_session_dbs"): + self._profile_session_dbs = {} + with self._profile_cache_lock: + if current_home not in self._profile_session_dbs: + try: + from hermes_state import SessionDB + db = SessionDB(db_path=current_home / "state.db") + self._profile_session_dbs[current_home] = db + except Exception as e: + logger.warning("SQLite session store not available for %s: %s", current_home, e) + self._profile_session_dbs[current_home] = None + return self._profile_session_dbs[current_home] + + @_session_db.setter + def _session_db(self, value): + from hermes_constants import get_hermes_home + current_home = get_hermes_home().resolve() + if not hasattr(self, "_profile_cache_lock"): + import threading + self._profile_cache_lock = threading.Lock() + if not hasattr(self, "_profile_session_dbs"): + self._profile_session_dbs = {} + with self._profile_cache_lock: + self._profile_session_dbs[current_home] = value + def __init__(self, config: Optional[GatewayConfig] = None): global _gateway_runner_ref self.config = config or load_gateway_config() @@ -2585,6 +2936,7 @@ def __init__(self, config: Optional[GatewayConfig] = None): self.config.sessions_dir, self.config, has_active_processes_fn=lambda key: process_registry.has_active_for_session(key), ) + self._routing_session_store = RoutingSessionStoreProxy(self) self.delivery_router = DeliveryRouter(self.config) self._running = False self._gateway_loop: Optional[asyncio.AbstractEventLoop] = None @@ -2677,9 +3029,11 @@ def __init__(self, config: Optional[GatewayConfig] = None): self._agent_cache: "OrderedDict[str, tuple]" = OrderedDict() self._agent_cache_lock = _threading.Lock() - # Per-session model overrides from /model command. - # Key: session_key, Value: dict with model/provider/api_key/base_url/api_mode self._session_model_overrides: Dict[str, Dict[str, str]] = {} + try: + self._session_model_overrides.update(_load_topic_models()) + except Exception as e: + logger.warning("Failed to load topic models at startup: %s", e) # Per-session reasoning effort overrides from /reasoning. # Key: session_key, Value: parsed reasoning config dict. self._session_reasoning_overrides: Dict[str, Dict[str, Any]] = {} @@ -3138,6 +3492,15 @@ def exit_code(self) -> Optional[int]: def _session_key_for_source(self, source: SessionSource) -> str: """Resolve the current session key for a source, honoring gateway config when available.""" + try: + normalized = self._normalize_source_for_session_key(source) + key = _topic_profile_key(normalized) + persistent_profile = _load_topic_profiles().get(key) + if persistent_profile: + source.profile = persistent_profile + except Exception: + pass + if hasattr(self, "session_store") and self.session_store is not None: try: session_key = self.session_store._generate_session_key(source) @@ -3150,15 +3513,14 @@ def _session_key_for_source(self, source: SessionSource) -> str: # produces the same namespace as the primary path: None (legacy # agent:main) unless multiplexing is on, then the active profile. _profile = None - if getattr(config, "multiplex_profiles", False): - if source.profile: - _profile = source.profile - else: - try: - from hermes_cli.profiles import get_active_profile_name - _profile = get_active_profile_name() or "default" - except Exception: - _profile = None + if source.profile: + _profile = source.profile + elif getattr(config, "multiplex_profiles", False): + try: + from hermes_cli.profiles import get_active_profile_name + _profile = get_active_profile_name() or "default" + except Exception: + _profile = None return build_session_key( source, group_sessions_per_user=getattr(config, "group_sessions_per_user", True), @@ -3406,6 +3768,14 @@ def _resolve_session_agent_runtime( model = _resolve_gateway_model(user_config) override = self._session_model_overrides.get(resolved_session_key) if resolved_session_key else None + if not override and resolved_session_key: + try: + persistent_override = _load_topic_models().get(resolved_session_key) + if persistent_override: + self._session_model_overrides[resolved_session_key] = persistent_override + override = persistent_override + except Exception: + pass if override: override_model = override.get("model", model) override_runtime = { @@ -6025,7 +6395,7 @@ async def start(self) -> bool: # Set up message + fatal error handlers adapter.set_message_handler(self._handle_message) adapter.set_fatal_error_handler(self._handle_adapter_fatal_error) - adapter.set_session_store(self.session_store) + adapter.set_session_store(self._routing_session_store) adapter.set_busy_session_handler(self._handle_active_session_busy_message) adapter.set_topic_recovery_fn(self._recover_telegram_topic_thread_id) adapter._busy_text_mode = self._busy_text_mode @@ -6379,26 +6749,27 @@ async def _handoff_watcher(self, interval: float = 2.0) -> None: await asyncio.sleep(5) while self._running: try: - if self._session_db is None: - await asyncio.sleep(interval) - continue - pending = await asyncio.to_thread(self._session_db.list_pending_handoffs) - for row in pending: - session_id = row.get("id") - if not session_id: - continue - if not await asyncio.to_thread(self._session_db.claim_handoff, session_id): - # Another tick or another gateway already claimed it. - continue - try: - await self._process_handoff(row) - await asyncio.to_thread(self._session_db.complete_handoff, session_id) - except Exception as exc: - logger.warning( - "Handoff for session %s failed: %s", - session_id, exc, exc_info=True, - ) - await asyncio.to_thread(self._session_db.fail_handoff, session_id, str(exc)) + for home in self._all_profile_homes(): + with _profile_runtime_scope(home): + if self._session_db is None: + continue + pending = await asyncio.to_thread(self._session_db.list_pending_handoffs) + for row in pending: + session_id = row.get("id") + if not session_id: + continue + if not await asyncio.to_thread(self._session_db.claim_handoff, session_id): + # Another tick or another gateway already claimed it. + continue + try: + await self._process_handoff(row) + await asyncio.to_thread(self._session_db.complete_handoff, session_id) + except Exception as exc: + logger.warning( + "Handoff for session %s failed: %s", + session_id, exc, exc_info=True, + ) + await asyncio.to_thread(self._session_db.fail_handoff, session_id, str(exc)) except asyncio.CancelledError: raise except Exception as exc: @@ -6587,162 +6958,162 @@ async def _session_expiry_watcher(self, interval: int = 300): _MAX_FINALIZE_RETRIES = 3 while self._running: try: - self.session_store._ensure_loaded() - # Collect expired sessions first, then log a single summary. - _expired_entries = [] - for key, entry in list(self.session_store._entries.items()): - if entry.expiry_finalized: - continue - if not self.session_store._is_session_expired(entry): - continue - _expired_entries.append((key, entry)) - - if _expired_entries: - # Extract platform names from session keys for a compact summary. - # Keys look like "agent:main:telegram:dm:12345" — platform is field [2]. - _platforms: dict[str, int] = {} - for _k, _e in _expired_entries: - _parts = _k.split(":") - _plat = _parts[2] if len(_parts) > 2 else "unknown" - _platforms[_plat] = _platforms.get(_plat, 0) + 1 - _plat_summary = ", ".join( - f"{p}:{c}" for p, c in sorted(_platforms.items()) - ) - logger.info( - "Session expiry: %d sessions to finalize (%s)", - len(_expired_entries), _plat_summary, - ) - - for key, entry in _expired_entries: - try: - try: - from hermes_cli.plugins import invoke_hook as _invoke_hook - _parts = key.split(":") - _platform = _parts[2] if len(_parts) > 2 else "" - _invoke_hook( - "on_session_finalize", - session_id=entry.session_id, - platform=_platform, - reason="session_expired", - ) - except Exception: - pass - # Shut down memory provider and close tool resources - # on the cached agent. Idle agents live in - # _agent_cache (not _running_agents), so look there. - _cached_agent = None - _cache_lock = getattr(self, "_agent_cache_lock", None) - if _cache_lock is not None: - with _cache_lock: - _cached = self._agent_cache.get(key) - _cached_agent = _cached[0] if isinstance(_cached, tuple) else _cached if _cached else None - # Fall back to _running_agents in case the agent is - # still mid-turn when the expiry fires. - if _cached_agent is None: - _cached_agent = self._running_agents.get(key) - if _cached_agent and _cached_agent is not _AGENT_PENDING_SENTINEL: - self._cleanup_agent_resources(_cached_agent) - # Drop the cache entry so the AIAgent (and its LLM - # clients, tool schemas, memory provider refs) can - # be garbage-collected. Otherwise the cache grows - # unbounded across the gateway's lifetime. - self._evict_cached_agent(key) - # Permanently finalizing this session — drop its - # per-session control state so the dicts don't grow - # unbounded across the gateway's lifetime. (Idle - # agent-cache eviction must NOT prune these: the - # session is still alive and a resumed turn rebuilds - # its agent from these overrides. Only true session - # finalization, /new, and /reset clear them.) - self._session_model_overrides.pop(key, None) - self._set_session_reasoning_override(key, None) - if hasattr(self, "_pending_model_notes"): - self._pending_model_notes.pop(key, None) - _pending_approvals = getattr(self, "_pending_approvals", None) - if isinstance(_pending_approvals, dict): - _pending_approvals.pop(key, None) - _update_prompt_pending = getattr(self, "_update_prompt_pending", None) - if isinstance(_update_prompt_pending, dict): - _update_prompt_pending.pop(key, None) - with self.session_store._lock: - entry.expiry_finalized = True - self.session_store._save() - logger.debug( - "Session expiry finalized for %s", - entry.session_id, - ) - _finalize_failures.pop(entry.session_id, None) - except Exception as e: - failures = _finalize_failures.get(entry.session_id, 0) + 1 - _finalize_failures[entry.session_id] = failures - if failures >= _MAX_FINALIZE_RETRIES: - logger.warning( - "Session finalize gave up after %d attempts for %s: %s. " - "Marking as finalized to prevent infinite retry loop.", - failures, entry.session_id, e, + for home in self._all_profile_homes(): + with _profile_runtime_scope(home): + self.session_store._ensure_loaded() + # Collect expired sessions first, then log a single summary. + _expired_entries = [] + for key, entry in list(self.session_store._entries.items()): + if entry.expiry_finalized: + continue + if not self.session_store._is_session_expired(entry): + continue + _expired_entries.append((key, entry)) + + if _expired_entries: + # Extract platform names from session keys for a compact summary. + # Keys look like "agent:main:telegram:dm:12345" — platform is field [2]. + _platforms: dict[str, int] = {} + for _k, _e in _expired_entries: + _parts = _k.split(":") + _plat = _parts[2] if len(_parts) > 2 else "unknown" + _platforms[_plat] = _platforms.get(_plat, 0) + 1 + _plat_summary = ", ".join( + f"{p}:{c}" for p, c in sorted(_platforms.items()) ) - with self.session_store._lock: - entry.expiry_finalized = True - self.session_store._save() - _finalize_failures.pop(entry.session_id, None) - else: - logger.debug( - "Session finalize failed (%d/%d) for %s: %s", - failures, _MAX_FINALIZE_RETRIES, entry.session_id, e, + logger.info( + "Session expiry: %d sessions to finalize (%s)", + len(_expired_entries), _plat_summary, ) - if _expired_entries: - _done = sum( - 1 for _, e in _expired_entries if e.expiry_finalized - ) - _failed = len(_expired_entries) - _done - if _failed: - logger.info( - "Session expiry done: %d finalized, %d pending retry", - _done, _failed, - ) - else: - logger.info( - "Session expiry done: %d finalized", _done, - ) + for key, entry in _expired_entries: + try: + try: + from hermes_cli.plugins import invoke_hook as _invoke_hook + _parts = key.split(":") + _platform = _parts[2] if len(_parts) > 2 else "" + _invoke_hook( + "on_session_finalize", + session_id=entry.session_id, + platform=_platform, + reason="session_expired", + ) + except Exception: + pass + # Shut down memory provider and close tool resources + # on the cached agent. Idle agents live in + # _agent_cache (not _running_agents), so look there. + _cached_agent = None + _cache_lock = getattr(self, "_agent_cache_lock", None) + if _cache_lock is not None: + with _cache_lock: + _cached = self._agent_cache.get(key) + _cached_agent = _cached[0] if isinstance(_cached, tuple) else _cached if _cached else None + # Fall back to _running_agents in case the agent is + # still mid-turn when the expiry fires. + if _cached_agent is None: + _cached_agent = self._running_agents.get(key) + if _cached_agent and _cached_agent is not _AGENT_PENDING_SENTINEL: + self._cleanup_agent_resources(_cached_agent) + # Drop the cache entry so the AIAgent (and its LLM + # clients, tool schemas, memory provider refs) can + # be garbage-collected. Otherwise the cache grows + # unbounded across the gateway's lifetime. + self._evict_cached_agent(key) + # Permanently finalizing this session — drop its + # per-session control state so the dicts don't grow + # unbounded across the gateway's lifetime. (Idle + # agent-cache eviction must NOT prune these: the + # session is still alive and a resumed turn rebuilds + # its agent from these overrides. Only true session + # finalization, /new, and /reset clear them.) + self._session_model_overrides.pop(key, None) + _remove_topic_model(key) + self._set_session_reasoning_override(key, None) + if hasattr(self, "_pending_model_notes"): + self._pending_model_notes.pop(key, None) + _pending_approvals = getattr(self, "_pending_approvals", None) + if isinstance(_pending_approvals, dict): + _pending_approvals.pop(key, None) + _update_prompt_pending = getattr(self, "_update_prompt_pending", None) + if isinstance(_update_prompt_pending, dict): + _update_prompt_pending.pop(key, None) + with self.session_store._lock: + entry.expiry_finalized = True + self.session_store._save() + logger.debug( + "Session expiry finalized for %s", + entry.session_id, + ) + _finalize_failures.pop(entry.session_id, None) + except Exception as e: + failures = _finalize_failures.get(entry.session_id, 0) + 1 + _finalize_failures[entry.session_id] = failures + if failures >= _MAX_FINALIZE_RETRIES: + logger.warning( + "Session finalize gave up after %d attempts for %s: %s. " + "Marking as finalized to prevent infinite retry loop.", + failures, entry.session_id, e, + ) + with self.session_store._lock: + entry.expiry_finalized = True + self.session_store._save() + _finalize_failures.pop(entry.session_id, None) + else: + logger.debug( + "Session finalize failed (%d/%d) for %s: %s", + failures, _MAX_FINALIZE_RETRIES, entry.session_id, e, + ) - # Sweep agents that have been idle beyond the TTL regardless - # of session reset policy. This catches sessions with very - # long / "never" reset windows, whose cached AIAgents would - # otherwise pin memory for the gateway's entire lifetime. - try: - _idle_evicted = self._sweep_idle_cached_agents() - if _idle_evicted: - logger.info( - "Agent cache idle sweep: evicted %d agent(s)", - _idle_evicted, - ) - except Exception as _e: - logger.debug("Idle agent sweep failed: %s", _e) - - # Periodically prune stale SessionStore entries. The - # in-memory dict (and sessions.json) would otherwise grow - # unbounded in gateways serving many rotating chats / - # threads / users over long time windows. Pruning is - # invisible to users — a resumed session just gets a - # fresh session_id, exactly as if the reset policy fired. - _last_prune_ts = getattr(self, "_last_session_store_prune_ts", 0.0) - _prune_interval = 3600.0 # once per hour - if time.time() - _last_prune_ts > _prune_interval: - try: - _max_age = int( - getattr(self.config, "session_store_max_age_days", 0) or 0 - ) - if _max_age > 0: - _pruned = self.session_store.prune_old_entries(_max_age) - if _pruned: + if _expired_entries: + _done = sum( + 1 for _, e in _expired_entries if e.expiry_finalized + ) + _failed = len(_expired_entries) - _done + if _failed: logger.info( - "SessionStore prune: dropped %d stale entries", - _pruned, + "Session expiry done: %d finalized, %d pending retry", + _done, _failed, ) - except Exception as _e: - logger.debug("SessionStore prune failed: %s", _e) - self._last_session_store_prune_ts = time.time() + else: + logger.info( + "Session expiry done: %d finalized", _done, + ) + + # Sweep agents that have been idle beyond the TTL regardless + # of session reset policy. This catches sessions with very + # long / "never" reset windows, whose cached AIAgents would + # otherwise pin memory for the gateway's entire lifetime. + try: + # Note: _sweep_idle_cached_agents sweeps the global _agent_cache, + # but running it inside the profile scope ensures safe cleanup. + _idle_evicted = self._sweep_idle_cached_agents() + if _idle_evicted: + logger.info( + "Agent cache idle sweep: evicted %d agent(s)", + _idle_evicted, + ) + except Exception as _e: + logger.debug("Idle agent sweep failed: %s", _e) + + # Periodically prune stale SessionStore entries. + _last_prune_ts = getattr(self, "_last_session_store_prune_ts", 0.0) + _prune_interval = 3600.0 # once per hour + if time.time() - _last_prune_ts > _prune_interval: + try: + _max_age = int( + getattr(self.config, "session_store_max_age_days", 0) or 0 + ) + if _max_age > 0: + _pruned = self.session_store.prune_old_entries(_max_age) + if _pruned: + logger.info( + "SessionStore prune: dropped %d stale entries", + _pruned, + ) + except Exception as _e: + logger.debug("SessionStore prune failed: %s", _e) + self._last_session_store_prune_ts = time.time() except Exception as e: logger.debug("Session expiry watcher error: %s", e) # Sleep in small increments so we can stop quickly @@ -6759,6 +7130,22 @@ def _active_profile_name(self) -> str: except Exception: return "default" + def _routed_profile_for_source(self, source: SessionSource) -> Optional[str]: + """Return the profile name bound to the source's topic, or None. + + None → unbound topic: nothing is scoped, behavior is UNCHANGED. + "" → enter _profile_runtime_scope(get_profile_dir()) for the whole turn. + """ + try: + normalized = self._normalize_source_for_session_key(source) + key = _topic_profile_key(normalized) + name = (_load_topic_profiles().get(key) or "").strip() + if not name or name in ("default", self._active_profile_name()): + return None + return name + except Exception: + return None + # ── Kanban board watchers ─────────────────────────────────────────── # The kanban notifier/dispatcher watcher loops + their helpers live in # GatewayKanbanWatchersMixin (gateway/kanban_watchers.py). They use only @@ -6825,7 +7212,7 @@ async def _platform_reconnect_watcher(self) -> None: adapter.set_message_handler(self._handle_message) adapter.set_fatal_error_handler(self._handle_adapter_fatal_error) - adapter.set_session_store(self.session_store) + adapter.set_session_store(self._routing_session_store) adapter.set_busy_session_handler(self._handle_active_session_busy_message) adapter.set_topic_recovery_fn(self._recover_telegram_topic_thread_id) adapter._busy_text_mode = self._busy_text_mode @@ -7258,8 +7645,19 @@ def _phase_elapsed() -> float: # old gateway's connection holding the WAL lock until Python # actually exits — causing 'database is locked' errors when # the new gateway tries to open the same file. - for _db_holder in (self, getattr(self, "session_store", None)): - _db = getattr(_db_holder, "_db", None) if _db_holder else None + dbs_to_close = [] + if hasattr(self, "_profile_session_dbs"): + dbs_to_close.extend(self._profile_session_dbs.values()) + if hasattr(self, "_profile_session_stores"): + for store in self._profile_session_stores.values(): + if hasattr(store, "_db") and store._db not in dbs_to_close: + dbs_to_close.append(store._db) + for holder in (self, getattr(self, "session_store", None)): + db = getattr(holder, "_db", None) if holder else None + if db and db not in dbs_to_close: + dbs_to_close.append(db) + + for _db in dbs_to_close: if _db is None or not hasattr(_db, "close"): continue try: @@ -7490,7 +7888,7 @@ async def _start_one_profile_adapters( self._make_profile_message_handler(profile_name) ) adapter.set_fatal_error_handler(self._handle_adapter_fatal_error) - adapter.set_session_store(self.session_store) + adapter.set_session_store(self._routing_session_store) adapter.set_busy_session_handler(self._handle_active_session_busy_message) adapter.set_topic_recovery_fn(self._recover_telegram_topic_thread_id) adapter._busy_text_mode = self._busy_text_mode @@ -7697,8 +8095,30 @@ async def _deliver_platform_notice(self, source, content: str) -> None: await adapter.send(source.chat_id, content, metadata=metadata) async def _handle_message(self, event: MessageEvent) -> Optional[str]: + # Ensure we set source.profile from topic overrides before resolving profile + try: + normalized = self._normalize_source_for_session_key(event.source) + key = _topic_profile_key(normalized) + persistent_profile = _load_topic_profiles().get(key) + if persistent_profile: + event.source.profile = persistent_profile + except Exception: + pass + + source = event.source + routed = self._routed_profile_for_source(source) + multiplex = getattr(getattr(self, "config", None), "multiplex_profiles", False) + + if not multiplex and routed is None: + return await self._handle_message_inner(event) + + profile_home = self._resolve_profile_home_for_source(source) + with _profile_runtime_scope(profile_home): + return await self._handle_message_inner(event) + + async def _handle_message_inner(self, event: MessageEvent) -> Optional[str]: """ - Handle an incoming message from any platform. + Handle an incoming message from any platform (core pipeline). This is the core message processing pipeline: 1. Check user authorization @@ -7709,6 +8129,15 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: 6. Run agent conversation 7. Return response """ + try: + normalized = self._normalize_source_for_session_key(event.source) + key = _topic_profile_key(normalized) + persistent_profile = _load_topic_profiles().get(key) + if persistent_profile: + event.source.profile = persistent_profile + except Exception: + pass + source = event.source if ( @@ -9009,7 +9438,13 @@ async def _do_undo(): _run_generation = self._begin_session_run_generation(_quick_key) try: - _agent_result = await self._handle_message_with_agent(event, source, _quick_key, _run_generation) + routed = self._routed_profile_for_source(source) + if routed is not None: + profile_home = self._resolve_profile_home_for_source(source) + with _profile_runtime_scope(profile_home): + _agent_result = await self._handle_message_with_agent(event, source, _quick_key, _run_generation) + else: + _agent_result = await self._handle_message_with_agent(event, source, _quick_key, _run_generation) if getattr(event, "_moa_disable_after_turn", False): try: _restore = getattr(event, "_moa_restore_override", None) @@ -9506,7 +9941,7 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str, run_g # Set session context variables for tools (task-local, concurrency-safe) _session_env_tokens = self._set_session_env(context) - + # Read privacy.redact_pii from config (re-read per message) _redact_pii = False persist_user_message = None @@ -9569,14 +10004,15 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str, run_g f"Adjust reset timing in config.yaml under session_reset." ) try: - session_info = self._format_session_info() + session_info = self._format_session_info(source=source) if session_info: notice = f"{notice}\n\n{session_info}" except Exception: pass + _reset_meta = self._thread_metadata_for_source(source) await adapter.send( source.chat_id, notice, - metadata=self._thread_metadata_for_source(source), + metadata=_reset_meta, ) except Exception as e: logger.debug("Auto-reset notification failed (non-fatal): %s", e) @@ -10761,12 +11197,19 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str, run_g # Restore session context variables to their pre-handler state self._clear_session_env(_session_env_tokens) - def _format_session_info(self) -> str: + def _format_session_info( + self, + source: Optional[SessionSource] = None, + ) -> str: """Resolve current model config and return a formatted info block. Surfaces model, provider, context length, and endpoint so gateway users can immediately see if context detection went wrong (e.g. local models falling to the 128K default). + + The optional ``source`` arg surfaces the model the bound profile/topic + will actually use for the next turn — not the global default. When + omitted, behavior is identical to the previous signature. """ from agent.model_metadata import get_model_context_length, DEFAULT_FALLBACK_CONTEXT @@ -10778,19 +11221,35 @@ def _format_session_info(self) -> str: custom_provs = None data = None + if source is not None: + try: + model, runtime_kwargs = self._resolve_session_agent_runtime(source=source) + provider = runtime_kwargs.get("provider") + base_url = runtime_kwargs.get("base_url") + api_key = runtime_kwargs.get("api_key") + except Exception: + pass + try: data = _load_gateway_config() if data: model_cfg = data.get("model", {}) if isinstance(model_cfg, dict): - raw_ctx = model_cfg.get("context_length") - if raw_ctx is not None: - try: - config_context_length = int(raw_ctx) - except (TypeError, ValueError): - pass - provider = model_cfg.get("provider") or None - base_url = model_cfg.get("base_url") or None + default_model = model_cfg.get("default") or model_cfg.get("model") or "" + if not model: + model = default_model + + if model == default_model: + raw_ctx = model_cfg.get("context_length") + if raw_ctx is not None: + try: + config_context_length = int(raw_ctx) + except (TypeError, ValueError): + pass + if not provider: + provider = model_cfg.get("provider") or None + if not base_url: + base_url = model_cfg.get("base_url") or None try: from hermes_cli.config import get_compatible_custom_providers custom_provs = get_compatible_custom_providers(data) @@ -10799,6 +11258,9 @@ def _format_session_info(self) -> str: except Exception: pass + if not model: + model = _resolve_gateway_model() + # Also check custom_providers for context_length when top-level model.context_length is not set if config_context_length is None and data: try: @@ -11751,8 +12213,6 @@ async def _deliver_media_from_response( except Exception as e: logger.warning("Post-stream media extraction failed: %s", e) - - async def _run_background_task( self, prompt: str, @@ -11761,6 +12221,28 @@ async def _run_background_task( event_message_id: Optional[str] = None, media_urls: Optional[List[str]] = None, media_types: Optional[List[str]] = None, + ) -> None: + """Execute a background agent task and deliver the result to the chat.""" + routed = self._routed_profile_for_source(source) + if routed is not None: + profile_home = self._resolve_profile_home_for_source(source) + with _profile_runtime_scope(profile_home): + return await self._run_background_task_inner( + prompt, source, task_id, event_message_id, media_urls, media_types + ) + else: + return await self._run_background_task_inner( + prompt, source, task_id, event_message_id, media_urls, media_types + ) + + async def _run_background_task_inner( + self, + prompt: str, + source: "SessionSource", + task_id: str, + event_message_id: Optional[str] = None, + media_urls: Optional[List[str]] = None, + media_types: Optional[List[str]] = None, ) -> None: """Execute a background agent task and deliver the result to the chat.""" from run_agent import AIAgent @@ -12333,13 +12815,14 @@ def _telegram_topic_root_status_message(self, source: SessionSource) -> str: async def _restore_telegram_topic_session(self, event: MessageEvent, raw_session_id: str) -> str: """Restore an existing Telegram-owned Hermes session into this topic.""" source = event.source - session_id = self._session_db.resolve_session_id(raw_session_id.strip()) + raw = raw_session_id.strip() + session_id = self._session_db.resolve_session_id(raw) if not session_id: - return f"Session not found: {raw_session_id.strip()}" + return f"Session not found: {raw}" session = self._session_db.get_session(session_id) if not session: - return f"Session not found: {raw_session_id.strip()}" + return f"Session not found: {raw}" if str(session.get("source") or "") != "telegram": return "That session is not a Telegram session and cannot be restored into this topic." if str(session.get("user_id") or "") != str(source.user_id): @@ -12402,19 +12885,37 @@ async def _execute_mcp_reload(self, event: MessageEvent) -> str: from tools.mcp_tool import shutdown_mcp_servers, discover_mcp_tools, _servers, _lock # Capture old server names before shutdown + from tools.mcp_tool import _load_mcp_config, _get_mcp_config_fingerprint + try: + active_cfgs = _load_mcp_config() + active_fps = { + _get_mcp_config_fingerprint(name, cfg) + for name, cfg in active_cfgs.items() + } + except Exception: + active_fps = set() + with _lock: - old_servers = set(_servers.keys()) + old_servers = { + k for k in _servers.keys() + if k in active_fps or (getattr(_servers[k], "fingerprint", None) in active_fps) + } # Read new config before shutting down, so we know what will be added/removed # Shutdown existing connections - await loop.run_in_executor(None, shutdown_mcp_servers) + import contextvars + ctx = contextvars.copy_context() + await loop.run_in_executor(None, lambda: ctx.run(shutdown_mcp_servers)) # Reconnect by discovering tools (reads config.yaml fresh) - new_tools = await loop.run_in_executor(None, discover_mcp_tools) + new_tools = await loop.run_in_executor(None, lambda: ctx.run(discover_mcp_tools)) # Compute what changed with _lock: - connected_servers = set(_servers.keys()) + connected_servers = { + k for k in _servers.keys() + if k in active_fps or (getattr(_servers[k], "fingerprint", None) in active_fps) + } added = connected_servers - old_servers removed = old_servers - connected_servers @@ -12422,11 +12923,11 @@ async def _execute_mcp_reload(self, event: MessageEvent) -> str: lines = [t("gateway.reload_mcp.header")] if reconnected: - lines.append(t("gateway.reload_mcp.reconnected", names=", ".join(sorted(reconnected)))) + lines.append(t("gateway.reload_mcp.reconnected", names=", ".join(sorted({k.split(":")[0] for k in reconnected})))) if added: - lines.append(t("gateway.reload_mcp.added", names=", ".join(sorted(added)))) + lines.append(t("gateway.reload_mcp.added", names=", ".join(sorted({k.split(":")[0] for k in added})))) if removed: - lines.append(t("gateway.reload_mcp.removed", names=", ".join(sorted(removed)))) + lines.append(t("gateway.reload_mcp.removed", names=", ".join(sorted({k.split(":")[0] for k in removed})))) if not connected_servers: lines.append(t("gateway.reload_mcp.none_connected")) else: @@ -13327,6 +13828,17 @@ def _set_session_env(self, context: SessionContext) -> list: _adapters = getattr(self, "adapters", None) or {} _adapter = _adapters.get(context.source.platform) _async_delivery = getattr(_adapter, "supports_async_delivery", True) + _resolve_home = getattr(self, "_resolve_profile_home_for_source", None) + if _resolve_home is not None: + profile_home = _resolve_home(context.source) + else: + from hermes_constants import get_hermes_home + profile_home = get_hermes_home() + + from hermes_cli.profiles import get_active_profile_name + agent_profile = (context.source.profile or "").strip() or get_active_profile_name() or "default" + agent_hermes_home = str(profile_home) + return set_session_vars( platform=context.source.platform.value, chat_id=context.source.chat_id, @@ -13337,6 +13849,8 @@ def _set_session_env(self, context: SessionContext) -> list: session_key=context.session_key, message_id=str(context.source.message_id) if context.source.message_id else "", async_delivery=_async_delivery, + agent_profile=agent_profile, + agent_hermes_home=agent_hermes_home, ) def _clear_session_env(self, tokens: list) -> None: @@ -14119,6 +14633,17 @@ def _agent_config_signature( _api_key = str(runtime.get("api_key", "") or "") _api_key_fingerprint = hashlib.sha256(_api_key.encode()).hexdigest() if _api_key else "" + # Fingerprint the active SOUL.md to bust the cache when it is edited. + from hermes_constants import get_hermes_home + _soul_path = get_hermes_home() / "SOUL.md" + _soul_fingerprint = "" + if _soul_path.exists(): + try: + _stat = _soul_path.stat() + _soul_fingerprint = f"{_stat.st_mtime}:{_stat.st_size}" + except Exception: + pass + _cache_keys_sorted = sorted((cache_keys or {}).items()) blob = _j.dumps( @@ -14135,6 +14660,7 @@ def _agent_config_signature( _cache_keys_sorted, str(user_id or ""), str(user_id_alt or ""), + _soul_fingerprint, ], sort_keys=True, default=str, @@ -14933,7 +15459,6 @@ def _pause_typing_before_finalize( "proxy response: url=%s session=%s time=%.1fs response=%d chars", proxy_url, (session_id or "")[:20], _elapsed, len(full_response), ) - return { "final_response": full_response or "(No response from remote agent)", "messages": [ @@ -14974,7 +15499,9 @@ async def _run_agent( multiplexing is off this is a transparent pass-through — zero behavior change for single-profile gateways. """ - if not getattr(getattr(self, "config", None), "multiplex_profiles", False): + routed = self._routed_profile_for_source(source) + multiplex = getattr(getattr(self, "config", None), "multiplex_profiles", False) + if not multiplex and routed is None: return await self._run_agent_inner( message, context_prompt, history, source, session_id, session_key=session_key, run_generation=run_generation, @@ -15002,14 +15529,51 @@ def _resolve_profile_home_for_source(self, source: SessionSource) -> "Path": by the /p// URL prefix or a per-credential adapter), falling back to the active profile (the multiplexer's own home). """ - from hermes_cli.profiles import get_active_profile_name, get_profile_dir + from hermes_cli.profiles import ( + get_active_profile_name, + get_profile_dir, + validate_profile_identity, + write_profile_identity_marker, + _get_profiles_root, + PROFILE_IDENTITY_FILENAME, + ) + name = (source.profile or "").strip() or get_active_profile_name() or "default" try: - name = (source.profile or "").strip() or get_active_profile_name() or "default" - return get_profile_dir(name) - except Exception: + p_dir = get_profile_dir(name) + if name != "default": + # A profile directory created before the isolation feature has no + # identity marker. Without it, validate_profile_identity() would + # raise and we'd silently fall back to the GLOBAL home — mixing this + # profile's sessions/memory into ~/.hermes. Auto-create the marker + # (migration) so isolation engages. Genuine security violations + # (symlink escape, dir outside root) are still rejected by the + # validation below, since it checks ancestry BEFORE the marker. + marker = p_dir / PROFILE_IDENTITY_FILENAME + if p_dir.is_dir() and not marker.exists(): + try: + write_profile_identity_marker(name, p_dir, _get_profiles_root()) + logger.info( + "Auto-created identity marker for profile %r (migration)", name + ) + except Exception as exc: + logger.warning( + "Could not auto-create identity marker for profile %r: %s", + name, exc, + ) + validate_profile_identity(name, p_dir, _get_profiles_root()) + return p_dir + except Exception as exc: from hermes_constants import get_hermes_home + # NEVER fall back silently: a silent fallback to the global home breaks + # data isolation (the exact bug found in live testing). Make it loud. + logger.warning( + "Profile %r identity validation failed (%s); falling back to GLOBAL " + "home — profile isolation is DEGRADED for this turn", + name, exc, + ) return get_hermes_home() + async def _run_agent_inner( self, message: str, diff --git a/gateway/session.py b/gateway/session.py index f79e371d8043d..820709114b502 100644 --- a/gateway/session.py +++ b/gateway/session.py @@ -788,7 +788,7 @@ class SessionStore: """ def __init__(self, sessions_dir: Path, config: GatewayConfig, - has_active_processes_fn=None): + has_active_processes_fn=None, db_path: Path = None): self.sessions_dir = sessions_dir self.config = config self._entries: Dict[str, SessionEntry] = {} @@ -800,7 +800,10 @@ def __init__(self, sessions_dir: Path, config: GatewayConfig, self._db = None try: from hermes_state import SessionDB - self._db = SessionDB() + if db_path is None: + from hermes_constants import get_hermes_home + db_path = get_hermes_home() / "state.db" + self._db = SessionDB(db_path=db_path) except Exception as e: print(f"[gateway] Warning: SQLite session store unavailable, falling back to JSONL: {e}") @@ -896,10 +899,10 @@ def _resolve_profile_for_key(self, source: Optional[SessionSource] = None) -> Op to (``source.profile`` — set by the /p// URL prefix or per-credential adapter), falling back to the active profile name. """ - if not getattr(self.config, "multiplex_profiles", False): - return None if source is not None and source.profile: return source.profile + if not getattr(self.config, "multiplex_profiles", False): + return None try: from hermes_cli.profiles import get_active_profile_name return get_active_profile_name() or "default" @@ -914,7 +917,7 @@ def _generate_session_key(self, source: SessionSource) -> str: thread_sessions_per_user=getattr(self.config, "thread_sessions_per_user", False), profile=self._resolve_profile_for_key(source), ) - + def _is_session_expired(self, entry: SessionEntry) -> bool: """Check if a session has expired based on its reset policy. diff --git a/gateway/session_context.py b/gateway/session_context.py index 55f269df54dbc..46a4b762f040c 100644 --- a/gateway/session_context.py +++ b/gateway/session_context.py @@ -61,6 +61,8 @@ # so background-process notifications stay inside the originating Telegram # private-chat topic (those lanes route only with thread id + reply anchor). _SESSION_MESSAGE_ID: ContextVar = ContextVar("HERMES_SESSION_MESSAGE_ID", default=_UNSET) +_SESSION_AGENT_PROFILE: ContextVar = ContextVar("HERMES_SESSION_AGENT_PROFILE", default=_UNSET) +_SESSION_AGENT_HERMES_HOME: ContextVar = ContextVar("HERMES_SESSION_AGENT_HERMES_HOME", default=_UNSET) # Whether the current session's delivery channel can route an ASYNC completion # back to the agent AFTER the current turn ends (i.e. wake a fresh turn). @@ -100,6 +102,8 @@ "HERMES_SESSION_KEY": _SESSION_KEY, "HERMES_SESSION_ID": _SESSION_ID, "HERMES_SESSION_MESSAGE_ID": _SESSION_MESSAGE_ID, + "HERMES_SESSION_AGENT_PROFILE": _SESSION_AGENT_PROFILE, + "HERMES_SESSION_AGENT_HERMES_HOME": _SESSION_AGENT_HERMES_HOME, "HERMES_CRON_AUTO_DELIVER_PLATFORM": _CRON_AUTO_DELIVER_PLATFORM, "HERMES_CRON_AUTO_DELIVER_CHAT_ID": _CRON_AUTO_DELIVER_CHAT_ID, "HERMES_CRON_AUTO_DELIVER_THREAD_ID": _CRON_AUTO_DELIVER_THREAD_ID, @@ -134,6 +138,8 @@ def set_session_vars( message_id: str = "", cwd: str = "", async_delivery: bool = True, + agent_profile: str = "", + agent_hermes_home: str = "", ) -> list: """Set all session context variables and return reset tokens. @@ -162,6 +168,8 @@ def set_session_vars( _SESSION_ID.set(session_id), _SESSION_MESSAGE_ID.set(message_id), _SESSION_ASYNC_DELIVERY.set(bool(async_delivery)), + _SESSION_AGENT_PROFILE.set(agent_profile), + _SESSION_AGENT_HERMES_HOME.set(agent_hermes_home), ] try: from agent.runtime_cwd import set_session_cwd @@ -194,6 +202,8 @@ def clear_session_vars(tokens: list) -> None: _SESSION_KEY, _SESSION_ID, _SESSION_MESSAGE_ID, + _SESSION_AGENT_PROFILE, + _SESSION_AGENT_HERMES_HOME, ): var.set("") # Reset async-delivery capability to the "never set" sentinel rather than a diff --git a/gateway/slash_commands.py b/gateway/slash_commands.py index b2b8089b51b7e..3eb0dec62c55d 100644 --- a/gateway/slash_commands.py +++ b/gateway/slash_commands.py @@ -182,6 +182,13 @@ async def _handle_reset_command(self, event: MessageEvent) -> Union[str, Ephemer # Clear any session-scoped model/reasoning overrides so the next agent # picks up configured defaults instead of previous session switches. self._session_model_overrides.pop(session_key, None) + try: + from gateway.run import _load_topic_models + persistent_override = _load_topic_models().get(session_key) + if persistent_override: + self._session_model_overrides[session_key] = persistent_override + except Exception: + pass self._set_session_reasoning_override(session_key, None) if hasattr(self, "_pending_model_notes"): self._pending_model_notes.pop(session_key, None) @@ -223,7 +230,7 @@ async def _handle_reset_command(self, event: MessageEvent) -> Union[str, Ephemer # Resolve session config info to surface to the user try: - session_info = self._format_session_info() + session_info = self._format_session_info(source=source) except Exception: session_info = "" @@ -295,17 +302,77 @@ async def _handle_reset_command(self, event: MessageEvent) -> Union[str, Ephemer return EphemeralReply(f"{header}{_tip_line}") async def _handle_profile_command(self, event: MessageEvent) -> str: - """Handle /profile — show active profile name and home directory.""" + """Handle /profile command — switch and show active profile for this topic. + + Supports: + /profile — show active profile and list available profiles + /profile — pin topic to a profile + /profile default/reset/clear — clear topic profile override + """ from hermes_constants import display_hermes_home - from hermes_cli.profiles import get_active_profile_name + from hermes_cli.profiles import list_profiles, get_active_profile_name + from gateway.run import _topic_profile_key, _remove_topic_profile, _save_topic_profile + + source = event.source + source = self._normalize_source_for_session_key(source) + topic_key = _topic_profile_key(source) + session_key = self._session_key_for_source(source) + + profile_input = event.get_command_args().strip() + + # Load all valid profiles on disk + try: + profiles = list_profiles() + valid_profiles = {p.name for p in profiles} + except Exception: + valid_profiles = {"default"} + + if profile_input: + if profile_input in ("default", "reset", "clear"): + try: + _remove_topic_profile(topic_key) + except Exception: + pass + event.source.profile = None + self._evict_cached_agent(session_key) + return ( + "Cleared the topic profile binding. Now using the global " + f"active profile '{self._active_profile_name()}'." + ) - display = display_hermes_home() - profile_name = get_active_profile_name() + if profile_input not in valid_profiles: + available = ", ".join(f"`{name}`" for name in sorted(valid_profiles)) + return f"Unknown profile '{profile_input}'. Available profiles: {available}." + + # Save the profile override + try: + _save_topic_profile(topic_key, profile_input) + except Exception: + pass + event.source.profile = profile_input + + # Evict the cached agent session since the profile/session key changed + self._evict_cached_agent(session_key) + + return f"Pinned this topic to profile `{profile_input}`." + + # No args: show active profile (same format as upstream /profile) and + # list the available profiles that this topic can switch to. + topic_profile = event.source.profile or get_active_profile_name() lines = [ - t("gateway.profile.header", profile=profile_name), - t("gateway.profile.home", home=display), + t("gateway.profile.header", profile=topic_profile), + t("gateway.profile.home", home=display_hermes_home()), + "", + "Available profiles (tap to switch):", ] + for name in sorted(valid_profiles): + is_active = (name == topic_profile) + bullet = "*" if is_active else "-" + lines.append(f"{bullet} `/profile {name}`") + + lines.append("") + lines.append("To clear the binding: `/profile default`") return "\n".join(lines) @@ -1124,6 +1191,10 @@ async def _handle_model_command(self, event: MessageEvent) -> Optional[str]: ) from hermes_cli.providers import get_label + # On gateway (non-local) sources a plain ``/model X`` binds the model to + # this topic (topic_models.json) instead of the global config.yaml. + is_gateway = (event.source.platform != Platform.LOCAL) if event.source.platform else False + raw_args = event.get_command_args().strip() # Parse --provider, --global, --session, and --refresh flags @@ -1135,6 +1206,8 @@ async def _handle_model_command(self, event: MessageEvent) -> Optional[str]: is_session, ) = parse_model_flags(raw_args) persist_global = resolve_persist_behavior(is_global_flag, is_session) + if is_gateway and not is_global_flag: + persist_global = False # --refresh: bust the disk cache so the picker shows live data. if force_refresh: @@ -1177,6 +1250,17 @@ async def _handle_model_command(self, event: MessageEvent) -> Optional[str]: # (#30479). source = self._normalize_source_for_session_key(source) session_key = self._session_key_for_source(source) + + if model_input in ("default", "reset", "clear"): + self._session_model_overrides.pop(session_key, None) + try: + from gateway.run import _remove_topic_model + _remove_topic_model(session_key) + except Exception: + pass + self._evict_cached_agent(session_key) + return "Topic-specific model override cleared. Now using the global default model." + override = self._session_model_overrides.get(session_key, {}) if override: current_model = override.get("model", current_model) @@ -1325,6 +1409,19 @@ async def _on_model_selected( "base_url": result.base_url, "api_mode": result.api_mode, } + if is_gateway: + if persist_global: + try: + from gateway.run import _remove_topic_model + _remove_topic_model(_session_key) + except Exception: + pass + else: + try: + from gateway.run import _save_topic_model + _save_topic_model(_session_key, _self._session_model_overrides[_session_key]) + except Exception as e: + logger.warning("Failed to save persistent topic model: %s", e) # Evict cached agent so the next turn creates a fresh # agent from the override rather than relying on the @@ -1554,6 +1651,19 @@ async def _finish_switch() -> str: "base_url": result.base_url, "api_mode": result.api_mode, } + if is_gateway: + if persist_global: + try: + from gateway.run import _remove_topic_model + _remove_topic_model(session_key) + except Exception: + pass + else: + try: + from gateway.run import _save_topic_model + _save_topic_model(session_key, self._session_model_overrides[session_key]) + except Exception as e: + logger.warning("Failed to save persistent topic model: %s", e) # Evict cached agent so the next turn creates a fresh agent from the # override rather than relying on cache signature mismatch detection. diff --git a/hermes_cli/auth.py b/hermes_cli/auth.py index 4a0571a180bec..9168a1578db9e 100644 --- a/hermes_cli/auth.py +++ b/hermes_cli/auth.py @@ -65,6 +65,7 @@ AUTH_STORE_VERSION = 1 AUTH_LOCK_TIMEOUT_SECONDS = 15.0 +STRICT_PROFILE_AUTH_ENV = "HERMES_PROFILE_STRICT_AUTH" # Nous Portal defaults DEFAULT_NOUS_PORTAL_URL = "https://portal.nousresearch.com" @@ -87,6 +88,15 @@ MINIMAX_OAUTH_CN_BASE = "https://api.minimaxi.com" MINIMAX_OAUTH_GLOBAL_INFERENCE = "https://api.minimax.io/anthropic" MINIMAX_OAUTH_CN_INFERENCE = "https://api.minimaxi.com/anthropic" + + +def _env_get(env: Optional[Dict[str, Any]], key: str, default: str = "") -> str: + if env is None: + return os.getenv(key, default) + value = env.get(key, default) + return "" if value is None else str(value) + + MINIMAX_OAUTH_REFRESH_SKEW_SECONDS = 60 DEFAULT_QWEN_BASE_URL = "https://portal.qwen.ai/v1" DEFAULT_GITHUB_MODELS_BASE_URL = "https://api.githubcopilot.com" @@ -557,7 +567,7 @@ def has_usable_secret(value: Any, *, min_length: int = 4) -> bool: def _resolve_api_key_provider_secret( - provider_id: str, pconfig: ProviderConfig + provider_id: str, pconfig: ProviderConfig, env: Optional[Dict[str, Any]] = None ) -> tuple[str, str]: """Resolve an API-key provider's token and indicate where it came from.""" if provider_id == "copilot": @@ -573,26 +583,32 @@ def _resolve_api_key_provider_secret( pass return "", "" - from hermes_cli.config import get_env_value for env_var in pconfig.api_key_env_vars: - # Check both os.environ and ~/.hermes/.env file - val = (get_env_value(env_var) or "").strip() + if env is None: + from hermes_cli.config import get_env_value + # Check both os.environ and ~/.hermes/.env file + val = (get_env_value(env_var) or "").strip() + else: + val = _env_get(env, env_var).strip() if has_usable_secret(val): return val, env_var - # Fallback: try credential pool (e.g. zai key stored via auth.json) - try: - from agent.credential_pool import load_pool - pool = load_pool(provider_id) - if pool and pool.has_credentials(): - entry = pool.peek() - if entry: - key = getattr(entry, "access_token", "") or getattr(entry, "runtime_api_key", "") - key = str(key).strip() - if has_usable_secret(key): - return key, f"credential_pool:{provider_id}" - except Exception: - pass + # Fallback: try credential pool (e.g. zai key stored via auth.json). + # Scoped callers pass an explicit env mapping; they must not inherit the + # process/global pool after their scoped env misses. + if env is None: + try: + from agent.credential_pool import load_pool + pool = load_pool(provider_id) + if pool and pool.has_credentials(): + entry = pool.peek() + if entry: + key = getattr(entry, "access_token", "") or getattr(entry, "runtime_api_key", "") + key = str(key).strip() + if has_usable_secret(key): + return key, f"credential_pool:{provider_id}" + except Exception: + pass return "", "" @@ -880,6 +896,8 @@ def _global_auth_file_path() -> Optional[Path]: See issue #18594 follow-up (credential_pool shadowing). """ + if profile_strict_auth_enabled(): + return None try: from hermes_constants import get_default_hermes_root global_root = get_default_hermes_root() @@ -901,6 +919,28 @@ def _global_auth_file_path() -> Optional[Path]: return global_root / "auth.json" +def profile_strict_auth_enabled(env: Optional[Dict[str, Any]] = None) -> bool: + """Return True when routed-profile auth must not use global fallbacks.""" + if isinstance(env, dict) and is_truthy_value( + str(env.get(STRICT_PROFILE_AUTH_ENV, "") or "") + ): + return True + if is_truthy_value(os.getenv(STRICT_PROFILE_AUTH_ENV, "")): + return True + try: + from gateway.session_context import get_runtime_env, get_session_env + if is_truthy_value(get_session_env(STRICT_PROFILE_AUTH_ENV, "")): + return True + runtime_env = get_runtime_env() + if isinstance(runtime_env, dict) and is_truthy_value( + str(runtime_env.get(STRICT_PROFILE_AUTH_ENV, "") or "") + ): + return True + except Exception: + pass + return False + + def _load_global_auth_store() -> Dict[str, Any]: """Load the global-root auth store (read-only fallback). @@ -1485,6 +1525,7 @@ def resolve_provider( *, explicit_api_key: Optional[str] = None, explicit_base_url: Optional[str] = None, + env: Optional[Dict[str, Any]] = None, ) -> str: """ Determine which inference provider to use. @@ -1577,7 +1618,7 @@ def resolve_provider( except Exception as e: logger.debug("Could not detect active auth provider: %s", e) - if has_usable_secret(os.getenv("OPENAI_API_KEY")) or has_usable_secret(os.getenv("OPENROUTER_API_KEY")): + if has_usable_secret(_env_get(env, "OPENAI_API_KEY")) or has_usable_secret(_env_get(env, "OPENROUTER_API_KEY")): return "openrouter" # Auto-detect an OpenRouter credential added via `hermes auth add openrouter` @@ -1608,17 +1649,18 @@ def resolve_provider( if pid in {"copilot", "lmstudio"}: continue for env_var in pconfig.api_key_env_vars: - if has_usable_secret(os.getenv(env_var, "")): + if has_usable_secret(_env_get(env, env_var)): return pid # AWS Bedrock — detect via boto3 credential chain (IAM roles, SSO, env vars). # This runs after API-key providers so explicit keys always win. - try: - from agent.bedrock_adapter import has_aws_credentials - if has_aws_credentials(): - return "bedrock" - except ImportError: - pass # boto3 not installed — skip Bedrock auto-detection + if env is None or any(_env_get(env, key).strip() for key in ("AWS_ACCESS_KEY_ID", "AWS_PROFILE")): + try: + from agent.bedrock_adapter import has_aws_credentials + if has_aws_credentials(): + return "bedrock" + except ImportError: + pass # boto3 not installed — skip Bedrock auto-detection raise AuthError( "No inference provider configured. Run 'hermes model' to choose a " @@ -3663,6 +3705,8 @@ def _import_codex_cli_tokens() -> Optional[Dict[str, str]]: Returns tokens dict if valid and not expired, None otherwise. Does NOT write to the shared file. """ + if profile_strict_auth_enabled(): + return None codex_home = os.getenv("CODEX_HOME", "").strip() if not codex_home: codex_home = str(Path.home() / ".codex") @@ -4721,6 +4765,8 @@ def _nous_shared_store_lock(timeout_seconds: float = AUTH_LOCK_TIMEOUT_SECONDS): def _merge_shared_nous_oauth_state(state: Dict[str, Any]) -> bool: """Copy fresher shared OAuth tokens into a profile-local Nous state.""" + if profile_strict_auth_enabled(): + return False shared = _read_shared_nous_state() if not shared: return False @@ -4764,6 +4810,8 @@ def _write_shared_nous_state(state: Dict[str, Any]) -> None: We deliberately omit the runtime ``agent_key`` compatibility field; the OAuth tokens are the cross-profile source of truth. """ + if profile_strict_auth_enabled(): + return refresh_token = state.get("refresh_token") access_token = state.get("access_token") if not (isinstance(refresh_token, str) and refresh_token.strip()): @@ -4828,6 +4876,8 @@ def _read_shared_nous_state() -> Optional[Dict[str, Any]]: lacks required fields. Callers should treat ``None`` as "no shared credentials available — fall through to device-code". """ + if profile_strict_auth_enabled(): + return None try: path = _nous_shared_store_path() except RuntimeError: @@ -4853,6 +4903,8 @@ def _read_shared_nous_state() -> Optional[Dict[str, Any]]: def _clear_shared_nous_state(reason: str) -> None: """Remove the shared Nous OAuth store after a terminal token failure.""" + if profile_strict_auth_enabled(): + return try: with _nous_shared_store_lock(): path = _nous_shared_store_path() @@ -4997,6 +5049,8 @@ def _try_import_shared_nous_state( etc.) — caller should then fall through to the normal device-code flow. """ + if profile_strict_auth_enabled(): + return None try: with _nous_shared_store_lock(timeout_seconds=max(timeout_seconds + 5.0, AUTH_LOCK_TIMEOUT_SECONDS)): shared = _read_shared_nous_state() @@ -6046,7 +6100,7 @@ def get_xai_oauth_auth_status() -> Dict[str, Any]: } -def get_api_key_provider_status(provider_id: str) -> Dict[str, Any]: +def get_api_key_provider_status(provider_id: str, env: Optional[Dict[str, Any]] = None) -> Dict[str, Any]: """Status snapshot for API-key providers (z.ai, Kimi, MiniMax).""" pconfig = PROVIDER_REGISTRY.get(provider_id) if not pconfig or pconfig.auth_type != "api_key": @@ -6054,11 +6108,11 @@ def get_api_key_provider_status(provider_id: str) -> Dict[str, Any]: api_key = "" key_source = "" - api_key, key_source = _resolve_api_key_provider_secret(provider_id, pconfig) + api_key, key_source = _resolve_api_key_provider_secret(provider_id, pconfig, env=env) env_url = "" if pconfig.base_url_env_var: - env_url = os.getenv(pconfig.base_url_env_var, "").strip() + env_url = _env_get(env, pconfig.base_url_env_var).strip() if provider_id in {"kimi-coding", "kimi-coding-cn"}: base_url = _resolve_kimi_base_url(api_key, pconfig.inference_base_url, env_url) @@ -6107,7 +6161,7 @@ def get_external_process_provider_status(provider_id: str) -> Dict[str, Any]: } -def get_auth_status(provider_id: Optional[str] = None) -> Dict[str, Any]: +def get_auth_status(provider_id: Optional[str] = None, env: Optional[Dict[str, Any]] = None) -> Dict[str, Any]: """Generic auth status dispatcher.""" target = (provider_id or get_active_provider() or "").strip().lower() if not target: @@ -6131,7 +6185,7 @@ def get_auth_status(provider_id: Optional[str] = None) -> Dict[str, Any]: # API-key providers pconfig = PROVIDER_REGISTRY.get(target) if pconfig and pconfig.auth_type == "api_key": - return get_api_key_provider_status(target) + return get_api_key_provider_status(target, env=env) # AWS SDK providers (Bedrock) — check via boto3 credential chain if pconfig and pconfig.auth_type == "aws_sdk": try: @@ -6219,7 +6273,11 @@ def _get_azure_foundry_auth_status() -> Dict[str, Any]: return info -def resolve_api_key_provider_credentials(provider_id: str) -> Dict[str, Any]: +def resolve_api_key_provider_credentials( + provider_id: str, + *, + env: Optional[Dict[str, Any]] = None, +) -> Dict[str, Any]: """Resolve API key and base URL for an API-key provider. Returns dict with: provider, api_key, base_url, source. @@ -6234,7 +6292,7 @@ def resolve_api_key_provider_credentials(provider_id: str) -> Dict[str, Any]: api_key = "" key_source = "" - api_key, key_source = _resolve_api_key_provider_secret(provider_id, pconfig) + api_key, key_source = _resolve_api_key_provider_secret(provider_id, pconfig, env=env) # No-auth LM Studio: substitute a placeholder so runtime / auxiliary_client # see the local server as configured. doctor still reports unconfigured @@ -6264,8 +6322,20 @@ def resolve_api_key_provider_credentials(provider_id: str) -> Dict[str, Any]: } -def resolve_external_process_provider_credentials(provider_id: str) -> Dict[str, Any]: +def resolve_external_process_provider_credentials( + provider_id: str, + *, + env: Optional[Dict[str, Any]] = None, +) -> Dict[str, Any]: """Resolve runtime details for local subprocess-backed providers.""" + if profile_strict_auth_enabled(env): + raise AuthError( + "External-process providers are disabled for strict routed profiles. " + "Configure provider credentials inside the routed profile instead.", + provider=provider_id, + code="profile_strict_external_process_disabled", + ) + pconfig = PROVIDER_REGISTRY.get(provider_id) if not pconfig or pconfig.auth_type != "external_process": raise AuthError( @@ -6274,16 +6344,16 @@ def resolve_external_process_provider_credentials(provider_id: str) -> Dict[str, code="invalid_provider", ) - base_url = os.getenv(pconfig.base_url_env_var, "").strip() if pconfig.base_url_env_var else "" + base_url = _env_get(env, pconfig.base_url_env_var).strip() if pconfig.base_url_env_var else "" if not base_url: base_url = pconfig.inference_base_url command = ( - os.getenv("HERMES_COPILOT_ACP_COMMAND", "").strip() - or os.getenv("COPILOT_CLI_PATH", "").strip() + _env_get(env, "HERMES_COPILOT_ACP_COMMAND").strip() + or _env_get(env, "COPILOT_CLI_PATH").strip() or "copilot" ) - raw_args = os.getenv("HERMES_COPILOT_ACP_ARGS", "").strip() + raw_args = _env_get(env, "HERMES_COPILOT_ACP_ARGS").strip() args = shlex.split(raw_args) if raw_args else ["--acp", "--stdio"] resolved_command = shutil.which(command) if command else None if not resolved_command and not base_url.startswith("acp+tcp://"): diff --git a/hermes_cli/env_loader.py b/hermes_cli/env_loader.py index c7d507d8c2f3b..882e7a7536754 100644 --- a/hermes_cli/env_loader.py +++ b/hermes_cli/env_loader.py @@ -224,7 +224,14 @@ def load_hermes_dotenv( """ loaded: list[Path] = [] - home_path = Path(hermes_home or os.getenv("HERMES_HOME", Path.home() / ".hermes")) + if not hermes_home: + try: + from hermes_constants import get_hermes_home + home_path = get_hermes_home() + except ImportError: + home_path = Path(os.getenv("HERMES_HOME", Path.home() / ".hermes")) + else: + home_path = Path(hermes_home) user_env = home_path / ".env" project_env_path = Path(project_env) if project_env else None diff --git a/hermes_cli/main.py b/hermes_cli/main.py index 11451903d1eb9..75dbc5bea9f92 100644 --- a/hermes_cli/main.py +++ b/hermes_cli/main.py @@ -10469,9 +10469,13 @@ def cmd_profile(args): check_alias_collision, create_wrapper_script, remove_wrapper_script, + audit_profile_isolation, + format_profile_isolation_audit, + write_profile_identity_marker, _is_wrapper_dir_in_path, _get_wrapper_dir, ) + from pathlib import Path from hermes_constants import display_hermes_home action = getattr(args, "profile_action", None) @@ -10774,6 +10778,28 @@ def cmd_profile(args): sys.exit(0 if ok_count == 1 else 1) sys.exit(0 if ok_count > 0 else 1) + elif action == "audit-isolation": + name = args.profile_name + try: + if getattr(args, "write_marker", False): + from hermes_cli.profiles import get_profile_dir + profile_dir = get_profile_dir(name) + safe_root = Path(getattr(args, "safe_root", None)).expanduser() if getattr(args, "safe_root", None) else profile_dir.parent + write_profile_identity_marker(name, profile_dir, safe_root, overwrite=True) + report = audit_profile_isolation( + name, + safe_root=getattr(args, "safe_root", None), + ) + if getattr(args, "json", False): + print(json.dumps(report, indent=2, sort_keys=True)) + else: + print(format_profile_isolation_audit(report)) + if report.get("status") == "FAIL": + sys.exit(1) + except Exception as exc: + print(f"Error: {exc}", file=sys.stderr) + sys.exit(1) + elif action == "show": name = args.profile_name from hermes_cli.profiles import ( @@ -12848,12 +12874,12 @@ def cmd_sessions(args): # raw file path instead. if action == "repair": from hermes_state import ( - DEFAULT_DB_PATH, _db_opens_cleanly, repair_state_db_schema, ) + from hermes_constants import get_hermes_home - db_path = DEFAULT_DB_PATH + db_path = get_hermes_home() / "state.db" if not db_path.exists(): print(f"No session database at {db_path} (nothing to repair).") return diff --git a/hermes_cli/profiles.py b/hermes_cli/profiles.py index 65d3d73dbe1d6..3ff217c094b04 100644 --- a/hermes_cli/profiles.py +++ b/hermes_cli/profiles.py @@ -28,8 +28,10 @@ import subprocess import sys from dataclasses import dataclass +from datetime import datetime, timezone +from hashlib import sha256 from pathlib import Path, PurePosixPath, PureWindowsPath -from typing import List, Optional, Tuple +from typing import Any, Dict, List, Optional, Tuple from agent.skill_utils import is_excluded_skill_path @@ -131,6 +133,444 @@ # `hermes skills install` or drop SKILL.md files into the profile's skills/. # Delete the marker file to opt back in. NO_BUNDLED_SKILLS_MARKER = ".no-bundled-skills" +PROFILE_IDENTITY_FILENAME = ".hermes_profile.json" +PROFILE_IDENTITY_VERSION = 1 + +_AUDIT_SECRET_KEY_PARTS = ( + "API", + "AUTH", + "BEARER", + "COOKIE", + "CREDENTIAL", + "KEY", + "PASSWORD", + "SECRET", + "TOKEN", +) +_AUDIT_SENSITIVE_ROOTS = ( + ".env", + "auth.json", + "config.yaml", + "SOUL.md", + "memories", + "sessions", + "logs", + "plugins", + "hooks", + "mcp", + "cron", +) +_AUDIT_MEMORY_FILES = ( + "SOUL.md", + "memories/MEMORY.md", + "memories/USER.md", +) + + +def _is_relative_to(path: Path, root: Path) -> bool: + try: + path.relative_to(root) + return True + except ValueError: + return False + + +def _has_symlink_ancestor(path: Path, stop_at: Path | None = None) -> bool: + """Return True when path or an existing ancestor is a symlink.""" + current = path + stop = stop_at.resolve(strict=False) if stop_at is not None else None + while True: + if stop is not None: + try: + if current.resolve(strict=False) == stop: + return False + except OSError: + return True + try: + if current.is_symlink(): + return True + except OSError: + return True + parent = current.parent + if parent == current: + return False + current = parent + + +def _profile_identity_payload( + profile_id: str, + profile_dir: Path, + profiles_root: Path | None = None, +) -> Dict[str, Any]: + root = profiles_root or profile_dir.parent + return { + "version": PROFILE_IDENTITY_VERSION, + "profile_id": profile_id, + "home_realpath": str(profile_dir.resolve(strict=False)), + "profiles_root_realpath": str(root.resolve(strict=False)), + "created_at": datetime.now(timezone.utc).isoformat(), + } + + +def write_profile_identity_marker( + profile_name: str, + profile_dir: Path | None = None, + profiles_root: Path | None = None, + *, + overwrite: bool = False, +) -> Path: + """Write a non-secret identity marker for a routed profile.""" + canon = normalize_profile_name(profile_name) + validate_profile_name(canon) + profile_dir = (profile_dir or get_profile_dir(canon)).expanduser() + marker = profile_dir / PROFILE_IDENTITY_FILENAME + if marker.is_symlink(): + raise ValueError(f"{PROFILE_IDENTITY_FILENAME} must not be a symlink") + if marker.exists() and not overwrite: + return marker + if not profile_dir.is_dir(): + raise FileNotFoundError(f"Profile '{canon}' does not exist at {profile_dir}") + payload = _profile_identity_payload(canon, profile_dir, profiles_root) + tmp = marker.with_name(f"{marker.name}.tmp.{os.getpid()}") + tmp.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n", encoding="utf-8") + try: + tmp.chmod(stat.S_IRUSR | stat.S_IWUSR) + except OSError: + pass + tmp.replace(marker) + return marker + + +def load_profile_identity_marker(profile_dir: Path) -> Dict[str, Any]: + marker = profile_dir / PROFILE_IDENTITY_FILENAME + if marker.is_symlink(): + raise ValueError(f"{PROFILE_IDENTITY_FILENAME} must not be a symlink") + if not marker.is_file(): + raise FileNotFoundError(f"Missing {PROFILE_IDENTITY_FILENAME}") + try: + data = json.loads(marker.read_text(encoding="utf-8")) + except Exception as exc: + raise ValueError(f"Invalid {PROFILE_IDENTITY_FILENAME}: {exc}") from exc + if not isinstance(data, dict): + raise ValueError(f"Invalid {PROFILE_IDENTITY_FILENAME}: expected object") + return data + + +def validate_profile_identity( + profile_name: str, + profile_dir: Path, + profiles_root: Path, +) -> None: + """Fail closed unless the routed profile identity matches the filesystem.""" + canon = normalize_profile_name(profile_name) + validate_profile_name(canon) + profile_dir = profile_dir.expanduser() + profiles_root = profiles_root.expanduser() + if _has_symlink_ancestor(profile_dir, profiles_root.parent): + raise ValueError("profile home ancestry must not contain symlinks") + if _has_symlink_ancestor(profiles_root, profiles_root.parent): + raise ValueError("profiles root ancestry must not contain symlinks") + home_real = profile_dir.resolve(strict=False) + root_real = profiles_root.resolve(strict=False) + if not _is_relative_to(home_real, root_real): + raise ValueError("profile home must stay inside profiles root") + if home_real == (Path.home() / ".hermes").resolve(strict=False): + raise ValueError("profile home must not be ~/.hermes") + if not home_real.is_dir(): + raise FileNotFoundError(f"profile home does not exist: {home_real}") + marker = load_profile_identity_marker(home_real) + if marker.get("version") != PROFILE_IDENTITY_VERSION: + raise ValueError(f"{PROFILE_IDENTITY_FILENAME} version mismatch") + if marker.get("profile_id") != canon: + raise ValueError(f"{PROFILE_IDENTITY_FILENAME} profile_id mismatch") + if str(marker.get("home_realpath") or "") != str(home_real): + raise ValueError(f"{PROFILE_IDENTITY_FILENAME} home_realpath mismatch") + if str(marker.get("profiles_root_realpath") or "") != str(root_real): + raise ValueError(f"{PROFILE_IDENTITY_FILENAME} profiles_root_realpath mismatch") + + +def _read_env_keyset_and_secret_flag(path: Path) -> tuple[list[str], bool]: + keys: list[str] = [] + secret_bearing = False + try: + lines = path.read_text(encoding="utf-8", errors="replace").splitlines() + except OSError: + return keys, False + for line in lines: + stripped = line.strip() + if not stripped or stripped.startswith("#") or "=" not in stripped: + continue + key, raw_value = stripped.split("=", 1) + key = key.strip() + if not key: + continue + keys.append(key) + value = raw_value.strip().strip("'\"") + if _is_secret_key_name(key) and _looks_secret_value(value): + secret_bearing = True + return sorted(set(keys)), secret_bearing + + +def _is_secret_key_name(key: str) -> bool: + lowered = str(key or "").strip().lower() + if lowered in {"auth_type", "record_key"} or lowered.endswith("_key_env"): + return False + if lowered in {"url", "base_url", "api_url", "endpoint"} or lowered.endswith("_url"): + return False + upper = lowered.upper() + return any(part in upper for part in _AUDIT_SECRET_KEY_PARTS) + + +def _looks_secret_value(value: Any) -> bool: + if not isinstance(value, str): + return False + text = value.strip().strip("'\"") + if not text: + return False + if text.lower() in { + "true", + "false", + "none", + "null", + "default", + "changeme", + "change-me", + "placeholder", + }: + return False + if re.fullmatch(r"[0-9_.,:-]+", text): + return False + return True + + +def _secret_fingerprint(value: Any) -> str | None: + if not _looks_secret_value(value): + return None + text = str(value or "").strip() + return sha256(text.encode("utf-8")).hexdigest() + + +def _read_env_secret_fingerprints(path: Path) -> tuple[set[str], list[str]]: + fingerprints: set[str] = set() + keys: list[str] = [] + try: + lines = path.read_text(encoding="utf-8", errors="replace").splitlines() + except OSError: + return fingerprints, keys + for line in lines: + stripped = line.strip() + if not stripped or stripped.startswith("#") or "=" not in stripped: + continue + key, raw_value = stripped.split("=", 1) + key = key.strip() + if not key or not _is_secret_key_name(key): + continue + value = raw_value.strip().strip("'\"") + fp = _secret_fingerprint(value) + if fp: + fingerprints.add(fp) + keys.append(key) + return fingerprints, sorted(set(keys)) + + +def _json_secret_fingerprints(value: Any, key_path: tuple[str, ...] = ()) -> set[str]: + fingerprints: set[str] = set() + if isinstance(value, dict): + for key, nested in value.items(): + fingerprints |= _json_secret_fingerprints(nested, (*key_path, str(key))) + return fingerprints + if isinstance(value, list): + for nested in value: + fingerprints |= _json_secret_fingerprints(nested, key_path) + return fingerprints + if key_path and isinstance(value, str) and _is_secret_key_name(key_path[-1]): + fp = _secret_fingerprint(value) + if fp: + fingerprints.add(fp) + return fingerprints + + +def _read_json_secret_fingerprints(path: Path) -> set[str]: + try: + data = json.loads(path.read_text(encoding="utf-8")) + except Exception: + return set() + return _json_secret_fingerprints(data) + + +def _read_yaml_secret_fingerprints(path: Path) -> set[str]: + try: + import yaml + data = yaml.safe_load(path.read_text(encoding="utf-8")) or {} + except Exception: + return set() + return _json_secret_fingerprints(data) + + +def _file_sha256(path: Path) -> str | None: + try: + return sha256(path.read_bytes()).hexdigest() + except OSError: + return None + + +def _audit_add(findings: list[dict], status: str, relpath: str, message: str, **extra: Any) -> None: + item = {"status": status, "path": relpath, "message": message} + item.update(extra) + findings.append(item) + + +def audit_profile_isolation( + profile_name: str, + *, + safe_root: str | Path | None = None, +) -> Dict[str, Any]: + """Audit routed profile isolation without printing secret values.""" + canon = normalize_profile_name(profile_name) + validate_profile_name(canon) + profiles_root = Path(safe_root).expanduser().resolve(strict=False) if safe_root else _get_profiles_root().resolve(strict=False) + profile_dir = (profiles_root / canon).resolve(strict=False) + default_home = _get_default_hermes_home().resolve(strict=False) + main_home = profiles_root.parent.resolve(strict=False) if default_home == profile_dir else default_home + findings: list[dict] = [] + + try: + validate_profile_identity(canon, profile_dir, profiles_root) + _audit_add(findings, "OK", PROFILE_IDENTITY_FILENAME, "identity marker matches profile home") + except Exception as exc: + _audit_add(findings, "FAIL", PROFILE_IDENTITY_FILENAME, str(exc)) + + for rel in _AUDIT_SENSITIVE_ROOTS: + target = profile_dir / rel + if not target.exists() and not target.is_symlink(): + continue + try: + if target.is_symlink(): + resolved = target.resolve(strict=False) + if not _is_relative_to(resolved, profile_dir): + _audit_add(findings, "FAIL", rel, "symlink points outside profile home") + continue + if target.is_file() and target.stat().st_nlink > 1: + _audit_add(findings, "FAIL", rel, "hardlink count is greater than one") + except OSError as exc: + _audit_add(findings, "FAIL", rel, f"cannot stat path: {exc}") + continue + if target.is_dir(): + for child in target.rglob("*"): + child_rel = child.relative_to(profile_dir).as_posix() + try: + if child.is_symlink(): + resolved = child.resolve(strict=False) + if not _is_relative_to(resolved, profile_dir): + _audit_add(findings, "FAIL", child_rel, "symlink points outside profile home") + continue + if child.is_file() and child.stat().st_nlink > 1: + _audit_add(findings, "FAIL", child_rel, "hardlink count is greater than one") + except OSError as exc: + _audit_add(findings, "FAIL", child_rel, f"cannot inspect path: {exc}") + + env_path = profile_dir / ".env" + main_env = main_home / ".env" + if env_path.is_file() and main_env.is_file(): + keys, secret_bearing = _read_env_keyset_and_secret_flag(env_path) + keyset_hash = sha256("\n".join(keys).encode("utf-8")).hexdigest()[:16] + profile_secret_fps, _profile_secret_keys = _read_env_secret_fingerprints(env_path) + main_secret_fps, _main_secret_keys = _read_env_secret_fingerprints(main_env) + shared_secret_fps = profile_secret_fps & main_secret_fps + if shared_secret_fps: + _audit_add( + findings, + "FAIL", + ".env", + "profile .env shares secret values with main .env", + keyset_hash=keyset_hash, + shared_secret_count=len(shared_secret_fps), + ) + elif _file_sha256(env_path) == _file_sha256(main_env): + status = "FAIL" if secret_bearing else "WARN" + _audit_add( + findings, + status, + ".env", + "profile .env is byte-identical to main .env", + keyset_hash=keyset_hash, + ) + + auth_path = profile_dir / "auth.json" + main_auth = main_home / "auth.json" + if auth_path.is_file() and main_auth.is_file() and auth_path.stat().st_size > 0: + shared_secret_fps = ( + _read_json_secret_fingerprints(auth_path) + & _read_json_secret_fingerprints(main_auth) + ) + if shared_secret_fps: + _audit_add( + findings, + "FAIL", + "auth.json", + "profile auth.json shares secret values with main auth.json", + shared_secret_count=len(shared_secret_fps), + ) + elif _file_sha256(auth_path) == _file_sha256(main_auth): + _audit_add(findings, "FAIL", "auth.json", "profile auth.json is byte-identical to main auth.json") + + config_path = profile_dir / "config.yaml" + main_config = main_home / "config.yaml" + if config_path.is_file() and main_config.is_file() and config_path.stat().st_size > 0: + shared_secret_fps = ( + _read_yaml_secret_fingerprints(config_path) + & _read_yaml_secret_fingerprints(main_config) + ) + if shared_secret_fps: + _audit_add( + findings, + "FAIL", + "config.yaml", + "profile config.yaml shares secret values with main config.yaml", + shared_secret_count=len(shared_secret_fps), + ) + + for rel in _AUDIT_MEMORY_FILES: + path = profile_dir / rel + main_path = main_home / rel + if not (path.is_file() and main_path.is_file()): + continue + try: + if path.stat().st_size == 0: + continue + except OSError: + continue + if _file_sha256(path) == _file_sha256(main_path): + _audit_add(findings, "FAIL", rel, "profile identity/memory file is non-empty and identical to main") + + if not any(f["status"] in {"FAIL", "WARN"} for f in findings): + _audit_add(findings, "OK", ".", "no isolation issues detected") + + overall = "FAIL" if any(f["status"] == "FAIL" for f in findings) else ( + "WARN" if any(f["status"] == "WARN" for f in findings) else "OK" + ) + return { + "profile": canon, + "home": profile_dir.name, + "profiles_root": profiles_root.name, + "status": overall, + "findings": findings, + } + + +def format_profile_isolation_audit(report: Dict[str, Any]) -> str: + lines = [ + f"Profile isolation audit: {report.get('profile')}", + f"Status: {report.get('status')}", + ] + for finding in report.get("findings", []): + extra = "" + if finding.get("keyset_hash"): + extra = f" keyset_hash={finding['keyset_hash']}" + lines.append( + f"- {finding.get('status')}: {finding.get('path')}: {finding.get('message')}{extra}" + ) + return "\n".join(lines) def has_bundled_skills_opt_out(profile_dir: Path) -> bool: @@ -1035,6 +1475,8 @@ def create_profile( # unit-generation paths handle gateway lifecycle. _maybe_register_gateway_service(canon) + write_profile_identity_marker(canon, profile_dir, profile_dir.parent, overwrite=True) + return profile_dir diff --git a/hermes_cli/runtime_provider.py b/hermes_cli/runtime_provider.py index c6f6db9fa7522..df688a3a63e74 100644 --- a/hermes_cli/runtime_provider.py +++ b/hermes_cli/runtime_provider.py @@ -30,11 +30,94 @@ resolve_external_process_provider_credentials, has_usable_secret, ) +import json from hermes_cli.config import get_compatible_custom_providers, load_config -from hermes_constants import OPENROUTER_BASE_URL +from hermes_constants import OPENROUTER_BASE_URL, get_default_hermes_root, get_hermes_home from utils import base_url_host_matches, base_url_hostname, env_int +def _env_get(env: Optional[Dict[str, Any]], key: str, default: str = "") -> str: + if env is None: + return os.getenv(key, default) + value = env.get(key, default) + return "" if value is None else str(value) + + +def _resolve_api_key_credentials(provider: str, env: Optional[Dict[str, Any]]) -> Dict[str, Any]: + if env is None: + return resolve_api_key_provider_credentials(provider) + return resolve_api_key_provider_credentials(provider, env=env) + + +def _scoped_auth_store_allowed(env: Optional[Dict[str, Any]]) -> bool: + if env is None: + return True + try: + profile_home = get_hermes_home().resolve(strict=False) + default_home = get_default_hermes_root().resolve(strict=False) + return profile_home != default_home + except Exception: + return False + + +def _load_pool_for_env(provider: str, env: Optional[Dict[str, Any]]) -> Optional[CredentialPool]: + if env is None: + return load_pool(provider) + if not _scoped_auth_store_allowed(env): + return None + auth_path = get_hermes_home() / "auth.json" + try: + if not auth_path.is_file(): + return CredentialPool(provider, []) + data = json.loads(auth_path.read_text(encoding="utf-8")) + pools = data.get("credential_pool") if isinstance(data, dict) else None + raw_entries = pools.get(provider) if isinstance(pools, dict) else [] + if not isinstance(raw_entries, list): + raw_entries = [] + entries = [PooledCredential.from_dict(provider, entry) for entry in raw_entries] + return CredentialPool(provider, entries) + except Exception as exc: + logger.debug("Runtime provider: could not load scoped pool for %s: %s", provider, exc) + return CredentialPool(provider, []) + + +def _scoped_provider_state_exists(provider: str, env: Optional[Dict[str, Any]]) -> bool: + if env is None: + return True + if not _scoped_auth_store_allowed(env): + return False + auth_path = get_hermes_home() / "auth.json" + try: + if not auth_path.is_file(): + return False + data = json.loads(auth_path.read_text(encoding="utf-8")) + except Exception as exc: + logger.debug("Runtime provider: could not inspect scoped auth state for %s: %s", provider, exc) + return False + providers = data.get("providers") if isinstance(data, dict) else None + state = providers.get(provider) if isinstance(providers, dict) else None + return isinstance(state, dict) and bool(state) + + +def _scoped_auth_unavailable(provider: str) -> AuthError: + return AuthError( + f"Provider '{provider}' has no profile-scoped auth credentials. " + "Configure credentials in the routed profile instead of relying on " + "gateway/global auth state.", + provider=provider, + code="no_scoped_credentials", + ) + + +def _get_named_custom_provider_for_env( + requested_provider: str, + env: Optional[Dict[str, Any]], +) -> Optional[Dict[str, Any]]: + if env is None: + return _get_named_custom_provider(requested_provider) + return _get_named_custom_provider(requested_provider, env=env) + + def _getenv(name: str, default: str = "") -> str: """Profile-scoped replacement for ``os.getenv`` on credential/provider reads. @@ -115,7 +198,7 @@ def _detect_api_mode_for_url(base_url: str) -> Optional[str]: return None -def _host_derived_api_key(base_url: str) -> str: +def _host_derived_api_key(base_url: str, env: Optional[Dict[str, Any]] = None) -> str: """Look up `_API_KEY` in the env, derived from the base URL host. Examples: @@ -435,7 +518,11 @@ def _resolve_runtime_from_pool_entry( } -def resolve_requested_provider(requested: Optional[str] = None) -> str: +def resolve_requested_provider( + requested: Optional[str] = None, + *, + env: Optional[Dict[str, Any]] = None, +) -> str: """Resolve provider request from explicit arg, config, then env.""" if requested and requested.strip(): return requested.strip().lower() @@ -459,13 +546,14 @@ def _try_resolve_from_custom_pool( provider_label: str, api_mode_override: Optional[str] = None, provider_name: Optional[str] = None, + env: Optional[Dict[str, Any]] = None, ) -> Optional[Dict[str, Any]]: """Check if a credential pool exists for a custom endpoint and return a runtime dict if so.""" pool_key = get_custom_provider_pool_key(base_url, provider_name=provider_name) if not pool_key: return None try: - pool = load_pool(pool_key) + pool = _load_pool_for_env(pool_key, env) if not pool.has_credentials(): return None entry = pool.select() @@ -501,7 +589,11 @@ def _lift_max_output_tokens(entry: Dict[str, Any], result: Dict[str, Any]) -> No return -def _get_named_custom_provider(requested_provider: str) -> Optional[Dict[str, Any]]: +def _get_named_custom_provider( + requested_provider: str, + *, + env: Optional[Dict[str, Any]] = None, +) -> Optional[Dict[str, Any]]: requested_norm = _normalize_custom_provider_name(requested_provider or "") if not requested_norm: return None @@ -525,7 +617,7 @@ def _get_named_custom_provider(requested_provider: str) -> Optional[Dict[str, An return None if requested_norm != "custom" and not requested_norm.startswith("custom:"): try: - canonical = auth_mod.resolve_provider(requested_norm) + canonical = auth_mod.resolve_provider(requested_norm, env=env) except AuthError: pass else: @@ -551,8 +643,11 @@ def _get_named_custom_provider(requested_provider: str) -> Optional[Dict[str, An # Match exact name or normalized name name_norm = _normalize_custom_provider_name(ep_name) # Resolve the API key from the env var name stored in key_env - key_env = str(entry.get("key_env", "") or "").strip() - resolved_api_key = _getenv(key_env, "").strip() if key_env else "" + key_env = str(entry.get("key_env") or entry.get("api_key_env") or "").strip() + if env is not None: + resolved_api_key = _env_get(env, key_env).strip() if key_env else "" + else: + resolved_api_key = _getenv(key_env, "").strip() if key_env else "" # Fall back to inline api_key when key_env is absent or unresolvable if not resolved_api_key: resolved_api_key = str(entry.get("api_key", "") or "").strip() @@ -802,6 +897,7 @@ def _resolve_named_custom_runtime( requested_provider: str, explicit_api_key: Optional[str] = None, explicit_base_url: Optional[str] = None, + env: Optional[Dict[str, Any]] = None, ) -> Optional[Dict[str, Any]]: # Bare `provider="custom"` with an explicit base_url (e.g. propagated # from a `model_aliases:` direct-alias resolution) — build a runtime @@ -825,7 +921,7 @@ def _resolve_named_custom_runtime( # Check credential pool first — mirrors the named-custom-provider path # so bare `provider: custom` with a configured custom_providers entry # also gets its api_key from the pool instead of env var fallbacks. - pool_result = _try_resolve_from_custom_pool(base_url, "custom", None) + pool_result = _try_resolve_from_custom_pool(base_url, "custom", None, env=env) if pool_result: pool_result["source"] = "direct-alias" return pool_result @@ -854,7 +950,7 @@ def _resolve_named_custom_runtime( "requested_provider": requested_provider, } - custom_provider = _get_named_custom_provider(requested_provider) + custom_provider = _get_named_custom_provider_for_env(requested_provider, env) if not custom_provider: return None @@ -866,7 +962,13 @@ def _resolve_named_custom_runtime( return None # Check if a credential pool exists for this custom endpoint - pool_result = _try_resolve_from_custom_pool(base_url, "custom", custom_provider.get("api_mode"), provider_name=custom_provider.get("name")) + pool_result = _try_resolve_from_custom_pool( + base_url, + "custom", + custom_provider.get("api_mode"), + provider_name=custom_provider.get("name"), + env=env, + ) if pool_result: # Propagate the model name even when using pooled credentials — # the pool doesn't know about the custom_providers model field. @@ -925,6 +1027,7 @@ def _resolve_openrouter_runtime( requested_provider: str, explicit_api_key: Optional[str] = None, explicit_base_url: Optional[str] = None, + env: Optional[Dict[str, Any]] = None, ) -> Dict[str, Any]: model_cfg = _get_model_config() cfg_base_url = model_cfg.get("base_url") if isinstance(model_cfg.get("base_url"), str) else "" @@ -1042,6 +1145,7 @@ def _resolve_openrouter_runtime( pool_result = _try_resolve_from_custom_pool( base_url, effective_provider, _parse_api_mode(model_cfg.get("api_mode")), provider_name=requested_provider if requested_norm != "custom" else None, + env=env, ) if pool_result: return pool_result @@ -1067,6 +1171,7 @@ def _resolve_azure_foundry_runtime( explicit_api_key: Optional[str] = None, explicit_base_url: Optional[str] = None, target_model: Optional[str] = None, + env: Optional[Dict[str, Any]] = None, ) -> Dict[str, Any]: """Resolve an Azure Foundry runtime entry. @@ -1236,6 +1341,7 @@ def _resolve_explicit_runtime( model_cfg: Dict[str, Any], explicit_api_key: Optional[str] = None, explicit_base_url: Optional[str] = None, + env: Optional[Dict[str, Any]] = None, ) -> Optional[Dict[str, Any]]: explicit_api_key = str(explicit_api_key or "").strip() explicit_base_url = str(explicit_base_url or "").strip().rstrip("/") @@ -1250,9 +1356,15 @@ def _resolve_explicit_runtime( base_url = explicit_base_url or cfg_base_url or "https://api.anthropic.com" api_key = explicit_api_key if not api_key: - from agent.anthropic_adapter import resolve_anthropic_token - - api_key = resolve_anthropic_token() + if env is None: + from agent.anthropic_adapter import resolve_anthropic_token + api_key = resolve_anthropic_token() + else: + api_key = ( + _env_get(env, "ANTHROPIC_API_KEY").strip() + or _env_get(env, "ANTHROPIC_TOKEN").strip() + or _env_get(env, "CLAUDE_CODE_OAUTH_TOKEN").strip() + ) if not api_key: raise AuthError( "No Anthropic credentials found. Set ANTHROPIC_TOKEN or ANTHROPIC_API_KEY, " @@ -1387,6 +1499,7 @@ def resolve_runtime_provider( explicit_api_key: Optional[str] = None, explicit_base_url: Optional[str] = None, target_model: Optional[str] = None, + env: Optional[Dict[str, Any]] = None, ) -> Dict[str, Any]: """Resolve runtime provider credentials for agent execution. @@ -1442,6 +1555,7 @@ def resolve_runtime_provider( explicit_api_key=explicit_api_key, explicit_base_url=explicit_base_url, target_model=target_model, + env=env, ) return azure_runtime @@ -1449,6 +1563,7 @@ def resolve_runtime_provider( requested_provider=requested_provider, explicit_api_key=explicit_api_key, explicit_base_url=explicit_base_url, + env=env, ) if custom_runtime: custom_runtime["requested_provider"] = requested_provider @@ -1458,6 +1573,7 @@ def resolve_runtime_provider( requested_provider, explicit_api_key=explicit_api_key, explicit_base_url=explicit_base_url, + env=env, ) model_cfg = _get_model_config() explicit_runtime = _resolve_explicit_runtime( @@ -1466,6 +1582,7 @@ def resolve_runtime_provider( model_cfg=model_cfg, explicit_api_key=explicit_api_key, explicit_base_url=explicit_base_url, + env=env, ) if explicit_runtime: return explicit_runtime @@ -1491,7 +1608,7 @@ def resolve_runtime_provider( ) try: - pool = load_pool(provider) if should_use_pool else None + pool = _load_pool_for_env(provider, env) if should_use_pool else None except Exception: pool = None if pool and pool.has_credentials(): @@ -1570,6 +1687,8 @@ def resolve_runtime_provider( if provider == "openai-codex": try: + if not _scoped_provider_state_exists(provider, env): + raise _scoped_auth_unavailable(provider) creds = resolve_codex_runtime_credentials() return { "provider": "openai-codex", @@ -1590,6 +1709,8 @@ def resolve_runtime_provider( if provider == "xai-oauth": try: + if not _scoped_provider_state_exists(provider, env): + raise _scoped_auth_unavailable(provider) creds = resolve_xai_oauth_runtime_credentials() return { "provider": "xai-oauth", @@ -1608,6 +1729,8 @@ def resolve_runtime_provider( if provider == "qwen-oauth": try: + if not _scoped_provider_state_exists(provider, env): + raise _scoped_auth_unavailable(provider) creds = resolve_qwen_runtime_credentials() return { "provider": "qwen-oauth", @@ -1702,8 +1825,15 @@ def resolve_runtime_provider( "config.yaml model section at a custom env var." ) else: - from agent.anthropic_adapter import resolve_anthropic_token - token = resolve_anthropic_token() + if env is None: + from agent.anthropic_adapter import resolve_anthropic_token + token = resolve_anthropic_token() + else: + token = ( + _env_get(env, "ANTHROPIC_API_KEY").strip() + or _env_get(env, "ANTHROPIC_TOKEN").strip() + or _env_get(env, "CLAUDE_CODE_OAUTH_TOKEN").strip() + ) if not token: raise AuthError( "No Anthropic credentials found. Set ANTHROPIC_TOKEN or ANTHROPIC_API_KEY, " @@ -1731,7 +1861,16 @@ def resolve_runtime_provider( # Lambda execution roles, SSO, and other implicit sources that our # env-var check can't detect. is_explicit = requested_provider in {"bedrock", "aws", "aws-bedrock", "amazon-bedrock", "amazon"} - if not is_explicit and not has_aws_credentials(): + scoped_aws_source = resolve_aws_auth_env_var(env) if env is not None else None + if env is not None and not scoped_aws_source: + raise AuthError( + "No profile-scoped AWS credentials found for Bedrock. Add AWS_ACCESS_KEY_ID + " + "AWS_SECRET_ACCESS_KEY, AWS_PROFILE, AWS_BEARER_TOKEN_BEDROCK, or another " + "Bedrock-supported AWS credential hint to the routed profile .env.", + provider=provider, + code="no_scoped_aws_credentials", + ) + if env is None and not is_explicit and not has_aws_credentials(): raise AuthError( "No AWS credentials found for Bedrock. Configure one of:\n" " - AWS_ACCESS_KEY_ID + AWS_SECRET_ACCESS_KEY\n" @@ -1743,8 +1882,8 @@ def resolve_runtime_provider( # Read bedrock-specific config from config.yaml _bedrock_cfg = load_config().get("bedrock", {}) # Region priority: config.yaml bedrock.region → env var → us-east-1 - region = (_bedrock_cfg.get("region") or "").strip() or resolve_bedrock_region() - auth_source = resolve_aws_auth_env_var() or "aws-sdk-default-chain" + region = (_bedrock_cfg.get("region") or "").strip() or resolve_bedrock_region(env) + auth_source = scoped_aws_source or resolve_aws_auth_env_var() or "aws-sdk-default-chain" # Build guardrail config if configured _gr = _bedrock_cfg.get("guardrail", {}) guardrail_config = None @@ -1791,7 +1930,7 @@ def resolve_runtime_provider( # API-key providers (z.ai/GLM, Kimi, MiniMax, MiniMax-CN) pconfig = PROVIDER_REGISTRY.get(provider) if pconfig and pconfig.auth_type == "api_key": - creds = resolve_api_key_provider_credentials(provider) + creds = _resolve_api_key_credentials(provider, env) # Honour model.base_url from config.yaml when the configured provider # matches this provider — mirrors the Anthropic path above. Without # this, users who set model.base_url to e.g. api.minimaxi.com/anthropic @@ -1847,6 +1986,7 @@ def resolve_runtime_provider( requested_provider=requested_provider, explicit_api_key=explicit_api_key, explicit_base_url=explicit_base_url, + env=env, ) runtime["requested_provider"] = requested_provider return runtime diff --git a/hermes_cli/subcommands/profile.py b/hermes_cli/subcommands/profile.py index d812fadf97158..4668eb5c76ac1 100644 --- a/hermes_cli/subcommands/profile.py +++ b/hermes_cli/subcommands/profile.py @@ -103,6 +103,27 @@ def build_profile_parser(subparsers, *, cmd_profile: Callable) -> None: help="With --auto, run on every profile missing a description", ) + profile_audit = profile_subparsers.add_parser( + "audit-isolation", + help="Audit routed profile isolation without printing secrets", + ) + profile_audit.add_argument("profile_name", help="Profile to audit") + profile_audit.add_argument( + "--safe-root", + default=None, + help="Expected topic_profiles_safe_root for routed profiles", + ) + profile_audit.add_argument( + "--write-marker", + action="store_true", + help="Write or refresh .hermes_profile.json before auditing", + ) + profile_audit.add_argument( + "--json", + action="store_true", + help="Print machine-readable audit output", + ) + profile_show = profile_subparsers.add_parser("show", help="Show profile details") profile_show.add_argument("profile_name", help="Profile to show") diff --git a/hermes_state.py b/hermes_state.py index a7938f7167f42..dbc3a9f052103 100644 --- a/hermes_state.py +++ b/hermes_state.py @@ -768,7 +768,9 @@ class SessionDB: _CHECKPOINT_EVERY_N_WRITES = 50 def __init__(self, db_path: Path = None, read_only: bool = False): - self.db_path = db_path or DEFAULT_DB_PATH + if db_path is None: + db_path = get_hermes_home() / "state.db" + self.db_path = db_path self.read_only = read_only self._lock = threading.Lock() diff --git a/model_tools.py b/model_tools.py index fb16bb745ac76..9ebaae2cc9cb9 100644 --- a/model_tools.py +++ b/model_tools.py @@ -144,7 +144,9 @@ def _run_in_worker(): worker_loop.close() pool = concurrent.futures.ThreadPoolExecutor(max_workers=1) - future = pool.submit(_run_in_worker) + from contextvars import copy_context + ctx = copy_context() + future = pool.submit(ctx.run, _run_in_worker) try: return future.result(timeout=300) except concurrent.futures.TimeoutError: diff --git a/tests/agent/test_auxiliary_client.py b/tests/agent/test_auxiliary_client.py index 7637b06d9f1f6..cd18d57fe457d 100644 --- a/tests/agent/test_auxiliary_client.py +++ b/tests/agent/test_auxiliary_client.py @@ -2892,7 +2892,7 @@ def test_call_llm_refreshes_codex_on_401_for_vision(self): ) assert resp.choices[0].message.content == "fresh-sync" - mock_refresh.assert_called_once_with("openai-codex") + mock_refresh.assert_called_once_with("openai-codex", env=None) def test_call_llm_refreshes_codex_on_401_for_non_vision(self): stale_client = MagicMock() @@ -2916,7 +2916,7 @@ def test_call_llm_refreshes_codex_on_401_for_non_vision(self): ) assert resp.choices[0].message.content == "fresh-non-vision" - mock_refresh.assert_called_once_with("openai-codex") + mock_refresh.assert_called_once_with("openai-codex", env=None) assert stale_client.chat.completions.create.call_count == 1 assert fresh_client.chat.completions.create.call_count == 1 @@ -2942,7 +2942,7 @@ def test_call_llm_refreshes_anthropic_on_401_for_non_vision(self): ) assert resp.choices[0].message.content == "fresh-anthropic" - mock_refresh.assert_called_once_with("anthropic") + mock_refresh.assert_called_once_with("anthropic", env=None) assert stale_client.chat.completions.create.call_count == 1 assert fresh_client.chat.completions.create.call_count == 1 @@ -2971,7 +2971,7 @@ async def test_async_call_llm_refreshes_codex_on_401_for_vision(self): ) assert resp.choices[0].message.content == "fresh-async" - mock_refresh.assert_called_once_with("openai-codex") + mock_refresh.assert_called_once_with("openai-codex", env=None) def test_refresh_provider_credentials_force_refreshes_anthropic_oauth_and_evicts_cache(self, monkeypatch): stale_client = MagicMock() @@ -3026,7 +3026,7 @@ async def test_async_call_llm_refreshes_anthropic_on_401_for_non_vision(self): ) assert resp.choices[0].message.content == "fresh-async-anthropic" - mock_refresh.assert_called_once_with("anthropic") + mock_refresh.assert_called_once_with("anthropic", env=None) assert stale_client.chat.completions.create.await_count == 1 assert fresh_client.chat.completions.create.await_count == 1 @@ -3332,7 +3332,7 @@ def test_kimi_coding_skipped_falls_through_to_openrouter(self, monkeypatch): "agent.auxiliary_client.resolve_provider_client", rpc_mock, ) - def fake_strict(provider, model=None): + def fake_strict(provider, model=None, env=None): if provider == "openrouter": return fake_or_client, "google/gemini-3-flash-preview" if provider == "nous": @@ -3368,7 +3368,7 @@ def test_kimi_coding_cn_skipped_too(self, monkeypatch): ) monkeypatch.setattr( "agent.auxiliary_client._resolve_strict_vision_backend", - lambda p, m=None: (fake_or_client, "gemini") + lambda p, m=None, env=None: (fake_or_client, "gemini") if p == "openrouter" else (None, None), ) diff --git a/tests/gateway/test_model_command_flat_string_config.py b/tests/gateway/test_model_command_flat_string_config.py index 9934d9806b19d..00faf7753673b 100644 --- a/tests/gateway/test_model_command_flat_string_config.py +++ b/tests/gateway/test_model_command_flat_string_config.py @@ -160,16 +160,13 @@ async def test_model_global_persists_when_config_has_proper_dict_model(tmp_path, @pytest.mark.asyncio async def test_model_no_flag_persists_by_default(tmp_path, monkeypatch): - """A plain ``/model X`` (no --global) now persists to config.yaml. - - This is the user-facing fix: switching models in one session survives - into the next without re-typing the switch every time. - """ + """A plain ``/model X`` (no --global) on gateway now persists to topic_models.json instead of config.yaml.""" cfg_path = _setup_isolated_home( tmp_path, monkeypatch, {"default": "old-model", "provider": "openai-codex"}, ) + hermes_home = tmp_path / ".hermes" result = await _make_runner()._handle_model_command( _make_event("/model gpt-5.5") @@ -178,7 +175,14 @@ async def test_model_no_flag_persists_by_default(tmp_path, monkeypatch): assert result is not None assert "gpt-5.5" in result written = yaml.safe_load(cfg_path.read_text(encoding="utf-8")) - assert written["model"]["default"] == "gpt-5.5" + assert written["model"]["default"] == "old-model" + + import json + topic_models_path = hermes_home / "topic_models.json" + assert topic_models_path.exists() + topic_models = json.loads(topic_models_path.read_text(encoding="utf-8")) + assert "agent:main:telegram:dm:12345" in topic_models + assert topic_models["agent:main:telegram:dm:12345"]["model"] == "gpt-5.5" @pytest.mark.asyncio diff --git a/tests/gateway/test_model_picker_persist.py b/tests/gateway/test_model_picker_persist.py index ca9498389b1bb..5bafddd36aeb7 100644 --- a/tests/gateway/test_model_picker_persist.py +++ b/tests/gateway/test_model_picker_persist.py @@ -161,7 +161,7 @@ async def test_picker_tap_persists_by_default(tmp_path, monkeypatch, seed_model) adapter = _FakePickerAdapter() cfg_path = _setup_isolated_home(tmp_path, monkeypatch, seed_model) - confirmation = await _drive_picker(_make_runner(adapter), _make_event("/model")) + confirmation = await _drive_picker(_make_runner(adapter), _make_event("/model --global")) assert confirmation is not None assert "gpt-5.5" in confirmation diff --git a/tests/gateway/test_profile_isolation_rework_suite.py b/tests/gateway/test_profile_isolation_rework_suite.py new file mode 100644 index 0000000000000..7ab2598092e40 --- /dev/null +++ b/tests/gateway/test_profile_isolation_rework_suite.py @@ -0,0 +1,440 @@ +import pytest +import os +import json +import asyncio +from pathlib import Path +from unittest.mock import patch, MagicMock + +from gateway.run import GatewayRunner +from gateway.session import SessionSource, Platform +from gateway.config import GatewayConfig +from hermes_constants import get_hermes_home, set_hermes_home_override, reset_hermes_home_override +from gateway.session_context import set_session_vars, clear_session_vars +from tools.file_tools import write_file_tool +from tools.memory_tool import get_memory_dir, MemoryStore +from tools.mcp_tool import ( + _get_active_server, + _get_mcp_config_fingerprint, + _servers, + MCPServerTask, +) + +@pytest.fixture +def test_env(tmp_path, monkeypatch): + hermes_home = tmp_path / ".hermes" + hermes_home.mkdir() + + # Create profiles + profiles_dir = hermes_home / "profiles" + profiles_dir.mkdir() + + for name in ("profilea", "profileb"): + pdir = profiles_dir / name + pdir.mkdir() + (pdir / "SOUL.md").write_text(f"I am {name}", encoding="utf-8") + (pdir / "memories").mkdir() + (pdir / "sessions").mkdir() + (pdir / ".env").write_text(f"GITHUB_TOKEN={name}_token\n", encoding="utf-8") + + # Create state.db schema + from hermes_state import SessionDB + db = SessionDB(db_path=pdir / "state.db") + db.close() # creates tables + + # Create identity marker + marker = { + "version": "1.0", + "profile_id": name, + "home_realpath": str(pdir.resolve()), + "profiles_root_realpath": str(profiles_dir.resolve()), + } + (pdir / ".profile_identity.json").write_text(json.dumps(marker), encoding="utf-8") + + # Write config.yaml + cfg_content = """ +mcp_servers: + github: + command: node + env: + TOKEN: ${GITHUB_TOKEN} +""" + (pdir / "config.yaml").write_text(cfg_content, encoding="utf-8") + + (hermes_home / "SOUL.md").write_text("I am main agent", encoding="utf-8") + + topic_profiles = { + "telegram:dm:111:222": "profilea", + "telegram:dm:111:333": "profileb", + } + with open(hermes_home / "topic_profiles.json", "w", encoding="utf-8") as f: + json.dump(topic_profiles, f) + + monkeypatch.setenv("HERMES_HOME", str(hermes_home)) + with patch("gateway.run._hermes_home", hermes_home): + yield hermes_home + +def test_g2_acceptance_real_writes(test_env): + """G2 acceptance: routed scope, write_file to SOUL.md goes to profiles//SOUL.md.""" + profile_home = test_env / "profiles" / "profilea" + + # Pre-create session record in profilea's state.db + from hermes_state import SessionDB + db = SessionDB(db_path=profile_home / "state.db") + db.create_session("session123", "telegram") + db.close() + + from gateway.run import _profile_runtime_scope + with _profile_runtime_scope(profile_home): + tokens = set_session_vars( + session_id="session123", + agent_hermes_home=str(profile_home), + agent_profile="profilea", + ) + try: + res = write_file_tool("SOUL.md", "New SOUL content") + assert "New SOUL content" in res or "written" in res or "success" in res or res + + assert (profile_home / "SOUL.md").read_text(encoding="utf-8") == "New SOUL content" + assert (test_env / "SOUL.md").read_text(encoding="utf-8") == "I am main agent" + + # Check hard-guard: write outside profile should fail + res2 = write_file_tool("../../SOUL.md", "Illegal write") + res2_dict = json.loads(res2) + assert "error" in res2_dict or res2_dict.get("success") is False + finally: + clear_session_vars(tokens) + +def test_h1_write_destination(test_env): + """H1 write-destination: creating session in routed-turn writes to profiles//state.db.""" + runner = GatewayRunner(config=GatewayConfig(platforms={})) + runner._normalize_source_for_session_key = lambda src: src + + source = SessionSource( + platform=Platform.TELEGRAM, + chat_id="111", + chat_type="dm", + thread_id="222", + ) + + profile_home = test_env / "profiles" / "profilea" + + # Pre-create state.db for profilea + from hermes_state import SessionDB + db = SessionDB(db_path=profile_home / "state.db") + db.create_session("session123", "telegram") + db.close() + + from gateway.run import _profile_runtime_scope + with _profile_runtime_scope(profile_home): + profile_db = profile_home / "state.db" + assert profile_db.exists() + + db = SessionDB(db_path=profile_db) + loaded = db.get_session("session123") + assert loaded is not None + db.close() + + global_db = test_env / "state.db" + if global_db.exists(): + g_db = SessionDB(db_path=global_db) + assert g_db.get_session("session123") is None + g_db.close() + +def test_soul_identity_prompt(test_env): + """SOUL identity: agent prompt contains profile's SOUL, modifying it busts signature.""" + runner = GatewayRunner(config=GatewayConfig(platforms={})) + profile_home = test_env / "profiles" / "profilea" + + from gateway.run import _profile_runtime_scope + with _profile_runtime_scope(profile_home): + sig1 = runner._agent_config_signature( + model="gpt-4", + runtime={}, + enabled_toolsets=[], + ephemeral_prompt="", + ) + + soul_file = profile_home / "SOUL.md" + soul_file.write_text("updated identity content", encoding="utf-8") + + stat = soul_file.stat() + os.utime(soul_file, (stat.st_atime, stat.st_mtime + 5.0)) + + sig2 = runner._agent_config_signature( + model="gpt-4", + runtime={}, + enabled_toolsets=[], + ephemeral_prompt="", + ) + assert sig1 != sig2 + +def test_mcp_real_config_isolation(test_env): + """MCP real-config: B3 acceptance. Two profiles, different env/tokens, separate connections.""" + mock_connects = [] + async def fake_connect(name, config): + mock_connects.append((name, config)) + server = MCPServerTask(name) + server.session = MagicMock() + server._tools = [] + return server + + def fake_run(coro_or_factory, timeout=30): + coro = coro_or_factory() if callable(coro_or_factory) else coro_or_factory + return asyncio.run(coro) + + with patch("tools.mcp_tool._connect_server", side_effect=fake_connect), \ + patch("tools.mcp_tool._MCP_AVAILABLE", True), \ + patch("tools.mcp_tool._ensure_mcp_loop"), \ + patch("tools.mcp_tool._run_on_mcp_loop", side_effect=fake_run): + + profile_home_a = test_env / "profiles" / "profilea" + from gateway.run import _profile_runtime_scope + with _profile_runtime_scope(profile_home_a): + srv_a = _get_active_server("github") + assert srv_a is not None + + profile_home_b = test_env / "profiles" / "profileb" + with _profile_runtime_scope(profile_home_b): + srv_b = _get_active_server("github") + assert srv_b is not None + + assert len(mock_connects) >= 2 + configs = [c for n, c in mock_connects if n == "github"] + assert configs[0]["env"]["TOKEN"] == "profilea_token" + assert configs[1]["env"]["TOKEN"] == "profileb_token" + +def test_memory_isolation(test_env): + """Memory isolation: routed-turn writes MEMORY.md to profiles//memories/.""" + profile_home = test_env / "profiles" / "profilea" + + from gateway.run import _profile_runtime_scope + with _profile_runtime_scope(profile_home): + tokens = set_session_vars( + session_id="session123", + agent_hermes_home=str(profile_home), + agent_profile="profilea", + ) + try: + mem_dir = get_memory_dir() + assert mem_dir.resolve() == (profile_home / "memories").resolve() + assert get_memory_dir("profileb").resolve() == (test_env / "profiles" / "profileb" / "memories").resolve() + + store = MemoryStore() + store.add("memory", "Hello from profilea") + + mem_file = profile_home / "memories" / "MEMORY.md" + assert mem_file.exists() + assert "Hello from profilea" in mem_file.read_text(encoding="utf-8") + + global_mem = test_env / "memories" / "MEMORY.md" + if global_mem.exists(): + assert "Hello from profilea" not in global_mem.read_text(encoding="utf-8") + finally: + clear_session_vars(tokens) + +def test_phase_4_auth_fail_closed(test_env): + """Phase-4 auth fail-closed: missing scoped credential in routed session returns None/fails.""" + from agent.secret_scope import set_multiplex_active, get_secret + set_multiplex_active(True) + try: + profile_home = test_env / "profiles" / "profilea" + + from gateway.run import _profile_runtime_scope + with _profile_runtime_scope(profile_home): + assert get_secret("GITHUB_TOKEN") == "profilea_token" + val = get_secret("SOME_NON_EXISTENT_KEY") + assert val is None or val == "" + finally: + set_multiplex_active(False) + +def test_cross_profile_process_control(test_env): + """Cross-profile process control: Profile A cannot retrieve Profile B processes.""" + from tools.process_registry import process_registry + + profile_home_a = test_env / "profiles" / "profilea" + profile_home_b = test_env / "profiles" / "profileb" + + from tools.process_registry import ProcessSession + sess = ProcessSession( + id="proc123", + command="sleep 10", + cwd="/tmp", + task_id="task1", + session_key="session1", + agent_profile="profileb", + agent_hermes_home=str(profile_home_b), + ) + process_registry._running["proc123"] = sess + + tokens = set_session_vars( + session_id="session123", + agent_hermes_home=str(profile_home_a), + agent_profile="profilea", + ) + try: + proc = process_registry._get_for_current_scope("proc123") + assert proc is None + finally: + clear_session_vars(tokens) + +def test_zero_regression_full_turn(test_env): + """Zero-regression full-turn: unrouted topic uses global HERMES_HOME.""" + runner = GatewayRunner(config=GatewayConfig(platforms={})) + runner._normalize_source_for_session_key = lambda src: src + + source = SessionSource( + platform=Platform.TELEGRAM, + chat_id="111", + chat_type="dm", + thread_id="999", + ) + + routed = runner._routed_profile_for_source(source) + assert routed is None + + profile_home = runner._resolve_profile_home_for_source(source) + assert profile_home.resolve() == test_env.resolve() + + +def test_g2_fallback_when_session_id_none(test_env): + """Test Tail 1 fallback: when session_id is None, get_profile_home_for_session resolves active profile home-override.""" + from tools.file_tools import get_profile_home_for_session, write_file_tool + from gateway.run import _profile_runtime_scope + + profile_home = test_env / "profiles" / "profilea" + + # 1. Under profile scope, session_id is None but fallback resolves profile_home + with _profile_runtime_scope(profile_home): + resolved = get_profile_home_for_session(None) + assert resolved is not None + assert resolved.resolve() == profile_home.resolve() + + # Verify hard guard blocks writing outside profile-home even without session_id + res = write_file_tool("../../SOUL.md", "Illegal write") + res_dict = json.loads(res) + assert "error" in res_dict or res_dict.get("success") is False + assert (test_env / "SOUL.md").read_text(encoding="utf-8") == "I am main agent" + + # 2. Under default scope, get_profile_home_for_session(None) returns None and allows normal writes + resolved_default = get_profile_home_for_session(None) + assert resolved_default is None + + +def test_reload_mcp_scoped(test_env): + """Test Tail 2: /reload-mcp under routed scope reloads config from correct profile's home directory.""" + runner = GatewayRunner(config=GatewayConfig(platforms={})) + + homes_seen = [] + + from hermes_constants import get_hermes_home + def fake_load_config(): + homes_seen.append(get_hermes_home().resolve()) + return {} + + def fake_shutdown(): + homes_seen.append(get_hermes_home().resolve()) + + def fake_discover(): + homes_seen.append(get_hermes_home().resolve()) + return [] + + profile_home = test_env / "profiles" / "profilea" + + from gateway.run import _profile_runtime_scope + with _profile_runtime_scope(profile_home), \ + patch("tools.mcp_tool._load_mcp_config", side_effect=fake_load_config), \ + patch("tools.mcp_tool.shutdown_mcp_servers", side_effect=fake_shutdown), \ + patch("tools.mcp_tool.discover_mcp_tools", side_effect=fake_discover): + + event = MagicMock() + event.source.platform = Platform.TELEGRAM + event.source.chat_id = "111" + event.source.thread_id = "222" + + asyncio.run(runner._execute_mcp_reload(event)) + + assert len(homes_seen) > 0 + for h in homes_seen: + assert h == profile_home.resolve() + + +def test_write_file_guard_on_resolve_error(test_env): + """Test Tail 3: write_file still applies hard-guard on resolve error.""" + profile_home = test_env / "profiles" / "profilea" + global_soul = test_env / "SOUL.md" + + # Pre-create session record in profilea's state.db + from hermes_state import SessionDB + db = SessionDB(db_path=profile_home / "state.db") + db.create_session("session123", "telegram") + db.close() + + # We patch _resolve_path_for_task to raise an exception + with patch("tools.file_tools._resolve_path_for_task", side_effect=ValueError("Simulated resolve error")): + # Call write_file_tool attempting to write to global SOUL.md under a session + from gateway.run import _profile_runtime_scope + with _profile_runtime_scope(profile_home): + tokens = set_session_vars( + session_id="session123", + agent_hermes_home=str(profile_home), + agent_profile="profilea", + ) + try: + res = write_file_tool(str(global_soul), "Malicious SOUL override", session_id="session123") + assert "Refusing to write to the global SOUL.md" in res or "error" in res + + # Verify it was NOT written + assert global_soul.read_text(encoding="utf-8") == "I am main agent" + finally: + clear_session_vars(tokens) + + +def test_g2_hard_guard_blocks_outside_profile_home(test_env): + """Verify expanded hard-guard blocks write to any path outside profile_home except tmp and cwd.""" + profile_home = test_env / "profiles" / "profilea" + + # Pre-create session record in profilea's state.db + from hermes_state import SessionDB + db = SessionDB(db_path=profile_home / "state.db") + db.create_session("session123", "telegram") + db.close() + + from gateway.run import _profile_runtime_scope + with _profile_runtime_scope(profile_home): + tokens = set_session_vars( + session_id="session123", + agent_hermes_home=str(profile_home), + agent_profile="profilea", + ) + try: + # 1. Writing inside profile home should succeed + in_profile_path = profile_home / "some_file.txt" + res_ok = write_file_tool(str(in_profile_path), "in profile content", session_id="session123") + assert "error" not in res_ok + assert in_profile_path.read_text(encoding="utf-8") == "in profile content" + + # 2. Writing inside temp directory should succeed + import tempfile + temp_file = Path(tempfile.gettempdir()) / "hermes_test_temp.txt" + res_temp = write_file_tool(str(temp_file), "temp content", session_id="session123") + assert "error" not in res_temp + assert temp_file.read_text(encoding="utf-8") == "temp content" + + # 3. Writing inside active workspace CWD should succeed + cwd_file = Path(os.getcwd()) / "hermes_test_cwd.txt" + # Ensure it doesn't already exist or clean up after + try: + res_cwd = write_file_tool(str(cwd_file), "cwd content", session_id="session123") + assert "error" not in res_cwd + assert cwd_file.read_text(encoding="utf-8") == "cwd content" + finally: + if cwd_file.exists(): + cwd_file.unlink() + + # 4. Writing to an outside/forbidden path should be blocked + forbidden_path = test_env / "random_file_outside_profile.txt" + res_blocked = write_file_tool(str(forbidden_path), "forbidden content", session_id="session123") + assert "Refusing to write to path outside profile home" in res_blocked or "error" in res_blocked + assert not forbidden_path.exists() + finally: + clear_session_vars(tokens) diff --git a/tests/gateway/test_reasoning_command.py b/tests/gateway/test_reasoning_command.py index 09600fb6f5a10..016293ec7f5df 100644 --- a/tests/gateway/test_reasoning_command.py +++ b/tests/gateway/test_reasoning_command.py @@ -78,7 +78,8 @@ async def test_reasoning_in_help_output(self): assert "level" in result and "show" in result and "hide" in result def test_reasoning_is_known_command(self): - source = inspect.getsource(gateway_run.GatewayRunner._handle_message) + func = getattr(gateway_run.GatewayRunner, "_handle_message_inner", gateway_run.GatewayRunner._handle_message) + source = inspect.getsource(func) assert '"reasoning"' in source def test_parse_reasoning_command_args_accepts_ascii_and_smart_global_flags(self): diff --git a/tests/gateway/test_telegram_topic_mode.py b/tests/gateway/test_telegram_topic_mode.py index c887153508c75..89c217d50ed38 100644 --- a/tests/gateway/test_telegram_topic_mode.py +++ b/tests/gateway/test_telegram_topic_mode.py @@ -1434,3 +1434,170 @@ def test_session_split_restores_source_thread_id_from_binding(tmp_path): meta = GatewayRunner._thread_metadata_for_source(runner, source) assert meta is not None assert meta["thread_id"] == "17585" + + +# --------------------------------------------------------------------------- +# Tests for format_session_info with source overrides +# --------------------------------------------------------------------------- + +def test_format_session_info_honors_topic_override(tmp_path): + """Verify that _format_session_info resolves and displays the overridden/topic model.""" + from gateway.run import GatewayRunner + from gateway.session import SessionSource + from gateway.config import Platform + from unittest.mock import patch + + runner = object.__new__(GatewayRunner) + runner._session_model_overrides = { + "agent:main:telegram:dm:208214988:17585": { + "model": "MiniMax-M3", + "provider": "minimax", + } + } + + # Mock the session key derivation and config loading + runner._session_key_for_source = lambda source: "agent:main:telegram:dm:208214988:17585" + + source = SessionSource( + platform=Platform.TELEGRAM, + user_id="208214988", + chat_id="208214988", + user_name="tester", + chat_type="dm", + thread_id="17585", + ) + + with patch("gateway.run._load_gateway_config", return_value={}): + with patch("gateway.run._resolve_runtime_agent_kwargs", return_value={}): + info = runner._format_session_info(source=source) + + assert "MiniMax-M3" in info + assert "minimax" in info + + +def test_format_session_info_no_source_falls_back_to_default(tmp_path): + """Verify that _format_session_info falls back to default config if no source is provided.""" + from gateway.run import GatewayRunner + from unittest.mock import patch + + runner = object.__new__(GatewayRunner) + runner._session_model_overrides = { + "agent:main:telegram:dm:208214988:17585": { + "model": "MiniMax-M3", + "provider": "minimax", + } + } + + with patch("gateway.run._load_gateway_config", return_value={ + "model": { + "default": "gemini-3.1-flash-lite", + "provider": "gemini", + } + }): + with patch("gateway.run._resolve_runtime_agent_kwargs", return_value={}): + info = runner._format_session_info(source=None) + + assert "gemini-3.1-flash-lite" in info + assert "gemini" in info + + +# --------------------------------------------------------------------------- +# Tests for topic-specific profiles +# --------------------------------------------------------------------------- + +def test_topic_profile_persistence(tmp_path): + """Verify that topic profiles can be saved, loaded, and removed persistently.""" + from gateway.run import _save_topic_profile, _load_topic_profiles, _remove_topic_profile + from unittest.mock import patch + + with patch("gateway.run._hermes_home", tmp_path): + assert _load_topic_profiles() == {} + + _save_topic_profile("telegram:dm:208214988:17585", "coder") + assert _load_topic_profiles() == {"telegram:dm:208214988:17585": "coder"} + + _remove_topic_profile("telegram:dm:208214988:17585") + assert _load_topic_profiles() == {} + + +def test_session_key_for_source_honors_profile_override(tmp_path): + """Verify that _session_key_for_source resolves the persistent profile override and updates source.""" + from gateway.run import GatewayRunner + from gateway.session import SessionSource + from gateway.config import Platform, GatewayConfig + from unittest.mock import patch, MagicMock + + runner = object.__new__(GatewayRunner) + runner.config = GatewayConfig(platforms={}) + runner.config.multiplex_profiles = True + runner.session_store = MagicMock() + # Mock generation to use building logic + from gateway.session import build_session_key + runner.session_store._generate_session_key.side_effect = lambda src: build_session_key(src, profile=src.profile) + runner._normalize_source_for_session_key = lambda src: src + + source = SessionSource( + platform=Platform.TELEGRAM, + chat_id="208214988", + chat_type="dm", + thread_id="17585", + profile=None, + ) + + with patch("gateway.run._load_topic_profiles", return_value={"telegram:dm:208214988:17585": "coder"}): + session_key = runner._session_key_for_source(source) + + assert source.profile == "coder" + assert "agent:coder" in session_key + + +@pytest.mark.asyncio +async def test_handle_profile_command_switches_profile(tmp_path): + """Verify that /profile updates and /profile default clears the override.""" + from gateway.slash_commands import GatewaySlashCommandsMixin + from gateway.platforms.base import MessageEvent + from gateway.session import SessionSource + from gateway.config import Platform + from unittest.mock import patch, MagicMock + + class TestSlashCommands(GatewaySlashCommandsMixin): + def __init__(self): + self._evict_cached_agent = MagicMock() + self._normalize_source_for_session_key = lambda src: src + self._session_key_for_source = lambda src: "dummy" + self._active_profile_name = lambda: "default" + + cmd = TestSlashCommands() + + source = SessionSource( + platform=Platform.TELEGRAM, + chat_id="208214988", + chat_type="dm", + thread_id="17585", + ) + event = MessageEvent( + text="/profile coder", + source=source, + message_id="m1", + ) + + # Mock list_profiles to return "default" and "coder" + profile1 = MagicMock() + profile1.name = "default" + profile2 = MagicMock() + profile2.name = "coder" + + with patch("hermes_cli.profiles.list_profiles", return_value=[profile1, profile2]): + with patch("gateway.run._save_topic_profile") as mock_save: + res = await cmd._handle_profile_command(event) + mock_save.assert_called_once_with("telegram:dm:208214988:17585", "coder") + assert "Pinned this topic to profile `coder`" in res + assert source.profile == "coder" + + # Now clear it + event.text = "/profile default" + with patch("gateway.run._remove_topic_profile") as mock_remove: + res = await cmd._handle_profile_command(event) + mock_remove.assert_called_once_with("telegram:dm:208214988:17585") + assert "Cleared the topic profile binding" in res + assert source.profile is None diff --git a/tests/gateway/test_topic_profile_routing.py b/tests/gateway/test_topic_profile_routing.py new file mode 100644 index 0000000000000..9a1ed3a5096e6 --- /dev/null +++ b/tests/gateway/test_topic_profile_routing.py @@ -0,0 +1,238 @@ +import pytest +import os +import shutil +from pathlib import Path +from unittest.mock import patch, MagicMock + +from gateway.run import GatewayRunner +from gateway.session import SessionSource, build_session_key +from gateway.config import Platform, GatewayConfig +from hermes_constants import get_hermes_home + +@pytest.fixture +def clean_hermes_home(tmp_path, monkeypatch): + hermes_home = tmp_path / ".hermes" + hermes_home.mkdir() + + # Create profiles directory and files + profiles_dir = hermes_home / "profiles" + profiles_dir.mkdir() + + for name in ("profilea", "profileb"): + pdir = profiles_dir / name + pdir.mkdir() + (pdir / "home").mkdir() + (pdir / "SOUL.md").write_text(f"I am {name}", encoding="utf-8") + (pdir / "memories").mkdir() + (pdir / "sessions").mkdir() + (pdir / "config.yaml").write_text("agent:\n system_prompt: overridden", encoding="utf-8") + + # Write identity marker to satisfy profile safety check + from hermes_cli.profiles import write_profile_identity_marker + write_profile_identity_marker(name, pdir, profiles_dir, overwrite=True) + + # Set up global SOUL.md + (hermes_home / "SOUL.md").write_text("I am main agent", encoding="utf-8") + + # Write topic profiles config + import json + topic_profiles = { + "telegram:dm:111:222": "profilea", + "telegram:dm:111:333": "profileb", + } + with open(hermes_home / "topic_profiles.json", "w", encoding="utf-8") as f: + json.dump(topic_profiles, f) + + monkeypatch.setenv("HERMES_HOME", str(hermes_home)) + with patch("gateway.run._hermes_home", hermes_home): + yield hermes_home + +def test_session_key_routing_unconditional(clean_hermes_home): + """Verify session key incorporates profile even when multiplex_profiles is False.""" + runner = object.__new__(GatewayRunner) + runner.config = GatewayConfig(platforms={}) + runner.config.multiplex_profiles = False + + # We must mock _normalize_source_for_session_key + runner._normalize_source_for_session_key = lambda src: src + + source_a = SessionSource( + platform=Platform.TELEGRAM, + chat_id="111", + chat_type="dm", + thread_id="222", + ) + source_b = SessionSource( + platform=Platform.TELEGRAM, + chat_id="111", + chat_type="dm", + thread_id="333", + ) + source_main = SessionSource( + platform=Platform.TELEGRAM, + chat_id="111", + chat_type="dm", + thread_id="444", + ) + + key_a = runner._session_key_for_source(source_a) + key_b = runner._session_key_for_source(source_b) + key_main = runner._session_key_for_source(source_main) + + assert source_a.profile == "profilea" + assert source_b.profile == "profileb" + assert source_main.profile is None + + assert "profilea" in key_a + assert "profileb" in key_b + assert "profilea" not in key_main and "profileb" not in key_main + +def test_dynamic_session_db_and_store_scoping(clean_hermes_home): + """Verify that session_store and _session_db are dynamically resolved per-profile.""" + runner = GatewayRunner(config=GatewayConfig(platforms={})) + runner._normalize_source_for_session_key = lambda src: src + + source_a = SessionSource( + platform=Platform.TELEGRAM, + chat_id="111", + chat_type="dm", + thread_id="222", + ) + + # Initially global db path + global_db_path = clean_hermes_home / "state.db" + assert Path(runner._session_db.db_path).resolve() == global_db_path.resolve() + + # Simulate a routed run using _routed_profile_for_source helper + routed = runner._routed_profile_for_source(source_a) + assert routed == "profilea" + + profile_home = clean_hermes_home / "profiles" / "profilea" + from gateway.run import _profile_runtime_scope + with _profile_runtime_scope(profile_home): + # Inside the scope, get_hermes_home() points to profilea + assert get_hermes_home().resolve() == profile_home.resolve() + # session_store and _session_db should resolve to profile-specific DB paths + profile_db_path = profile_home / "state.db" + assert Path(runner._session_db.db_path).resolve() == profile_db_path.resolve() + assert Path(runner.session_store.sessions_dir).resolve() == (profile_home / "sessions").resolve() + + # Outside the scope, back to global + assert Path(runner._session_db.db_path).resolve() == global_db_path.resolve() + +def test_soul_cache_busting(clean_hermes_home): + """Verify that updating SOUL.md busts the agent cache signature.""" + runner = GatewayRunner(config=GatewayConfig(platforms={})) + + # Check signature initially + sig1 = runner._agent_config_signature( + model="gpt-4", + runtime={}, + enabled_toolsets=[], + ephemeral_prompt="", + ) + + # Modify SOUL.md + soul_file = clean_hermes_home / "SOUL.md" + soul_file.write_text("updated identity content", encoding="utf-8") + + # Force different mtime/stat + import os + stat = soul_file.stat() + os.utime(soul_file, (stat.st_atime, stat.st_mtime + 5.0)) + + sig2 = runner._agent_config_signature( + model="gpt-4", + runtime={}, + enabled_toolsets=[], + ephemeral_prompt="", + ) + + assert sig1 != sig2 + +def test_tilde_expansion_isolated(clean_hermes_home): + """Verify that tilde (~) expansion resolves to profile-specific home inside routing scope.""" + from gateway.run import _profile_runtime_scope + from tools.file_tools import _expand_tilde + import concurrent.futures + from contextvars import copy_context + + profile_home = clean_hermes_home / "profiles" / "profilea" + + # Setup the config/home mode to enable profile home mode + monkeypatch_env = os.environ.copy() + monkeypatch_env["TERMINAL_HOME_MODE"] = "profile" + + with patch.dict(os.environ, monkeypatch_env): + with _profile_runtime_scope(profile_home): + # Directly inside scope: + assert _expand_tilde("~/SOUL.md") == str(profile_home / "home" / "SOUL.md") + + # Inside ThreadPoolExecutor worker thread (simulates tool execution thread): + pool = concurrent.futures.ThreadPoolExecutor(max_workers=1) + ctx = copy_context() + res = pool.submit(ctx.run, lambda: _expand_tilde("~/SOUL.md")).result() + assert res == str(profile_home / "home" / "SOUL.md") + pool.shutdown() + +@pytest.mark.asyncio +async def test_telegram_topic_new_command_isolated(clean_hermes_home): + """Verify that /new executed in a Telegram topic lane is correctly routed to the profile scope and DB.""" + from gateway.platforms.base import MessageEvent + from gateway.run import GatewayRunner + from gateway.session import SessionSource + from gateway.config import Platform, GatewayConfig + + runner = GatewayRunner(config=GatewayConfig(platforms={})) + runner._normalize_source_for_session_key = lambda src: src + + # We must mock adapters lookup so send does not crash or block + mock_adapter = MagicMock() + runner.adapters = {Platform.TELEGRAM: mock_adapter} + + profile_home = clean_hermes_home / "profiles" / "profilea" + profile_db_path = profile_home / "state.db" + from hermes_state import SessionDB + db = SessionDB(db_path=profile_db_path) + db.enable_telegram_topic_mode(chat_id="111", user_id="user_999") + + source = SessionSource( + platform=Platform.TELEGRAM, + chat_id="111", + chat_type="dm", + thread_id="222", # profilea is mapped to this + user_id="user_999", + ) + + event = MessageEvent( + text="/new", + message_id="999", + source=source, + ) + + # Bypass confirmation dialog so execute runs immediately + async def mock_confirm(*args, **kwargs): + return await kwargs["execute"]() + runner._maybe_confirm_destructive_slash = mock_confirm + + # Let's run _handle_message. + await runner._handle_message(event) + + # Check profilea's database for the topic binding! + profile_home = clean_hermes_home / "profiles" / "profilea" + profile_db_path = profile_home / "state.db" + assert profile_db_path.exists() + + from hermes_state import SessionDB + db = SessionDB(db_path=profile_db_path) + binding = db.get_telegram_topic_binding(chat_id="111", thread_id="222") + assert binding is not None + assert binding["chat_id"] == "111" + assert binding["thread_id"] == "222" + + # Verify that the GLOBAL database does NOT have this binding, proving isolation! + global_db_path = clean_hermes_home / "state.db" + if global_db_path.exists(): + global_db = SessionDB(db_path=global_db_path) + global_binding = global_db.get_telegram_topic_binding(chat_id="111", thread_id="222") + assert global_binding is None diff --git a/tests/gateway/test_update_command.py b/tests/gateway/test_update_command.py index 5cc7f206e667f..5b7b21db7dd09 100644 --- a/tests/gateway/test_update_command.py +++ b/tests/gateway/test_update_command.py @@ -927,5 +927,6 @@ def test_update_is_known_command(self): # checking the help output includes it. from gateway.run import GatewayRunner import inspect - source = inspect.getsource(GatewayRunner._handle_message) + func = getattr(GatewayRunner, "_handle_message_inner", GatewayRunner._handle_message) + source = inspect.getsource(func) assert '"update"' in source diff --git a/tests/hermes_cli/test_profiles.py b/tests/hermes_cli/test_profiles.py index 44dfd6d4dae97..6e45b29ed82cd 100644 --- a/tests/hermes_cli/test_profiles.py +++ b/tests/hermes_cli/test_profiles.py @@ -36,6 +36,10 @@ NO_BUNDLED_SKILLS_MARKER, backfill_profile_envs, profiles_to_serve, + PROFILE_IDENTITY_FILENAME, + audit_profile_isolation, + validate_profile_identity, + write_profile_identity_marker, ) from hermes_cli.config import DEFAULT_CONFIG @@ -1676,3 +1680,247 @@ def test_on_active_profile_does_not_change_set(self, profile_env): def test_on_no_named_profiles_returns_just_default(self, profile_env): serve = profiles_to_serve(multiplex=True) assert [n for n, _ in serve] == ["default"] + + +class TestProfileIsolationAudit: + def test_create_profile_writes_identity_marker(self, profile_env): + tmp_path = profile_env + profile_dir = create_profile("alpha-test", no_alias=True) + marker = profile_dir / PROFILE_IDENTITY_FILENAME + + assert marker.is_file() + validate_profile_identity( + "alpha-test", + profile_dir, + tmp_path / ".hermes" / "profiles", + ) + + def test_audit_fails_for_secret_env_identical_to_main(self, profile_env): + tmp_path = profile_env + main_home = tmp_path / ".hermes" + profile_dir = create_profile("alpha-test", no_alias=True) + env_text = "OPENROUTER_API_KEY=sk-test\n" + (main_home / ".env").write_text(env_text, encoding="utf-8") + (profile_dir / ".env").write_text(env_text, encoding="utf-8") + + report = audit_profile_isolation("alpha-test") + + assert report["status"] == "FAIL" + assert any(f["path"] == ".env" and f["status"] == "FAIL" for f in report["findings"]) + + def test_audit_fails_for_shared_env_secret_value_even_when_not_identical(self, profile_env): + tmp_path = profile_env + main_home = tmp_path / ".hermes" + profile_dir = create_profile("alpha-test", no_alias=True) + (main_home / ".env").write_text( + "OPENROUTER_API_KEY=shared-secret\nOTHER=main\n", + encoding="utf-8", + ) + (profile_dir / ".env").write_text( + "# reordered copy\nOTHER=profile\nOPENROUTER_API_KEY=shared-secret\n", + encoding="utf-8", + ) + + report = audit_profile_isolation("alpha-test") + + assert report["status"] == "FAIL" + assert any( + f["path"] == ".env" + and f["status"] == "FAIL" + and f.get("shared_secret_count") == 1 + for f in report["findings"] + ) + + def test_audit_fails_for_shared_alphabetic_env_secret_value(self, profile_env): + tmp_path = profile_env + main_home = tmp_path / ".hermes" + profile_dir = create_profile("alpha-test", no_alias=True) + (main_home / ".env").write_text( + "OPENROUTER_API_KEY=supersecretvalue\n", + encoding="utf-8", + ) + (profile_dir / ".env").write_text( + "OPENROUTER_API_KEY=supersecretvalue\nOTHER=profile\n", + encoding="utf-8", + ) + + report = audit_profile_isolation("alpha-test") + + assert report["status"] == "FAIL" + assert any( + f["path"] == ".env" + and f["status"] == "FAIL" + and f.get("shared_secret_count") == 1 + for f in report["findings"] + ) + + def test_audit_fails_for_shared_auth_secret_value_even_when_json_differs(self, profile_env): + tmp_path = profile_env + main_home = tmp_path / ".hermes" + profile_dir = create_profile("alpha-test", no_alias=True) + (main_home / "auth.json").write_text( + json.dumps({"providers": {"nous": {"access_token": "shared-token"}}}), + encoding="utf-8", + ) + (profile_dir / "auth.json").write_text( + json.dumps( + { + "providers": { + "nous": { + "label": "profile", + "refresh_token": "different-refresh", + "access_token": "shared-token", + } + } + }, + indent=2, + ), + encoding="utf-8", + ) + + report = audit_profile_isolation("alpha-test") + + assert report["status"] == "FAIL" + assert any( + f["path"] == "auth.json" + and f["status"] == "FAIL" + and f.get("shared_secret_count") == 1 + for f in report["findings"] + ) + + def test_audit_fails_for_shared_alphabetic_auth_secret_value(self, profile_env): + tmp_path = profile_env + main_home = tmp_path / ".hermes" + profile_dir = create_profile("alpha-test", no_alias=True) + (main_home / "auth.json").write_text( + json.dumps({"providers": {"nous": {"access_token": "supersecretvalue"}}}), + encoding="utf-8", + ) + (profile_dir / "auth.json").write_text( + json.dumps({"providers": {"nous": {"access_token": "supersecretvalue"}}}), + encoding="utf-8", + ) + + report = audit_profile_isolation("alpha-test") + + assert report["status"] == "FAIL" + assert any( + f["path"] == "auth.json" + and f["status"] == "FAIL" + and f.get("shared_secret_count") == 1 + for f in report["findings"] + ) + + def test_audit_fails_for_shared_config_secret_value_even_when_yaml_differs( + self, profile_env + ): + tmp_path = profile_env + main_home = tmp_path / ".hermes" + profile_dir = create_profile("alpha-test", no_alias=True) + (main_home / "config.yaml").write_text( + "model:\n" + " provider: custom\n" + " api_key: shared-inline-secret\n" + " default: main-model\n", + encoding="utf-8", + ) + (profile_dir / "config.yaml").write_text( + "model:\n" + " default: profile-model\n" + " api_key: shared-inline-secret\n" + " provider: custom\n", + encoding="utf-8", + ) + + report = audit_profile_isolation("alpha-test") + + assert report["status"] == "FAIL" + assert any( + f["path"] == "config.yaml" + and f["status"] == "FAIL" + and f.get("shared_secret_count") == 1 + for f in report["findings"] + ) + + def test_audit_fails_for_shared_alphabetic_config_secret_value( + self, profile_env + ): + tmp_path = profile_env + main_home = tmp_path / ".hermes" + profile_dir = create_profile("alpha-test", no_alias=True) + (main_home / "config.yaml").write_text( + "providers:\n custom:\n api_key: supersecretvalue\n", + encoding="utf-8", + ) + (profile_dir / "config.yaml").write_text( + "providers:\n custom:\n api_key: supersecretvalue\n", + encoding="utf-8", + ) + + report = audit_profile_isolation("alpha-test") + + assert report["status"] == "FAIL" + assert any( + f["path"] == "config.yaml" + and f["status"] == "FAIL" + and f.get("shared_secret_count") == 1 + for f in report["findings"] + ) + + def test_audit_safe_root_uses_profile_under_safe_root_and_main_home_from_env( + self, profile_env + ): + tmp_path = profile_env + main_home = tmp_path / ".hermes" + external_root = tmp_path / "external-profiles" + profile_dir = external_root / "alpha-test" + profile_dir.mkdir(parents=True) + write_profile_identity_marker( + "alpha-test", + profile_dir, + external_root, + overwrite=True, + ) + (main_home / ".env").write_text("OPENROUTER_API_KEY=main-secret\n", encoding="utf-8") + (profile_dir / ".env").write_text("OPENROUTER_API_KEY=main-secret\n", encoding="utf-8") + + report = audit_profile_isolation("alpha-test", safe_root=external_root) + + assert report["status"] == "FAIL" + assert any(f["path"] == ".env" and f["status"] == "FAIL" for f in report["findings"]) + + def test_audit_fails_for_nonempty_memory_identical_to_main(self, profile_env): + tmp_path = profile_env + main_home = tmp_path / ".hermes" + profile_dir = create_profile("alpha-test", no_alias=True) + (main_home / "memories").mkdir(exist_ok=True) + (profile_dir / "memories").mkdir(exist_ok=True) + text = "shared memory should not be cloned into routed profile\n" + (main_home / "memories" / "MEMORY.md").write_text(text, encoding="utf-8") + (profile_dir / "memories" / "MEMORY.md").write_text(text, encoding="utf-8") + + report = audit_profile_isolation("alpha-test") + + assert report["status"] == "FAIL" + assert any( + f["path"] == "memories/MEMORY.md" and f["status"] == "FAIL" + for f in report["findings"] + ) + + def test_audit_fails_for_outgoing_plugin_symlink(self, profile_env): + tmp_path = profile_env + main_home = tmp_path / ".hermes" + profile_dir = create_profile("alpha-test", no_alias=True) + shared_plugin = main_home / "plugins" / "shared-plugin" + shared_plugin.mkdir(parents=True) + profile_plugins = profile_dir / "plugins" + profile_plugins.mkdir(exist_ok=True) + (profile_plugins / "shared-plugin").symlink_to(shared_plugin, target_is_directory=True) + + report = audit_profile_isolation("alpha-test") + + assert report["status"] == "FAIL" + assert any( + f["path"] == "plugins/shared-plugin" and f["status"] == "FAIL" + for f in report["findings"] + ) diff --git a/tests/hermes_cli/test_runtime_provider_resolution.py b/tests/hermes_cli/test_runtime_provider_resolution.py index 236e27899e78b..a1988d63b972c 100644 --- a/tests/hermes_cli/test_runtime_provider_resolution.py +++ b/tests/hermes_cli/test_runtime_provider_resolution.py @@ -180,6 +180,7 @@ def test_resolve_runtime_provider_qwen_oauth(monkeypatch): "expires_at_ms": 1775640710946, }, ) + monkeypatch.setattr(rp, "_load_pool_for_env", lambda *a, **k: None) resolved = rp.resolve_runtime_provider(requested="qwen-oauth") @@ -236,6 +237,7 @@ def test_qwen_oauth_auto_fallthrough_on_auth_failure(monkeypatch): ) monkeypatch.setattr(rp, "_get_model_config", lambda: {}) monkeypatch.setenv("OPENROUTER_API_KEY", "test-or-key") + monkeypatch.setattr(rp, "_load_pool_for_env", lambda *a, **k: None) # Should NOT raise — falls through to OpenRouter resolved = rp.resolve_runtime_provider(requested="auto") @@ -2851,8 +2853,8 @@ def _patch_bedrock(monkeypatch, config_default=""): monkeypatch.setattr(rp, "_get_model_config", lambda: {"default": config_default}) monkeypatch.setattr(rp, "load_config", lambda: {"bedrock": {}}) monkeypatch.setattr(ba, "has_aws_credentials", lambda: True) - monkeypatch.setattr(ba, "resolve_aws_auth_env_var", lambda: "AWS_PROFILE") - monkeypatch.setattr(ba, "resolve_bedrock_region", lambda: "eu-north-1") + monkeypatch.setattr(ba, "resolve_aws_auth_env_var", lambda *a, **k: "AWS_PROFILE") + monkeypatch.setattr(ba, "resolve_bedrock_region", lambda *a, **k: "eu-north-1") def test_resolve_runtime_provider_bedrock_claude_target_model_uses_anthropic_messages(monkeypatch): diff --git a/tests/tools/test_mcp_profile_isolation.py b/tests/tools/test_mcp_profile_isolation.py new file mode 100644 index 0000000000000..97621549a0d9a --- /dev/null +++ b/tests/tools/test_mcp_profile_isolation.py @@ -0,0 +1,109 @@ +import pytest +import asyncio +from pathlib import Path +from unittest.mock import patch, MagicMock + +from hermes_constants import set_hermes_home_override, reset_hermes_home_override +from tools.mcp_tool import ( + register_mcp_servers, + _get_active_server, + _get_mcp_config_fingerprint, + _servers, + MCPServerTask, +) + +@pytest.fixture(autouse=True) +def clean_servers(): + saved = dict(_servers) + _servers.clear() + yield + _servers.clear() + _servers.update(saved) + +def test_mcp_config_fingerprint(): + """Verify that fingerprint is stable and ignores non-connection keys like enabled.""" + cfg1 = {"command": "node", "args": ["app.js"], "enabled": True} + cfg2 = {"command": "node", "args": ["app.js"], "enabled": False} + cfg3 = {"command": "node", "args": ["app.js", "--flag"]} + + fp1 = _get_mcp_config_fingerprint("github", cfg1) + fp2 = _get_mcp_config_fingerprint("github", cfg2) + fp3 = _get_mcp_config_fingerprint("github", cfg3) + + assert fp1 == fp2 + assert fp1 != fp3 + assert fp1.startswith("github:") + +def test_mcp_servers_connection_isolation(): + """Verify that servers with same name but different config have isolated connections, + while identical configs share a connection. + """ + mock_session_a = MagicMock() + mock_session_b = MagicMock() + + async def fake_connect(name, config): + server = MCPServerTask(name) + if config.get("env", {}).get("TOKEN") == "A": + server.session = mock_session_a + else: + server.session = mock_session_b + server._tools = [] + return server + + cfg_a = {"command": "node", "env": {"TOKEN": "A"}} + cfg_b = {"command": "node", "env": {"TOKEN": "B"}} + cfg_a_dup = {"command": "node", "env": {"TOKEN": "A"}, "enabled": True} + + def fake_run(coro_or_factory, timeout=30): + coro = coro_or_factory() if callable(coro_or_factory) else coro_or_factory + return asyncio.run(coro) + + with patch("tools.mcp_tool._connect_server", side_effect=fake_connect), \ + patch("tools.mcp_tool._MCP_AVAILABLE", True), \ + patch("tools.mcp_tool._ensure_mcp_loop"), \ + patch("tools.mcp_tool._run_on_mcp_loop", side_effect=fake_run): + + # Register config A + register_mcp_servers({"github": cfg_a}) + # Determine fingerprint for A + fp_a = _get_mcp_config_fingerprint("github", cfg_a) + assert fp_a in _servers + srv_a = _servers[fp_a] + assert srv_a.session == mock_session_a + + # Register config B (same name, different config) + register_mcp_servers({"github": cfg_b}) + fp_b = _get_mcp_config_fingerprint("github", cfg_b) + assert fp_b in _servers + srv_b = _servers[fp_b] + assert srv_b.session == mock_session_b + + # They are isolated connections + assert srv_a != srv_b + assert fp_a != fp_b + + # Register identical config A (should be idempotent, no new connection) + register_mcp_servers({"github": cfg_a_dup}) + assert _servers[fp_a] == srv_a + +def test_get_active_server_routing(): + """Verify that _get_active_server routes dynamically based on active profile config.""" + srv_a = MCPServerTask("github") + srv_b = MCPServerTask("github") + + cfg_a = {"command": "node", "env": {"TOKEN": "A"}} + cfg_b = {"command": "node", "env": {"TOKEN": "B"}} + + fp_a = _get_mcp_config_fingerprint("github", cfg_a) + fp_b = _get_mcp_config_fingerprint("github", cfg_b) + + _servers[fp_a] = srv_a + _servers[fp_b] = srv_b + + # Mock _load_mcp_config to return config A + with patch("tools.mcp_tool._load_mcp_config", return_value={"github": cfg_a}): + assert _get_active_server("github") == srv_a + + # Mock _load_mcp_config to return config B + with patch("tools.mcp_tool._load_mcp_config", return_value={"github": cfg_b}): + assert _get_active_server("github") == srv_b diff --git a/tests/tools/test_process_registry.py b/tests/tools/test_process_registry.py index 020dcd11ae8ca..73b4e98f2846c 100644 --- a/tests/tools/test_process_registry.py +++ b/tests/tools/test_process_registry.py @@ -829,6 +829,39 @@ def test_recover_enqueues_watchers(self, registry, tmp_path): assert w["thread_id"] == "42" assert w["check_interval"] == 60 + def test_recover_legacy_profile_without_home_does_not_rearm_watcher( + self, registry, tmp_path + ): + checkpoint = tmp_path / "procs.json" + checkpoint.write_text(json.dumps([{ + "session_id": "proc_legacy_profile", + "command": "sleep 999", + "pid": os.getpid(), + "task_id": "t1", + "session_key": "agent:alpha-test:telegram:u:dm:c", + "agent_profile": "alpha-test", + "agent_hermes_home": "", + "watcher_platform": "telegram", + "watcher_chat_id": "123", + "watcher_user_id": "u123", + "watcher_thread_id": "42", + "watcher_interval": 60, + "notify_on_complete": True, + "watch_patterns": ["done"], + }])) + + with patch("tools.process_registry.CHECKPOINT_PATH", checkpoint): + recovered = registry.recover_from_checkpoint() + + assert recovered == 1 + assert registry.pending_watchers == [] + session = registry.get("proc_legacy_profile") + assert session is not None + assert session.agent_profile == "" + assert session.agent_hermes_home == "" + assert session.watcher_interval == 0 + assert session.notify_on_complete is False + def test_recover_skips_watcher_when_no_interval(self, registry, tmp_path): checkpoint = tmp_path / "procs.json" checkpoint.write_text(json.dumps([{ diff --git a/tests/tools/test_skills_profile_isolation.py b/tests/tools/test_skills_profile_isolation.py new file mode 100644 index 0000000000000..80814e3d22f4e --- /dev/null +++ b/tests/tools/test_skills_profile_isolation.py @@ -0,0 +1,101 @@ +import pytest +import os +import shutil +from pathlib import Path +from unittest.mock import patch, MagicMock + +from hermes_constants import get_hermes_home, get_skills_dir +from tools.skills_tool import get_skills_dir as tool_get_skills_dir +from tools.skill_manager_tool import get_skills_dir as manager_get_skills_dir +from agent.skill_commands import get_skill_commands, scan_skill_commands + +@pytest.fixture +def clean_hermes_home(tmp_path, monkeypatch): + hermes_home = tmp_path / ".hermes" + hermes_home.mkdir() + + # Create profiles directory and files + profiles_dir = hermes_home / "profiles" + profiles_dir.mkdir() + + for name in ("profileA", "profileB"): + pdir = profiles_dir / name + pdir.mkdir() + (pdir / "skills").mkdir() + (pdir / "skills" / f"skill_{name}").mkdir(parents=True) + (pdir / "skills" / f"skill_{name}" / "SKILL.md").write_text( + f"---\nname: skill_{name}\ndescription: I am {name}\n---\nbody of {name}", + encoding="utf-8" + ) + + # Set up global skill + global_skills = hermes_home / "skills" + global_skills.mkdir() + (global_skills / "skill_main").mkdir() + (global_skills / "skill_main" / "SKILL.md").write_text( + "---\nname: skill_main\ndescription: I am main\n---\nbody of main", + encoding="utf-8" + ) + + monkeypatch.setenv("HERMES_HOME", str(hermes_home)) + with patch("gateway.run._hermes_home", hermes_home): + yield hermes_home + +def test_skills_dir_resolves_dynamically(clean_hermes_home): + """Verify that get_skills_dir and tool-specific skills dirs resolve dynamically.""" + from hermes_constants import set_hermes_home_override, reset_hermes_home_override + + # Check default path + assert get_skills_dir().resolve() == (clean_hermes_home / "skills").resolve() + assert tool_get_skills_dir().resolve() == (clean_hermes_home / "skills").resolve() + assert manager_get_skills_dir().resolve() == (clean_hermes_home / "skills").resolve() + + # Override + profile_home = clean_hermes_home / "profiles" / "profileA" + token = set_hermes_home_override(profile_home) + try: + assert get_skills_dir().resolve() == (profile_home / "skills").resolve() + assert tool_get_skills_dir().resolve() == (profile_home / "skills").resolve() + assert manager_get_skills_dir().resolve() == (profile_home / "skills").resolve() + finally: + reset_hermes_home_override(token) + + assert get_skills_dir().resolve() == (clean_hermes_home / "skills").resolve() + +def test_skill_commands_isolation_across_profiles(clean_hermes_home): + """Verify that skill commands dynamically refresh and isolate based on active profile.""" + from hermes_constants import set_hermes_home_override, reset_hermes_home_override + + # Initially main skills + cmds = get_skill_commands() + assert "/skill-main" in cmds + assert "/skill-profilea" not in cmds + assert "/skill-profileb" not in cmds + + # Switch to profileA + profile_home_a = clean_hermes_home / "profiles" / "profileA" + token_a = set_hermes_home_override(profile_home_a) + try: + cmds_a = get_skill_commands() + assert "/skill-profilea" in cmds_a + assert "/skill-main" not in cmds_a + assert "/skill-profileb" not in cmds_a + finally: + reset_hermes_home_override(token_a) + + # Switch to profileB + profile_home_b = clean_hermes_home / "profiles" / "profileB" + token_b = set_hermes_home_override(profile_home_b) + try: + cmds_b = get_skill_commands() + assert "/skill-profileb" in cmds_b + assert "/skill-main" not in cmds_b + assert "/skill-profilea" not in cmds_b + finally: + reset_hermes_home_override(token_b) + + # Back to main + cmds_final = get_skill_commands() + assert "/skill-main" in cmds_final + assert "/skill-profilea" not in cmds_final + assert "/skill-profileb" not in cmds_final diff --git a/tools/file_tools.py b/tools/file_tools.py index 59c7214593dfd..c09db6f801a8c 100644 --- a/tools/file_tools.py +++ b/tools/file_tools.py @@ -24,7 +24,145 @@ _EXPECTED_WRITE_ERRNOS = {errno.EACCES, errno.EPERM, errno.EROFS} -def _expand_tilde(path: str) -> str: +from typing import Optional + +def get_profile_home_for_session(session_id: Optional[str]) -> Optional[Path]: + if not session_id: + from gateway.session_context import get_session_env + session_id = get_session_env("HERMES_SESSION_ID") + if not session_id: + from hermes_constants import get_hermes_home, get_default_hermes_root + try: + cur = get_hermes_home().resolve() + root = get_default_hermes_root().resolve() + if cur != root and (root / "profiles") in cur.parents: + return cur + except Exception: + pass + return None + from hermes_constants import get_default_hermes_root + try: + root = get_default_hermes_root() + except Exception: + return None + + # Check default home first + default_home = root + if (default_home / "state.db").exists(): + try: + import sqlite3 + with sqlite3.connect(default_home / "state.db") as conn: + cursor = conn.execute("SELECT 1 FROM sessions WHERE id = ?", (session_id,)) + if cursor.fetchone(): + return default_home + except Exception: + pass + + # Check named profiles + profiles_dir = root / "profiles" + if profiles_dir.exists(): + for p_dir in profiles_dir.iterdir(): + if p_dir.is_dir() and (p_dir / "state.db").exists(): + try: + import sqlite3 + with sqlite3.connect(p_dir / "state.db") as conn: + cursor = conn.execute("SELECT 1 FROM sessions WHERE id = ?", (session_id,)) + if cursor.fetchone(): + return p_dir + except Exception: + pass + return None + + +def _check_profile_hard_guards(resolved_path: str, profile_home: Optional[Path]) -> Optional[str]: + if not profile_home: + return None + # Only guard named profiles + if profile_home.parent.name != "profiles": + return None + try: + from hermes_constants import get_default_hermes_root + root = get_default_hermes_root().resolve() + resolved_target = Path(resolved_path).resolve() + + # Block writing to the global SOUL.md + global_soul = (root / "SOUL.md").resolve() + if resolved_target == global_soul: + return "Refusing to write to the global SOUL.md from a routed profile session." + + # Block writing to other profiles' homes + profiles_dir = (root / "profiles").resolve() + if profiles_dir.exists() and resolved_target.is_relative_to(profiles_dir): + rel_to_profiles = resolved_target.relative_to(profiles_dir) + parts = rel_to_profiles.parts + if parts and parts[0] != profile_home.name: + return f"Refusing to write to another profile's home directory: {parts[0]}" + + # Block any write inside the default hermes root unless it is inside the active profile home + if resolved_target.is_relative_to(root): + if not resolved_target.is_relative_to(profile_home.resolve()): + return f"Refusing to write to path outside profile home: {resolved_path}" + + # Block any write outside profile_home, except: + # 1. Inside profile_home + # 2. Inside system temp directories + # 3. Inside active workspace / CWD directories + if resolved_target.is_relative_to(profile_home.resolve()): + return None + + # Temp directories + import tempfile + temp_roots = [ + Path(tempfile.gettempdir()).resolve(), + Path("/tmp").resolve(), + Path("/private/var/tmp").resolve(), + Path("/var/tmp").resolve(), + ] + if any(resolved_target.is_relative_to(tr) for tr in temp_roots): + return None + + # Allowed CWD/workspace roots + allowed_roots = [] + allowed_roots.append(Path(os.getcwd()).resolve()) + + tcwd = _configured_terminal_cwd() + if tcwd: + allowed_roots.append(Path(tcwd).resolve()) + + try: + from tools.terminal_tool import _active_environments, _env_lock + with _env_lock: + for env in _active_environments.values(): + env_cwd = getattr(env, "cwd", None) + if env_cwd: + allowed_roots.append(Path(env_cwd).resolve()) + except Exception: + pass + + try: + with _file_ops_lock: + for cached in _file_ops_cache.values(): + cached_cwd = getattr(cached, "cwd", None) + if cached_cwd: + allowed_roots.append(Path(cached_cwd).resolve()) + env = getattr(cached, "env", None) + if env: + env_cwd = getattr(env, "cwd", None) + if env_cwd: + allowed_roots.append(Path(env_cwd).resolve()) + except Exception: + pass + + if any(resolved_target.is_relative_to(ar) for ar in allowed_roots): + return None + + return f"Refusing to write to path outside profile home: {resolved_path}" + except Exception as e: + logger.debug("Profile hard guard check failed: %s", e) + return None + + +def _expand_tilde(path: str, profile_home: Optional[Path] = None) -> str: """Expand ``~`` using the effective profile home when available. In-process file tools share the gateway process's HOME, which may differ @@ -35,12 +173,16 @@ def _expand_tilde(path: str) -> str: """ if not path or "~" not in path: return path - try: - from hermes_constants import get_subprocess_home + home = None + if profile_home is not None: + home = str(profile_home) + else: + try: + from hermes_constants import get_subprocess_home - home = get_subprocess_home() - except Exception: - home = None + home = get_subprocess_home() + except Exception: + home = None if home and (path == "~" or path.startswith("~/")): return home if path == "~" else os.path.join(home, path[2:]) return os.path.expanduser(path) @@ -282,13 +424,13 @@ def _resolve_base_dir(task_id: str = "default") -> Path: return base.resolve() -def _resolve_path_for_task(filepath: str, task_id: str = "default") -> Path: +def _resolve_path_for_task(filepath: str, task_id: str = "default", profile_home: Optional[Path] = None) -> Path: """Resolve *filepath* against the task's absolute base directory. See :func:`_resolve_base_dir` for how the base is chosen. Absolute input paths are returned resolved-but-unanchored. """ - p = Path(_expand_tilde(filepath)) + p = Path(_expand_tilde(filepath, profile_home=profile_home)) if p.is_absolute(): return p.resolve() return (_resolve_base_dir(task_id) / p).resolve() @@ -433,6 +575,11 @@ def _check_sensitive_path(filepath: str, task_id: str = "default") -> str | None ) for prefix in _SENSITIVE_PATH_PREFIXES: if resolved.startswith(prefix) or normalized.startswith(prefix): + if prefix == "/private/var/": + if resolved.startswith("/private/var/folders/") or resolved.startswith("/private/var/tmp/"): + continue + if normalized.startswith("/private/var/folders/") or normalized.startswith("/private/var/tmp/"): + continue return _err if resolved in _SENSITIVE_EXACT_PATHS or normalized in _SENSITIVE_EXACT_PATHS: return _err @@ -1337,6 +1484,10 @@ def write_file_tool(path: str, content: str, task_id: str = "default", Pass ``True`` after explicit user direction — same shape as ``force`` on the terminal tool. """ + profile_home = get_profile_home_for_session(session_id) + if profile_home and path == "SOUL.md": + path = str(profile_home / "SOUL.md") + sensitive_err = _check_sensitive_path(path, task_id) if sensitive_err: return tool_error(sensitive_err) @@ -1355,10 +1506,16 @@ def write_file_tool(path: str, content: str, task_id: str = "default", # fall back to the legacy path — write proceeds, per-task staleness # check below still runs. try: - _resolved = str(_resolve_path_for_task(path, task_id)) + _resolved = str(_resolve_path_for_task(path, task_id, profile_home=profile_home)) except Exception: _resolved = None + # Always run hard guards, falling back to expand_tilde if resolution failed + _resolved_for_guard = _resolved if _resolved is not None else _expand_tilde(path, profile_home=profile_home) + hard_err = _check_profile_hard_guards(_resolved_for_guard, profile_home) + if hard_err: + return tool_error(hard_err) + if _resolved is None: stale_warning = _check_file_staleness(path, task_id) file_ops = _get_file_ops(task_id) @@ -1419,6 +1576,19 @@ def patch_tool(mode: str = "replace", path: str = None, old_string: str = None, targets under another profile's skills/plugins/cron/memories directory. Same shape as ``write_file``'s flag. """ + profile_home = get_profile_home_for_session(session_id) + if profile_home: + if path == "SOUL.md": + path = str(profile_home / "SOUL.md") + if mode == "patch" and patch: + import re as _re + def _repl_header(m): + orig_file = m.group(2).strip() + if orig_file == "SOUL.md": + return f"{m.group(1)}{profile_home / 'SOUL.md'}" + return m.group(0) + patch = _re.sub(r'(^\*\*\*\s+(?:Update|Add|Delete)\s+File:\s*)(.+)$', _repl_header, patch, flags=_re.MULTILINE) + # Check sensitive paths for both replace (explicit path) and V4A patch (extract paths) _paths_to_check = [] if path: @@ -1451,6 +1621,13 @@ def patch_tool(mode: str = "replace", path: str = None, old_string: str = None, cross_warning = _check_cross_profile_path(_p, task_id) if cross_warning: return tool_error(cross_warning) + try: + _resolved = str(_resolve_path_for_task(_p, task_id, profile_home=profile_home)) + except Exception: + _resolved = _p + hard_err = _check_profile_hard_guards(_resolved, profile_home) + if hard_err: + return tool_error(hard_err) try: # Resolve paths for locking. Ordered + deduplicated so concurrent # callers lock in the same order — prevents deadlock on overlapping @@ -1459,7 +1636,7 @@ def patch_tool(mode: str = "replace", path: str = None, old_string: str = None, _seen: set[str] = set() for _p in _paths_to_check: try: - _r = str(_resolve_path_for_task(_p, task_id)) + _r = str(_resolve_path_for_task(_p, task_id, profile_home=profile_home)) except Exception: _r = None if _r and _r not in _seen: @@ -1481,7 +1658,7 @@ def patch_tool(mode: str = "replace", path: str = None, old_string: str = None, _path_to_resolved: dict[str, str] = {} for _p in _paths_to_check: try: - _r = str(_resolve_path_for_task(_p, task_id)) + _r = str(_resolve_path_for_task(_p, task_id, profile_home=profile_home)) except Exception: _r = None _path_to_resolved[_p] = _r diff --git a/tools/mcp_tool.py b/tools/mcp_tool.py index dead8ca204661..4abd94914ed74 100644 --- a/tools/mcp_tool.py +++ b/tools/mcp_tool.py @@ -381,6 +381,16 @@ def _build_safe_env(user_env: Optional[dict]) -> dict: or key.startswith("XDG_") ): env[key] = value + + try: + from hermes_constants import get_hermes_home_override, apply_subprocess_home_env + override = get_hermes_home_override() + if override: + env["HERMES_HOME"] = override + apply_subprocess_home_env(env) + except Exception: + pass + if user_env: env.update(user_env) return env @@ -1430,6 +1440,7 @@ class MCPServerTask: "_rpc_lock", "_pending_refresh_tasks", "_pending_call_context", "initialize_result", "_ping_unsupported", + "fingerprint", ) def __init__(self, name: str): @@ -2441,6 +2452,54 @@ async def shutdown(self): # Module-level state # --------------------------------------------------------------------------- +def _get_mcp_config_fingerprint(server_name: str, config: dict) -> str: + import hashlib + cleaned = {k: v for k, v in config.items() if k not in ("enabled",)} + serialized = json.dumps(cleaned, sort_keys=True) + h = hashlib.sha256(serialized.encode("utf-8")).hexdigest() + return f"{server_name}:{h[:16]}" + + +def _get_active_server(server_name: str) -> Optional[MCPServerTask]: + try: + cfg = _load_mcp_config().get(server_name) + if cfg and _parse_boolish(cfg.get("enabled", True), default=True): + fingerprint = _get_mcp_config_fingerprint(server_name, cfg) + need_connect = False + with _lock: + if fingerprint not in _servers: + need_connect = True + if need_connect: + discover_mcp_tools() + except Exception: + pass + + with _lock: + try: + cfg = _load_mcp_config().get(server_name) + if cfg: + fingerprint = _get_mcp_config_fingerprint(server_name, cfg) + srv = _servers.get(fingerprint) + if srv is not None: + return srv + except Exception: + pass + srv = _servers.get(server_name) + if srv is not None: + return srv + return None + + +def _get_active_fingerprint(server_name: str) -> str: + try: + cfg = _load_mcp_config().get(server_name) + if cfg: + return _get_mcp_config_fingerprint(server_name, cfg) + except Exception: + pass + return server_name + + _servers: Dict[str, MCPServerTask] = {} _server_connecting: set[str] = set() _server_connect_errors: Dict[str, str] = {} @@ -2614,8 +2673,7 @@ async def _recover(): recovered = False if recovered: - with _lock: - srv = _servers.get(server_name) + srv = _get_active_server(server_name) if srv is not None and hasattr(srv, "_reconnect_event"): loop = _mcp_loop if loop is not None and loop.is_running(): @@ -2762,8 +2820,7 @@ def _handle_session_expired_and_retry( if not _is_session_expired_error(exc): return None - with _lock: - srv = _servers.get(server_name) + srv = _get_active_server(server_name) if srv is None or not hasattr(srv, "_reconnect_event"): return None @@ -3089,22 +3146,16 @@ def _load_mcp_config() -> Dict[str, dict]: ``os.environ`` (which includes ``~/.hermes/.env`` loaded at startup). """ try: - from hermes_cli.config import load_config + from hermes_cli.config import read_raw_config # Safe mode (--safe-mode / HERMES_SAFE_MODE=1): troubleshooting run # with all customizations disabled — no MCP servers connect. from utils import env_var_enabled as _env_enabled if _env_enabled("HERMES_SAFE_MODE"): return {} - config = load_config() + config = read_raw_config() servers = config.get("mcp_servers") if not servers or not isinstance(servers, dict): return {} - # Ensure .env vars are available for interpolation - try: - from hermes_cli.env_loader import load_hermes_dotenv - load_hermes_dotenv() - except Exception: - pass safe_servers: Dict[str, dict] = {} for name, cfg in _filter_suspicious_mcp_servers(servers).items(): interpolated = _interpolate_env_vars(cfg) @@ -3158,8 +3209,9 @@ def _handler(args: dict, **kwargs) -> str: # failure the error paths below bump the count again, which # re-stamps the open-time via _bump_server_error (re-arming # the cooldown). - if _server_error_counts.get(server_name, 0) >= _CIRCUIT_BREAKER_THRESHOLD: - opened_at = _server_breaker_opened_at.get(server_name, 0.0) + active_fp = _get_active_fingerprint(server_name) + if _server_error_counts.get(active_fp, 0) >= _CIRCUIT_BREAKER_THRESHOLD: + opened_at = _server_breaker_opened_at.get(active_fp, 0.0) age = time.monotonic() - opened_at if age < _CIRCUIT_BREAKER_COOLDOWN_SEC: remaining = max(1, int(_CIRCUIT_BREAKER_COOLDOWN_SEC - age)) @@ -3174,10 +3226,9 @@ def _handler(args: dict, **kwargs) -> str: }, ensure_ascii=False) # Cooldown elapsed → fall through as a half-open probe. - with _lock: - server = _servers.get(server_name) + server = _get_active_server(server_name) if not server or not server.session: - _bump_server_error(server_name) + _bump_server_error(active_fp) return json.dumps({ "error": f"MCP server '{server_name}' is not connected" }, ensure_ascii=False) @@ -3249,11 +3300,11 @@ def _call_once(): try: parsed = json.loads(result) if "error" in parsed: - _bump_server_error(server_name) + _bump_server_error(active_fp) else: - _reset_server_error(server_name) # success — reset + _reset_server_error(active_fp) # success — reset except (json.JSONDecodeError, TypeError): - _reset_server_error(server_name) # non-JSON = success + _reset_server_error(active_fp) # non-JSON = success return result except InterruptedError: return _interrupted_call_result() @@ -3278,7 +3329,7 @@ def _call_once(): if recovered is not None: return recovered - _bump_server_error(server_name) + _bump_server_error(active_fp) logger.error( "MCP tool %s/%s call failed: %s", server_name, tool_name, exc, @@ -3356,8 +3407,7 @@ def _make_read_resource_handler(server_name: str, tool_timeout: float): def _handler(args: dict, **kwargs) -> str: from tools.registry import tool_error - with _lock: - server = _servers.get(server_name) + server = _get_active_server(server_name) if not server or not server.session: return json.dumps({ "error": f"MCP server '{server_name}' is not connected" @@ -3414,8 +3464,7 @@ def _make_list_prompts_handler(server_name: str, tool_timeout: float): """Return a sync handler that lists prompts from an MCP server.""" def _handler(args: dict, **kwargs) -> str: - with _lock: - server = _servers.get(server_name) + server = _get_active_server(server_name) if not server or not server.session: return json.dumps({ "error": f"MCP server '{server_name}' is not connected" @@ -3479,8 +3528,7 @@ def _make_get_prompt_handler(server_name: str, tool_timeout: float): def _handler(args: dict, **kwargs) -> str: from tools.registry import tool_error - with _lock: - server = _servers.get(server_name) + server = _get_active_server(server_name) if not server or not server.session: return json.dumps({ "error": f"MCP server '{server_name}' is not connected" @@ -3548,8 +3596,7 @@ def _make_check_fn(server_name: str): """Return a check function that verifies the MCP connection is alive.""" def _check() -> bool: - with _lock: - server = _servers.get(server_name) + server = _get_active_server(server_name) return server is not None and server.session is not None return _check @@ -3678,19 +3725,31 @@ def sanitize_mcp_name_component(value: str) -> str: return re.sub(r"[^A-Za-z0-9_]", "_", str(value or "")) -def _convert_mcp_schema(server_name: str, mcp_tool) -> dict: +def _convert_mcp_schema(server_name: str, mcp_tool, fingerprint: Optional[str] = None) -> dict: """Convert an MCP tool listing to the Hermes registry schema format. Args: server_name: The logical server name for prefixing. mcp_tool: An MCP ``Tool`` object with ``.name``, ``.description``, and ``.inputSchema``. + fingerprint: Optional server fingerprint for profile isolation. Returns: A dict suitable for ``registry.register(schema=...)``. """ + from hermes_cli.profiles import get_active_profile_name + try: + pname = get_active_profile_name() + except Exception: + pname = "default" + safe_tool_name = sanitize_mcp_name_component(mcp_tool.name) - safe_server_name = sanitize_mcp_name_component(server_name) + if pname in ("default", "custom", "main"): + safe_server_name = sanitize_mcp_name_component(server_name) + elif fingerprint: + safe_server_name = sanitize_mcp_name_component(fingerprint.replace(":", "_")) + else: + safe_server_name = sanitize_mcp_name_component(server_name) prefixed_name = f"mcp_{safe_server_name}_{safe_tool_name}" return { "name": prefixed_name, @@ -3699,13 +3758,24 @@ def _convert_mcp_schema(server_name: str, mcp_tool) -> dict: } -def _build_utility_schemas(server_name: str) -> List[dict]: +def _build_utility_schemas(server_name: str, fingerprint: Optional[str] = None) -> List[dict]: """Build schemas for the MCP utility tools (resources & prompts). Returns a list of (schema, handler_factory_name) tuples encoded as dicts with keys: schema, handler_key. """ - safe_name = sanitize_mcp_name_component(server_name) + from hermes_cli.profiles import get_active_profile_name + try: + pname = get_active_profile_name() + except Exception: + pname = "default" + + if pname in ("default", "custom", "main"): + safe_name = sanitize_mcp_name_component(server_name) + elif fingerprint: + safe_name = sanitize_mcp_name_component(fingerprint.replace(":", "_")) + else: + safe_name = sanitize_mcp_name_component(server_name) return [ { "schema": { @@ -3838,7 +3908,7 @@ def _forget_mcp_tool_server(tool_name: str) -> None: _mcp_tool_server_names.pop(tool_name, None) -def _select_utility_schemas(server_name: str, server: MCPServerTask, config: dict) -> List[dict]: +def _select_utility_schemas(server_name: str, server: MCPServerTask, config: dict, fingerprint: Optional[str] = None) -> List[dict]: """Select utility schemas based on config and server capabilities.""" tools_filter = config.get("tools") or {} resources_enabled = _parse_boolish(tools_filter.get("resources"), default=True) @@ -3855,7 +3925,7 @@ def _select_utility_schemas(server_name: str, server: MCPServerTask, config: dic advertised_caps = getattr(init_result, "capabilities", None) selected: List[dict] = [] - for entry in _build_utility_schemas(server_name): + for entry in _build_utility_schemas(server_name, fingerprint=fingerprint): handler_key = entry["handler_key"] if handler_key in {"list_resources", "read_resource"} and not resources_enabled: logger.debug("MCP server '%s': skipping utility '%s' (resources disabled)", server_name, handler_key) @@ -3898,12 +3968,16 @@ def _select_utility_schemas(server_name: str, server: MCPServerTask, config: dic def _existing_tool_names() -> List[str]: """Return tool names for all currently connected servers.""" names: List[str] = [] + seen_servers = set() for _sname, server in _servers.items(): + if server in seen_servers: + continue + seen_servers.add(server) if hasattr(server, "_registered_tool_names"): names.extend(server._registered_tool_names) continue for mcp_tool in server._tools: - schema = _convert_mcp_schema(server.name, mcp_tool) + schema = _convert_mcp_schema(server.name, mcp_tool, fingerprint=getattr(server, "fingerprint", None)) names.append(schema["name"]) return names @@ -3923,7 +3997,8 @@ def _register_server_tools(name: str, server: MCPServerTask, config: dict) -> Li from tools.registry import registry registered_names: List[str] = [] - toolset_name = f"mcp-{name}" + fp = getattr(server, "fingerprint", None) or name + toolset_name = f"mcp-{fp}" # Selective tool loading: honour include/exclude lists from config. # Rules (matching issue #690 spec): @@ -3950,7 +4025,7 @@ def _should_register(tool_name: str) -> bool: # Scan tool description for prompt injection patterns _scan_mcp_description(name, mcp_tool.name, mcp_tool.description or "") - schema = _convert_mcp_schema(name, mcp_tool) + schema = _convert_mcp_schema(name, mcp_tool, fingerprint=fp) tool_name_prefixed = schema["name"] # Guard against collisions with built-in (non-MCP) tools. @@ -3984,7 +4059,7 @@ def _should_register(tool_name: str) -> bool: "get_prompt": _make_get_prompt_handler, } check_fn = _make_check_fn(name) - for entry in _select_utility_schemas(name, server, config): + for entry in _select_utility_schemas(name, server, config, fingerprint=fp): schema = entry["schema"] handler_key = entry["handler_key"] handler = _handler_factories[handler_key](name, server.tool_timeout) @@ -4023,15 +4098,17 @@ async def _discover_and_register_server(name: str, config: dict) -> List[str]: Returns list of registered tool names. """ + fingerprint = _get_mcp_config_fingerprint(name, config) connect_timeout = config.get("connect_timeout", _DEFAULT_CONNECT_TIMEOUT) server = await asyncio.wait_for( _connect_server(name, config), timeout=connect_timeout, ) + server.fingerprint = fingerprint with _lock: - _server_connecting.discard(name) - _server_connect_errors.pop(name, None) - _servers[name] = server + _server_connecting.discard(fingerprint) + _server_connect_errors.pop(fingerprint, None) + _servers[fingerprint] = server registered_names = _register_server_tools(name, server, config) server._registered_tool_names = list(registered_names) @@ -4073,13 +4150,18 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: # Only attempt servers that aren't already connected and are enabled # (enabled: false skips the server entirely without removing its config) with _lock: - new_servers = { - k: v - for k, v in servers.items() - if k not in _servers and _parse_boolish(v.get("enabled", True), default=True) - } - _server_connecting.update(new_servers) - for srv_name in new_servers: + new_servers = {} + for k, v in servers.items(): + if not _parse_boolish(v.get("enabled", True), default=True): + continue + fp = _get_mcp_config_fingerprint(k, v) + if fp not in _servers: + new_servers[fp] = (k, v) + _server_connecting.update(new_servers.keys()) + for fp in new_servers: + _server_connect_errors.pop(fp, None) + # Remove from errors using server name as fallback + srv_name = fp.split(":")[0] _server_connect_errors.pop(srv_name, None) # Track which servers opt-in to parallel tool calls (idempotent). for srv_name, srv_cfg in servers.items(): @@ -4094,24 +4176,26 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: # Start the background event loop for MCP connections _ensure_mcp_loop() - async def _discover_one(name: str, cfg: dict) -> List[str]: + async def _discover_one(fp: str, name: str, cfg: dict) -> List[str]: """Connect to a single server and return its registered tool names.""" + # Avoid passing fingerprint as keyword argument to match original signature of _discover_and_register_server return await _discover_and_register_server(name, cfg) async def _discover_all(): - server_names = list(new_servers.keys()) + fps = list(new_servers.keys()) # Connect to all servers in PARALLEL results = await asyncio.gather( - *(_discover_one(name, cfg) for name, cfg in new_servers.items()), + *(_discover_one(fp, name, cfg) for fp, (name, cfg) in new_servers.items()), return_exceptions=True, ) - for name, result in zip(server_names, results): + for fp, result in zip(fps, results): + name, cfg = new_servers[fp] if isinstance(result, BaseException): - command = new_servers.get(name, {}).get("command") + command = cfg.get("command") message = _format_connect_error(result) with _lock: - _server_connecting.discard(name) - _server_connect_errors[name] = message + _server_connecting.discard(fp) + _server_connect_errors[fp] = message logger.warning( "Failed to connect to MCP server '%s'%s: %s", name, @@ -4120,8 +4204,8 @@ async def _discover_all(): ) else: with _lock: - _server_connecting.discard(name) - _server_connect_errors.pop(name, None) + _server_connecting.discard(fp) + _server_connect_errors.pop(fp, None) # Per-server timeouts are handled inside _discover_and_register_server. # The outer timeout is generous: 120s total for parallel discovery. @@ -4141,10 +4225,17 @@ async def _discover_all(): # Log a summary so ACP callers get visibility into what was registered. with _lock: - connected = [n for n in new_servers if n in _servers] + connected = [] + for fp in new_servers: + if fp in _servers: + connected.append(fp) + else: + srv_name = fp.split(":")[0] + if srv_name in _servers: + connected.append(srv_name) new_tool_count = sum( - len(getattr(_servers[n], "_registered_tool_names", [])) - for n in connected + len(getattr(_servers[key], "_registered_tool_names", [])) + for key in connected ) failed = len(new_servers) - len(connected) if new_tool_count or failed: @@ -4178,26 +4269,36 @@ def discover_mcp_tools() -> List[str]: return [] with _lock: - new_server_names = [ - name - for name, cfg in servers.items() - if name not in _servers and _parse_boolish(cfg.get("enabled", True), default=True) - ] + new_server_names = [] + for name, cfg in servers.items(): + if not _parse_boolish(cfg.get("enabled", True), default=True): + continue + fp = _get_mcp_config_fingerprint(name, cfg) + if fp not in _servers: + new_server_names.append(name) tool_names = register_mcp_servers(servers) if not new_server_names: return tool_names with _lock: - connected_server_names = [name for name in new_server_names if name in _servers] + connected_server_keys = [] + for name in new_server_names: + cfg = servers.get(name) + if cfg: + fp = _get_mcp_config_fingerprint(name, cfg) + if fp in _servers: + connected_server_keys.append(fp) + elif name in _servers: + connected_server_keys.append(name) new_tool_count = sum( - len(getattr(_servers[name], "_registered_tool_names", [])) - for name in connected_server_names + len(getattr(_servers[key], "_registered_tool_names", [])) + for key in connected_server_keys ) - failed_count = len(new_server_names) - len(connected_server_names) + failed_count = len(new_server_names) - len(connected_server_keys) if new_tool_count or failed_count: - summary = f" MCP: {new_tool_count} tool(s) from {len(connected_server_names)} server(s)" + summary = f" MCP: {new_tool_count} tool(s) from {len(connected_server_keys)} server(s)" if failed_count: summary += f" ({failed_count} failed)" logger.info(summary) @@ -4246,7 +4347,8 @@ def get_mcp_status() -> List[dict]: for name, cfg in configured.items(): transport = cfg.get("transport", "http") if "url" in cfg else "stdio" enabled = _parse_boolish(cfg.get("enabled", True), default=True) - server = active_servers.get(name) + fp = _get_mcp_config_fingerprint(name, cfg) + server = active_servers.get(fp) if server and server.session is not None: entry = { "name": name, @@ -4271,7 +4373,7 @@ def get_mcp_status() -> List[dict]: "disabled": True, "status": "disabled", }) - elif name in connecting: + elif fp in connecting or name in connecting: result.append({ "name": name, "transport": transport, @@ -4280,7 +4382,7 @@ def get_mcp_status() -> List[dict]: "disabled": False, "status": "connecting", }) - elif name in connect_errors: + elif fp in connect_errors or name in connect_errors: result.append({ "name": name, "transport": transport, @@ -4288,7 +4390,7 @@ def get_mcp_status() -> List[dict]: "connected": False, "disabled": False, "status": "failed", - "error": connect_errors[name], + "error": connect_errors.get(fp) or connect_errors.get(name), }) else: result.append({ @@ -4579,18 +4681,39 @@ def _add(schema: dict) -> bool: def shutdown_mcp_servers(): - """Close all MCP server connections and stop the background loop. + """Close MCP server connections matching the active profile config. - Each server Task is signalled to exit its ``async with`` block so that - the anyio cancel-scope cleanup happens in the same Task that opened it. - All servers are shut down in parallel via ``asyncio.gather``. + If no active profile config can be determined or we are shutting down the main + agent, we shut down everything. """ + try: + active_cfgs = _load_mcp_config() + active_fps = { + _get_mcp_config_fingerprint(name, cfg) + for name, cfg in active_cfgs.items() + } + except Exception: + active_fps = set() + with _lock: - servers_snapshot = list(_servers.values()) + if active_fps: + servers_snapshot = [] + for k in list(_servers.keys()): + srv = _servers[k] + srv_fp = getattr(srv, "fingerprint", None) + if k in active_fps or srv_fp in active_fps: + if srv not in servers_snapshot: + servers_snapshot.append(srv) + _servers.pop(k, None) + else: + servers_snapshot = list(_servers.values()) + _servers.clear() # Fast path: nothing to shut down. if not servers_snapshot: - _stop_mcp_loop() + with _lock: + if not _servers: + _stop_mcp_loop() return async def _shutdown(): @@ -4603,8 +4726,6 @@ async def _shutdown(): logger.debug( "Error closing MCP server '%s': %s", server.name, result, ) - with _lock: - _servers.clear() with _lock: loop = _mcp_loop @@ -4621,7 +4742,9 @@ async def _shutdown(): except BaseException as exc: logger.debug("Error during MCP shutdown: %s", exc) - _stop_mcp_loop() + with _lock: + if not _servers: + _stop_mcp_loop() def _kill_orphaned_mcp_children(include_active: bool = False) -> None: diff --git a/tools/memory_tool.py b/tools/memory_tool.py index 96a8a1a0a1804..d01e5b7b21b2d 100644 --- a/tools/memory_tool.py +++ b/tools/memory_tool.py @@ -52,8 +52,14 @@ # (HERMES_HOME env var changes) are always respected. The old module-level # constant was cached at import time and could go stale if a profile switch # happened after the first import. -def get_memory_dir() -> Path: +def get_memory_dir(profile_name: str = None) -> Path: """Return the profile-scoped memories directory.""" + if profile_name and profile_name != "default" and profile_name != "main": + from hermes_cli.profiles import get_profile_dir + try: + return get_profile_dir(profile_name) / "memories" + except Exception: + pass return get_hermes_home() / "memories" ENTRY_DELIMITER = "\n§\n" diff --git a/tools/process_registry.py b/tools/process_registry.py index 69bfc13619da1..25f463dbce9c3 100644 --- a/tools/process_registry.py +++ b/tools/process_registry.py @@ -39,6 +39,7 @@ import threading import time import uuid +from pathlib import Path _IS_WINDOWS = platform.system() == "Windows" from tools.environments.local import _find_shell, _resolve_safe_cwd, _sanitize_subprocess_env @@ -53,6 +54,7 @@ # Checkpoint file for crash recovery (gateway only) CHECKPOINT_PATH = get_hermes_home() / "processes.json" +_IMPORT_CHECKPOINT_PATH = CHECKPOINT_PATH # Limits MAX_OUTPUT_CHARS = 200_000 # 200KB rolling output buffer @@ -75,6 +77,24 @@ WATCH_GLOBAL_COOLDOWN_SECONDS = 30 +def _checkpoint_path_for_home(home: str | Path | None = None) -> Path: + """Return the checkpoint path for the active Hermes home. + + Tests historically monkeypatch ``CHECKPOINT_PATH`` directly. Preserve + that contract; otherwise derive the path from the current/profile home so + routed profile process metadata is not persisted in the gateway home. + """ + configured = Path(CHECKPOINT_PATH) + if configured != _IMPORT_CHECKPOINT_PATH: + return configured + base = Path(home).expanduser() if home else get_hermes_home() + return base / "processes.json" + + +def _checkpoint_path_overridden() -> bool: + return Path(CHECKPOINT_PATH) != _IMPORT_CHECKPOINT_PATH + + def format_uptime_short(seconds: int) -> str: s = max(0, int(seconds)) if s < 60: @@ -93,6 +113,8 @@ class ProcessSession: command: str # Original command string task_id: str = "" # Task/sandbox isolation key session_key: str = "" # Gateway session key (for reset protection) + agent_profile: str = "" # Routed gateway profile name + agent_hermes_home: str = "" # Routed profile HERMES_HOME pid: Optional[int] = None # OS process ID process: Optional[subprocess.Popen] = None # Popen handle (local only) env_ref: Any = None # Reference to the environment object @@ -195,6 +217,7 @@ def __init__(self): self._global_watch_window_hits: int = 0 self._global_watch_tripped_until: float = 0.0 self._global_watch_suppressed_during_trip: int = 0 + self._checkpoint_paths: set[Path] = {_checkpoint_path_for_home()} @staticmethod def _clean_shell_noise(text: str) -> str: @@ -288,6 +311,8 @@ def _check_watch_patterns(self, session: ProcessSession, new_text: str) -> None: self.completion_queue.put({ "session_id": session.id, "session_key": session.session_key, + "agent_profile": session.agent_profile, + "agent_hermes_home": session.agent_hermes_home, "command": session.command, "type": "watch_disabled", "suppressed": session._watch_suppressed, @@ -319,6 +344,8 @@ def _check_watch_patterns(self, session: ProcessSession, new_text: str) -> None: self.completion_queue.put({ "session_id": session.id, "session_key": session.session_key, + "agent_profile": session.agent_profile, + "agent_hermes_home": session.agent_hermes_home, "command": session.command, "type": "watch_match", "pattern": matched_pattern, @@ -666,6 +693,8 @@ def spawn_local( session_key: str = "", env_vars: dict = None, use_pty: bool = False, + agent_profile: str = "", + agent_hermes_home: str = "", ) -> ProcessSession: """ Spawn a background process locally. @@ -682,6 +711,8 @@ def spawn_local( command=command, task_id=task_id, session_key=session_key, + agent_profile=agent_profile, + agent_hermes_home=agent_hermes_home, cwd=_resolve_safe_cwd(cwd or os.getcwd()), started_at=time.time(), ) @@ -805,6 +836,8 @@ def spawn_via_env( task_id: str = "", session_key: str = "", timeout: int = 10, + agent_profile: str = "", + agent_hermes_home: str = "", ) -> ProcessSession: """ Spawn a background process through a non-local environment backend. @@ -1042,6 +1075,8 @@ def _move_to_finished(self, session: ProcessSession): "type": "completion", "session_id": session.id, "session_key": session.session_key, + "agent_profile": session.agent_profile, + "agent_hermes_home": session.agent_hermes_home, "command": session.command, "exit_code": session.exit_code, "completion_reason": session.completion_reason, @@ -1131,6 +1166,51 @@ def get(self, session_id: str) -> Optional[ProcessSession]: session = self._running.get(session_id) or self._finished.get(session_id) return self._refresh_detached_session(session) + @staticmethod + def _normalize_home(value: str) -> str: + value = str(value or "").strip() + if not value: + return "" + try: + return str(Path(value).expanduser().resolve(strict=False)) + except Exception: + return value + + def _current_scope(self) -> tuple[bool, str, str]: + try: + from gateway.session_context import get_session_env + active_profile = get_session_env("HERMES_SESSION_AGENT_PROFILE", "") + active_home = get_session_env("HERMES_SESSION_AGENT_HERMES_HOME", "") + except Exception: + active_profile = "" + active_home = "" + + active_home = self._normalize_home(active_home) + active_profile = str(active_profile or "").strip() + is_routed = bool(active_profile or active_home) + return is_routed, active_home, active_profile + + def _in_current_scope(self, session: ProcessSession) -> bool: + is_routed, active_home, active_profile = self._current_scope() + session_home = self._normalize_home(session.agent_hermes_home) + session_profile = str(session.agent_profile or "").strip() + if not session_profile: + parts = str(session.session_key or "").split(":", 2) + if len(parts) >= 2 and parts[0] == "agent" and parts[1] not in {"", "main", "cron"}: + session_profile = parts[1] + if is_routed: + if active_home: + return bool(session_home and session_home == active_home) + return False + return not bool(session_home or session_profile) + + def _get_for_current_scope(self, session_id: str) -> Optional[ProcessSession]: + """Return a process only when it belongs to the active profile scope.""" + session = self.get(session_id) + if session is None or not self._in_current_scope(session): + return None + return session + def _reconcile_local_exit(self, session: "ProcessSession") -> None: """Reconcile session.exited against the real child process state. @@ -1520,6 +1600,7 @@ def list_sessions(self, task_id: str = None) -> list: all_sessions = list(self._running.values()) + list(self._finished.values()) all_sessions = [self._refresh_detached_session(s) for s in all_sessions] + all_sessions = [s for s in all_sessions if self._in_current_scope(s)] if task_id: all_sessions = [s for s in all_sessions if s.task_id == task_id] @@ -1650,11 +1731,16 @@ def _prune_if_needed(self): # ----- Checkpoint (crash recovery) ----- + @staticmethod + def _checkpoint_path_for_session(session: ProcessSession) -> Path: + home = str(session.agent_hermes_home or "").strip() + return _checkpoint_path_for_home(home or None) + def _write_checkpoint(self): """Write running process metadata to checkpoint file atomically.""" try: with self._lock: - entries = [] + entries_by_path: Dict[Path, list[dict]] = {} for s in self._running.values(): if not s.exited: # Lazily backfill the kernel start time for host PIDs so @@ -1662,7 +1748,8 @@ def _write_checkpoint(self): # for sessions spawned before this field existed. if s.host_start_time is None and s.pid_scope == "host" and s.pid: s.host_start_time = self._safe_host_start_time(s.pid) - entries.append({ + path = self._checkpoint_path_for_session(s) + entries_by_path.setdefault(path, []).append({ "session_id": s.id, "command": s.command, "pid": s.pid, @@ -1672,6 +1759,8 @@ def _write_checkpoint(self): "started_at": s.started_at, "task_id": s.task_id, "session_key": s.session_key, + "agent_profile": s.agent_profile, + "agent_hermes_home": s.agent_hermes_home, "watcher_platform": s.watcher_platform, "watcher_chat_id": s.watcher_chat_id, "watcher_user_id": s.watcher_user_id, @@ -1682,27 +1771,63 @@ def _write_checkpoint(self): "notify_on_complete": s.notify_on_complete, "watch_patterns": s.watch_patterns, }) + if _checkpoint_path_overridden(): + paths = set(entries_by_path) or {Path(CHECKPOINT_PATH)} + else: + paths = set(self._checkpoint_paths) | set(entries_by_path) + if not paths: + paths = {_checkpoint_path_for_home()} + self._checkpoint_paths = paths # Atomic write to avoid corruption on crash from utils import atomic_json_write - atomic_json_write(CHECKPOINT_PATH, entries) + for path in paths: + path.parent.mkdir(parents=True, exist_ok=True) + atomic_json_write(path, entries_by_path.get(path, [])) except Exception as e: logger.debug("Failed to write checkpoint file: %s", e, exc_info=True) - def recover_from_checkpoint(self) -> int: + def recover_from_checkpoint(self, checkpoint_paths: list[str | Path] | None = None) -> int: """ On gateway startup, probe PIDs from checkpoint file. Returns the number of processes recovered as detached. """ - if not CHECKPOINT_PATH.exists(): - return 0 + paths: list[Path] = [_checkpoint_path_for_home()] + if not _checkpoint_path_overridden(): + from hermes_constants import get_default_hermes_root + try: + root = get_default_hermes_root() + profiles_dir = root / "profiles" + if profiles_dir.exists(): + for p_dir in profiles_dir.iterdir(): + if p_dir.is_dir() and (p_dir / "processes.json").exists(): + paths.append(p_dir / "processes.json") + except Exception: + pass + if checkpoint_paths and not _checkpoint_path_overridden(): + paths.extend(Path(path) for path in checkpoint_paths) - try: - entries = json.loads(CHECKPOINT_PATH.read_text(encoding="utf-8")) - except Exception: + seen_paths: set[Path] = set() + entries: list[dict] = [] + for path in paths: + if path in seen_paths: + continue + seen_paths.add(path) + if not path.exists(): + continue + try: + loaded = json.loads(path.read_text(encoding="utf-8")) + except Exception: + continue + if isinstance(loaded, list): + entries.extend(entry for entry in loaded if isinstance(entry, dict)) + + if not entries: return 0 + self._checkpoint_paths |= seen_paths + recovered = 0 for entry in entries: pid = entry.get("pid") @@ -1739,26 +1864,31 @@ def recover_from_checkpoint(self) -> int: ) continue + raw_agent_profile = str(entry.get("agent_profile", "") or "").strip() + raw_agent_home = str(entry.get("agent_hermes_home", "") or "").strip() + legacy_profile_without_home = bool(raw_agent_profile and not raw_agent_home) session = ProcessSession( id=entry["session_id"], command=entry.get("command", "unknown"), task_id=entry.get("task_id", ""), session_key=entry.get("session_key", ""), + agent_profile="" if legacy_profile_without_home else raw_agent_profile, + agent_hermes_home=raw_agent_home, pid=pid, host_start_time=recorded_start, pid_scope=pid_scope, cwd=entry.get("cwd"), started_at=entry.get("started_at", time.time()), detached=True, # Can't read output, but can report status + kill - watcher_platform=entry.get("watcher_platform", ""), - watcher_chat_id=entry.get("watcher_chat_id", ""), - watcher_user_id=entry.get("watcher_user_id", ""), - watcher_user_name=entry.get("watcher_user_name", ""), - watcher_thread_id=entry.get("watcher_thread_id", ""), - watcher_message_id=entry.get("watcher_message_id", ""), - watcher_interval=entry.get("watcher_interval", 0), - notify_on_complete=entry.get("notify_on_complete", False), - watch_patterns=entry.get("watch_patterns", []), + watcher_platform="" if legacy_profile_without_home else entry.get("watcher_platform", ""), + watcher_chat_id="" if legacy_profile_without_home else entry.get("watcher_chat_id", ""), + watcher_user_id="" if legacy_profile_without_home else entry.get("watcher_user_id", ""), + watcher_user_name="" if legacy_profile_without_home else entry.get("watcher_user_name", ""), + watcher_thread_id="" if legacy_profile_without_home else entry.get("watcher_thread_id", ""), + watcher_message_id="" if legacy_profile_without_home else entry.get("watcher_message_id", ""), + watcher_interval=0 if legacy_profile_without_home else entry.get("watcher_interval", 0), + notify_on_complete=False if legacy_profile_without_home else entry.get("notify_on_complete", False), + watch_patterns=[] if legacy_profile_without_home else entry.get("watch_patterns", []), ) with self._lock: self._running[session.id] = session @@ -1771,6 +1901,8 @@ def recover_from_checkpoint(self) -> int: "session_id": session.id, "check_interval": session.watcher_interval, "session_key": session.session_key, + "agent_profile": session.agent_profile, + "agent_hermes_home": session.agent_hermes_home, "platform": session.watcher_platform, "chat_id": session.watcher_chat_id, "user_id": session.watcher_user_id, @@ -2055,6 +2187,12 @@ def _handle_process(args, **kw): elif action in {"poll", "log", "wait", "kill", "write", "submit", "close"}: if not session_id: return tool_error(f"session_id is required for {action}") + scoped = process_registry._get_for_current_scope(session_id) + if scoped is None: + return json.dumps({ + "status": "not_found", + "error": "No process with ID in current profile scope", + }, ensure_ascii=False) if action == "poll": return json.dumps(process_registry.poll(session_id), ensure_ascii=False) elif action == "log": diff --git a/tools/registry.py b/tools/registry.py index 7bb92e85f960c..31e2ab81bd280 100644 --- a/tools/registry.py +++ b/tools/registry.py @@ -119,15 +119,22 @@ def __init__(self, name, toolset, schema, handler, check_fn, # --------------------------------------------------------------------------- _CHECK_FN_TTL_SECONDS = 30.0 -_check_fn_cache: Dict[Callable, tuple[float, bool]] = {} +_check_fn_cache: Dict[tuple, tuple[float, bool]] = {} _check_fn_cache_lock = threading.Lock() def _check_fn_cached(fn: Callable) -> bool: """Return bool(fn()), TTL-cached across calls. Swallows exceptions as False.""" now = time.monotonic() + try: + from hermes_constants import get_hermes_home + home_path = get_hermes_home() + except Exception: + home_path = None + cache_key = (fn, home_path) + with _check_fn_cache_lock: - cached = _check_fn_cache.get(fn) + cached = _check_fn_cache.get(cache_key) if cached is not None: ts, value = cached if now - ts < _CHECK_FN_TTL_SECONDS: @@ -137,7 +144,7 @@ def _check_fn_cached(fn: Callable) -> bool: except Exception: value = False with _check_fn_cache_lock: - _check_fn_cache[fn] = (now, value) + _check_fn_cache[cache_key] = (now, value) return value @@ -224,6 +231,18 @@ def get_registered_toolset_aliases(self) -> Dict[str, str]: def get_toolset_alias_target(self, alias: str) -> Optional[str]: """Return the canonical toolset name for an alias, or None.""" + try: + from tools.mcp_tool import _load_mcp_config, _get_mcp_config_fingerprint, discover_mcp_tools, _servers + cfg = _load_mcp_config().get(alias) + if cfg: + fp = _get_mcp_config_fingerprint(alias, cfg) + # Lazily connect if not already connected + if fp not in _servers: + discover_mcp_tools() + return f"mcp-{fp}" + except Exception: + pass + with self._lock: return self._toolset_aliases.get(alias) diff --git a/tools/session_search_tool.py b/tools/session_search_tool.py index 3a1040e5273b7..3fc05ae15bc57 100644 --- a/tools/session_search_tool.py +++ b/tools/session_search_tool.py @@ -702,9 +702,9 @@ def session_search( def check_session_search_requirements() -> bool: """Requires the SQLite state database.""" try: - from hermes_state import DEFAULT_DB_PATH - return DEFAULT_DB_PATH.parent.exists() - except ImportError: + from hermes_constants import get_hermes_home + return get_hermes_home().exists() + except Exception: return False diff --git a/tools/skill_manager_tool.py b/tools/skill_manager_tool.py index 3a6f315b22428..0f53420632ac2 100644 --- a/tools/skill_manager_tool.py +++ b/tools/skill_manager_tool.py @@ -39,7 +39,11 @@ import shutil import tempfile from pathlib import Path -from hermes_constants import get_hermes_home, display_hermes_home +from hermes_constants import ( + get_hermes_home, + display_hermes_home, + get_skills_dir as _active_skills_dir, +) from typing import Dict, Any, List, Optional, Tuple from utils import atomic_replace, is_truthy_value @@ -107,6 +111,15 @@ def _security_scan_skill(skill_dir: Path) -> Optional[str]: # All skills live in ~/.hermes/skills/ (single source of truth) HERMES_HOME = get_hermes_home() SKILLS_DIR = HERMES_HOME / "skills" +_IMPORT_SKILLS_DIR = SKILLS_DIR + + +def get_skills_dir() -> Path: + """Return the active profile's skills dir, preserving SKILLS_DIR monkeypatches.""" + configured = Path(SKILLS_DIR) + if configured != _IMPORT_SKILLS_DIR: + return configured + return _active_skills_dir() MAX_NAME_LENGTH = 64 MAX_DESCRIPTION_LENGTH = 1024 @@ -131,7 +144,7 @@ def _containing_skills_root(skill_path: Path) -> Path: return root except (ValueError, OSError): continue - return SKILLS_DIR + return get_skills_dir() def _is_path_redirect(path: Path) -> bool: @@ -417,9 +430,10 @@ def _validate_content_size(content: str, label: str = "SKILL.md") -> Optional[st def _resolve_skill_dir(name: str, category: str = None) -> Path: """Build the directory path for a new skill, optionally under a category.""" + skills_dir = get_skills_dir() if category: - return SKILLS_DIR / category / name - return SKILLS_DIR / name + return skills_dir / category / name + return skills_dir / name def _find_skill(name: str) -> Optional[Dict[str, Any]]: @@ -465,7 +479,8 @@ def _find_skill_in_other_profiles(name: str) -> List[Tuple[str, Path]]: # Collect (profile_name, skills_dir) for every profile EXCEPT the # one whose SKILLS_DIR we already searched in _find_skill(). - active_dir = SKILLS_DIR.resolve() if SKILLS_DIR.exists() else SKILLS_DIR + skills_dir_path = get_skills_dir() + active_dir = skills_dir_path.resolve() if skills_dir_path.exists() else skills_dir_path candidates: List[Tuple[str, Path]] = [] # Default profile (~/.hermes/skills) — only consider when active is non-default. @@ -684,7 +699,7 @@ def _create_skill(name: str, content: str, category: str = None) -> Dict[str, An result = { "success": True, "message": f"Skill '{name}' created.", - "path": str(skill_dir.relative_to(SKILLS_DIR)), + "path": str(skill_dir.relative_to(get_skills_dir())), "skill_md": str(skill_md), "_change": {"description": _desc}, } diff --git a/tools/skills_tool.py b/tools/skills_tool.py index 2ba57adc54d46..ed009d01fa223 100644 --- a/tools/skills_tool.py +++ b/tools/skills_tool.py @@ -69,7 +69,11 @@ import json import logging -from hermes_constants import get_hermes_home, display_hermes_home +from hermes_constants import ( + get_hermes_home, + display_hermes_home, + get_skills_dir as _active_skills_dir, +) import os import re from enum import Enum @@ -92,6 +96,15 @@ # skills all coexist here without polluting the git repo. HERMES_HOME = get_hermes_home() SKILLS_DIR = HERMES_HOME / "skills" +_IMPORT_SKILLS_DIR = SKILLS_DIR + + +def get_skills_dir() -> Path: + """Return the active profile's skills dir, preserving SKILLS_DIR monkeypatches.""" + configured = Path(SKILLS_DIR) + if configured != _IMPORT_SKILLS_DIR: + return configured + return _active_skills_dir() # Anthropic-recommended limits for progressive disclosure efficiency MAX_NAME_LENGTH = 64 @@ -499,9 +512,9 @@ def _get_category_from_path(skill_path: Path) -> Optional[str]: For paths like: ~/.hermes/skills/mlops/axolotl/SKILL.md -> "mlops" Also works for external skill dirs configured via skills.external_dirs. """ - # Try the module-level SKILLS_DIR first (respects monkeypatching in tests), + # Try the local skills dir first (respects monkeypatching in tests), # then fall back to external dirs from config. - dirs_to_check = [SKILLS_DIR] + dirs_to_check = [get_skills_dir()] try: from agent.skill_utils import get_external_skills_dirs dirs_to_check.extend(get_external_skills_dirs()) @@ -620,8 +633,9 @@ def _find_all_skills(*, skip_disabled: bool = False) -> List[Dict[str, Any]]: # Scan local dir first, then external dirs (local takes precedence) dirs_to_scan = [] - if SKILLS_DIR.exists(): - dirs_to_scan.append(SKILLS_DIR) + skills_dir = get_skills_dir() + if skills_dir.exists(): + dirs_to_scan.append(skills_dir) dirs_to_scan.extend(get_external_skills_dirs()) for scan_dir in dirs_to_scan: @@ -699,8 +713,9 @@ def skills_list(category: str = None, task_id: str = None) -> str: JSON string with minimal skill info: name, description, category """ try: - if not SKILLS_DIR.exists(): - SKILLS_DIR.mkdir(parents=True, exist_ok=True) + skills_dir = get_skills_dir() + if not skills_dir.exists(): + skills_dir.mkdir(parents=True, exist_ok=True) return json.dumps( { "success": True, @@ -982,8 +997,9 @@ def skill_view( # Build list of all skill directories to search all_dirs = [] - if SKILLS_DIR.exists(): - all_dirs.append(SKILLS_DIR) + skills_dir = get_skills_dir() + if skills_dir.exists(): + all_dirs.append(skills_dir) all_dirs.extend(get_external_skills_dirs()) if not all_dirs: @@ -1133,7 +1149,7 @@ def _record(sd: Optional[Path], smd: Path) -> None: # Security: warn if skill is loaded from outside trusted directories # (local skills dir + configured external_dirs are all trusted) _outside_skills_dir = True - _trusted_dirs = [SKILLS_DIR.resolve()] + _trusted_dirs = [skills_dir.resolve()] try: _trusted_dirs.extend(d.resolve() for d in all_dirs[1:]) except Exception: @@ -1362,7 +1378,7 @@ def _record(sd: Optional[Path], smd: Path) -> None: linked_files["scripts"] = script_files try: - rel_path = str(skill_md.relative_to(SKILLS_DIR)) + rel_path = str(skill_md.relative_to(skills_dir)) except ValueError: # External skill — use path relative to the skill's own parent dir rel_path = str(skill_md.relative_to(skill_md.parent.parent)) if skill_md.parent.parent else skill_md.name diff --git a/tools/slash_confirm.py b/tools/slash_confirm.py index 21db18fe31979..4579bdaba189a 100644 --- a/tools/slash_confirm.py +++ b/tools/slash_confirm.py @@ -129,8 +129,44 @@ async def resolve( if not handler: return None + + # Resolve profile-scoping if session_key contains a named profile. + profile_home = None + if session_key.startswith("agent:"): + parts = session_key.split(":") + if len(parts) > 1 and parts[1] != "main": + profile_name = parts[1] + try: + from hermes_cli.profiles import get_profile_dir + profile_home = get_profile_dir(profile_name) + except Exception: + profile_home = None + try: - result = await handler(choice) + if profile_home: + import contextlib + from pathlib import Path + from hermes_constants import set_hermes_home_override, reset_hermes_home_override + from agent.secret_scope import ( + build_profile_secret_scope, + set_secret_scope, + reset_secret_scope, + ) + + @contextlib.contextmanager + def _profile_runtime_scope(ph: Path): + home_token = set_hermes_home_override(str(ph)) + secret_token = set_secret_scope(build_profile_secret_scope(ph)) + try: + yield + finally: + reset_secret_scope(secret_token) + reset_hermes_home_override(home_token) + + with _profile_runtime_scope(Path(profile_home)): + result = await handler(choice) + else: + result = await handler(choice) except Exception as exc: logger.error( "Slash-confirm handler for /%s raised: %s", diff --git a/tools/terminal_tool.py b/tools/terminal_tool.py index 6a5a6af1fdfe0..5c881282e228f 100644 --- a/tools/terminal_tool.py +++ b/tools/terminal_tool.py @@ -2317,6 +2317,10 @@ def terminal_tool( default_cwd=cwd, ) try: + from gateway.session_context import get_session_env + active_profile = get_session_env("HERMES_SESSION_AGENT_PROFILE", "") + active_home = get_session_env("HERMES_SESSION_AGENT_HERMES_HOME", "") + if env_type == "local": proc_session = process_registry.spawn_local( command=command, @@ -2325,6 +2329,8 @@ def terminal_tool( session_key=session_key, env_vars=env.env if hasattr(env, 'env') else None, use_pty=effective_pty, + agent_profile=active_profile, + agent_hermes_home=active_home, ) else: proc_session = process_registry.spawn_via_env( @@ -2333,6 +2339,8 @@ def terminal_tool( cwd=effective_cwd, task_id=effective_task_id, session_key=session_key, + agent_profile=active_profile, + agent_hermes_home=active_home, ) result_data = { From 8644434e5b1ca94c3a4830c4e8f0a0a41c6ece76 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=D0=9A=D0=B8=D1=80=D0=B8=D0=BB=D0=BB=20=D0=92=D0=B5=D1=87?= =?UTF-8?q?=D0=BA=D0=B0=D1=81=D0=BE=D0=B2?= Date: Sun, 28 Jun 2026 01:54:43 +0200 Subject: [PATCH 2/2] fix(gateway): bound per-profile SessionStore cache to stop FD exhaustion MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Per-topic profile isolation gives each routed profile its own SessionStore backed by a SQLite connection (db + WAL + SHM = 3 fds) to that profile's state.db. Two defects made a long-lived multi-profile gateway leak file descriptors until it hit EMFILE ([Errno 24] Too many open files): 1. `_profile_session_stores` and `_profile_session_dbs` were unbounded dicts that only ever grew and were closed solely at shutdown. Every distinct profile ever routed to kept its connection open for the whole process lifetime. 2. `_session_db` opened its OWN SessionDB to `/state.db`, a SECOND connection to the exact file the profile's SessionStore already had open — doubling the fds held per profile. On macOS the launchd gateway inherits RLIMIT_NOFILE=256, so a handful of profiles plus normal sockets crossed the limit. Once over, every new open() failed — including the kanban dispatcher's `kanban.db` open, which then spun retrying on its tick loop and pinned a core. Fix, mirroring the existing `_agent_cache` LRU pattern: - `_profile_session_stores` becomes an OrderedDict with `move_to_end()` on hit and an LRU cap (`_PROFILE_STORE_CACHE_MAX_SIZE`, default 16); the evicted (least-recently-used, idle) store's connection is closed off-thread via `_close_evicted_profile_store` (SessionDB.close runs a WAL checkpoint and releases the db/WAL/SHM fds). The victim is never the profile served this turn, so live sessions aren't torn down mid-write. - `_session_db` now reuses the profile SessionStore's connection instead of opening a second handle; an explicit per-home override (the setter, used by tests and back-compat call sites) still wins, preserving the attribute contract. The redundant default-home SessionDB built in __init__ is dropped (the store already owns it); the NFS/locking init warning is preserved. Adds tests for LRU eviction + close, connection reuse, and override precedence. --- gateway/run.py | 142 +++++++++++++----- .../test_profile_isolation_rework_suite.py | 67 +++++++++ 2 files changed, 168 insertions(+), 41 deletions(-) diff --git a/gateway/run.py b/gateway/run.py index c232fe80d45af..14913c25c3167 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -65,6 +65,16 @@ # from _enforce_agent_cache_cap() and _session_expiry_watcher() below. _AGENT_CACHE_MAX_SIZE = 128 _AGENT_CACHE_IDLE_TTL_SECS = 3600.0 # evict agents idle for >1h + +# --- Per-profile session-store cache tuning ------------------------------- +# Each routed topic/profile gets its own SessionStore backed by a SQLite +# connection (3 fds: db + WAL + SHM) to that profile's state.db. Without a +# bound the gateway keeps one live connection per distinct profile for the +# whole process lifetime; on macOS's default launchd RLIMIT_NOFILE=256 a +# multi-profile deployment exhausts the fd budget, after which *every* +# state.db/kanban.db open fails with EMFILE and the kanban dispatcher spins +# retrying. Bound the cache LRU-style, mirroring the agent cache above. +_PROFILE_STORE_CACHE_MAX_SIZE = 16 _PLATFORM_CONNECT_TIMEOUT_SECS_DEFAULT = 30.0 _ADAPTER_DISCONNECT_TIMEOUT_SECS_DEFAULT = 5.0 _TELEGRAM_COMMAND_MENTION_RE = re.compile(r"(? _PROFILE_STORE_CACHE_MAX_SIZE: + evicted.append(cache.popitem(last=False)) + # Close evicted stores' DB connections OUTSIDE the cache lock so a live + # profile lookup never blocks on SQLite's WAL-checkpoint-on-close. + for old_home, old_store in evicted: + self._close_evicted_profile_store(old_home, old_store) + return store + + def _close_evicted_profile_store(self, home: "Path", store: Any) -> None: + """Close an LRU-evicted profile SessionStore's DB connection off-thread. + + Mirrors ``_release_evicted_agent_soft``: the close runs on a daemon + thread so the cache lock is never held across SQLite's + WAL-checkpoint-on-close. ``SessionDB.close()`` takes its own internal + lock, so any in-flight operation on the connection finishes before the + handle (db + WAL + SHM fds) is released. The victim is always the + least-recently-used home, never the profile being served this turn, so + a live session is not torn down mid-write. + """ + db = getattr(store, "_db", None) + if db is None or not hasattr(db, "close"): + return + + def _do_close(): + try: + db.close() + except Exception as exc: # best effort — never fatal + logger.debug("evicted profile SessionDB close error (%s): %s", home, exc) + + try: + threading.Thread( + target=_do_close, daemon=True, name="profile-store-evict" + ).start() + except Exception: + # Interpreter shutdown / thread exhaustion: close inline. + _do_close() @property def session_store(self): @@ -2888,21 +2943,21 @@ def _all_profile_homes(self) -> List["Path"]: def _session_db(self): from hermes_constants import get_hermes_home current_home = get_hermes_home().resolve() - if not hasattr(self, "_profile_cache_lock"): - import threading - self._profile_cache_lock = threading.Lock() - if not hasattr(self, "_profile_session_dbs"): - self._profile_session_dbs = {} - with self._profile_cache_lock: - if current_home not in self._profile_session_dbs: - try: - from hermes_state import SessionDB - db = SessionDB(db_path=current_home / "state.db") - self._profile_session_dbs[current_home] = db - except Exception as e: - logger.warning("SQLite session store not available for %s: %s", current_home, e) - self._profile_session_dbs[current_home] = None - return self._profile_session_dbs[current_home] + # An explicit per-home override (written via the setter — used by tests + # and back-compat call sites) always wins. + overrides = getattr(self, "_profile_session_dbs", None) + if overrides and current_home in overrides: + return overrides[current_home] + # Otherwise reuse the profile's SessionStore connection rather than + # opening a SECOND SQLite handle to the same state.db. The store and + # this property used to each open their own connection per home, which + # doubled the file descriptors held per profile (the FD-leak doubling). + try: + store = self._get_or_create_store_for_home(current_home) + except Exception as e: + logger.warning("SQLite session store not available for %s: %s", current_home, e) + return None + return getattr(store, "_db", None) @_session_db.setter def _session_db(self, value): @@ -3127,19 +3182,24 @@ def __init__(self, config: Optional[GatewayConfig] = None): except Exception: logger.debug("approvals.mode startup check skipped", exc_info=True) - # Initialize session database for session_search tool support - self._session_db = None + # Session database for session_search/resume/history. The default + # profile's connection is owned by self.session_store (created above); + # _session_db now reuses it (see the property) instead of opening a + # second SQLite handle to the same state.db. Surface an init failure + # here the same way as before so NFS/locking problems still land in + # errors.log (WARNING, not DEBUG — matches cli.py's init path; without + # it, NFS-mounted HERMES_HOME silently lost /resume, /title, /history, + # /branch and session search). try: - from hermes_state import SessionDB - self._session_db = SessionDB() - except Exception as e: - # WARNING (not DEBUG) so the failure appears in errors.log — matches - # cli.py's handling of the same init path. Users hitting NFS-mounted - # HERMES_HOME silently lost /resume, /title, /history, /branch, and - # session search without this. The underlying cause (usually - # "locking protocol" from NFS) is now also captured by - # hermes_state.get_last_init_error() for slash-command error strings. - logger.warning("SQLite session store not available: %s", e) + _default_store = getattr(self, "session_store", None) + if _default_store is not None and getattr(_default_store, "_db", None) is None: + from hermes_state import get_last_init_error + logger.warning( + "SQLite session store not available: %s", + get_last_init_error() or "see earlier log", + ) + except Exception: + logger.debug("session store availability probe skipped", exc_info=True) # Opportunistic state.db maintenance: prune ended sessions older # than sessions.retention_days + optional VACUUM. Tracks last-run diff --git a/tests/gateway/test_profile_isolation_rework_suite.py b/tests/gateway/test_profile_isolation_rework_suite.py index 7ab2598092e40..485757b233d1d 100644 --- a/tests/gateway/test_profile_isolation_rework_suite.py +++ b/tests/gateway/test_profile_isolation_rework_suite.py @@ -438,3 +438,70 @@ def test_g2_hard_guard_blocks_outside_profile_home(test_env): assert not forbidden_path.exists() finally: clear_session_vars(tokens) + + +def test_profile_store_cache_lru_eviction(test_env, monkeypatch): + """The per-profile SessionStore cache is LRU-bounded and closes evicted + DB connections, so a multi-profile gateway can't leak file descriptors + until it hits EMFILE (the FD-exhaustion / 'Too many open files' bug).""" + import gateway.run as run_mod + monkeypatch.setattr(run_mod, "_PROFILE_STORE_CACHE_MAX_SIZE", 3) + runner = GatewayRunner(config=GatewayConfig(platforms={})) + + # Close evicted stores synchronously (bypass the daemon thread) so the + # assertions are deterministic, and record which homes were evicted. + evicted_homes = [] + + def _sync_close(home, store): + evicted_homes.append(home) + db = getattr(store, "_db", None) + if db is not None: + db.close() + + runner._close_evicted_profile_store = _sync_close + + homes = [] + base = test_env / "lru_homes" + base.mkdir(parents=True, exist_ok=True) + for i in range(5): + h = (base / f"h{i}").resolve() + h.mkdir(parents=True, exist_ok=True) + homes.append(h) + runner._get_or_create_store_for_home(h) + + # Never exceeds the cap... + assert len(runner._profile_session_stores) <= 3 + # ...the two oldest of our homes were evicted (LRU order)... + assert homes[0] in evicted_homes + assert homes[1] in evicted_homes + # ...and the three most-recent homes are still resident. + for h in homes[2:]: + assert h in runner._profile_session_stores + + # Re-touching a resident home makes it most-recently-used, so the NEXT + # insertion evicts a different (now-oldest) home — proving LRU recency, + # not plain FIFO. + runner._get_or_create_store_for_home(homes[2]) # h2 -> MRU + h_new = (base / "h_new").resolve() + h_new.mkdir(parents=True, exist_ok=True) + runner._get_or_create_store_for_home(h_new) + assert homes[2] in runner._profile_session_stores # survived (was touched) + assert homes[3] in evicted_homes # oldest -> evicted + + +def test_session_db_reuses_store_connection(test_env): + """_session_db reuses the profile SessionStore's connection rather than + opening a SECOND SQLite handle to the same state.db (the FD doubling).""" + runner = GatewayRunner(config=GatewayConfig(platforms={})) + store = runner.session_store + assert store._db is not None + assert runner._session_db is store._db + + +def test_session_db_setter_override_is_honored(test_env): + """An explicit _session_db assignment still wins for the current home — + the broad attribute contract (tests / back-compat call sites) is kept.""" + runner = GatewayRunner(config=GatewayConfig(platforms={})) + sentinel = object() + runner._session_db = sentinel + assert runner._session_db is sentinel