refactor(hermes_cli): group D — is_relative_to/closing/suppress collapses, dedupe external paths via dict, drop redundant snap_dir/guards, compact bang_shell/archive_safe/__init__ layout

This commit is contained in:
Teknium
2026-09-02 22:51:30 -07:00
parent 0ebee818e1
commit bf1300b27c
5 changed files with 51 additions and 123 deletions
+6 -12
View File
@@ -15,32 +15,26 @@ def _ensure_utf8():
`hermes setup` on a fresh Pi).
"""
repaired = False
for stream_name in ("stdout", "stderr"):
stream = getattr(sys, stream_name, None)
if stream is None:
continue
try:
encoding = (getattr(stream, "encoding", "") or "").lower().replace("-", "")
if encoding == "utf8":
if (getattr(stream, "encoding", "") or "").lower().replace("-", "") == "utf8":
continue
# Preferred: reconfigure in place, preserving object identity so code already holding
# a reference to the old sys.stdout benefits from the repair too.
reconfigure = getattr(stream, "reconfigure", None)
if callable(reconfigure):
reconfigure(encoding="utf-8", errors="replace")
repaired = True
continue
# Fallback for streams without reconfigure(): reopen the fd as UTF-8 (closefd=False
# keeps the original fd open).
new_stream = open(stream.fileno(), "w", encoding="utf-8", errors="replace", buffering=1, closefd=False)
setattr(sys, stream_name, new_stream)
else:
# No reconfigure(): reopen the fd as UTF-8 (closefd=False keeps the original fd open).
new_stream = open(stream.fileno(), "w", encoding="utf-8", errors="replace",
buffering=1, closefd=False)
setattr(sys, stream_name, new_stream)
repaired = True
except (AttributeError, OSError, ValueError):
pass
# Only nudge child processes toward UTF-8 when a non-UTF-8 locale was actually detected; on a
# healthy UTF-8 host children inherit it from the locale already.
if repaired:
+11 -25
View File
@@ -6,6 +6,7 @@ import os
import shutil
import tarfile
import tempfile
from contextlib import suppress
from pathlib import Path, PurePosixPath, PureWindowsPath
@@ -19,12 +20,9 @@ def normalize_archive_parts(member_name: str) -> list[str]:
normalized_name = member_name.replace("\\", "/")
posix_path = PurePosixPath(normalized_name)
windows_path = PureWindowsPath(member_name)
if not normalized_name or posix_path.is_absolute() or windows_path.is_absolute() or windows_path.drive:
raise ValueError(f"Unsafe archive member path: {member_name}")
parts = [part for part in posix_path.parts if part not in {"", "."}]
if not parts or ".." in parts:
if (not normalized_name or posix_path.is_absolute() or windows_path.is_absolute()
or windows_path.drive or not parts or ".." in parts):
raise ValueError(f"Unsafe archive member path: {member_name}")
return parts
@@ -38,15 +36,13 @@ def make_targz(base: str, root_dir: str, base_dir: str) -> str:
dest_dir = os.path.dirname(archive_path) or "."
fd, tmp_path = tempfile.mkstemp(dir=dest_dir, prefix=".archive_", suffix=".tar.gz.tmp")
try:
with os.fdopen(fd, "wb") as f:
with tarfile.open(fileobj=f, mode="w:gz", format=tarfile.GNU_FORMAT) as tf:
tf.add(str(Path(root_dir) / base_dir), arcname=base_dir)
with os.fdopen(fd, "wb") as f, \
tarfile.open(fileobj=f, mode="w:gz", format=tarfile.GNU_FORMAT) as tf:
tf.add(str(Path(root_dir) / base_dir), arcname=base_dir)
os.replace(tmp_path, archive_path)
except BaseException:
try:
with suppress(OSError):
os.unlink(tmp_path)
except OSError:
pass
raise
return archive_path
@@ -61,26 +57,19 @@ def safe_extract_targz(archive: Path, destination: Path) -> None:
with tarfile.open(archive, "r:gz") as tf:
for member in tf.getmembers():
target = destination.joinpath(*normalize_archive_parts(member.name))
if member.isdir():
target.mkdir(parents=True, exist_ok=True)
continue
if not member.isfile():
raise ValueError(f"Unsupported archive member type: {member.name}")
target.parent.mkdir(parents=True, exist_ok=True)
extracted = tf.extractfile(member)
if extracted is None:
raise ValueError(f"Cannot read archive member: {member.name}")
with extracted, open(target, "wb") as dst:
shutil.copyfileobj(extracted, dst)
try:
with suppress(OSError):
os.chmod(target, member.mode & 0o777)
except OSError:
pass
def archive_root_dirs(archive: Path) -> set[str]:
@@ -91,12 +80,9 @@ def archive_root_dirs(archive: Path) -> set[str]:
archive) without first mutating a live tree.
"""
with tarfile.open(archive, "r:gz") as tf:
return {
parts[0]
for member in tf.getmembers()
for parts in [normalize_archive_parts(member.name)]
if len(parts) > 1 or member.isdir()
}
return {parts[0] for member in tf.getmembers()
for parts in [normalize_archive_parts(member.name)]
if len(parts) > 1 or member.isdir()}
def copy_regular_files(src: Path, dst: Path) -> int:
+21 -54
View File
@@ -11,7 +11,7 @@ import tempfile
import threading
import time
import zipfile
from contextlib import contextmanager, suppress
from contextlib import closing, contextmanager, suppress
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
@@ -170,11 +170,7 @@ def _atomic_output_path(final_path: Path):
def _is_within(path: Path, root: Path) -> bool:
"""True when *path* resolves inside the already-resolved *root* (traversal / symlink guard)."""
try:
path.resolve().relative_to(root)
except ValueError:
return False
return True
return path.resolve().is_relative_to(root)
def _collect_memory_provider_external_paths() -> List[Path]:
@@ -193,18 +189,13 @@ def _collect_memory_provider_external_paths() -> List[Path]:
except Exception as exc:
logger.warning("backup_paths() failed for memory provider %r: %s", active, exc)
return []
out: List[Path] = []
seen: set = set()
out: Dict[Path, Path] = {} # resolved -> first declared spelling
for raw in declared:
try:
with suppress(Exception):
p = Path(raw).expanduser()
resolved = p.resolve() if p.exists() else None
except Exception:
continue
if resolved is not None and resolved not in seen:
seen.add(resolved)
out.append(p)
return out
if p.exists():
out.setdefault(p.resolve(), p)
return list(out.values())
def _iter_external_files(base: Path) -> List[Path]:
@@ -215,13 +206,10 @@ def _iter_external_files(base: Path) -> List[Path]:
return []
files: List[Path] = []
for dirpath, dirnames, filenames in os.walk(base, followlinks=False):
dp = Path(dirpath)
dirnames[:] = [d for d in dirnames if d not in _EXCLUDED_DIRS]
for fname in filenames:
fpath = dp / fname
if fpath.is_symlink() or fname in _EXCLUDED_NAMES or fname.endswith(_EXCLUDED_SUFFIXES):
continue
files.append(fpath)
files.extend(fp for fp in (Path(dirpath) / f for f in filenames)
if not (fp.is_symlink() or fp.name in _EXCLUDED_NAMES
or fp.name.endswith(_EXCLUDED_SUFFIXES)))
return files
@@ -448,11 +436,8 @@ def _safe_restore_db(src: Path, dst: Path) -> bool:
# Checkpoint first so the backup starts clean rather than writing on top of a deep WAL.
with suppress(Exception):
dst_conn.execute("PRAGMA wal_checkpoint(TRUNCATE)")
src_conn = sqlite3.connect(f"file:{src}?mode=ro", uri=True)
try:
with closing(sqlite3.connect(f"file:{src}?mode=ro", uri=True)) as src_conn:
src_conn.backup(dst_conn)
finally:
src_conn.close()
dst_conn.close()
with suppress(Exception):
dst.chmod(src.stat().st_mode)
@@ -474,10 +459,8 @@ def _unlink_move_restore_db(src: Path, dst: Path) -> bool:
try:
holders = _foreign_db_holder_pids(dst)
if holders:
logger.error(
"Refusing unlink+move restore of %s: process(es) %s still "
"hold the database or its WAL open. Stop them and retry.",
dst, holders)
logger.error("Refusing unlink+move restore of %s: process(es) %s still "
"hold the database or its WAL open. Stop them and retry.", dst, holders)
return False
with offline_file_access(dst, what="unlink+move restore of"):
tmp = dst.parent / f".{dst.name}.snap_restore"
@@ -491,10 +474,8 @@ def _unlink_move_restore_db(src: Path, dst: Path) -> bool:
shutil.move(str(tmp), str(dst))
return True
except LiveConnectionError as exc2:
logger.error(
"Refusing unlink+move restore of %s: %s Close the in-process "
"database handles (or restart Hermes) and retry.",
dst, exc2)
logger.error("Refusing unlink+move restore of %s: %s Close the in-process "
"database handles (or restart Hermes) and retry.", dst, exc2)
return False
except Exception as exc2:
logger.error("Fallback restore also failed for %s -> %s: %s", src, dst, exc2)
@@ -592,11 +573,9 @@ def _collect_external_entries() -> tuple[list[tuple[Path, str]], list[str]]:
skipped_external.append(str(base))
continue
for fpath in _iter_external_files(base):
try:
with suppress(ValueError, OSError):
rel_to_home = fpath.resolve().relative_to(home_dir)
except (ValueError, OSError):
continue
external_to_add.append((fpath, _EXTERNAL_PREFIX + rel_to_home.as_posix()))
external_to_add.append((fpath, _EXTERNAL_PREFIX + rel_to_home.as_posix()))
return external_to_add, skipped_external
@@ -694,8 +673,6 @@ def _validate_backup_zip(zf: zipfile.ZipFile) -> tuple[bool, str]:
def _detect_prefix(zf: zipfile.ZipFile) -> str:
"""Detect if the zip has a common directory prefix wrapping all entries."""
names = [n for n in zf.namelist() if not n.endswith("/")]
if not names:
return ""
first_parts = {Path(n).parts[0] for n in names if len(Path(n).parts) > 1}
if len(first_parts) == 1 and first_parts <= {".hermes", "hermes"}:
return first_parts.pop() + "/"
@@ -954,16 +931,8 @@ def _revive_gateway_after_import(hermes_root: Path) -> None:
# directories (recursive); missing entries are skipped. Pairing data lives in platform JSON blobs
# outside state.db, so it is listed explicitly — ``hermes update`` snapshots this set (#15733).
_QUICK_STATE_FILES = (
"state.db",
"config.yaml",
".env",
"auth.json",
"cron/jobs.json",
"cron/executions.db",
"gateway_state.json",
"channel_directory.json",
"channel_aliases.json",
"processes.json",
"state.db", "config.yaml", ".env", "auth.json", "cron/jobs.json", "cron/executions.db",
"gateway_state.json", "channel_directory.json", "channel_aliases.json", "processes.json",
"gateway/discord_message_recovery.db", # Discord reconnect replay ledger
# Per-profile user stores, destroyed if the update flow replaces the file and the post-update
# schema-init re-creates an empty one (#52889). Skipped when outside HERMES_HOME.
@@ -1058,14 +1027,13 @@ def _copy_quick_snapshot_files(
def _create_quick_snapshot_locked(
label: Optional[str], hermes_home: Optional[Path], keep: Optional[int], max_file_size: Optional[int]
label: Optional[str], home: Path, keep: Optional[int], max_file_size: Optional[int]
) -> Optional[str]:
"""Copy the quick-snapshot set to a timestamped dir under state-snapshots/ and prune old ones.
``max_file_size`` skips (with a warning) larger files: the pre-update snapshot uses it so a
multi-GB ``state.db`` never stalls ``hermes update`` while the small files are always captured.
"""
home = hermes_home or get_hermes_home()
root = _quick_snapshot_root(home)
ts = datetime.now(timezone.utc).strftime("%Y%m%d-%H%M%S")
base_snap_id = f"{ts}-{label}" if label else ts
@@ -1073,7 +1041,6 @@ def _create_quick_snapshot_locked(
while (root / snap_id).exists():
snap_id = f"{base_snap_id}-{suffix}"
suffix += 1
snap_dir = root / snap_id
staging_dir = root / f".{snap_id}.{os.getpid()}.partial"
shutil.rmtree(staging_dir, ignore_errors=True)
staging_dir.mkdir(parents=True, exist_ok=False)
@@ -1098,7 +1065,7 @@ def _create_quick_snapshot_locked(
}
with open(staging_dir / "manifest.json", "w", encoding="utf-8") as f:
json.dump(meta, f, indent=2)
os.replace(staging_dir, snap_dir)
os.replace(staging_dir, root / snap_id)
# Auto-prune (pre-update callers pass a smaller keep so state.db copies don't accumulate).
# Skip when a DB failed to capture OR was skipped for size (#68805): the snapshot is
# incomplete and the older one may hold the only recoverable database.
+9 -24
View File
@@ -10,6 +10,7 @@ from __future__ import annotations
import os
import subprocess
from contextlib import suppress
from typing import Optional
USAGE_HINT = "Usage: !<command> — run a shell command without spending a model turn (e.g. !git status)"
@@ -34,9 +35,7 @@ def parse_bang_command(text: str) -> str:
``! ls -la`` -> ``ls -la``; ``!!`` -> ``!`` — a literal second bang belongs to the user's shell
(history expansion), not to Hermes.
"""
if not is_bang_command(text):
return ""
return text.strip()[1:].strip()
return text.strip()[1:].strip() if is_bang_command(text) else ""
def bang_shell_enabled() -> bool:
@@ -52,11 +51,8 @@ def bang_shell_enabled() -> bool:
def env_var_enabled(name, default=""): # type: ignore[misc]
return str(os.getenv(name, default)).strip().lower() in {"1", "true", "yes", "on"}
return not (
env_var_enabled("HERMES_GATEWAY_SESSION")
or env_var_enabled("HERMES_CRON_SESSION")
or (os.getenv("HERMES_SESSION_PLATFORM") or "").strip()
)
return not (env_var_enabled("HERMES_GATEWAY_SESSION") or env_var_enabled("HERMES_CRON_SESSION")
or (os.getenv("HERMES_SESSION_PLATFORM") or "").strip())
def resolve_bang_cwd(session_key: Optional[str] = None) -> Optional[str]:
@@ -67,7 +63,6 @@ def resolve_bang_cwd(session_key: Optional[str] = None) -> Optional[str]:
"""
try:
from tools.terminal_tool import _get_env_config, get_session_cwd
return get_session_cwd(session_key) or (_get_env_config() or {}).get("cwd") or None
except Exception:
return None
@@ -99,7 +94,6 @@ def _bang_env() -> dict:
"""
try:
from tools.environments.local import _sanitize_subprocess_env
return _sanitize_subprocess_env(os.environ.copy())
except Exception:
return os.environ.copy()
@@ -111,29 +105,23 @@ def run_bang_command(command: str, *, cwd: Optional[str] = None, timeout: int =
Output exists only on the user's terminal — nothing is returned for insertion into history.
"""
emit = writer or (lambda line: print(line, end="" if line.endswith("\n") else "\n"))
run_cwd = os.path.expanduser(cwd) if cwd else None
if run_cwd and not os.path.isdir(run_cwd):
run_cwd = None
try:
from hermes_cli._subprocess_compat import windows_hide_flags
creationflags = windows_hide_flags()
except Exception:
creationflags = 0
try:
# shell=True is intentional (matches quick_commands): the human typed this, not the model.
proc = subprocess.Popen(
command, shell=True, stdout=subprocess.PIPE, stderr=subprocess.STDOUT,
text=True, encoding="utf-8", errors="replace",
cwd=run_cwd, env=_bang_env(), creationflags=creationflags,
)
command, shell=True, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True,
encoding="utf-8", errors="replace", cwd=run_cwd, env=_bang_env(),
creationflags=creationflags)
except Exception as exc:
emit(f"!: failed to run command: {exc}")
return 127
try:
if proc.stdout is not None:
for line in proc.stdout:
@@ -149,10 +137,7 @@ def run_bang_command(command: str, *, cwd: Optional[str] = None, timeout: int =
emit("!: interrupted")
return 130
finally:
try:
if proc.stdout is not None:
if proc.stdout is not None:
with suppress(Exception):
proc.stdout.close()
except Exception:
pass
return int(proc.returncode or 0)
+4 -8
View File
@@ -152,9 +152,8 @@ def _canonical_github_remote(url: str | None) -> str:
def _is_official_ssh_remote(url: str | None) -> bool:
if not url or not url.strip().lower().startswith(("git@", "ssh://")):
return False
return _canonical_github_remote(url) == _OFFICIAL_REPO_CANONICAL
return bool(url) and url.strip().lower().startswith(("git@", "ssh://")) and (
_canonical_github_remote(url) == _OFFICIAL_REPO_CANONICAL)
_GIT_TEXT_KW = {"text": True, "encoding": "utf-8", "errors": "replace"}
@@ -464,11 +463,8 @@ def prefetch_banner_data():
if _banner_data_prefetch_started:
return
_banner_data_prefetch_started = True
def _run() -> None:
for warm in (get_git_banner_state, get_latest_release_tag, get_available_skills):
_quiet(warm)
_daemon("banner-data-prefetch", _run)
_daemon("banner-data-prefetch", lambda: [_quiet(warm) for warm in (
get_git_banner_state, get_latest_release_tag, get_available_skills)])
def get_update_result(timeout: float = 0.5) -> Optional[int]: