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:
+6
-12
@@ -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
@@ -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
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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]:
|
||||
|
||||
Reference in New Issue
Block a user