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
51 changes: 51 additions & 0 deletions gateway/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -1126,6 +1126,45 @@ def _resolve_profile_for_key(self, source: Optional[SessionSource] = None) -> Op
except Exception:
return None

@staticmethod
def _profile_from_session_key(session_key: Optional[str]) -> Optional[str]:
"""Extract the profile namespace encoded in a gateway session key."""
if not session_key:
return None
parts = str(session_key).split(":")
if len(parts) < 2 or parts[0] != "agent":
return None
namespace = parts[1] or "main"
return "default" if namespace == "main" else namespace

@staticmethod
def _active_profile_name() -> str:
try:
from hermes_cli.profiles import get_active_profile_name
return get_active_profile_name() or "default"
except Exception:
return "default"

def _recovered_row_allowed_for_active_profile(
self,
*,
requested_session_key: str,
recovered: Dict[str, Any],
) -> bool:
"""Prevent non-multiplexed gateways from reviving another profile's row."""
if getattr(self.config, "multiplex_profiles", False):
return True

recovered_key = str(recovered.get("session_key") or "")
if not recovered_key or recovered_key == requested_session_key:
return True

recovered_profile = self._profile_from_session_key(recovered_key)
if recovered_profile is None:
return True

return recovered_profile == self._active_profile_name()

def _generate_session_key(self, source: SessionSource) -> str:
"""Generate a session key from a source."""
return build_session_key(
Expand Down Expand Up @@ -1186,6 +1225,18 @@ def _recover_session_from_db(
return None
if not recovered:
return None
if not self._recovered_row_allowed_for_active_profile(
requested_session_key=session_key,
recovered=recovered,
):
logger.warning(
"Gateway session DB recovery ignored %s for %s because "
"multiplex_profiles is disabled and the row belongs to a "
"different profile",
recovered.get("session_key"),
session_key,
)
return None
try:
self._db.reopen_session(str(recovered["id"]))
except Exception as exc:
Expand Down
63 changes: 63 additions & 0 deletions tests/gateway/test_multiplex_phase0.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
on.
"""
import pytest
from datetime import datetime
from unittest.mock import patch

from gateway.config import GatewayConfig, Platform
Expand Down Expand Up @@ -163,3 +164,65 @@ def test_flag_on_default_profile_stays_legacy(self, tmp_path):
assert store._generate_session_key(s) == "agent:main:telegram:dm:99"


class _RecoveringDB:
def __init__(self, row):
self.row = row
self.reopened = []

def find_latest_gateway_session_for_peer(self, **_kwargs):
return self.row

def reopen_session(self, session_id):
self.reopened.append(session_id)


class TestSessionStoreUnmultiplexedRecovery:
"""Turning multiplexing off must not recover another profile's session."""

def _store_with_row(self, tmp_path, row, **cfg_kw):
config = GatewayConfig(**cfg_kw)
with patch("gateway.session.SessionStore._ensure_loaded"):
store = SessionStore(sessions_dir=tmp_path, config=config)
store._db = _RecoveringDB(row)
store._loaded = True
return store

def test_flag_off_rejects_other_profile_peer_fallback(self, tmp_path):
row = {
"id": "sess-coder",
"started_at": 1700000000,
"session_key": "agent:coder:telegram:dm:99",
}
store = self._store_with_row(tmp_path, row)
source = _src(chat_id="99", chat_type="dm")

with patch("hermes_cli.profiles.get_active_profile_name", return_value="default"):
recovered = store._recover_session_from_db(
session_key="agent:main:telegram:dm:99",
source=source,
now=datetime.fromtimestamp(1700000001),
)

assert recovered is None
assert store._db.reopened == []

def test_flag_off_allows_active_profile_peer_fallback(self, tmp_path):
row = {
"id": "sess-coder",
"started_at": 1700000000,
"session_key": "agent:coder:telegram:dm:99",
}
store = self._store_with_row(tmp_path, row)
source = _src(chat_id="99", chat_type="dm")

with patch("hermes_cli.profiles.get_active_profile_name", return_value="coder"):
recovered = store._recover_session_from_db(
session_key="agent:main:telegram:dm:99",
source=source,
now=datetime.fromtimestamp(1700000001),
)

assert recovered is not None
assert recovered.session_id == "sess-coder"
assert recovered.session_key == "agent:main:telegram:dm:99"
assert store._db.reopened == ["sess-coder"]