From b8f85f18e860ca2d03d1bf0e68f421fcf1ef63ec Mon Sep 17 00:00:00 2001 From: kshitij <82637225+kshitijk4poor@users.noreply.github.com> Date: Fri, 14 Aug 2026 01:26:22 +0530 Subject: [PATCH] refactor: shared /model persist helper, canonical gateway_runtime reader, cross-surface route persistence, tests - Extract the two duplicated /model session-persist blocks into _persist_model_switch_to_session; persist the route BOTH nested (gateway_runtime, CLI reader) and top-level (TUI gateway's _stored_session_runtime_overrides reader) so a CLI switch also survives a desktop/TUI session.resume. - Add SessionDB.session_gateway_runtime as the canonical tolerant row-level route reader (session_yolo_enabled precedent); use it in _restore_session_model instead of hand-rolled JSON parsing. - Clear stale launch-time _explicit_api_key/_explicit_base_url when resume restores a different provider (same leak guard _apply_model_switch_result already has). - 12 new tests incl. a real-SessionDB round trip; mutation-checked. --- cli.py | 131 ++++++++---------- hermes_state.py | 33 +++++ tests/cli/test_resume_model_restore.py | 183 +++++++++++++++++++++++++ 3 files changed, 274 insertions(+), 73 deletions(-) create mode 100644 tests/cli/test_resume_model_restore.py diff --git a/cli.py b/cli.py index 078d934dbb..6b7b9e7a4e 100644 --- a/cli.py +++ b/cli.py @@ -7690,6 +7690,40 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): else: self._console_print(f"[dim]{_escape(msg)}[/dim]") + def _persist_model_switch_to_session(self, result) -> None: + """Persist a session-scoped /model switch to the session DB row. + + Writes the model column plus the runtime route so ``--resume`` + (CLI, reads ``gateway_runtime``) and ``session.resume`` (TUI/desktop, + reads top-level ``model_config`` keys via + ``_stored_session_runtime_overrides``) both restore the switched + provider instead of recombining the model with the ambient default + (#79536). Mirrors the gateway's ``update_session_model()`` call. + getattr: tests drive the switch paths with ``object.__new__`` stubs. + """ + db = getattr(self, "_session_db", None) + sid = getattr(self, "session_id", None) + if not db or not sid: + return + route = { + k: v + for k, v in { + "provider": result.target_provider, + "base_url": result.base_url, + "api_mode": result.api_mode, + }.items() + if v + } + try: + db.update_session_model(sid, result.new_model) + # Both shapes: nested for the CLI reader, top-level for the + # TUI gateway's resume path. + db.patch_session_model_config(sid, {"gateway_runtime": route, **route}) + except Exception: + logger.debug( + "Failed to persist model switch to session DB", exc_info=True + ) + def _restore_session_model(self, session_meta: dict, *, quiet: bool = False) -> None: """Restore model/provider from the session DB row on resume. @@ -7720,20 +7754,14 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): # An explicit -m / --model on the command line overrides resume. if getattr(self, "_explicit_model_override", False): return - # Stored provider/endpoint from model_config.gateway_runtime - # (written by gateway turns and CLI /model switches alike). - stored_provider = stored_base_url = stored_api_mode = None - try: - import json as _json - raw_config = (session_meta or {}).get("model_config") - config = _json.loads(raw_config) if isinstance(raw_config, str) and raw_config else (raw_config if isinstance(raw_config, dict) else {}) - runtime = config.get("gateway_runtime") or {} if isinstance(config, dict) else {} - if isinstance(runtime, dict): - stored_provider = runtime.get("provider") or None - stored_base_url = runtime.get("base_url") or None - stored_api_mode = runtime.get("api_mode") or None - except Exception: - pass + # Stored provider/endpoint via the canonical row-level reader + # (prefers model_config.gateway_runtime, falls back to the TUI + # gateway's top-level keys). + from hermes_state import SessionDB as _SessionDB + _stored_runtime = _SessionDB.session_gateway_runtime(session_meta) + stored_provider = _stored_runtime.get("provider") or None + stored_base_url = _stored_runtime.get("base_url") or None + stored_api_mode = _stored_runtime.get("api_mode") or None model_changed = stored_model != self.model provider_changed = bool(stored_provider) and stored_provider != self.provider if not model_changed and not provider_changed: @@ -7747,6 +7775,13 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): if stored_api_mode: self.api_mode = stored_api_mode if provider_changed: + # Stale launch-time explicit overrides belong to the AMBIENT + # provider; carrying them into the restored provider's + # resolution poisons _ensure_runtime_credentials on startup + # resume (same leak _apply_model_switch_result guards against + # by overwriting _explicit_* on every switch). + self._explicit_api_key = None + self._explicit_base_url = stored_base_url # Re-resolve credentials for the restored provider. api_key is # never persisted to the session DB (by design) — the normal # runtime provider resolution owns credentials. @@ -9663,35 +9698,10 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): else: _cprint(" (session only — add --global to persist)") - # Persist the model change to the session DB row so --resume - # restores the correct model instead of falling back to the - # config default. Skipped for --global (config.yaml is the source - # of truth). Mirrors the gateway's update_session_model() call - # after a /model switch. The provider/endpoint go into - # model_config.gateway_runtime — same key the gateway writes and - # _restore_session_model reads; persisting only the model name - # recombines it with the ambient provider on resume (#79536). - # getattr: tests build bare stubs via object.__new__ without - # __init__ attributes. - _picker_session_db = getattr(self, "_session_db", None) - _picker_session_id = getattr(self, "session_id", None) - if not persist_global and _picker_session_db and _picker_session_id: - try: - _picker_session_db.update_session_model( - _picker_session_id, result.new_model - ) - _picker_session_db.patch_session_model_config( - _picker_session_id, - {"gateway_runtime": {k: v for k, v in { - "provider": result.target_provider, - "base_url": result.base_url, - "api_mode": result.api_mode, - }.items() if v}}, - ) - except Exception: - logger.debug( - "Failed to persist model switch to session DB", exc_info=True - ) + # Persist session-scoped switches so --resume / session.resume + # restore them; --global switches live in config.yaml instead. + if not persist_global: + HermesCLI._persist_model_switch_to_session(self, result) def _handle_model_picker_selection(self, persist_global: bool = False) -> None: state = self._model_picker_state @@ -10044,36 +10054,11 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): else: _cprint(" (session only — add --global to persist)") - # Persist the model change to the session DB row so --resume - # restores the correct model instead of falling back to the - # config default. Skipped for --global (config.yaml is the source - # of truth) and --once (ephemeral, restored after one turn). - # Mirrors the gateway's update_session_model() call after a - # /model switch. The provider/endpoint go into - # model_config.gateway_runtime — same key the gateway writes and - # _restore_session_model reads; persisting only the model name - # recombines it with the ambient provider on resume (#79536). - # getattr: tests build bare stubs via object.__new__ without - # __init__ attributes. - _switch_session_db = getattr(self, "_session_db", None) - _switch_session_id = getattr(self, "session_id", None) - if not persist_global and not one_turn and _switch_session_db and _switch_session_id: - try: - _switch_session_db.update_session_model( - _switch_session_id, result.new_model - ) - _switch_session_db.patch_session_model_config( - _switch_session_id, - {"gateway_runtime": {k: v for k, v in { - "provider": result.target_provider, - "base_url": result.base_url, - "api_mode": result.api_mode, - }.items() if v}}, - ) - except Exception: - logger.debug( - "Failed to persist model switch to session DB", exc_info=True - ) + # Persist session-scoped switches so --resume / session.resume + # restore them; --global lives in config.yaml, --once is ephemeral + # (restored after one turn). + if not persist_global and not one_turn: + HermesCLI._persist_model_switch_to_session(self, result) def _handle_codex_runtime(self, cmd_original: str) -> None: """Handle /codex-runtime — toggle the codex app-server runtime opt-in. diff --git a/hermes_state.py b/hermes_state.py index 800a7828e1..4fe015d4f3 100644 --- a/hermes_state.py +++ b/hermes_state.py @@ -6109,6 +6109,39 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) return False return bool(raw.get("yolo_mode")) + @staticmethod + def session_gateway_runtime(session_meta: Optional[Dict[str, Any]]) -> Dict[str, Any]: + """Read the persisted runtime route off a session row dict. + + Accepts the dict returned by ``get_session`` (``model_config`` is a + JSON string) or an already-parsed dict. Prefers the nested + ``gateway_runtime`` key (written by the gateway's + ``_sync_session_model_from_agent`` and the CLI ``/model`` persist), + falling back to the top-level ``provider``/``base_url``/``api_mode`` + keys the TUI gateway's ``_runtime_model_config`` writes. Returns an + empty dict on any parse failure — resume falls back to ambient + config resolution. + """ + raw = (session_meta or {}).get("model_config") + if isinstance(raw, str): + try: + raw = json.loads(raw) + except Exception: + return {} + if not isinstance(raw, dict): + return {} + runtime = raw.get("gateway_runtime") + if isinstance(runtime, dict) and runtime.get("provider"): + return dict(runtime) + top_level = { + key: raw.get(key) + for key in ("provider", "base_url", "api_mode") + if raw.get(key) + } + if top_level: + return top_level + return dict(runtime) if isinstance(runtime, dict) else {} + def update_session_billing_route( self, session_id: str, diff --git a/tests/cli/test_resume_model_restore.py b/tests/cli/test_resume_model_restore.py new file mode 100644 index 0000000000..dae3ea6afe --- /dev/null +++ b/tests/cli/test_resume_model_restore.py @@ -0,0 +1,183 @@ +"""Tests for CLI resume model restoration and /model session persistence. + +Covers _restore_session_model, _persist_model_switch_to_session (cli.py) and +SessionDB.session_gateway_runtime (hermes_state.py) — the round trip that +makes `hermes --resume` reopen a session on the model/provider it actually +used instead of the ambient config default (#57588-class, #79536). +""" + +import json + +import pytest + +import cli as cli_mod +from hermes_state import SessionDB + + +def _make_stub(**overrides): + """Bare HermesCLI the way resume paths see it (no __init__).""" + stub = object.__new__(cli_mod.HermesCLI) + stub.model = "ambient-model" + stub.provider = "openrouter" + stub.requested_provider = "openrouter" + stub.base_url = "https://openrouter.ai/api/v1" + stub.api_key = "ambient-key" + stub.api_mode = "" + stub.agent = None + stub._console_print = lambda s: None + for key, value in overrides.items(): + setattr(stub, key, value) + return stub + + +def _row(model="glm-4.7", model_config=None): + return { + "model": model, + "model_config": json.dumps(model_config) if model_config else None, + } + + +# ── SessionDB.session_gateway_runtime ─────────────────────────────── + + +def test_session_gateway_runtime_prefers_nested_key(): + meta = _row(model_config={ + "gateway_runtime": {"provider": "custom:feather", "base_url": "https://f/v1"}, + "provider": "openrouter", + }) + runtime = SessionDB.session_gateway_runtime(meta) + assert runtime["provider"] == "custom:feather" + assert runtime["base_url"] == "https://f/v1" + + +def test_session_gateway_runtime_falls_back_to_top_level_keys(): + # The TUI gateway's _runtime_model_config writes top-level keys only. + meta = _row(model_config={"provider": "nous", "api_mode": "chat_completions"}) + runtime = SessionDB.session_gateway_runtime(meta) + assert runtime == {"provider": "nous", "api_mode": "chat_completions"} + + +def test_session_gateway_runtime_tolerates_garbage(): + assert SessionDB.session_gateway_runtime(None) == {} + assert SessionDB.session_gateway_runtime({}) == {} + assert SessionDB.session_gateway_runtime({"model_config": "{not json"}) == {} + assert SessionDB.session_gateway_runtime({"model_config": json.dumps([1, 2])}) == {} + + +# ── _restore_session_model ────────────────────────────────────────── + + +def test_restore_session_model_restores_model_and_provider(): + stub = _make_stub() + stub._restore_session_model(_row(model_config={ + "gateway_runtime": {"provider": "custom:feather", "base_url": "https://f/v1"}, + })) + assert stub.model == "glm-4.7" + assert stub.provider == "custom:feather" + assert stub.requested_provider == "custom:feather" + assert stub.base_url == "https://f/v1" + # Stale launch-time explicit overrides must not leak into the restored + # provider's credential resolution. + assert stub._explicit_api_key is None + assert stub._explicit_base_url == "https://f/v1" + + +def test_restore_session_model_explicit_cli_flag_wins(): + stub = _make_stub(model="cli-flag-model", _explicit_model_override=True) + stub._restore_session_model(_row()) + assert stub.model == "cli-flag-model" + assert stub.provider == "openrouter" + + +def test_restore_session_model_no_stored_model_is_noop(): + stub = _make_stub() + stub._restore_session_model(_row(model=None)) + assert stub.model == "ambient-model" + + +def test_restore_session_model_matching_state_is_silent_noop(): + notes = [] + stub = _make_stub(model="glm-4.7", provider="custom:feather", + requested_provider="custom:feather", + _console_print=lambda s: notes.append(s)) + stub._restore_session_model(_row(model_config={ + "gateway_runtime": {"provider": "custom:feather"}, + })) + assert not notes + + +def test_restore_session_model_swaps_running_agent_in_place(): + calls = {} + + class _Agent: + def switch_model(self, **kwargs): + calls.update(kwargs) + + stub = _make_stub(agent=_Agent()) + stub._restore_session_model(_row()) + assert calls["new_model"] == "glm-4.7" + + +# ── _persist_model_switch_to_session ──────────────────────────────── + + +class _Result: + new_model = "deepseek-v4-flash-free" + target_provider = "custom:opencode-zen" + base_url = "https://oz/v1" + api_mode = "" + + +def test_persist_model_switch_writes_model_and_both_route_shapes(): + written = {} + + class _DB: + def update_session_model(self, sid, model): + written["model"] = (sid, model) + + def patch_session_model_config(self, sid, patch): + written["patch"] = (sid, patch) + + stub = _make_stub(_session_db=_DB(), session_id="s1") + stub._persist_model_switch_to_session(_Result()) + assert written["model"] == ("s1", "deepseek-v4-flash-free") + sid, patch = written["patch"] + # Nested shape for the CLI reader... + assert patch["gateway_runtime"]["provider"] == "custom:opencode-zen" + # ...and top-level for the TUI gateway's _stored_session_runtime_overrides. + assert patch["provider"] == "custom:opencode-zen" + assert patch["base_url"] == "https://oz/v1" + assert "api_mode" not in patch["gateway_runtime"] # empty values dropped + + +def test_persist_model_switch_noop_without_db_or_session(): + stub = _make_stub() # no _session_db / session_id attributes at all + stub._persist_model_switch_to_session(_Result()) # must not raise + + +def test_persist_model_switch_swallows_db_errors(): + class _DB: + def update_session_model(self, *a): + raise RuntimeError("disk full") + + stub = _make_stub(_session_db=_DB(), session_id="s1") + stub._persist_model_switch_to_session(_Result()) # must not raise + + +# ── round trip: persist → get_session shape → restore ─────────────── + + +def test_round_trip_persist_then_restore(tmp_path, monkeypatch): + monkeypatch.setenv("HERMES_HOME", str(tmp_path / ".hermes")) + db = SessionDB(db_path=tmp_path / "state.db") + db.create_session(session_id="rt1", source="cli", model="ambient-model") + + stub = _make_stub(_session_db=db, session_id="rt1") + stub._persist_model_switch_to_session(_Result()) + + meta = db.get_session("rt1") + restored = _make_stub() + restored._restore_session_model(meta) + assert restored.model == "deepseek-v4-flash-free" + assert restored.provider == "custom:opencode-zen" + assert restored.base_url == "https://oz/v1"