refactor(agent/E_session): final code compaction — inline single-use predicates, tail() keyword helper, dict literals, elif chains

This commit is contained in:
Teknium
2026-09-02 23:08:45 -07:00
parent 63f1950b4f
commit f2fceda404
6 changed files with 152 additions and 311 deletions
+11 -21
View File
@@ -31,9 +31,7 @@ def is_interrupted_tool_result(content: Any) -> bool:
if not isinstance(content, str):
return False
lowered = content.lower()
if "[command interrupted]" in lowered:
return True
return "exit_code" in lowered and ("130" in lowered or "-1" in lowered) and "interrupt" in lowered
return "[command interrupted]" in lowered or ("exit_code" in lowered and ("130" in lowered or "-1" in lowered) and "interrupt" in lowered)
def _call_name(call: Dict[str, Any]) -> str:
@@ -62,8 +60,7 @@ def strip_interrupted_tool_tails(agent_history: List[Dict[str, Any]]) -> List[Di
if not agent_history:
return agent_history
cleaned: List[Dict[str, Any]] = []
i = 0
n = len(agent_history)
i, n = 0, len(agent_history)
while i < n:
msg = agent_history[i]
if msg.get("role") == "assistant" and "tool_calls" in msg:
@@ -71,24 +68,20 @@ def strip_interrupted_tool_tails(agent_history: List[Dict[str, Any]]) -> List[Di
while j < n and agent_history[j].get("role") == "tool":
j += 1
tool_results = agent_history[i + 1:j]
if tool_results and any(is_interrupted_tool_result(m.get("content", "")) for m in tool_results):
if any(is_interrupted_tool_result(m.get("content", "")) for m in tool_results):
calls = msg.get("tool_calls") or []
if _any_side_effecting(calls):
call_names = {_call_id(call): _call_name(call) for call in calls}
cleaned.append(msg)
for tool_result in tool_results:
if not is_interrupted_tool_result(tool_result.get("content", "")):
cleaned.append(tool_result)
continue
recovered = dict(tool_result)
name = call_names.get(str(tool_result.get("tool_call_id") or ""), "")
recovered["effect_disposition"], recovered["content"] = _orphan_recovery(name, _INTERRUPTED_NOTICES)
cleaned.append(recovered)
if is_interrupted_tool_result(tool_result.get("content", "")):
name = call_names.get(str(tool_result.get("tool_call_id") or ""), "")
disposition, content = _orphan_recovery(name, _INTERRUPTED_NOTICES)
tool_result = {**tool_result, "effect_disposition": disposition, "content": content}
cleaned.append(tool_result)
else:
logger.debug(
"Stripping interrupted read-only assistant→tool replay block (indices %d–%d, tool_results=%d)",
i, j - 1, len(tool_results),
)
logger.debug("Stripping interrupted read-only assistant→tool replay block (indices %d–%d, tool_results=%d)",
i, j - 1, len(tool_results))
i = j
continue
if msg.get("role") == "tool" and is_interrupted_tool_result(msg.get("content", "")):
@@ -151,10 +144,7 @@ _EXPIRED_CONFIRMATION_SENTINEL = (
def is_dangerous_confirmation(content: Any) -> bool:
"""True if user-message text contains a known dangerous confirmation phrase."""
if not isinstance(content, str):
return False
text = content.strip().lower()
return any(pattern in text for pattern in _DANGEROUS_CONFIRMATION_PATTERNS)
return isinstance(content, str) and any(pattern in content.strip().lower() for pattern in _DANGEROUS_CONFIRMATION_PATTERNS)
def strip_stale_dangerous_confirmations(
+24 -61
View File
@@ -103,14 +103,6 @@ def _durable_content(content: Any) -> Any:
return "\n".join(txt) if txt else None
def _tool_calls_data(msg: Dict) -> Any:
if hasattr(msg, "tool_calls") and isinstance(msg.tool_calls, list) and msg.tool_calls:
return [{"name": tc.function.name, "arguments": tc.function.arguments} for tc in msg.tool_calls]
if isinstance(msg.get("tool_calls"), list):
return msg["tool_calls"]
return None
def _persist_lock(agent):
"""Close and turn-start persistence can run on separate CLI threads: one critical section."""
return getattr(agent, "_session_persist_lock", None) or nullcontext()
@@ -122,9 +114,8 @@ def _db_flush_seed_ids(agent) -> set:
"""One-shot ``_flushed_db_message_ids`` seed (same session, after a non-empty flush); the scan translates
it to markers and the flush clears it."""
current_session_id = getattr(agent, "session_id", None)
seed_ids = None
if getattr(agent, "_flushed_db_message_session_id", None) == current_session_id and agent._last_flushed_db_idx != 0:
seed_ids = getattr(agent, "_flushed_db_message_ids", None)
same_session = getattr(agent, "_flushed_db_message_session_id", None) == current_session_id
seed_ids = getattr(agent, "_flushed_db_message_ids", None) if same_session and agent._last_flushed_db_idx != 0 else None
agent._flushed_db_message_session_id = current_session_id
return seed_ids if isinstance(seed_ids, set) else set()
@@ -146,7 +137,7 @@ def _db_flush_row(agent, msg: Dict, is_current_turn_user: bool) -> Dict[str, Any
# api_content sidecar: exact bytes sent to the API when they differ from clean content (replay parity).
api_content = msg.get("api_content") if isinstance(msg.get("api_content"), str) else None
timestamp = msg.get("timestamp")
if is_current_turn_user and msg.get("role") == "user":
if is_current_turn_user and role == "user":
override = getattr(agent, "_persist_user_message_override", None)
if _override_replaces_content(msg, content, override):
# Live content is what the wire sent, the override is the clean transcript; keep the sent bytes.
@@ -154,8 +145,7 @@ def _db_flush_row(agent, msg: Dict, is_current_turn_user: bool) -> Dict[str, Any
api_content = content
content = override
ov_timestamp = getattr(agent, "_persist_user_message_timestamp", None)
if ov_timestamp is not None:
timestamp = ov_timestamp
timestamp = timestamp if ov_timestamp is None else ov_timestamp
if api_content == content:
api_content = None
# get_messages_as_conversation replays rows through sanitize_context().strip(); capture the sent bytes
@@ -167,20 +157,14 @@ def _db_flush_row(agent, msg: Dict, is_current_turn_user: bool) -> Dict[str, Any
api_content = content
# Key order is the divert-JSONL wire order (divert_session_transcript_jsonl).
row = {
"role": role,
"content": _durable_content(content),
"tool_name": msg.get("tool_name"),
"tool_calls": _tool_calls_data(msg),
"tool_call_id": msg.get("tool_call_id"),
"finish_reason": msg.get("finish_reason"),
"role": role, "content": _durable_content(content), "tool_name": msg.get("tool_name"),
"tool_calls": msg["tool_calls"] if isinstance(msg.get("tool_calls"), list) else None,
"tool_call_id": msg.get("tool_call_id"), "finish_reason": msg.get("finish_reason"),
**{k: msg.get(k) for k in _ROW_REASONING_KEYS},
"_compressed_summary": bool(msg.get(COMPRESSED_SUMMARY_METADATA_KEY)),
"timestamp": timestamp,
"api_content": api_content,
"display_kind": _summary_display_kind(msg),
"display_metadata": msg.get("display_metadata"),
# Load-bearing for restart drain-window recovery dedup.
"platform_message_id": msg.get("platform_message_id"),
"timestamp": timestamp, "api_content": api_content,
"display_kind": _summary_display_kind(msg), "display_metadata": msg.get("display_metadata"),
"platform_message_id": msg.get("platform_message_id"), # load-bearing for restart drain-window recovery dedup
}
if isinstance(msg.get("_row_id"), int):
row["_row_id"] = msg["_row_id"]
@@ -217,8 +201,7 @@ def _db_flush_write(agent, batch_rows: List[Dict[str, Any]], batch_msgs: List[Di
if not batch_rows:
return
agent._session_db.append_messages_batch(
session_id=agent.session_id,
messages=batch_rows,
session_id=agent.session_id, messages=batch_rows,
compression_lock_holder=getattr(agent, "_active_compression_lock_holder", None),
turn_lease_holder=getattr(agent, "_active_session_turn_lease_holder", None),
turn_lease_ttl_seconds=getattr(agent, "_active_session_turn_lease_ttl_seconds", 300.0) or 300.0,
@@ -244,9 +227,7 @@ def _db_flush_adopt_compression_tip(agent) -> bool:
if tip_row is None or tip_row.get("ended_at") is not None:
return False
logger.warning("Adopted live compression tip %s for closed session %s; retrying flush once", tip, old_id)
agent.session_id = tip
agent._flushed_db_message_ids = set()
agent._last_flushed_db_idx = 0
agent.session_id, agent._flushed_db_message_ids, agent._last_flushed_db_idx = tip, set(), 0
agent._compression_adoption_failed = False
return True
@@ -257,23 +238,17 @@ def _db_flush_failed(agent, e: Exception, batch_rows: List[Dict[str, Any]], adop
# The only place the SQLite error is visible before it becomes a bare False — classify it so the turn-end
# explanation can distinguish lock contention from disk-full/read-only.
from hermes_state import (
CompressionSessionClosedError,
StateDbCorruptError,
StateDbReplacedError,
classify_persistence_error,
CompressionSessionClosedError, StateDbCorruptError, StateDbReplacedError, classify_persistence_error,
divert_session_transcript_jsonl,
)
agent._last_persistence_error_cause = classify_persistence_error(e)
if isinstance(e, (StateDbReplacedError, StateDbCorruptError)):
# A replaced/quarantined handle will not take this batch again — keep it on disk.
try:
divert_session_transcript_jsonl(getattr(agent, "session_id", "") or "", batch_rows)
except Exception:
logger.warning(
"JSONL divert failed after state.db %s for %s",
agent._last_persistence_error_cause, getattr(agent, "session_id", None), exc_info=True,
)
logger.warning("JSONL divert failed after state.db %s for %s",
agent._last_persistence_error_cause, getattr(agent, "session_id", None), exc_info=True)
if isinstance(e, CompressionSessionClosedError):
# Compression race: another path rotated this session mid-write. Retry exactly once on the live tip; a
# second closed-parent write fails closed.
@@ -329,8 +304,7 @@ class SessionPersistenceMixin:
msg["content"] = override
if timestamp is not None:
msg["timestamp"] = timestamp
# Load-bearing for restart drain-window recovery dedup (has_platform_message_id).
if platform_id is not None:
if platform_id is not None: # load-bearing for restart drain-window recovery dedup (has_platform_message_id)
msg["platform_message_id"] = platform_id
def _persist_session(self, messages: List[Dict], conversation_history: List[Dict] = None):
@@ -429,10 +403,8 @@ class SessionPersistenceMixin:
"""Convert REASONING_SCRATCHPAD to think tags and clean up whitespace."""
if not content:
return content
content = convert_scratchpad_to_think(content)
content = re.sub(r'\n+(<think>)', r'\n\1', content)
content = re.sub(r'(</think>)\n+', r'\1\n', content)
return content.strip()
content = re.sub(r'\n+(<think>)', r'\n\1', convert_scratchpad_to_think(content))
return re.sub(r'(</think>)\n+', r'\1\n', content).strip()
@staticmethod
def _redact_message_content(content):
@@ -441,11 +413,8 @@ class SessionPersistenceMixin:
return redact_sensitive_text(content)
if not isinstance(content, list):
return content
return [
{**p, **{k: redact_sensitive_text(p[k]) for k in ("text", "content") if isinstance(p.get(k), str)}}
if isinstance(p, dict) else p
for p in content
]
return [{**p, **{k: redact_sensitive_text(p[k]) for k in ("text", "content") if isinstance(p.get(k), str)}}
if isinstance(p, dict) else p for p in content]
def _save_session_log(self, messages: List[Dict[str, Any]] = None):
"""Optional per-session JSON snapshot (``sessions.write_json_snapshots``, default False) for external
@@ -465,16 +434,10 @@ class SessionPersistenceMixin:
if _existing_log_is_larger(log_file, len(cleaned)):
return
entry = {
"session_id": self.session_id,
"model": self.model,
"base_url": self.base_url,
"platform": self.platform,
"session_start": self.session_start.isoformat(),
"last_updated": datetime.now().isoformat(),
"system_prompt": redact_sensitive_text(self._cached_system_prompt or ""),
"tools": self.tools or [],
"message_count": len(cleaned),
"messages": cleaned,
"session_id": self.session_id, "model": self.model, "base_url": self.base_url, "platform": self.platform,
"session_start": self.session_start.isoformat(), "last_updated": datetime.now().isoformat(),
"system_prompt": redact_sensitive_text(self._cached_system_prompt or ""), "tools": self.tools or [],
"message_count": len(cleaned), "messages": cleaned,
}
atomic_json_write(log_file, entry, indent=2, default=str)
except Exception as e:
+44 -99
View File
@@ -142,7 +142,6 @@ def register_from_config(cfg: Optional[Dict[str, Any]], *, accept_hooks: bool =
if not isinstance(cfg, dict):
return []
from utils import env_var_enabled
if env_var_enabled("HERMES_SAFE_MODE"): # hooks are user customizations too — fire zero user-configured code
logger.info("HERMES_SAFE_MODE=1 — shell-hook registration skipped")
return []
@@ -150,10 +149,8 @@ def register_from_config(cfg: Optional[Dict[str, Any]], *, accept_hooks: bool =
specs = _parse_hooks_block(cfg.get("hooks"))
if not specs:
return []
registered: List[ShellHookSpec] = []
from hermes_cli.plugins import get_plugin_manager # lazy: avoids import cycle
manager = get_plugin_manager()
home_key = _home_key()
manager, home_key, registered = get_plugin_manager(), _home_key(), []
# Idempotence + allowlist read under the lock; TTY prompt outside it; mutation re-takes the lock and re-checks.
for spec in specs:
key = (home_key, spec.event, spec.matcher, spec.command)
@@ -162,11 +159,9 @@ def register_from_config(cfg: Optional[Dict[str, Any]], *, accept_hooks: bool =
continue
already_allowlisted = _is_allowlisted(spec.event, spec.command)
if not already_allowlisted and not _prompt_and_record(spec.event, spec.command, accept_hooks=effective_accept):
logger.warning(
"shell hook for %s (%s) not allowlisted — skipped. Use --accept-hooks / "
"HERMES_ACCEPT_HOOKS=1 / hooks_auto_accept: true, or approve at the TTY prompt next run.",
spec.event, spec.command,
)
logger.warning("shell hook for %s (%s) not allowlisted — skipped. Use --accept-hooks / "
"HERMES_ACCEPT_HOOKS=1 / hooks_auto_accept: true, or approve at the TTY prompt next run.",
spec.event, spec.command)
continue
with _registered_lock:
if key in _registered:
@@ -174,10 +169,8 @@ def register_from_config(cfg: Optional[Dict[str, Any]], *, accept_hooks: bool =
manager._hooks.setdefault(spec.event, []).append(_make_callback(spec))
_registered.add(key)
registered.append(spec)
logger.info(
"shell hook registered: %s -> %s (matcher=%s, timeout=%ds, fail_closed=%s)",
spec.event, spec.command, spec.matcher, spec.timeout, spec.fail_closed,
)
logger.info("shell hook registered: %s -> %s (matcher=%s, timeout=%ds, fail_closed=%s)",
spec.event, spec.command, spec.matcher, spec.timeout, spec.fail_closed)
return registered
@@ -205,20 +198,15 @@ def reset_for_tests() -> None:
def _parse_hooks_block(hooks_cfg: Any) -> List[ShellHookSpec]:
"""Normalise ``hooks:`` into specs; malformed entries warn-and-skip, never raise."""
from hermes_cli.plugins import SHELL_UNSUPPORTED_HOOKS, VALID_HOOKS
if not isinstance(hooks_cfg, dict):
return []
specs: List[ShellHookSpec] = []
for event_name, entries in hooks_cfg.items():
if event_name in ("output_spill", "outbound"): # reserved non-event sub-sections under `hooks:`
continue
if event_name in SHELL_UNSUPPORTED_HOOKS:
# _parse_response has no channel for these directives — refuse loudly.
logger.warning(
"hook event %r is Python-plugin-only: shell hooks cannot return its directive, "
"so this registration is refused rather than silently ignored",
event_name,
)
if event_name in SHELL_UNSUPPORTED_HOOKS: # _parse_response has no channel for these directives — refuse loudly
logger.warning("hook event %r is Python-plugin-only: shell hooks cannot return its directive, "
"so this registration is refused rather than silently ignored", event_name)
continue
if event_name not in VALID_HOOKS:
suggestion = difflib.get_close_matches(str(event_name), VALID_HOOKS, n=1, cutoff=0.6)
@@ -252,11 +240,8 @@ def _parse_single_entry(event: str, index: int, raw: Any) -> Optional[ShellHookS
warn(".matcher must be a string regex; ignoring")
matcher = None
if matcher is not None and event not in _TOOL_EVENTS:
warn(
".matcher=%r will be ignored at runtime — the matcher field is only honored for "
"pre_tool_call / post_tool_call. The hook will fire on every %s event.",
matcher, event,
)
warn(".matcher=%r will be ignored at runtime — the matcher field is only honored for "
"pre_tool_call / post_tool_call. The hook will fire on every %s event.", matcher, event)
matcher = None
try:
timeout = int(raw.get("timeout", DEFAULT_TIMEOUT_SECONDS))
@@ -266,7 +251,7 @@ def _parse_single_entry(event: str, index: int, raw: Any) -> Optional[ShellHookS
if timeout < 1:
warn(".timeout must be >=1; using default %ds", DEFAULT_TIMEOUT_SECONDS)
timeout = DEFAULT_TIMEOUT_SECONDS
if timeout > MAX_TIMEOUT_SECONDS:
elif timeout > MAX_TIMEOUT_SECONDS:
warn(".timeout=%ds exceeds max %ds; clamping", timeout, MAX_TIMEOUT_SECONDS)
timeout = MAX_TIMEOUT_SECONDS
# ``fail_closed`` (canonical) wins over ``failClosed`` (Cursor/Claude-Code compat).
@@ -275,11 +260,8 @@ def _parse_single_entry(event: str, index: int, raw: Any) -> Optional[ShellHookS
warn(".fail_closed must be a boolean (got %r); using default false (fail open)", fail_closed)
fail_closed = False
if fail_closed and event not in _BLOCKING_EVENTS:
warn(
".fail_closed=true will be ignored at runtime — fail_closed only applies to blocking-capable "
"events (%s). The hook will fail open on %s like any other hook.",
", ".join(sorted(_BLOCKING_EVENTS)), event,
)
warn(".fail_closed=true will be ignored at runtime — fail_closed only applies to blocking-capable "
"events (%s). The hook will fail open on %s like any other hook.", ", ".join(sorted(_BLOCKING_EVENTS)), event)
fail_closed = False
return ShellHookSpec(event=event, command=command.strip(), matcher=matcher, timeout=timeout, fail_closed=fail_closed)
@@ -292,9 +274,7 @@ _POPEN_ERRORS = ((FileNotFoundError, "command not found"), (PermissionError, "co
def _spawn(spec: ShellHookSpec, stdin_json: str) -> Dict[str, Any]:
"""The single subprocess site: run ``spec.command`` with ``stdin_json`` on stdin. Same result keys for every outcome."""
result: Dict[str, Any] = {
"returncode": None, "stdout": "", "stderr": "", "timed_out": False, "elapsed_seconds": 0.0, "error": None,
}
result: Dict[str, Any] = {"returncode": None, "stdout": "", "stderr": "", "timed_out": False, "elapsed_seconds": 0.0, "error": None}
def failed(error: str) -> Dict[str, Any]:
result["error"] = error
@@ -311,27 +291,21 @@ def _spawn(spec: ShellHookSpec, stdin_json: str) -> Dict[str, Any]:
# / taskkill /T). Hooks that finish in time keep detached helpers alive.
popen_kwargs: Dict[str, Any] = {"creationflags": windows_hide_flags()} if IS_WINDOWS else {"process_group": 0}
try:
proc = subprocess.Popen(
argv, stdin=subprocess.PIPE, stdout=subprocess.PIPE, stderr=subprocess.PIPE,
text=True, encoding='utf-8', errors='replace', shell=False, **popen_kwargs,
)
proc = subprocess.Popen(argv, stdin=subprocess.PIPE, stdout=subprocess.PIPE, stderr=subprocess.PIPE,
text=True, encoding='utf-8', errors='replace', shell=False, **popen_kwargs)
except Exception as exc:
return failed(next((msg for cls, msg in _POPEN_ERRORS if isinstance(exc, cls)), str(exc)))
try:
stdout, stderr = proc.communicate(input=stdin_json, timeout=spec.timeout)
except Exception as exc:
# Kill the whole tree — forked helpers holding the pipes would stall the drain.
kill_process_tree(proc)
kill_process_tree(proc) # the whole tree — forked helpers holding the pipes would stall the drain
with suppress(Exception):
proc.communicate(timeout=1)
if not isinstance(exc, subprocess.TimeoutExpired): # pragma: no cover — defensive
return failed(str(exc))
result["timed_out"] = True
result["elapsed_seconds"] = round(time.monotonic() - t0, 3)
result.update(timed_out=True, elapsed_seconds=round(time.monotonic() - t0, 3))
return result
result.update(
returncode=proc.returncode, stdout=stdout or "", stderr=stderr or "", elapsed_seconds=round(time.monotonic() - t0, 3),
)
result.update(returncode=proc.returncode, stdout=stdout or "", stderr=stderr or "", elapsed_seconds=round(time.monotonic() - t0, 3))
return result
@@ -359,10 +333,10 @@ def _evaluate_result(spec: ShellHookSpec, r: Dict[str, Any]) -> Optional[Dict[st
fail_closed = spec.fail_closed and blocking_event
if r["error"]:
logger.warning("shell hook failed (event=%s command=%s): %s", spec.event, spec.command, r["error"])
return _fail_closed_block(spec, r["error"]) if fail_closed else None
if r["timed_out"]:
elif r["timed_out"]:
logger.warning("shell hook timed out after %.2fs (event=%s command=%s)", r["elapsed_seconds"], spec.event, spec.command)
return _fail_closed_block(spec, f"timed out after {spec.timeout}s") if fail_closed else None
if r["error"] or r["timed_out"]:
return _fail_closed_block(spec, r["error"] or f"timed out after {spec.timeout}s") if fail_closed else None
stderr = r["stderr"].strip()
if stderr:
logger.debug("shell hook stderr (event=%s command=%s): %s", spec.event, spec.command, stderr[:_STDERR_MESSAGE_LIMIT])
@@ -375,10 +349,8 @@ def _evaluate_result(spec: ShellHookSpec, r: Dict[str, Any]) -> Optional[Dict[st
return {"action": "block", "message": message}
# Other non-zero exits: still parse stdout so exit-code failures can carry a block directive.
if r["returncode"] != 0:
logger.warning(
"shell hook exited %d (event=%s command=%s); stderr=%s",
r["returncode"], spec.event, spec.command, stderr[:_STDERR_MESSAGE_LIMIT],
)
logger.warning("shell hook exited %d (event=%s command=%s); stderr=%s",
r["returncode"], spec.event, spec.command, stderr[:_STDERR_MESSAGE_LIMIT])
stdout = (r["stdout"] or "").strip()
parsed = _parse_response(spec.event, stdout)
if parsed is None and fail_closed and stdout and not _is_json_object(stdout):
@@ -434,10 +406,7 @@ def _parse_context(data: Dict[str, Any]) -> Optional[Dict[str, Any]]:
return {"context": context} if isinstance(context, str) and context.strip() else None
_RESPONSE_PARSERS: Dict[str, Callable[[Dict[str, Any]], Optional[Dict[str, Any]]]] = {
"pre_tool_call": _parse_pre_tool_call,
"pre_verify": _parse_pre_verify,
}
_RESPONSE_PARSERS: Dict[str, Callable[[Dict[str, Any]], Optional[Dict[str, Any]]]] = {"pre_tool_call": _parse_pre_tool_call, "pre_verify": _parse_pre_verify}
def _parse_response(event: str, stdout: str) -> Optional[Dict[str, Any]]:
@@ -488,12 +457,9 @@ def save_allowlist(data: Dict[str, Any]) -> None:
os.unlink(tmp_path)
raise
except OSError as exc:
logger.warning(
"Failed to persist shell hook allowlist to %s: %s. The approval is in-memory for this run, "
"but the next startup will re-prompt (or skip registration on non-TTY runs without "
"--accept-hooks / HERMES_ACCEPT_HOOKS).",
p, exc,
)
logger.warning("Failed to persist shell hook allowlist to %s: %s. The approval is in-memory for this run, "
"but the next startup will re-prompt (or skip registration on non-TTY runs without "
"--accept-hooks / HERMES_ACCEPT_HOOKS).", p, exc)
def _is_allowlisted(event: str, command: str) -> bool:
@@ -531,31 +497,22 @@ def _prompt_and_record(event: str, command: str, *, accept_hooks: bool) -> bool:
if not sys.stdin.isatty():
return False
print(
f"\n⚠ Hermes is about to register a shell hook that will run a\n"
f" command on your behalf.\n\n"
f" Event: {event}\n"
f" Command: {command}\n\n"
f" Commands run with your full user credentials. Only approve\n"
f" commands you trust."
f"\n⚠ Hermes is about to register a shell hook that will run a\n command on your behalf.\n\n"
f" Event: {event}\n Command: {command}\n\n"
f" Commands run with your full user credentials. Only approve\n commands you trust."
)
try:
answer = input("Allow this hook to run? [y/N]: ").strip().lower()
except (EOFError, KeyboardInterrupt):
print() # keep the terminal tidy after ^C
return False
if answer not in {"y", "yes"}:
return False
_record_approval(event, command)
return True
if answer in {"y", "yes"}:
_record_approval(event, command)
return answer in {"y", "yes"}
def _record_approval(event: str, command: str) -> None:
entry = {
"event": event,
"command": command,
"approved_at": _utc_now_iso(),
"script_mtime_at_approval": script_mtime_iso(command),
}
entry = {"event": event, "command": command, "approved_at": _utc_now_iso(), "script_mtime_at_approval": script_mtime_iso(command)}
with _locked_update_approvals() as data:
data["approvals"] = [e for e in data.get("approvals", []) if not _entry_matches(e, event, command)] + [entry]
@@ -568,9 +525,7 @@ def revoke(command: str) -> int:
return before - len(data["approvals"])
_SCRIPT_EXTENSIONS: Tuple[str, ...] = (
".sh", ".bash", ".zsh", ".fish", ".py", ".pyw", ".rb", ".pl", ".lua", ".js", ".mjs", ".cjs", ".ts",
)
_SCRIPT_EXTENSIONS: Tuple[str, ...] = (".sh", ".bash", ".zsh", ".fish", ".py", ".pyw", ".rb", ".pl", ".lua", ".js", ".mjs", ".cjs", ".ts")
def _command_script_path(command: str) -> str:
@@ -579,11 +534,8 @@ def _command_script_path(command: str) -> str:
parts = split_command_line(command) or [command]
except ValueError:
return command
return (
next((p for p in parts if p.lower().endswith(_SCRIPT_EXTENSIONS)), None)
or next((p for p in parts if "/" in p or p.startswith("~")), None)
or parts[0]
)
return (next((p for p in parts if p.lower().endswith(_SCRIPT_EXTENSIONS)), None)
or next((p for p in parts if "/" in p or p.startswith("~")), None) or parts[0])
def _resolve_effective_accept(cfg: Dict[str, Any], accept_hooks_arg: bool) -> bool:
@@ -591,9 +543,7 @@ def _resolve_effective_accept(cfg: Dict[str, Any], accept_hooks_arg: bool) -> bo
if accept_hooks_arg or os.environ.get("HERMES_ACCEPT_HOOKS", "").strip().lower() in _TRUTHY:
return True
cfg_val = cfg.get("hooks_auto_accept", False)
if isinstance(cfg_val, bool):
return cfg_val
return isinstance(cfg_val, str) and cfg_val.strip().lower() in _TRUTHY
return cfg_val if isinstance(cfg_val, bool) else isinstance(cfg_val, str) and cfg_val.strip().lower() in _TRUTHY
# --- Introspection (used by `hermes hooks` CLI) ---
@@ -606,27 +556,22 @@ def allowlist_entry_for(event: str, command: str) -> Optional[Dict[str, Any]]:
def script_mtime_iso(command: str) -> Optional[str]:
"""ISO-8601 mtime of the resolved script path, or ``None`` if missing."""
path = _command_script_path(command)
if not path:
return None
try:
mtime = os.path.getmtime(os.path.expanduser(path))
mtime = os.path.getmtime(os.path.expanduser(path)) if path else None
except OSError:
return None
return datetime.fromtimestamp(mtime, tz=timezone.utc).isoformat().replace("+00:00", "Z")
return None if mtime is None else datetime.fromtimestamp(mtime, tz=timezone.utc).isoformat().replace("+00:00", "Z")
def script_is_executable(command: str) -> bool:
"""Runnable as configured: a bare script needs X_OK, an interpreter-prefixed one only R_OK (as ``_spawn`` does)."""
path = _command_script_path(command)
expanded = os.path.expanduser(path)
if not path or not os.path.isfile(expanded):
return False
try:
argv = split_command_line(command)
argv = split_command_line(command) if path and os.path.isfile(expanded) else None
except ValueError:
return False
is_bare_invocation = bool(argv) and argv[0] == path
return os.access(expanded, os.X_OK if is_bare_invocation else os.R_OK)
return argv is not None and os.access(expanded, os.X_OK if argv and argv[0] == path else os.R_OK)
def run_once(spec: ShellHookSpec, kwargs: Dict[str, Any]) -> Dict[str, Any]:
+22 -31
View File
@@ -13,6 +13,7 @@ import math
import secrets
import threading
import time
import contextlib
from contextlib import contextmanager
from concurrent.futures import Future, TimeoutError
from typing import Any, Callable, Mapping, Optional
@@ -182,10 +183,6 @@ def _session_id_of(agent: Any) -> Optional[str]:
return str(getattr(agent, "session_id", "") or "") or None
def _finite_number(value: Any) -> bool:
return not isinstance(value, bool) and isinstance(value, (int, float)) and math.isfinite(value)
def _clip(value: Any) -> Optional[str]:
return str(value)[:_MAX_RESULT_CHARS] if value is not None else None
@@ -196,7 +193,7 @@ _HANDLE_FIELD_CHECKS: tuple[tuple[str, Callable[[Any], bool]], ...] = (
("subagent_id", lambda v: isinstance(v, str) and bool(v)),
("parent_session_id", _opt_str),
("correlation_id", _opt_str),
("created_at", _finite_number),
("created_at", lambda v: not isinstance(v, bool) and isinstance(v, (int, float)) and math.isfinite(v)),
("provider", _opt_str),
("model", _opt_str),
("role", lambda v: isinstance(v, str)),
@@ -280,13 +277,13 @@ class SubagentLifecycleService:
record = self._record(handle)
if record is None:
return SubagentTerminalState(handle, SubagentState.UNKNOWN, True, diagnostic="UNKNOWN_HANDLE")
if record.future is not None:
try:
try:
if record.future is not None:
record.future.result(timeout=timeout_seconds)
except TimeoutError:
return SubagentTerminalState(record.handle, record.state, False, True)
except Exception:
pass
except TimeoutError:
return SubagentTerminalState(record.handle, record.state, False, True)
except Exception:
pass
with _REGISTRY.lock:
return SubagentTerminalState(record.handle, record.state, record.result is not None)
@@ -302,12 +299,10 @@ class SubagentLifecycleService:
record.updated_at = time.time()
accepted = False
if agent is not None:
try:
with contextlib.suppress(Exception):
accepted = request_hard_interrupt(
agent, f"Lifecycle cancellation requested: {reason[:500]}", tool_reason="subagent cancellation requested",
)
except Exception:
accepted = False
return SubagentCancelResult(bool(accepted), unsupported=not accepted, state=SubagentState.CANCEL_REQUESTED)
def result(self, handle: SubagentHandle) -> SubagentResult:
@@ -365,9 +360,8 @@ class SubagentLifecycleService:
else:
state = SubagentState.SUCCEEDED if status == "completed" else SubagentState.FAILED
fields: dict[str, Any] = dict(
summary=_clip(raw.get("summary")),
summary=_clip(raw.get("summary")), error_message=_clip(raw.get("error") or None),
error_classification=None if state == SubagentState.SUCCEEDED else status.upper(),
error_message=_clip(raw.get("error") or None),
usage_metadata={"api_calls": raw.get("api_calls", 0)} if is_dict else {},
tool_execution_summary={"duration_seconds": raw.get("duration_seconds", 0)} if is_dict else {},
)
@@ -377,14 +371,10 @@ class SubagentLifecycleService:
result = SubagentResult(record.handle, state, True, started_at=record.started_at, completed_at=time.time(), **fields)
payload = dataclasses.asdict(result)
payload.pop("result_hash", None)
digest = hashlib.sha256(json.dumps(payload, sort_keys=True, default=str).encode()).hexdigest()
result = dataclasses.replace(result, result_hash=digest)
result = dataclasses.replace(result, result_hash=hashlib.sha256(json.dumps(payload, sort_keys=True, default=str).encode()).hexdigest())
with _REGISTRY.lock:
record.agent = None
record.result = result
record.state = result.terminal_state
record.completed_at = result.completed_at
record.updated_at = result.completed_at or time.time()
record.agent, record.result, record.state = None, result, result.terminal_state
record.completed_at = record.updated_at = result.completed_at
@staticmethod
def _capability(subagent_id: str, parent_session_id: Optional[str], created_at: float) -> str:
@@ -402,11 +392,12 @@ class SubagentLifecycleService:
raise SubagentLifecycleError("metadata must be JSON-serializable.") from exc
if metadata_bytes > _MAX_METADATA_BYTES:
raise SubagentLifecycleError("metadata exceeds 8192 bytes.")
if request.allowed_toolsets:
from toolsets import TOOLSETS
unknown = set(request.allowed_toolsets) - set(TOOLSETS)
if unknown:
raise SubagentLifecycleError(f"Unknown toolsets: {', '.join(sorted(unknown))}.")
enabled = getattr(parent, "enabled_toolsets", None)
if enabled is not None and not set(request.allowed_toolsets).issubset(set(enabled)):
raise SubagentLifecycleError("Requested toolsets would broaden parent permissions.")
if not request.allowed_toolsets:
return
from toolsets import TOOLSETS
unknown = set(request.allowed_toolsets) - set(TOOLSETS)
if unknown:
raise SubagentLifecycleError(f"Unknown toolsets: {', '.join(sorted(unknown))}.")
enabled = getattr(parent, "enabled_toolsets", None)
if enabled is not None and not set(request.allowed_toolsets).issubset(set(enabled)):
raise SubagentLifecycleError("Requested toolsets would broaden parent permissions.")
+30 -45
View File
@@ -61,11 +61,8 @@ _LANGUAGE_RULE_PINNED = "- Write the title in {language}."
# Constrains the response to a single title field ("model answered instead of titling" failure class).
_TITLE_RESPONSE_FORMAT = {
"type": "json_schema",
"json_schema": {
"name": "session_title",
"strict": True,
"schema": {"type": "object", "properties": {"title": {"type": "string"}}, "required": ["title"], "additionalProperties": False},
},
"json_schema": {"name": "session_title", "strict": True, "schema": {
"type": "object", "properties": {"title": {"type": "string"}}, "required": ["title"], "additionalProperties": False}},
}
# Control-tag wrappers around machine-authored content inside a nominal "user" message (Codex CLI's
@@ -147,9 +144,8 @@ def _summarize_user_message(user_message: str) -> str:
def is_titleable_user_message(user_message: str) -> bool:
"""False for machine-authored openers and turns that reduce to nothing once scaffolding is stripped."""
if not isinstance(user_message, str) or not user_message.strip() or user_message.lstrip().startswith(_MACHINE_PREFIXES):
return False
return bool(_summarize_user_message(user_message).strip())
return (isinstance(user_message, str) and bool(user_message.strip()) and not user_message.lstrip().startswith(_MACHINE_PREFIXES)
and bool(_summarize_user_message(user_message).strip()))
def derive_title(user_message: str) -> Optional[str]:
@@ -186,10 +182,9 @@ def _extract_title_text(content: str) -> str:
pass
match = re.search(r'"title\"\s*:\s*"((?:[^"\\]|\\.)*)"', raw)
if match:
try:
with suppress(ValueError):
return json.loads(f'"{match.group(1)}"').strip()
except ValueError:
return match.group(1).strip()
return match.group(1).strip()
# Prose fallback: scrub <think> blocks so reasoning can't leak into a title.
try:
from agent.agent_runtime_helpers import strip_think_blocks
@@ -209,10 +204,9 @@ def _clean_title(text: str) -> Optional[str]:
def _safe_callback(callback: Optional[Callable], args: tuple, log_fmt: str, label: str) -> None:
"""Invoke an optional consumer callback, never raising."""
if callback is None:
return
try:
callback(*args)
if callback is not None:
callback(*args)
except Exception:
logger.debug(log_fmt, label, exc_info=True)
@@ -237,20 +231,20 @@ def generate_title(
if not _auto_title_enabled():
logger.debug("Auto-title skipped: auxiliary.title_generation.enabled=false")
return None
if runtime_validator is not None:
try:
if not runtime_validator():
logger.debug("Title generation skipped: runtime validator returned False")
return None
except Exception: # fail open: a broken validator must not disable titling
logger.debug("Title runtime validator raised; proceeding", exc_info=True)
try:
if runtime_validator is not None and not runtime_validator():
logger.debug("Title generation skipped: runtime validator returned False")
return None
except Exception: # fail open: a broken validator must not disable titling
logger.debug("Title runtime validator raised; proceeding", exc_info=True)
user_snippet = _summarize_user_message(user_message)[:MAX_TITLE_INPUT_CHARS]
if not user_snippet.strip():
return None
language = _title_language()
language_rule = _LANGUAGE_RULE_PINNED.format(language=language) if language else _LANGUAGE_RULE_MATCH_USER
# str.replace, not str.format: the prompt embeds literal JSON braces.
prompt = _TITLE_PROMPT_TEMPLATE.replace("__LANGUAGE_RULE__", language_rule)
prompt = _TITLE_PROMPT_TEMPLATE.replace(
"__LANGUAGE_RULE__", _LANGUAGE_RULE_PINNED.format(language=language) if language else _LANGUAGE_RULE_MATCH_USER,
)
try:
response = call_llm(
task="title_generation",
@@ -260,8 +254,8 @@ def generate_title(
extra_body={"response_format": _TITLE_RESPONSE_FORMAT},
)
title = _clean_title(_extract_title_text(response.choices[0].message.content or ""))
# Answer-shaped output: reject (not truncate) so the caller retries next exchange.
if title is not None and len(title.split()) > _MAX_TITLE_WORDS:
# Answer-shaped output: reject (not truncate) so the caller retries next exchange.
logger.debug("Rejecting answer-shaped title output (%d words > %d)", len(title.split()), _MAX_TITLE_WORDS)
return None
return title
@@ -309,9 +303,7 @@ def _persist_session_title(session_db, session_id, title, *, source, dedupe=True
return _set(title)
except ValueError:
next_title_fn = getattr(session_db, "get_next_title_in_lineage", None)
if not dedupe or next_title_fn is None:
raise
deduped = next_title_fn(title)
deduped = next_title_fn(title) if dedupe and next_title_fn is not None else None
if not deduped or deduped == title:
raise
return _set(deduped)
@@ -323,9 +315,7 @@ def apply_instant_title(session_db, session_id: str, user_message: str, title_ca
return None
try:
title = derive_title(user_message) if is_titleable_user_message(user_message) else None
if not title:
return None
persisted = _persist_session_title(session_db, session_id, title, source="derived", dedupe=False)
persisted = _persist_session_title(session_db, session_id, title, source="derived", dedupe=False) if title else None
if persisted:
_notify_title(title_callback, persisted, "derived", "Instant-title")
return persisted
@@ -359,22 +349,21 @@ def auto_title_session(
conversation_id = session_db.get_conversation_root(session_id) or session_id
set_conversation_context(conversation_id)
set_accounting_context(session_db, session_id)
title = generate_title(
title, source = generate_title(
user_message, failure_callback=failure_callback, main_runtime=main_runtime, runtime_validator=runtime_validator,
)
source = "llm"
), "llm"
if not title: # the inline attempt declined collisions; off the critical path the lineage scan is affordable
title, source = derive_title(user_message), "derived"
if not title:
return
if not title:
return
try:
persisted = _persist_session_title(session_db, session_id, title, source=source)
if persisted is None:
return
logger.debug("Auto-generated session title: %s", persisted)
_notify_title(title_callback, persisted, source, "Auto-title")
except Exception as e:
logger.debug("Failed to set auto-generated title: %s", e)
return
if persisted is not None:
logger.debug("Auto-generated session title: %s", persisted)
_notify_title(title_callback, persisted, source, "Auto-title")
except Exception as e:
# WARNING so operators see it in agent.log; names the likely cause.
logger.warning("Auto-title failed (harmless; if this started after an update, restart the running Hermes process): %s", e)
@@ -393,10 +382,8 @@ def _is_real_user_turn(message: Any) -> bool:
def _session_is_untitled(session_db, session_id: str) -> bool:
"""No title of any provenance; False when it can't tell (no model call per turn for an unreadable title)."""
getter = getattr(session_db, "get_session_title", None)
if not callable(getter):
return False
try:
return not str(getter(session_id) or "").strip()
return callable(getter) and not str(getter(session_id) or "").strip()
except Exception:
logger.debug("Untitled check failed for %s", session_id, exc_info=True)
return False
@@ -418,9 +405,7 @@ def maybe_auto_title(
# History may be pre- or post-message. Skip only when BOTH past the opening turn AND named: count alone
# left a machinery-opened session nameless; title alone never titles on an old store.
user_msg_count = sum(1 for m in (conversation_history or []) if _is_real_user_turn(m))
if user_msg_count > 1 and not _session_is_untitled(session_db, session_id):
return
if not is_titleable_user_message(user_message):
if (user_msg_count > 1 and not _session_is_untitled(session_db, session_id)) or not is_titleable_user_message(user_message):
return
if not _auto_title_enabled(): # config read after the cheap guards so the file isn't touched every turn
logger.debug("Auto-title skipped: auxiliary.title_generation.enabled=false")
+21 -54
View File
@@ -65,21 +65,18 @@ def _text_block(text: Any, redact: bool) -> Dict[str, Any]:
def _part_to_block(part: Any, redact: bool) -> Dict[str, Any]:
if not isinstance(part, dict):
return _text_block(str(part), redact)
ptype = part.get("type")
if ptype == "text":
if part.get("type") == "text":
return _text_block(part.get("text", ""), redact)
if ptype in ("image_url", "image"):
if part.get("type") in ("image_url", "image"):
return {"type": "text", "text": "[image omitted]"} # the viewer renders text turns; no base64
return _text_block(json.dumps(part), redact)
def _content_to_blocks(content: Any, redact: bool) -> List[Dict[str, Any]]:
"""Normalize a message ``content`` field into Anthropic content blocks."""
if content is None:
return []
if isinstance(content, list):
return [_part_to_block(part, redact) for part in content]
return [_text_block(content if isinstance(content, str) else json.dumps(content), redact)]
return [] if content is None else [_text_block(content if isinstance(content, str) else json.dumps(content), redact)]
def _parse_tool_args(raw_args: Any) -> Dict[str, Any]:
@@ -105,12 +102,8 @@ def _tool_calls_to_blocks(tool_calls: Any, redact: bool) -> List[Dict[str, Any]]
except (json.JSONDecodeError, ValueError):
logger.warning("Trace upload redacted tool arguments are not valid JSON; refusing upload")
raise TraceRedactionError(_REDACTION_BLOCKED_MESSAGE)
blocks.append({
"type": "tool_use",
"id": tc.get("id") or f"toolu_{uuid.uuid4().hex[:16]}",
"name": fn.get("name") or tc.get("name") or "tool",
"input": parsed,
})
blocks.append({"type": "tool_use", "id": tc.get("id") or f"toolu_{uuid.uuid4().hex[:16]}",
"name": fn.get("name") or tc.get("name") or "tool", "input": parsed})
return blocks
@@ -119,10 +112,8 @@ def _git_branch(cwd: str) -> str:
return ""
try:
import subprocess
r = subprocess.run(
["git", "rev-parse", "--abbrev-ref", "HEAD"],
capture_output=True, text=True, encoding="utf-8", errors="replace", timeout=3, cwd=cwd,
)
r = subprocess.run(["git", "rev-parse", "--abbrev-ref", "HEAD"],
capture_output=True, text=True, encoding="utf-8", errors="replace", timeout=3, cwd=cwd)
except Exception:
return ""
return r.stdout.strip() if r.returncode == 0 else ""
@@ -135,14 +126,10 @@ def _assistant_message(msg: Dict[str, Any], model: str, redact: bool) -> Dict[st
def _tool_result_message(msg: Dict[str, Any], model: str, redact: bool) -> Dict[str, Any]:
content = msg.get("content")
return {
"role": "user",
"content": [{
"type": "tool_result",
"tool_use_id": msg.get("tool_call_id") or msg.get("tool_name") or "tool",
"content": _redact(content if isinstance(content, str) else json.dumps(content), redact),
}],
}
return {"role": "user", "content": [{
"type": "tool_result", "tool_use_id": msg.get("tool_call_id") or msg.get("tool_name") or "tool",
"content": _redact(content if isinstance(content, str) else json.dumps(content), redact),
}]}
def _user_message(msg: Dict[str, Any], model: str, redact: bool) -> Dict[str, Any]:
@@ -151,10 +138,7 @@ def _user_message(msg: Dict[str, Any], model: str, redact: bool) -> Dict[str, An
# role -> (Claude Code line type, message builder). Unknown roles render as user.
_ROLE_RENDERERS: Dict[Any, Tuple[str, Any]] = {
"assistant": ("assistant", _assistant_message),
"tool": ("user", _tool_result_message),
}
_ROLE_RENDERERS: Dict[Any, Tuple[str, Any]] = {"assistant": ("assistant", _assistant_message), "tool": ("user", _tool_result_message)}
def build_trace_jsonl(messages: List[Dict[str, Any]], *, session_id: str, model: str = "", cwd: str = "", redact: bool = True) -> str:
@@ -170,18 +154,10 @@ def build_trace_jsonl(messages: List[Dict[str, Any]], *, session_id: str, model:
continue
turn_uuid = str(uuid.uuid4())
line_type, render = _ROLE_RENDERERS.get(role, ("user", _user_message))
entry = {
"parentUuid": parent,
"isSidechain": False,
"userType": "external",
"cwd": cwd or os.getcwd(),
"sessionId": session_id,
"version": _HERMES_VERSION,
"gitBranch": git_branch,
"uuid": turn_uuid,
"timestamp": base_ts,
"type": line_type,
"message": render(msg, model, redact),
entry = { # key order is the wire order
"parentUuid": parent, "isSidechain": False, "userType": "external", "cwd": cwd or os.getcwd(),
"sessionId": session_id, "version": _HERMES_VERSION, "gitBranch": git_branch, "uuid": turn_uuid,
"timestamp": base_ts, "type": line_type, "message": render(msg, model, redact),
}
lines.append(json.dumps(entry, ensure_ascii=False))
parent = turn_uuid
@@ -207,10 +183,10 @@ def _do_upload(jsonl: str, *, token: str, session_id: str, dataset_name: str = D
api = HfApi(token=token)
try:
who = api.whoami()
user = who.get("name") if isinstance(who, dict) else None
except Exception as e:
logger.warning("HF whoami failed: %s", e)
return "Your Hugging Face token was rejected (whoami failed). Make sure it has WRITE access and isn't expired."
user = who.get("name") if isinstance(who, dict) else None
if not user:
return "Could not resolve your Hugging Face username from the token."
repo_id = f"{user}/{dataset_name}"
@@ -221,10 +197,8 @@ def _do_upload(jsonl: str, *, token: str, session_id: str, dataset_name: str = D
return f"Could not create/access dataset {repo_id}: {e}"
path_in_repo = f"sessions/{session_id}.jsonl"
try:
api.upload_file(
path_or_fileobj=jsonl.encode("utf-8"), path_in_repo=path_in_repo, repo_id=repo_id, repo_type="dataset",
commit_message=f"add session trace {session_id}",
)
api.upload_file(path_or_fileobj=jsonl.encode("utf-8"), path_in_repo=path_in_repo, repo_id=repo_id,
repo_type="dataset", commit_message=f"add session trace {session_id}")
except Exception as e:
logger.warning("HF upload_file failed for %s: %s", repo_id, e)
return f"Upload to Hugging Face failed: {e}"
@@ -249,15 +223,8 @@ def load_session_messages(session_id: str, db_path=None) -> Tuple[List[Dict[str,
def upload_session_trace(
session_id: str,
*,
model: str = "",
cwd: str = "",
redact: bool = True,
private: bool = True,
dataset_name: str = DEFAULT_DATASET_NAME,
db_path=None,
token: Optional[str] = None,
session_id: str, *, model: str = "", cwd: str = "", redact: bool = True, private: bool = True,
dataset_name: str = DEFAULT_DATASET_NAME, db_path=None, token: Optional[str] = None,
) -> str:
"""CLI/gateway entry point: load, convert, upload to ``{user}/hermes-traces``. Status string, never raises."""
if not session_id: