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:
Teknium
2026-09-02 22:14:55 -07:00
parent 4a37db2269
commit b001c4ab07
5 changed files with 69 additions and 105 deletions

View File

@@ -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]:

View File

@@ -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

View File

@@ -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:

View File

@@ -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

View File

@@ -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"