refactor(approval): share activity_heartbeat across gateway/transport waits, fold prompt inner wrapper, compact readers and block-result literals (1310->1274)
This commit is contained in:
@@ -50,9 +50,7 @@ def reset_hermes_interactive_context(token: contextvars.Token) -> None:
|
||||
def _is_interactive_cli() -> bool:
|
||||
"""True for an interactive CLI/ACP session (contextvar first, env fallback)."""
|
||||
ctx_val = _hermes_interactive_ctx.get()
|
||||
if ctx_val is not None:
|
||||
return is_truthy_value(ctx_val)
|
||||
return env_var_enabled("HERMES_INTERACTIVE")
|
||||
return is_truthy_value(ctx_val) if ctx_val is not None else env_var_enabled("HERMES_INTERACTIVE")
|
||||
|
||||
|
||||
def _fire_approval_hook(hook_name: str, **kwargs) -> None:
|
||||
@@ -90,9 +88,8 @@ def reset_current_session_key(token: contextvars.Token[str]) -> None:
|
||||
_Tokens = tuple[contextvars.Token[str], contextvars.Token[str], contextvars.Token[str]]
|
||||
|
||||
|
||||
def set_current_observability_context(
|
||||
*, turn_id: str = "", tool_call_id: str = "", session_id: str = "",
|
||||
) -> _Tokens:
|
||||
def set_current_observability_context(*, turn_id: str = "", tool_call_id: str = "",
|
||||
session_id: str = "") -> _Tokens:
|
||||
"""Bind active tool correlation IDs to approval hooks."""
|
||||
return (_approval_turn_id.set(turn_id or ""), _approval_tool_call_id.set(tool_call_id or ""),
|
||||
_approval_session_id.set(session_id or ""))
|
||||
@@ -108,8 +105,7 @@ def reset_current_observability_context(tokens: _Tokens) -> None:
|
||||
|
||||
def get_current_session_key(default: str = "default") -> str:
|
||||
"""Return the active session key: approval contextvar → session_context → os.environ."""
|
||||
session_key = _approval_session_key.get()
|
||||
if session_key:
|
||||
if session_key := _approval_session_key.get():
|
||||
return session_key
|
||||
from gateway.session_context import get_session_env
|
||||
return get_session_env("HERMES_SESSION_KEY", default)
|
||||
@@ -229,8 +225,7 @@ def _get_approval_mode() -> str:
|
||||
from tools import approval as _a
|
||||
try:
|
||||
from gateway.hosted_room_execution_policy import current_room_execution_policy
|
||||
room_policy = current_room_execution_policy()
|
||||
if room_policy is not None:
|
||||
if (room_policy := current_room_execution_policy()) is not None:
|
||||
return room_policy.approval_mode
|
||||
except Exception:
|
||||
pass
|
||||
@@ -253,13 +248,11 @@ def _get_approval_timeout() -> int:
|
||||
from agent.deadline import MAX_SAFE_TIMEOUT_S
|
||||
safe_cap = int(MAX_SAFE_TIMEOUT_S)
|
||||
except Exception:
|
||||
# Fail CLOSED: the raw value would re-open the overflow this prevents.
|
||||
safe_cap = 365 * 24 * 3600
|
||||
safe_cap = 365 * 24 * 3600 # fail CLOSED: the raw value would re-open the overflow
|
||||
if raw > safe_cap:
|
||||
logger.warning("approvals.timeout=%s exceeds the platform-safe maximum; "
|
||||
"clamping to %ss", raw, safe_cap)
|
||||
return safe_cap
|
||||
return raw
|
||||
return min(raw, safe_cap)
|
||||
|
||||
|
||||
def _binary_approval_mode(key: str) -> str:
|
||||
@@ -294,13 +287,11 @@ def _tirith_fail_open() -> bool:
|
||||
False means the operator opted into fail-closed: an un-importable scanner
|
||||
must not silently grant access."""
|
||||
try:
|
||||
from hermes_cli.config import load_config_readonly as _load_cfg
|
||||
_sec = (_load_cfg() or {}).get("security", {}) or {}
|
||||
if _sec.get("tirith_enabled", True):
|
||||
return bool(_sec.get("tirith_fail_open", True))
|
||||
from hermes_cli.config import load_config_readonly
|
||||
_sec = (load_config_readonly() or {}).get("security", {}) or {}
|
||||
return bool(_sec.get("tirith_fail_open", True)) if _sec.get("tirith_enabled", True) else True
|
||||
except Exception:
|
||||
pass
|
||||
return True
|
||||
return True
|
||||
|
||||
|
||||
def _get_approval_transport_config() -> tuple[str, str | None]:
|
||||
|
||||
@@ -7,6 +7,7 @@ allowlist runs after. Session state stays in ``tools.approval`` and is read
|
||||
through it at call time so tests that rebind it keep working.
|
||||
"""
|
||||
|
||||
import contextlib
|
||||
import fnmatch
|
||||
import logging
|
||||
import re
|
||||
@@ -42,17 +43,12 @@ def _match_user_deny_rule(command: str) -> str | None:
|
||||
|
||||
def _user_deny_block_result(pattern: str) -> dict:
|
||||
"""Build the standard block result for an ``approvals.deny`` match."""
|
||||
return {
|
||||
"approved": False,
|
||||
"user_deny": True,
|
||||
"message": (
|
||||
f"BLOCKED: this command matches the user-defined deny rule "
|
||||
f"'{pattern}' (approvals.deny in config.yaml). It cannot be "
|
||||
"executed via the agent — not even with --yolo, /yolo, or "
|
||||
"approvals.mode=off. Do NOT retry or rephrase this command; "
|
||||
"the user has explicitly forbidden it."
|
||||
),
|
||||
}
|
||||
return {"approved": False, "user_deny": True, "message": (
|
||||
f"BLOCKED: this command matches the user-defined deny rule "
|
||||
f"'{pattern}' (approvals.deny in config.yaml). It cannot be "
|
||||
"executed via the agent — not even with --yolo, /yolo, or "
|
||||
"approvals.mode=off. Do NOT retry or rephrase this command; "
|
||||
"the user has explicitly forbidden it.")}
|
||||
|
||||
|
||||
def _save_blocked_payload(command: str) -> str | None:
|
||||
@@ -70,11 +66,9 @@ def _save_blocked_payload(command: str) -> str | None:
|
||||
# Opportunistic cleanup: blocked payloads older than 7 days.
|
||||
cutoff = time.time() - 7 * 86400
|
||||
for old in script_dir.glob("blocked-*.sh"):
|
||||
try:
|
||||
with contextlib.suppress(OSError):
|
||||
if old.stat().st_mtime < cutoff:
|
||||
old.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
path = script_dir / f"blocked-{int(time.time())}-{uuid.uuid4().hex[:8]}.sh"
|
||||
path.write_text(
|
||||
"#!/bin/bash\n"
|
||||
@@ -130,16 +124,12 @@ def _hardline_block_result(description: str, command: str = "") -> dict:
|
||||
|
||||
def _sudo_stdin_block_result(description: str) -> dict:
|
||||
"""Build the standard block result for sudo stdin guard."""
|
||||
return {
|
||||
"approved": False,
|
||||
"message": (
|
||||
f"BLOCKED: {description}. "
|
||||
"Do not pipe passwords to 'sudo -S' — this is a brute-force "
|
||||
"attack vector. Set SUDO_PASSWORD in your .env file if the "
|
||||
"agent needs passwordless sudo, or run the sudo command "
|
||||
"manually in your own terminal."
|
||||
),
|
||||
}
|
||||
return {"approved": False, "message": (
|
||||
f"BLOCKED: {description}. "
|
||||
"Do not pipe passwords to 'sudo -S' — this is a brute-force "
|
||||
"attack vector. Set SUDO_PASSWORD in your .env file if the "
|
||||
"agent needs passwordless sudo, or run the sudo command "
|
||||
"manually in your own terminal.")}
|
||||
|
||||
|
||||
# Shell control characters that make a command compound when they appear OUTSIDE
|
||||
@@ -209,10 +199,7 @@ def _command_matches_permanent_allowlist(command: str) -> bool:
|
||||
patterns = tuple(_a._permanent_approved)
|
||||
for pattern in patterns:
|
||||
pattern = pattern.strip() if isinstance(pattern, str) else ""
|
||||
if not pattern:
|
||||
continue
|
||||
if command == pattern:
|
||||
return True
|
||||
if any(ch in pattern for ch in "*?[") and fnmatch.fnmatchcase(command, pattern):
|
||||
if pattern and (command == pattern or (any(ch in pattern for ch in "*?[")
|
||||
and fnmatch.fnmatchcase(command, pattern))):
|
||||
return True
|
||||
return False
|
||||
|
||||
@@ -16,6 +16,7 @@ import time
|
||||
import uuid
|
||||
|
||||
from tools.interrupt import is_interrupted
|
||||
from tools.approval_human_wait import activity_heartbeat
|
||||
|
||||
logger = logging.getLogger("tools.approval")
|
||||
|
||||
@@ -49,14 +50,8 @@ def _poll_event(event: threading.Event, session_key: str, *, interrupt_log: str)
|
||||
dedicated interrupt-cause channel, not string matching."""
|
||||
from tools.approval import _get_approval_timeout, human_wait_window
|
||||
|
||||
timeout = _get_approval_timeout()
|
||||
try:
|
||||
from tools.environments.base import touch_activity_if_due
|
||||
except Exception: # pragma: no cover
|
||||
touch_activity_if_due = None
|
||||
now = time.monotonic()
|
||||
deadline = now + max(timeout, 0)
|
||||
activity_state = {"last_touch": now, "start": now}
|
||||
deadline = time.monotonic() + max(_get_approval_timeout(), 0)
|
||||
heartbeat = activity_heartbeat("waiting for user approval")
|
||||
with human_wait_window(session_key):
|
||||
while True:
|
||||
if is_interrupted():
|
||||
@@ -67,8 +62,7 @@ def _poll_event(event: threading.Event, session_key: str, *, interrupt_log: str)
|
||||
return "timeout"
|
||||
if event.wait(timeout=min(1.0, remaining)):
|
||||
return "set"
|
||||
if touch_activity_if_due is not None:
|
||||
touch_activity_if_due(activity_state, "waiting for user approval")
|
||||
heartbeat()
|
||||
|
||||
|
||||
def _finish(payload: dict, resolved: bool, choice: str | None, reason, **extra) -> dict:
|
||||
|
||||
@@ -82,6 +82,19 @@ def _resolve_key(session_key: str | None) -> str:
|
||||
return get_current_session_key()
|
||||
|
||||
|
||||
def activity_heartbeat(label: str):
|
||||
"""Callable that pings the agent's inactivity tracker (at most every ~10s)
|
||||
while a human wait is parked, so the gateway watchdog does not kill the agent
|
||||
while the user is still answering. No-op in minimal tool-only environments."""
|
||||
try:
|
||||
from tools.environments.base import touch_activity_if_due
|
||||
except Exception: # pragma: no cover - minimal tool-only environments
|
||||
return lambda: None
|
||||
now = time.monotonic()
|
||||
state = {"last_touch": now, "start": now}
|
||||
return lambda: touch_activity_if_due(state, label)
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def human_wait_window(session_key: str | None = None):
|
||||
"""Mark the enclosed block as time spent blocked on a human prompt. Wrap ONLY
|
||||
|
||||
@@ -10,8 +10,7 @@ import logging
|
||||
import os
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from tools.approval_human_wait import human_wait_window
|
||||
from tools.approval_human_wait import activity_heartbeat, human_wait_window
|
||||
from tools.interrupt import is_interrupted
|
||||
|
||||
logger = logging.getLogger("tools.approval")
|
||||
@@ -43,10 +42,8 @@ def prompt_dangerous_approval(command: str, description: str, timeout_seconds: i
|
||||
# bounded by the approval deadline): record it as human-wait time so the
|
||||
# concurrent batch deadline excludes it.
|
||||
with human_wait_window():
|
||||
return _prompt_dangerous_approval_inner(
|
||||
command, description, timeout_seconds, allow_permanent,
|
||||
approval_callback, allow_session=allow_session, smart_denied=smart_denied,
|
||||
)
|
||||
return _ask_human(command, description, timeout_seconds, allow_permanent,
|
||||
approval_callback, allow_session, smart_denied)
|
||||
|
||||
|
||||
_CLI_CHOICE_ALIASES = {
|
||||
@@ -79,9 +76,8 @@ def _read_choice(prompt: str, timeout_seconds: int) -> str | None:
|
||||
return None if thread.is_alive() else result["choice"]
|
||||
|
||||
|
||||
def _prompt_dangerous_approval_inner(command: str, description: str, timeout_seconds: int,
|
||||
allow_permanent: bool = True, approval_callback=None,
|
||||
*, allow_session: bool = True, smart_denied: bool = False) -> str:
|
||||
def _ask_human(command: str, description: str, timeout_seconds: int, allow_permanent: bool,
|
||||
approval_callback, allow_session: bool, smart_denied: bool) -> str:
|
||||
# Redact before any user-visible rendering; the original `command` still
|
||||
# executes after approval. Same redactor as memory/log sanitization so
|
||||
# tokens mask consistently across surfaces.
|
||||
@@ -93,11 +89,10 @@ def _prompt_dangerous_approval_inner(command: str, description: str, timeout_sec
|
||||
|
||||
if approval_callback is not None:
|
||||
try:
|
||||
callback_kwargs = {"allow_permanent": allow_permanent}
|
||||
if not allow_session:
|
||||
callback_kwargs["allow_session"] = False
|
||||
if smart_denied:
|
||||
callback_kwargs["smart_denied"] = True
|
||||
# Non-default scopes only: legacy callbacks lack the newer keywords.
|
||||
callback_kwargs = {"allow_permanent": allow_permanent,
|
||||
**({"allow_session": False} if not allow_session else {}),
|
||||
**({"smart_denied": True} if smart_denied else {})}
|
||||
return approval_callback(display_command, display_description, **callback_kwargs)
|
||||
except Exception as e:
|
||||
logger.error("Approval callback failed: %s", e, exc_info=True)
|
||||
@@ -111,11 +106,9 @@ def _prompt_dangerous_approval_inner(command: str, description: str, timeout_sec
|
||||
try:
|
||||
from prompt_toolkit.application.current import get_app_or_none
|
||||
if get_app_or_none() is not None:
|
||||
logger.warning(
|
||||
"Dangerous-command approval requested on a thread with no "
|
||||
"approval callback while prompt_toolkit is active; denying "
|
||||
"to avoid stdin deadlock. command=%r description=%r", command, description,
|
||||
)
|
||||
logger.warning("Dangerous-command approval requested on a thread with no "
|
||||
"approval callback while prompt_toolkit is active; denying "
|
||||
"to avoid stdin deadlock. command=%r description=%r", command, description)
|
||||
return "deny"
|
||||
except Exception:
|
||||
pass # prompt_toolkit absent or detection failed: legacy input() path is safe
|
||||
@@ -124,11 +117,8 @@ def _prompt_dangerous_approval_inner(command: str, description: str, timeout_sec
|
||||
try:
|
||||
from agent.i18n import t
|
||||
# (prompt key, menu key) by menu shape: once/deny, full, or no [a]lways.
|
||||
prompt_key, menu_key = (
|
||||
("approval.prompt_smart_deny", "approval.choose_smart_deny") if once_only
|
||||
else ("approval.prompt_long", "approval.choose_long") if allow_permanent
|
||||
else ("approval.prompt_short", "approval.choose_short")
|
||||
)
|
||||
shape = "smart_deny" if once_only else "long" if allow_permanent else "short"
|
||||
prompt_key, menu_key = f"approval.prompt_{shape}", f"approval.choose_{shape}"
|
||||
print(f"\n {t('approval.dangerous_header', description=display_description)}"
|
||||
f"\n {display_command}\n\n{t(menu_key)}\n")
|
||||
sys.stdout.flush()
|
||||
@@ -137,10 +127,9 @@ def _prompt_dangerous_approval_inner(command: str, description: str, timeout_sec
|
||||
print("\n" + t("approval.timeout"))
|
||||
return "timeout" # distinct from deny: the user never answered
|
||||
if once_only:
|
||||
decision = {
|
||||
**dict.fromkeys(t("approval.smart_deny_once_inputs").split(","), "once"),
|
||||
**dict.fromkeys(t("approval.smart_deny_deny_inputs").split(","), "deny"),
|
||||
}.get(choice, "deny")
|
||||
decision = {**dict.fromkeys(t("approval.smart_deny_once_inputs").split(","), "once"),
|
||||
**dict.fromkeys(t("approval.smart_deny_deny_inputs").split(","), "deny"),
|
||||
}.get(choice, "deny")
|
||||
else:
|
||||
decision = _CLI_CHOICE_ALIASES.get(choice, "deny")
|
||||
if decision == "always" and not allow_permanent:
|
||||
@@ -217,21 +206,11 @@ def _present_with_selected_transport(*, command: str, description: str, pattern_
|
||||
request_id=request.request_id, request_digest=request.digest,
|
||||
)
|
||||
_a._fire_approval_hook("pre_approval_request", **hook_kwargs)
|
||||
try:
|
||||
from tools.environments.base import touch_activity_if_due
|
||||
except Exception: # pragma: no cover - minimal tool-only environments
|
||||
touch_activity_if_due = None
|
||||
now = time.monotonic()
|
||||
activity_state = {"last_touch": now, "start": now}
|
||||
|
||||
def _poll() -> None:
|
||||
if touch_activity_if_due is not None:
|
||||
touch_activity_if_due(activity_state, "waiting for plugin approval transport")
|
||||
|
||||
with human_wait_window(session_key):
|
||||
result = invoke_approval_transport(
|
||||
registered.present, request, timeout_seconds=timeout_seconds,
|
||||
on_poll=_poll, is_interrupted=is_interrupted,
|
||||
on_poll=activity_heartbeat("waiting for plugin approval transport"),
|
||||
is_interrupted=is_interrupted,
|
||||
)
|
||||
hook_choice = result.choice if result.failure is None else f"transport_{result.failure}"
|
||||
_a._fire_approval_hook("post_approval_response", **hook_kwargs, choice=hook_choice)
|
||||
@@ -291,10 +270,11 @@ def request_elicitation_consent(message: str, description: str, *,
|
||||
logger.warning("Elicitation requested in gateway session %s but no "
|
||||
"notify_cb is registered — failing closed", session_key)
|
||||
return "decline"
|
||||
approval_data = {"command": message, "description": description,
|
||||
"pattern_key": "mcp_elicitation", "pattern_keys": ["mcp_elicitation"]}
|
||||
try:
|
||||
decision = _a._await_gateway_decision(session_key, notify_cb, approval_data, surface=surface)
|
||||
decision = _a._await_gateway_decision(
|
||||
session_key, notify_cb, {"command": message, "description": description,
|
||||
"pattern_key": "mcp_elicitation",
|
||||
"pattern_keys": ["mcp_elicitation"]}, surface=surface)
|
||||
except Exception as exc:
|
||||
logger.error("Elicitation gateway dispatch failed: %s", exc, exc_info=True)
|
||||
return "decline"
|
||||
@@ -306,9 +286,8 @@ def request_elicitation_consent(message: str, description: str, *,
|
||||
|
||||
# allow_permanent=False: elicitation is a per-call confirmation — no pattern to remember.
|
||||
try:
|
||||
choice = _a.prompt_dangerous_approval(
|
||||
message, description, timeout_seconds=timeout_seconds, allow_permanent=False,
|
||||
)
|
||||
choice = _a.prompt_dangerous_approval(message, description, timeout_seconds=timeout_seconds,
|
||||
allow_permanent=False)
|
||||
except Exception as exc:
|
||||
logger.error("Elicitation CLI prompt failed: %s", exc, exc_info=True)
|
||||
return "decline"
|
||||
|
||||
Reference in New Issue
Block a user