diff --git a/gateway/control_socket.py b/gateway/control_socket.py index b45608a9dd..1845936afc 100644 --- a/gateway/control_socket.py +++ b/gateway/control_socket.py @@ -12,6 +12,7 @@ from __future__ import annotations import asyncio import contextlib import hashlib +import inspect import json import logging import os @@ -124,7 +125,7 @@ class GatewayControlServer: because its control socket couldn't bind; consumers fall back to the scan layer.""" def __init__(self, home: Optional[Path] = None, *, - verb_handlers: Optional[dict[str, Callable[[], dict[str, Any]]]] = None) -> None: + verb_handlers: Optional[dict[str, Callable[..., dict[str, Any]]]] = None) -> None: if home is None: from gateway.status import _get_process_hermes_home home = _get_process_hermes_home() @@ -133,7 +134,7 @@ class GatewayControlServer: self._pipe_server: Any = None # Windows proactor pipe server self._bind_path: Optional[Path] = None self._pointer_file: Optional[Path] = None - self._handlers: dict[str, Callable[[], dict[str, Any]]] = { + self._handlers: dict[str, Callable[..., dict[str, Any]]] = { "identify": build_identify_payload, "status": build_status_payload, **(verb_handlers or {})} async def start(self) -> bool: @@ -210,7 +211,12 @@ class GatewayControlServer: response: dict[str, Any] = {"ok": False, "error": f"unknown verb: {verb!r}", "protocol": CONTROL_PROTOCOL_VERSION, "supported_verbs": sorted(self._handlers)} else: - response = {"ok": True, "protocol": CONTROL_PROTOCOL_VERSION, "result": handler()} + # Verbs that carry arguments (e.g. migrate-profile-identity) declare a ``params`` + # parameter; argument-less verbs (identify/status/rescan) keep their bare signature. + params = request.get("params") if isinstance(request.get("params"), dict) else {} + wants_params = "params" in inspect.signature(handler).parameters + response = {"ok": True, "protocol": CONTROL_PROTOCOL_VERSION, + "result": handler(params) if wants_params else handler()} except Exception as exc: response = {"ok": False, "error": f"{type(exc).__name__}: {exc}", "protocol": CONTROL_PROTOCOL_VERSION} if request_id is not None: @@ -264,11 +270,15 @@ class _PipeControlProtocol(asyncio.Protocol): self._transport.close() -def query_gateway_control(home: Path, verb: str, *, timeout: float = _DEFAULT_CLIENT_TIMEOUT) -> Optional[dict[str, Any]]: +def query_gateway_control(home: Path, verb: str, *, params: Optional[dict[str, Any]] = None, + timeout: float = _DEFAULT_CLIENT_TIMEOUT) -> Optional[dict[str, Any]]: """Ask the gateway serving ``home`` a control verb; returns its ``result`` payload. Any failure (no/stale socket, timeout, malformed answer, ``ok: false``) returns None so callers fall back to the scan layer. - Never raises.""" - request = json.dumps({"verb": verb, "id": 1, "protocol": CONTROL_PROTOCOL_VERSION}).encode("utf-8") + b"\n" + ``params`` carries verb arguments (e.g. ``{"old": ..., "new": ...}``). Never raises.""" + payload: dict[str, Any] = {"verb": verb, "id": 1, "protocol": CONTROL_PROTOCOL_VERSION} + if params: + payload["params"] = params + request = json.dumps(payload).encode("utf-8") + b"\n" query = _query_windows_pipe if _IS_WINDOWS else _query_unix_socket try: raw = query(Path(home), request, timeout) @@ -348,3 +358,13 @@ def rescan_gateway_profiles(home: Path, *, timeout: float = 8.0) -> Optional[dic when no gateway answers / the gateway predates the verb — callers then rely on the periodic rescan (or the restart reminder).""" return query_gateway_control(home, "rescan-profiles", timeout=timeout) + + +def migrate_gateway_profile_identity(home: Path, old_name: str, new_name: str, *, + timeout: float = 8.0) -> Optional[dict[str, Any]]: + """Ask the multiplexer serving ``home`` to rekey a renamed profile's in-memory + on-disk routing + from ``agent::`` to ``agent::`` now. Returns its ``{"rekeyed": N, ...}`` answer, or None + when no gateway answers / the gateway predates the verb — the CLI's durable DB rewrite still lands, + and a restart reconciles the in-memory copy.""" + return query_gateway_control(home, "migrate-profile-identity", + params={"old": old_name, "new": new_name}, timeout=timeout) diff --git a/gateway/run.py b/gateway/run.py index 339079fb83..26a3a67f59 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -5116,9 +5116,43 @@ async def _start_gateway_start_control_socket(runner): except concurrent.futures.TimeoutError: return {"multiplex": True, "pending": True, "served_profiles": runner.served_profile_names()} + def _migrate_profile_identity_handler(params: dict) -> dict: + """Migrate both durable stores and the routing index owned by this live gateway.""" + old, new = str(params.get("old") or "").strip(), str(params.get("new") or "").strip() + if not old or not new or old == new: + return {"ok": False, "error": "old/new required and must differ"} + store = getattr(runner, "session_store", None) + if store is None: + return {"ok": False, "error": "live gateway has no session store"} + acquired = [] + try: + from hermes_state_registry import acquire, release_or_close + db_counts: dict[str, dict[str, int]] = {} + routing_db = getattr(store, "_routing_db", None) + if routing_db is not None and hasattr(routing_db, "rekey_profile_state"): + db_counts["routing"] = routing_db.rekey_profile_state(old, new) + routing_home = getattr(store, "_routing_home", None) + profile_path = Path(routing_home) / "profiles" / new / "state.db" if routing_home else None + if profile_path is not None and profile_path.exists(): + profile_db = acquire(profile_path) + acquired.append(profile_db) + db_counts["profile"] = profile_db.rekey_profile_state(old, new) + rekeyed = store.rekey_profile_routing(old, new) + return {"ok": True, "rekeyed": rekeyed, "db": db_counts} + except Exception as exc: + logger.warning("Profile identity migration failed for %r->%r: %s", old, new, exc) + return {"ok": False, "error": f"{type(exc).__name__}: {exc}"} + finally: + for db in acquired: + try: + release_or_close(db) + except Exception: + logger.debug("Failed to release renamed profile state DB", exc_info=True) + _control_server = GatewayControlServer( verb_handlers={"pause-for-update": _pause_for_update_handler, - "rescan-profiles": _rescan_profiles_handler}) + "rescan-profiles": _rescan_profiles_handler, + "migrate-profile-identity": _migrate_profile_identity_handler}) if not await _control_server.start(): _control_server = None else: diff --git a/gateway/session.py b/gateway/session.py index fb074d0cc6..59cda21e99 100644 --- a/gateway/session.py +++ b/gateway/session.py @@ -1121,6 +1121,33 @@ class SessionStore( self._save() return new_entry + def rekey_profile_routing(self, old_name: str, new_name: str) -> int: + """Rekey the live routing index and reject target collisions before mutation.""" + from dataclasses import replace as _dc_replace + old, new = (old_name or "").strip(), (new_name or "").strip() + if not old or not new or old == new: + return 0 + old_ns, new_ns = f"agent:{old}:", f"agent:{new}:" + with self._lock: + moving = [key for key in self._entries if key.startswith(old_ns)] + collisions = [ + new_ns + key[len(old_ns):] for key in moving + if new_ns + key[len(old_ns):] in self._entries] + if collisions: + raise ValueError( + f"profile routing collision while renaming {old!r} to {new!r}: " + f"{collisions[0]!r} already exists") + for key in moving: + new_key = new_ns + key[len(old_ns):] + entry = self._entries.pop(key) + origin = entry.origin + if origin is not None and getattr(origin, "profile", None) == old: + origin = _dc_replace(origin, profile=new) + self._entries[new_key] = _dc_replace(entry, session_key=new_key, origin=origin) + if moving: + self._save() + return len(moving) + # Compression repoint is store bookkeeping, not user activity — leave ``updated_at`` alone so a # background compression on an idle session cannot make it look fresh to the # restart-resume freshness gate (#85709). diff --git a/hermes_cli/profiles.py b/hermes_cli/profiles.py index aa7cf4feaf..3cf198d06e 100644 --- a/hermes_cli/profiles.py +++ b/hermes_cli/profiles.py @@ -1864,12 +1864,58 @@ def rename_profile(old_name: str, new_name: str) -> Path: # 5. Update active_profile if it pointed to old name _retarget_active_profile(old_canon, new_canon, f"✓ Active profile updated: {new_canon}") - # 6. Hot-serve the renamed profile now (mirrors create; a missed signal only delays it). + # 6. Migrate profile-name-keyed session/routing state (session keys, profile_name, heartbeats, + # delivery + routing index) from the old name to the new one. A stale ``agent::*`` routing + # key otherwise resolves to a profile that no longer exists on every inbound event. + _migrate_profile_identity(old_canon, new_canon, live_mux) + + # 7. Hot-serve the renamed profile now (mirrors create; a missed signal only delays it). if live_mux: _notify_multiplexer(new_canon) return new_dir +def _migrate_profile_identity(old_canon: str, new_canon: str, live_mux: bool) -> None: + """Rekey renamed-profile identity without racing a live gateway's in-memory routing index.""" + if live_mux: + try: + from hermes_constants import get_default_hermes_root + from gateway.control_socket import migrate_gateway_profile_identity + answer = migrate_gateway_profile_identity( + get_default_hermes_root(), old_canon, new_canon) + except Exception as exc: + answer, failure = None, f"{type(exc).__name__}: {exc}" + else: + failure = answer.get("error") if isinstance(answer, dict) else None + if isinstance(answer, dict) and answer.get("ok") is True: + return + detail = f" ({failure})" if failure else "" + print( + "⚠ Profile was renamed, but the live gateway could not migrate session identity" + f"{detail}. Restart the gateway, then retry the identity migration.", + file=sys.stderr) + return + + from hermes_state_registry import acquire, release_or_close + from hermes_constants import get_default_hermes_root + root = get_default_hermes_root() + for db_path in (root / "state.db", get_profile_dir(new_canon) / "state.db"): + if not db_path.exists(): + continue + db = None + try: + db = acquire(db_path) + db.rekey_profile_state(old_canon, new_canon) + except Exception as exc: + click.echo( + f"⚠ Profile was renamed, but identity migration failed for {db_path}: " + f"{type(exc).__name__}: {exc}", err=True) + finally: + if db is not None: + with contextlib.suppress(Exception): + release_or_close(db) + + # Profile env resolution (called from _apply_profile_override) def resolve_profile_env(profile_name: str) -> str: diff --git a/hermes_state_gateway.py b/hermes_state_gateway.py index 29905eddd6..e12600b52f 100644 --- a/hermes_state_gateway.py +++ b/hermes_state_gateway.py @@ -546,6 +546,113 @@ class SessionGatewayMixin: return self._write_sql("DELETE FROM gateway_hygiene_state WHERE session_key = ?", (session_key,)) + def rekey_profile_state(self, old_name: str, new_name: str) -> Dict[str, int]: + """Atomically rewrite exact profile identity in this state database.""" + old, new = (old_name or "").strip(), (new_name or "").strip() + counts: Dict[str, int] = {} + if not old or not new or old == new: + return counts + old_ns, new_ns = f"agent:{old}:", f"agent:{new}:" + ns_len = len(old_ns) + + def _do(conn): + existing = {row[0] for row in conn.execute( + "SELECT name FROM sqlite_master WHERE type='table'").fetchall()} + collision = conn.execute( + "SELECT old.scope, ? || substr(old.session_key, ?) " + "FROM gateway_routing AS old JOIN gateway_routing AS target " + "ON target.scope = old.scope " + "AND target.session_key = ? || substr(old.session_key, ?) " + "WHERE substr(old.session_key, 1, ?) = ? LIMIT 1", + (new_ns, ns_len + 1, new_ns, ns_len + 1, ns_len, old_ns), + ).fetchone() + if collision is not None: + raise ValueError( + f"profile routing collision in scope {collision[0]!r}: {collision[1]!r}") + for table, columns in ( + ("telegram_dm_topic_mode", ("chat_id",)), + ("telegram_dm_topic_bindings", ("chat_id", "thread_id")), + ): + if table not in existing: + continue + equality = " AND ".join( + f"target.{column} = old.{column}" for column in columns) + collision = conn.execute( + f"SELECT 1 FROM {table} AS old JOIN {table} AS target " + f"ON target.profile_name = ? AND {equality} " + "WHERE old.profile_name = ? LIMIT 1", (new, old)).fetchone() + if collision is not None: + raise ValueError(f"profile identity collision in {table}") + + counts["sessions_profile_name"] = conn.execute( + "UPDATE sessions SET profile_name = ? WHERE profile_name = ?", (new, old)).rowcount + counts["gateway_heartbeats_profile"] = conn.execute( + "UPDATE gateway_heartbeats SET profile = ? WHERE profile = ?", (new, old)).rowcount + counts["sessions_session_key"] = conn.execute( + "UPDATE sessions SET session_key = ? || substr(session_key, ?) " + "WHERE substr(session_key, 1, ?) = ?", + (new_ns, ns_len + 1, ns_len, old_ns)).rowcount + + origin_count = 0 + for session_id, origin_json in conn.execute( + "SELECT id, origin_json FROM sessions WHERE origin_json IS NOT NULL").fetchall(): + try: + payload = json.loads(origin_json) + except (ValueError, TypeError): + continue + if isinstance(payload, dict) and payload.get("profile") == old: + payload["profile"] = new + conn.execute("UPDATE sessions SET origin_json = ? WHERE id = ?", + (json.dumps(payload, ensure_ascii=False), session_id)) + origin_count += 1 + counts["sessions_origin_json"] = origin_count + + if "delivery_obligations" in existing: + counts["delivery_obligations_adapter_profile"] = conn.execute( + "UPDATE delivery_obligations SET adapter_profile = ? WHERE adapter_profile = ?", + (new, old)).rowcount + counts["delivery_obligations_session_key"] = conn.execute( + "UPDATE delivery_obligations SET session_key = ? || substr(session_key, ?) " + "WHERE substr(session_key, 1, ?) = ?", + (new_ns, ns_len + 1, ns_len, old_ns)).rowcount + for table in ("telegram_dm_topic_mode", "telegram_dm_topic_bindings"): + if table in existing: + counts[f"{table}_profile_name"] = conn.execute( + f"UPDATE {table} SET profile_name = ? WHERE profile_name = ?", + (new, old)).rowcount + if "telegram_dm_topic_bindings" in existing: + counts["telegram_dm_topic_bindings_session_key"] = conn.execute( + "UPDATE telegram_dm_topic_bindings " + "SET session_key = ? || substr(session_key, ?) " + "WHERE substr(session_key, 1, ?) = ?", + (new_ns, ns_len + 1, ns_len, old_ns)).rowcount + + routing = conn.execute( + "SELECT rowid, session_key, entry_json FROM gateway_routing " + "WHERE substr(session_key, 1, ?) = ?", (ns_len, old_ns)).fetchall() + for rowid, session_key, entry_json in routing: + new_session_key = new_ns + session_key[ns_len:] + new_json = entry_json + if entry_json: + try: + payload = json.loads(entry_json) + except (ValueError, TypeError): + payload = None + if isinstance(payload, dict): + if isinstance(payload.get("session_key"), str) and payload["session_key"].startswith(old_ns): + payload["session_key"] = new_ns + payload["session_key"][ns_len:] + origin = payload.get("origin") + if isinstance(origin, dict) and origin.get("profile") == old: + origin["profile"] = new + new_json = json.dumps(payload, ensure_ascii=False) + conn.execute( + "UPDATE gateway_routing SET session_key = ?, entry_json = ? WHERE rowid = ?", + (new_session_key, new_json, rowid)) + counts["gateway_routing"] = len(routing) + + self._execute_write(_do) + return counts + @staticmethod def session_gateway_runtime(session_meta: Optional[Dict[str, Any]]) -> Dict[str, Any]: """Read the persisted runtime route off a session row dict (``model_config`` as diff --git a/tests/gateway/test_control_socket.py b/tests/gateway/test_control_socket.py index b4034ebb0c..825f73b4f3 100644 --- a/tests/gateway/test_control_socket.py +++ b/tests/gateway/test_control_socket.py @@ -161,6 +161,39 @@ def test_unknown_verb_and_malformed_request(home: Path): assert payload["protocol"] == CONTROL_PROTOCOL_VERSION +def test_verb_handler_receives_params(home: Path): + """A handler declaring a ``params`` argument is called with the request's params dict; a bare + handler is still called with no args (backward compat for identify/status/rescan).""" + received = {} + + def with_params(params): + received.update(params) + return {"echo": params} + + def bare(): + return {"ok": 1} + + async def scenario(): + server = GatewayControlServer( + home, verb_handlers={"with-params": with_params, "bare": bare}) + assert await server.start() + try: + loop = asyncio.get_running_loop() + got = await loop.run_in_executor( + None, lambda: query_gateway_control( + home, "with-params", params={"old": "a", "new": "b"})) + bare_ok = await loop.run_in_executor( + None, lambda: query_gateway_control(home, "bare")) + return got, bare_ok + finally: + await server.stop() + + got, bare_ok = _run(scenario()) + assert got == {"echo": {"old": "a", "new": "b"}} + assert received == {"old": "a", "new": "b"} + assert bare_ok == {"ok": 1} + + def test_stop_removes_socket_and_pointer(home: Path): async def scenario(): server = GatewayControlServer( diff --git a/tests/gateway/test_rekey_profile_routing.py b/tests/gateway/test_rekey_profile_routing.py new file mode 100644 index 0000000000..daef3d6d97 --- /dev/null +++ b/tests/gateway/test_rekey_profile_routing.py @@ -0,0 +1,77 @@ +"""In-memory routing rekey for `hermes profile rename`. + +The routing index lives in ``SessionStore._entries`` and is written back periodically, so a durable +DB rewrite alone is clobbered — the live store must rekey its in-memory copy too. This is why a +renamed profile's old namespace kept resurfacing until the gateway restarted. +""" +from __future__ import annotations + + +def _make_store(tmp_path): + from gateway.config import GatewayConfig + from gateway.session import SessionStore + sessions_dir = tmp_path / "sessions" + sessions_dir.mkdir() + store = SessionStore( + sessions_dir, + GatewayConfig(sessions_dir=sessions_dir, write_sessions_json=False, + multiplex_profiles=True), + ) + store._ensure_loaded() + return store + + +def _entry(session_key, chat_id, profile): + from gateway.session import SessionEntry, SessionSource, Platform + from gateway.session_lifecycle import _now + now = _now() + return SessionEntry( + session_key=session_key, session_id=f"sid-{chat_id}", + platform=Platform.FEISHU, chat_type="dm", created_at=now, updated_at=now, + origin=SessionSource(platform=Platform.FEISHU, chat_id=chat_id, profile=profile), + ) + + +def test_rekeys_old_namespace_and_origin_profile(tmp_path): + store = _make_store(tmp_path) + with store._lock: + store._entries["agent:oldname:feishu:dm:chatA"] = _entry( + "agent:oldname:feishu:dm:chatA", "chatA", "oldname") + store._entries["agent:keepme:feishu:dm:chatB"] = _entry( + "agent:keepme:feishu:dm:chatB", "chatB", "keepme") + + moved = store.rekey_profile_routing("oldname", "newname") + assert moved == 1 + + assert "agent:oldname:feishu:dm:chatA" not in store._entries + new_entry = store._entries["agent:newname:feishu:dm:chatA"] + assert new_entry.session_key == "agent:newname:feishu:dm:chatA" + assert new_entry.origin.profile == "newname" + # Bystander namespace untouched. + assert store._entries["agent:keepme:feishu:dm:chatB"].origin.profile == "keepme" + + +def test_noop_for_equal_or_empty_names(tmp_path): + store = _make_store(tmp_path) + with store._lock: + store._entries["agent:oldname:feishu:dm:chatA"] = _entry( + "agent:oldname:feishu:dm:chatA", "chatA", "oldname") + assert store.rekey_profile_routing("x", "x") == 0 + assert store.rekey_profile_routing("", "y") == 0 + assert "agent:oldname:feishu:dm:chatA" in store._entries + + +def test_does_not_overwrite_existing_new_namespace_key(tmp_path): + store = _make_store(tmp_path) + with store._lock: + store._entries["agent:oldname:feishu:dm:chatA"] = _entry( + "agent:oldname:feishu:dm:chatA", "chatA", "oldname") + # A collision on the target key (should not happen in practice) is left alone. + store._entries["agent:newname:feishu:dm:chatA"] = _entry( + "agent:newname:feishu:dm:chatA", "chatA", "newname") + + import pytest + with pytest.raises(ValueError, match="routing collision"): + store.rekey_profile_routing("oldname", "newname") + assert "agent:oldname:feishu:dm:chatA" in store._entries + assert "agent:newname:feishu:dm:chatA" in store._entries diff --git a/tests/hermes_cli/test_profiles.py b/tests/hermes_cli/test_profiles.py index 010d32d776..b8d8203228 100644 --- a/tests/hermes_cli/test_profiles.py +++ b/tests/hermes_cli/test_profiles.py @@ -867,6 +867,77 @@ class TestRenameProfile: assert not (tmp_path / ".hermes" / "profiles" / ".deleted").exists() assert not old_dir.exists() and new_dir.is_dir() + def test_rename_migrates_session_identity_without_live_gateway(self, profile_env): + """No live gateway → the CLI performs the durable rekey itself so a renamed profile's session + keys / profile_name / routing rows follow the new name (else inbound events on the old name's + chats resolve to a nonexistent profile and flood errors.log).""" + from hermes_state import SessionDB + tmp_path = profile_env + create_profile("oldname", no_alias=True) + old_dir = tmp_path / ".hermes" / "profiles" / "oldname" + # Seed a session owned by the old profile in the profile's own store + the root routing index. + pdb = SessionDB(old_dir / "state.db") + pdb.create_session( + "sess1", "feishu", session_key="agent:oldname:feishu:dm:chatA", + profile_name="oldname", chat_id="chatA", chat_type="dm") + pdb.close() + root_db = SessionDB(tmp_path / ".hermes" / "state.db") + root_db.save_gateway_routing_entry( + "agent:oldname:feishu:dm:chatA", + json.dumps({"session_key": "agent:oldname:feishu:dm:chatA", "session_id": "sess1", + "origin": {"platform": "feishu", "chat_id": "chatA", "profile": "oldname"}}), + scope=str(tmp_path / ".hermes" / "sessions")) + root_db.close() + + with patch("hermes_cli.profiles.check_alias_collision", return_value="skip"), \ + patch("hermes_cli.profiles._live_default_multiplexer", return_value=False): + rename_profile("oldname", "newname") + + new_dir = tmp_path / ".hermes" / "profiles" / "newname" + moved_db = SessionDB(new_dir / "state.db") + row = moved_db._read_one( + "SELECT session_key, profile_name FROM sessions WHERE id = ?", ("sess1",)) + assert row["session_key"] == "agent:newname:feishu:dm:chatA" + assert row["profile_name"] == "newname" + moved_db.close() + root_db2 = SessionDB(tmp_path / ".hermes" / "state.db") + routing = root_db2.load_gateway_routing_entries( + scope=str(tmp_path / ".hermes" / "sessions")) + assert "agent:oldname:feishu:dm:chatA" not in routing + assert "agent:newname:feishu:dm:chatA" in routing + root_db2.close() + + def test_rename_delegates_identity_migration_to_live_gateway(self, profile_env): + """Under a live multiplexer the CLI must NOT rewrite the routing DB directly (the gateway holds + it in memory and would clobber the write); it delegates to the control verb instead.""" + tmp_path = profile_env + create_profile("oldname", no_alias=True) + + with patch("hermes_cli.profiles.check_alias_collision", return_value="skip"), \ + patch("hermes_cli.profiles._live_default_multiplexer", return_value=True), \ + patch("hermes_cli.profiles._notify_multiplexer"), \ + patch("gateway.control_socket.migrate_gateway_profile_identity", + return_value={"ok": True, "rekeyed": 1, "db": {}}) as verb, \ + patch("hermes_state_registry.acquire") as acquire: + rename_profile("oldname", "newname") + + # Delegated to the gateway; the CLI's own durable-rewrite branch never ran. + assert verb.call_count == 1 + assert verb.call_args.args[1:] == ("oldname", "newname") + acquire.assert_not_called() + + + def test_live_gateway_failure_does_not_rewrite_db_directly(self, profile_env, capsys): + create_profile("oldname", no_alias=True) + with patch("hermes_cli.profiles.check_alias_collision", return_value="skip"), \ + patch("hermes_cli.profiles._live_default_multiplexer", return_value=True), \ + patch("hermes_cli.profiles._notify_multiplexer"), \ + patch("gateway.control_socket.migrate_gateway_profile_identity", return_value=None), \ + patch("hermes_state_registry.acquire") as acquire: + rename_profile("oldname", "newname") + acquire.assert_not_called() + assert "Restart the gateway" in capsys.readouterr().err + # =================================================================== # TestExportImport diff --git a/tests/hermes_state/test_rekey_profile_state.py b/tests/hermes_state/test_rekey_profile_state.py new file mode 100644 index 0000000000..307338fbd2 --- /dev/null +++ b/tests/hermes_state/test_rekey_profile_state.py @@ -0,0 +1,157 @@ +"""Regression tests for profile-name-keyed state migration on `hermes profile rename`. + +When a profile is renamed, the directory move carries the row data, but the profile name is also +baked into session keys (``agent::*``), ``sessions.profile_name``, ``gateway_heartbeats.profile``, +``delivery_obligations`` and the ``gateway_routing`` index. Left stale, an inbound event on a chat keyed +to the old name resolves to a profile that no longer exists and floods errors.log. See the profile-rename +identity-migration fix. +""" +import json +import time + +import pytest + +from hermes_state import SessionDB + + +@pytest.fixture +def db(tmp_path): + database = SessionDB(tmp_path / "state.db") + yield database + database.close() + + +class TestRekeyProfileState: + def test_rekeys_session_key_namespace_and_profile_columns(self, db): + # A session owned by the old profile, keyed in its namespace. + db.create_session( + "sess_old", "feishu", session_key="agent:oldname:feishu:dm:chatA", + profile_name="oldname", chat_id="chatA", chat_type="dm", + ) + # An unrelated profile's session must be left untouched. + db.create_session( + "sess_other", "feishu", session_key="agent:keepme:feishu:dm:chatB", + profile_name="keepme", chat_id="chatB", chat_type="dm", + ) + + counts = db.rekey_profile_state("oldname", "newname") + + assert counts["sessions_session_key"] == 1 + assert counts["sessions_profile_name"] == 1 + # Renamed row now lives under the new namespace + owner. + row = db._read_one( + "SELECT session_key, profile_name FROM sessions WHERE id = ?", ("sess_old",)) + assert row["session_key"] == "agent:newname:feishu:dm:chatA" + assert row["profile_name"] == "newname" + # Bystander untouched. + other = db._read_one( + "SELECT session_key, profile_name FROM sessions WHERE id = ?", ("sess_other",)) + assert other["session_key"] == "agent:keepme:feishu:dm:chatB" + assert other["profile_name"] == "keepme" + + def test_rekeys_routing_index_key_and_embedded_profile(self, db): + entry = { + "session_key": "agent:oldname:feishu:dm:chatA", + "session_id": "sess_old", + "origin": {"platform": "feishu", "chat_id": "chatA", "profile": "oldname"}, + } + db.save_gateway_routing_entry( + "agent:oldname:feishu:dm:chatA", json.dumps(entry), scope="/root/sessions") + + counts = db.rekey_profile_state("oldname", "newname") + assert counts["gateway_routing"] == 1 + + rows = db.load_gateway_routing_entries(scope="/root/sessions") + assert "agent:oldname:feishu:dm:chatA" not in rows + assert "agent:newname:feishu:dm:chatA" in rows + payload = json.loads(rows["agent:newname:feishu:dm:chatA"]) + assert payload["session_key"] == "agent:newname:feishu:dm:chatA" + assert payload["origin"]["profile"] == "newname" + + def test_rekeys_heartbeats_and_delivery_obligations(self, db, tmp_path, monkeypatch): + db.register_backend_heartbeat( + backend_id="be1", pid=123, started_at=time.time(), + profile="oldname", host="h") + # delivery_obligations is created lazily by the delivery ledger against the same state.db. + monkeypatch.setenv("HERMES_HOME", str(db.db_path.parent)) + from gateway import delivery_ledger + monkeypatch.setattr(delivery_ledger, "_db_path", lambda: db.db_path) + with delivery_ledger._connect() as conn: + now = time.time() + conn.execute( + "INSERT INTO delivery_obligations (obligation_id, session_key, platform, chat_id, " + "content, state, created_at, updated_at, adapter_profile) " + "VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)", + ("ob1", "agent:oldname:feishu:dm:chatA", "feishu", "chatA", + "hi", "pending", now, now, "oldname"), + ) + conn.commit() + + counts = db.rekey_profile_state("oldname", "newname") + + assert counts["gateway_heartbeats_profile"] == 1 + assert counts["delivery_obligations_adapter_profile"] == 1 + assert counts["delivery_obligations_session_key"] == 1 + hb = db._read_one("SELECT profile FROM gateway_heartbeats WHERE backend_id = ?", ("be1",)) + assert hb["profile"] == "newname" + ob = db._read_one( + "SELECT session_key, adapter_profile FROM delivery_obligations WHERE obligation_id = ?", + ("ob1",)) + assert ob["session_key"] == "agent:newname:feishu:dm:chatA" + assert ob["adapter_profile"] == "newname" + + def test_noop_when_names_equal_or_empty(self, db): + assert db.rekey_profile_state("x", "x") == {} + assert db.rekey_profile_state("", "y") == {} + assert db.rekey_profile_state("x", "") == {} + + def test_idempotent(self, db): + db.create_session( + "sess_old", "feishu", session_key="agent:oldname:feishu:dm:chatA", + profile_name="oldname", chat_id="chatA", chat_type="dm", + ) + first = db.rekey_profile_state("oldname", "newname") + assert first["sessions_session_key"] == 1 + second = db.rekey_profile_state("oldname", "newname") + # Nothing left under the old name. + assert second["sessions_session_key"] == 0 + assert second["sessions_profile_name"] == 0 + + def test_underscore_in_profile_name_is_not_a_like_wildcard(self, db): + db.create_session( + "literal", "telegram", session_key="agent:foo_bar:telegram:dm:a", + profile_name="foo_bar", chat_id="a", chat_type="dm") + db.create_session( + "bystander", "telegram", session_key="agent:fooXbar:telegram:dm:b", + profile_name="fooXbar", chat_id="b", chat_type="dm") + db.rekey_profile_state("foo_bar", "renamed") + assert db._read_one( + "SELECT session_key FROM sessions WHERE id = ?", ("literal",) + )["session_key"] == "agent:renamed:telegram:dm:a" + assert db._read_one( + "SELECT session_key FROM sessions WHERE id = ?", ("bystander",) + )["session_key"] == "agent:fooXbar:telegram:dm:b" + + def test_rekeys_origin_json_and_telegram_topic_state(self, db): + db.create_session( + "sess_old", "telegram", session_key="agent:oldname:telegram:dm:chatA", + profile_name="oldname", chat_id="chatA", chat_type="dm") + db._write_sql( + "UPDATE sessions SET origin_json = ? WHERE id = ?", + (json.dumps({"platform": "telegram", "profile": "oldname"}), "sess_old")) + db.enable_telegram_topic_mode( + chat_id="chatA", user_id="userA", profile_name="oldname") + db.bind_telegram_topic( + chat_id="chatA", thread_id="threadA", user_id="userA", + session_key="agent:oldname:telegram:dm:chatA", session_id="sess_old", + profile_name="oldname") + counts = db.rekey_profile_state("oldname", "newname") + row = db._read_one("SELECT origin_json FROM sessions WHERE id = ?", ("sess_old",)) + assert json.loads(row["origin_json"])["profile"] == "newname" + assert counts["sessions_origin_json"] == 1 + binding = db._read_one( + "SELECT profile_name, session_key FROM telegram_dm_topic_bindings " + "WHERE chat_id = ? AND thread_id = ?", ("chatA", "threadA")) + assert binding["profile_name"] == "newname" + assert binding["session_key"] == "agent:newname:telegram:dm:chatA" +