Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 33 additions & 5 deletions gateway/sticker_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,12 +12,38 @@
import os
import tempfile
import time
from pathlib import Path
from typing import Optional

from hermes_cli.config import get_hermes_home


CACHE_PATH = get_hermes_home() / "sticker_cache.json"
# Import-time default, used by ``_resolve_cache_path`` below to detect a test
# monkeypatch of the constant (test seam, see tests/gateway/test_sticker_cache.py)
# vs. an unmodified import-time value (in which case it re-resolves through
# the active profile override).
_CACHE_PATH_IMPORT_DEFAULT = CACHE_PATH


def _resolve_cache_path() -> Path:
"""Resolve the sticker cache path, honoring the active profile.

``CACHE_PATH`` is frozen at import time, which pins every later
read/write to whichever profile's ``HERMES_HOME`` was active when this
module was first imported — a cross-profile leak under the multiplexed
gateway, where multiple profiles share one process. Re-resolve through
``get_hermes_home()`` on every call so the context-local profile
override (``set_hermes_home_override``) is honored, mirroring the
cache-dir profile-isolation fix.

A test that monkeypatches the module constant away from its import-time
default is respected (test seam preserved).
"""
current = CACHE_PATH
if current != _CACHE_PATH_IMPORT_DEFAULT:
return current
return get_hermes_home() / "sticker_cache.json"

# Vision prompt for describing stickers -- kept concise to save tokens
STICKER_VISION_PROMPT = (
Expand All @@ -28,26 +54,28 @@

def _load_cache() -> dict:
"""Load the sticker cache from disk."""
if CACHE_PATH.exists():
cache_path = _resolve_cache_path()
if cache_path.exists():
try:
return json.loads(CACHE_PATH.read_text(encoding="utf-8"))
return json.loads(cache_path.read_text(encoding="utf-8"))
except (json.JSONDecodeError, OSError):
return {}
return {}


def _save_cache(cache: dict) -> None:
"""Save the sticker cache to disk atomically."""
CACHE_PATH.parent.mkdir(parents=True, exist_ok=True)
cache_path = _resolve_cache_path()
cache_path.parent.mkdir(parents=True, exist_ok=True)
fd, tmp_path = tempfile.mkstemp(
dir=str(CACHE_PATH.parent), suffix=".tmp"
dir=str(cache_path.parent), suffix=".tmp"
)
try:
with os.fdopen(fd, "w", encoding="utf-8") as f:
json.dump(cache, f, indent=2, ensure_ascii=False)
f.flush()
os.fsync(f.fileno())
os.replace(tmp_path, str(CACHE_PATH))
os.replace(tmp_path, str(cache_path))
except BaseException:
try:
os.unlink(tmp_path)
Expand Down
64 changes: 64 additions & 0 deletions tests/test_profile_isolation_runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -143,6 +143,70 @@ def test_store_path_follows_override(self, two_profiles, monkeypatch):
assert b_seen.endswith("state/rich_sent_index.json")


class TestCheckpointManagerPathResolution:
"""tools/checkpoint_manager.py's checkpoint store root must honor the
active profile — otherwise one profile's CheckpointManager instance can
read/write code-edit checkpoints into a different profile's store under
the multiplexed gateway."""

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Please add a two-profile checkpoint operation test in addition to this resolver assertion: persist under profile A, switch the same manager to B, then verify B reads/writes only B's store. The affected production sinks are _take() and list_checkpoints().

def test_checkpoint_base_follows_override(self, two_profiles):
prof_a, prof_b = two_profiles
import tools.checkpoint_manager as cm

a_seen = _under_override(prof_a, lambda: cm._resolve_checkpoint_base())
b_seen = _under_override(prof_b, lambda: cm._resolve_checkpoint_base())

assert a_seen == prof_a / "checkpoints"
assert b_seen == prof_b / "checkpoints"
assert a_seen != b_seen

def test_store_path_follows_override(self, two_profiles):
prof_a, prof_b = two_profiles
import tools.checkpoint_manager as cm

b_seen = _under_override(prof_b, lambda: cm._store_path())
assert b_seen == prof_b / "checkpoints" / "store"

def test_monkeypatched_constant_still_wins(self, two_profiles, monkeypatch, tmp_path):
"""The existing test seam (monkeypatch the module constant, see
tests/tools/test_checkpoint_manager.py) is preserved."""
_prof_a, prof_b = two_profiles
import tools.checkpoint_manager as cm

forced = tmp_path / "forced_checkpoints"
monkeypatch.setattr("tools.checkpoint_manager.CHECKPOINT_BASE", forced)
seen = _under_override(prof_b, lambda: cm._resolve_checkpoint_base())
assert seen == forced


class TestStickerCachePathResolution:
"""gateway/sticker_cache.py's cache file must honor the active profile —
otherwise one profile's Telegram sticker-description cache leaks into a
different profile's under the multiplexed gateway."""

def test_cache_path_follows_override(self, two_profiles):
prof_a, prof_b = two_profiles
import gateway.sticker_cache as sc

a_seen = _under_override(prof_a, lambda: sc._resolve_cache_path())
b_seen = _under_override(prof_b, lambda: sc._resolve_cache_path())

assert a_seen == prof_a / "sticker_cache.json"
assert b_seen == prof_b / "sticker_cache.json"
assert a_seen != b_seen

def test_monkeypatched_constant_still_wins(self, two_profiles, monkeypatch, tmp_path):
"""The existing test seam (monkeypatch the module constant, see
tests/gateway/test_sticker_cache.py) is preserved."""
_prof_a, prof_b = two_profiles
import gateway.sticker_cache as sc

forced = tmp_path / "forced_sticker_cache.json"
monkeypatch.setattr("gateway.sticker_cache.CACHE_PATH", forced)
seen = _under_override(prof_b, lambda: sc._resolve_cache_path())
assert seen == forced


# ---------------------------------------------------------------------------
# M2 — thread / executor context propagation
# ---------------------------------------------------------------------------
Expand Down
47 changes: 37 additions & 10 deletions tools/checkpoint_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,33 @@
# ---------------------------------------------------------------------------

CHECKPOINT_BASE = get_hermes_home() / "checkpoints"
# Import-time default, used by ``_resolve_checkpoint_base`` below to detect a
# test monkeypatch of the constant (test seam — many call sites in
# tests/tools/test_checkpoint_manager.py do
# ``monkeypatch.setattr("tools.checkpoint_manager.CHECKPOINT_BASE", ...)``) vs.
# an unmodified import-time value (in which case it re-resolves through the
# active profile override).
_CHECKPOINT_BASE_IMPORT_DEFAULT = CHECKPOINT_BASE


def _resolve_checkpoint_base() -> Path:
"""Resolve the checkpoint store root, honoring the active profile.

``CHECKPOINT_BASE`` is frozen at import time, which pins every later
checkpoint read/write to whichever profile's ``HERMES_HOME`` was active
when this module was first imported — a cross-profile leak under the
multiplexed gateway, where multiple profiles (each owning their own
``CheckpointManager`` instance per ``AIAgent``) share one process. The
manager's methods call this per-invocation instead of reading the module
constant directly, mirroring the cache-dir profile-isolation fix.

A test that monkeypatches the module constant away from its import-time
default is respected (test seam preserved).
"""
current = CHECKPOINT_BASE
if current != _CHECKPOINT_BASE_IMPORT_DEFAULT:
return current
return get_hermes_home() / "checkpoints"

# Single shared store directory under CHECKPOINT_BASE.
_STORE_DIRNAME = "store"
Expand Down Expand Up @@ -206,7 +233,7 @@ def _project_hash(working_dir: str) -> str:

def _store_path(base: Optional[Path] = None) -> Path:
"""Return the single shared shadow store path."""
return (base or CHECKPOINT_BASE) / _STORE_DIRNAME
return (base or _resolve_checkpoint_base()) / _STORE_DIRNAME


def _shadow_repo_path(working_dir: str) -> Path: # pragma: no cover — kept for BC
Expand Down Expand Up @@ -690,7 +717,7 @@ def ensure_checkpoint(self, working_dir: str, reason: str = "auto") -> bool:
def list_checkpoints(self, working_dir: str) -> List[Dict]:
"""List available checkpoints for a directory (most recent first)."""
abs_dir = str(_normalize_path(working_dir))
store = _store_path(CHECKPOINT_BASE)
store = _store_path()

if not (store / "HEAD").exists():
return []
Expand Down Expand Up @@ -748,7 +775,7 @@ def diff(self, working_dir: str, commit_hash: str) -> Dict:
return {"success": False, "error": hash_err}

abs_dir = str(_normalize_path(working_dir))
store = _store_path(CHECKPOINT_BASE)
store = _store_path()

if not (store / "HEAD").exists():
return {"success": False, "error": "No checkpoints exist for this directory"}
Expand Down Expand Up @@ -804,7 +831,7 @@ def restore(self, working_dir: str, commit_hash: str, file_path: str = None) ->
if path_err:
return {"success": False, "error": path_err}

store = _store_path(CHECKPOINT_BASE)
store = _store_path()

if not (store / "HEAD").exists():
return {"success": False, "error": "No checkpoints exist for this directory"}
Expand Down Expand Up @@ -872,7 +899,7 @@ def get_working_dir_for_path(self, file_path: str) -> str:

def _take(self, working_dir: str, reason: str) -> bool:
"""Take a snapshot. Returns True on success."""
store = _store_path(CHECKPOINT_BASE)
store = _store_path()

err = _init_store(store, working_dir)
if err:
Expand Down Expand Up @@ -1281,7 +1308,7 @@ def prune_checkpoints(

Never raises — maintenance must never block interactive startup.
"""
base = checkpoint_base or CHECKPOINT_BASE
base = checkpoint_base or _resolve_checkpoint_base()
result = {
"scanned": 0,
"deleted_orphan": 0,
Expand Down Expand Up @@ -1511,7 +1538,7 @@ def maybe_auto_prune_checkpoints(
Returns ``{"skipped": bool, "result": prune_checkpoints-dict,
"error": optional str}``.
"""
base = checkpoint_base or CHECKPOINT_BASE
base = checkpoint_base or _resolve_checkpoint_base()
out: Dict[str, object] = {"skipped": False}

try:
Expand Down Expand Up @@ -1574,7 +1601,7 @@ def store_status(checkpoint_base: Optional[Path] = None) -> Dict:
"total_size_bytes": N, "project_count": N, "projects": [...],
"legacy_archives": [...]}``
"""
base = checkpoint_base or CHECKPOINT_BASE
base = checkpoint_base or _resolve_checkpoint_base()
out: Dict = {
"base": str(base),
"store_size_bytes": 0,
Expand Down Expand Up @@ -1639,7 +1666,7 @@ def clear_all(checkpoint_base: Optional[Path] = None) -> Dict[str, int]:

Returns ``{"bytes_freed": N, "deleted": bool}``.
"""
base = checkpoint_base or CHECKPOINT_BASE
base = checkpoint_base or _resolve_checkpoint_base()
out = {"bytes_freed": 0, "deleted": False}
if not base.exists():
return out
Expand All @@ -1658,7 +1685,7 @@ def clear_legacy(checkpoint_base: Optional[Path] = None) -> Dict[str, int]:

Returns ``{"bytes_freed": N, "deleted": count}``.
"""
base = checkpoint_base or CHECKPOINT_BASE
base = checkpoint_base or _resolve_checkpoint_base()
out = {"bytes_freed": 0, "deleted": 0}
if not base.exists():
return out
Expand Down
Loading