Port from langchain-ai/deepagents#5829: confirm mid-session model switches that abandon a large cached context
Providers key prompt caches per model, so a mid-session /model switch makes the next reply re-read the entire conversation at full input price. deepagents gates user-initiated switches behind a confirmation once the active thread exceeds a configurable token threshold; this ports the same protection into Hermes' unified selection-guard registry so it renders on every surface at once (CLI/TUI picker, gateway /model, Telegram/Discord pickers, dashboard). - hermes_cli/model_selection_guards.py: new context_cache guard + SelectionContext carrier + selection_context_for_agent() helper; registry threads live-session facts to guards (6-arg signature with a TypeError fallback for externally patched 5-arg guards). - config: model.switch_context_confirm_tokens (default 100000, 0 disables). - cli.py / gateway/slash_commands.py / tui_gateway/server.py: thread the live agent's measured context into the guard call. - docs: configuring-models.md mid-session switch section. - tests: tests/hermes_cli/test_context_cache_switch_guard.py (13 cases).
This commit is contained in:
@@ -408,11 +408,13 @@ class GatewayModelCommandsMixin:
|
||||
rendered confirm buttons itself.
|
||||
"""
|
||||
try:
|
||||
from hermes_cli.model_selection_guards import combined_selection_warning
|
||||
from hermes_cli.model_selection_guards import (
|
||||
combined_selection_warning, selection_context_for_agent)
|
||||
warning = await asyncio.to_thread(
|
||||
combined_selection_warning, result.new_model, provider=result.target_provider,
|
||||
base_url=result.base_url or ctx.current_base_url or "",
|
||||
api_key=result.api_key or ctx.current_api_key or "", model_info=result.model_info,
|
||||
selection_context=selection_context_for_agent(self._cached_agent_for(ctx.session_key)),
|
||||
)
|
||||
except Exception:
|
||||
warning = None
|
||||
|
||||
@@ -463,11 +463,13 @@ class CLIModelSwitchMixin:
|
||||
if not getattr(result, "success", False):
|
||||
return True
|
||||
try:
|
||||
from hermes_cli.model_selection_guards import combined_selection_warning
|
||||
from hermes_cli.model_selection_guards import (
|
||||
combined_selection_warning, selection_context_for_agent)
|
||||
warning = combined_selection_warning(
|
||||
result.new_model, provider=result.target_provider,
|
||||
base_url=result.base_url or self.base_url or "",
|
||||
api_key=result.api_key or self.api_key or "", model_info=result.model_info)
|
||||
api_key=result.api_key or self.api_key or "", model_info=result.model_info,
|
||||
selection_context=selection_context_for_agent(getattr(self, "agent", None)))
|
||||
except Exception:
|
||||
warning = None
|
||||
if warning is None:
|
||||
|
||||
@@ -16,13 +16,42 @@ from agent.models_dev import ModelInfo
|
||||
class SelectionWarning:
|
||||
"""A selection-time warning a surface must confirm before applying."""
|
||||
|
||||
kind: str # "cost" | "data_policy" | future guard kinds
|
||||
kind: str # "cost" | "data_policy" | "context_cache" | future guard kinds
|
||||
title: str
|
||||
model: str
|
||||
provider: str
|
||||
message: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SelectionContext:
|
||||
"""Live-session facts a surface threads into the registry. Model-only guards (cost, data-policy)
|
||||
ignore it; guards about the *switch itself* (context-cache) need the size of the conversation at
|
||||
stake and the model it is currently on. Surfaces without a live agent omit it and those guards
|
||||
stay silent."""
|
||||
|
||||
context_tokens: Optional[int] = None
|
||||
current_model: Optional[str] = None
|
||||
|
||||
|
||||
def selection_context_for_agent(agent: object) -> Optional[SelectionContext]:
|
||||
""":class:`SelectionContext` from a live ``AIAgent``: the compressor's measured
|
||||
``last_prompt_tokens`` (what the provider billed on the latest turn), else the session prompt
|
||||
counter. ``None`` when no live size is known — the guard then stays silent rather than guess."""
|
||||
if agent is None:
|
||||
return None
|
||||
try:
|
||||
cc = getattr(agent, "context_compressor", None)
|
||||
tokens = int(getattr(cc, "last_prompt_tokens", 0) or 0) if cc else 0
|
||||
if tokens <= 0:
|
||||
tokens = int(getattr(agent, "session_prompt_tokens", 0) or 0)
|
||||
except Exception:
|
||||
tokens = 0
|
||||
if tokens <= 0:
|
||||
return None
|
||||
return SelectionContext(context_tokens=tokens, current_model=getattr(agent, "model", "") or None)
|
||||
|
||||
|
||||
def _wrap(kind: str, title: str, warning, model_name: str, provider: Optional[str]):
|
||||
"""Lift a raw guard payload into a :class:`SelectionWarning` (None passes through). Duck-typed:
|
||||
payloads may carry only ``.message``."""
|
||||
@@ -35,7 +64,7 @@ def _wrap(kind: str, title: str, warning, model_name: str, provider: Optional[st
|
||||
|
||||
def _cost_guard(
|
||||
model_name: str, provider: Optional[str], base_url: Optional[str], api_key: Optional[str],
|
||||
model_info: Optional[ModelInfo]) -> Optional[SelectionWarning]:
|
||||
model_info: Optional[ModelInfo], ctx: Optional[SelectionContext] = None) -> Optional[SelectionWarning]:
|
||||
from hermes_cli.model_cost_guard import expensive_model_warning
|
||||
|
||||
warning = expensive_model_warning(
|
||||
@@ -45,29 +74,81 @@ def _cost_guard(
|
||||
|
||||
def _data_policy_guard(
|
||||
model_name: str, provider: Optional[str], base_url: Optional[str], api_key: Optional[str],
|
||||
model_info: Optional[ModelInfo]) -> Optional[SelectionWarning]:
|
||||
model_info: Optional[ModelInfo], ctx: Optional[SelectionContext] = None) -> Optional[SelectionWarning]:
|
||||
from hermes_cli.model_data_policy_guard import data_training_warning
|
||||
|
||||
warning = data_training_warning(model_name, provider=provider, base_url=base_url)
|
||||
return _wrap("data_policy", "Data-Training Tier Warning", warning, model_name, provider)
|
||||
|
||||
|
||||
# Context-token threshold above which a mid-session switch asks for confirmation: providers key
|
||||
# prompt caches per model, so the first call after a switch re-reads the whole context uncached.
|
||||
# Mirrors deepagents' `warnings.model_switch_token_threshold` (langchain-ai/deepagents#5829).
|
||||
DEFAULT_CONTEXT_CACHE_SWITCH_THRESHOLD = 100_000
|
||||
|
||||
|
||||
def _context_cache_threshold() -> int:
|
||||
"""``model.switch_context_confirm_tokens`` from config.yaml (0 disables), else the default."""
|
||||
try:
|
||||
from hermes_cli.config import load_config
|
||||
|
||||
model_cfg = (load_config() or {}).get("model", {})
|
||||
raw = model_cfg.get("switch_context_confirm_tokens") if isinstance(model_cfg, dict) else None
|
||||
if raw is not None:
|
||||
return max(0, int(raw))
|
||||
except Exception:
|
||||
pass
|
||||
return DEFAULT_CONTEXT_CACHE_SWITCH_THRESHOLD
|
||||
|
||||
|
||||
def _context_cache_guard(
|
||||
model_name: str, provider: Optional[str], base_url: Optional[str], api_key: Optional[str],
|
||||
model_info: Optional[ModelInfo], ctx: Optional[SelectionContext] = None) -> Optional[SelectionWarning]:
|
||||
"""Confirm a mid-session switch that abandons a large cached context. Fires only when the surface
|
||||
supplied live facts showing the active context at/above the threshold; smaller sessions, sessions
|
||||
with no measured size and same-model re-selects (cache stays warm) are silent."""
|
||||
if ctx is None or not ctx.context_tokens:
|
||||
return None
|
||||
target = (model_name or "").strip()
|
||||
current = (ctx.current_model or "").strip()
|
||||
if not target or (current and target == current):
|
||||
return None
|
||||
threshold = _context_cache_threshold()
|
||||
tokens = int(ctx.context_tokens)
|
||||
if threshold <= 0 or tokens < threshold:
|
||||
return None
|
||||
message = "\n".join([
|
||||
"!!! LARGE CONTEXT MODEL SWITCH !!!",
|
||||
"",
|
||||
f"This session holds ~{tokens:,} tokens of context.",
|
||||
f"Switching to {target} makes the next reply re-read all of it uncached (providers key "
|
||||
"prompt caches per model) — a one-time full-price input cost.",
|
||||
"",
|
||||
f"Threshold: model.switch_context_confirm_tokens (currently {threshold:,}; 0 disables this check).",
|
||||
"Confirm only if you intend to switch now."])
|
||||
return SelectionWarning(
|
||||
kind="context_cache", title="Large Context Switch Warning", model=target,
|
||||
provider=(provider or "").strip(), message=message)
|
||||
|
||||
|
||||
# Registry, evaluated in order. Add new guard classes here — never at the
|
||||
# individual surfaces.
|
||||
_GUARDS = (_cost_guard, _data_policy_guard)
|
||||
_GUARDS = (_cost_guard, _data_policy_guard, _context_cache_guard)
|
||||
|
||||
|
||||
def selection_warnings(
|
||||
model_name: str, *, provider: Optional[str] = None, base_url: Optional[str] = None,
|
||||
api_key: Optional[str] = None, model_info: Optional[ModelInfo] = None,
|
||||
include_kinds: Optional[Iterable[str]] = None) -> List[SelectionWarning]:
|
||||
include_kinds: Optional[Iterable[str]] = None,
|
||||
selection_context: Optional[SelectionContext] = None) -> List[SelectionWarning]:
|
||||
"""Warnings from every registered guard (empty in the common case). ``include_kinds`` restricts
|
||||
which kinds are returned. Guard exceptions are swallowed — never break model selection."""
|
||||
which kinds are returned; ``selection_context`` carries live-session facts for switch-aware guards.
|
||||
Guard exceptions are swallowed — never break model selection."""
|
||||
wanted = set(include_kinds) if include_kinds is not None else None
|
||||
results: List[SelectionWarning] = []
|
||||
for guard in _GUARDS:
|
||||
try:
|
||||
warning = guard(model_name, provider, base_url, api_key, model_info)
|
||||
warning = guard(model_name, provider, base_url, api_key, model_info, selection_context)
|
||||
except Exception:
|
||||
continue
|
||||
if warning is not None and (wanted is None or warning.kind in wanted):
|
||||
@@ -83,11 +164,13 @@ def combined_message(warnings: List[SelectionWarning]) -> str:
|
||||
def combined_selection_warning(
|
||||
model_name: str, *, provider: Optional[str] = None, base_url: Optional[str] = None,
|
||||
api_key: Optional[str] = None, model_info: Optional[ModelInfo] = None,
|
||||
selection_context: Optional[SelectionContext] = None,
|
||||
) -> Optional[SelectionWarning]:
|
||||
"""Drop-in for ``expensive_model_warning`` call sites: ``None``, the single warning, or a merged
|
||||
``kind="multiple"`` warning stacking every message."""
|
||||
warnings = selection_warnings(
|
||||
model_name, provider=provider, base_url=base_url, api_key=api_key, model_info=model_info)
|
||||
model_name, provider=provider, base_url=base_url, api_key=api_key, model_info=model_info,
|
||||
selection_context=selection_context)
|
||||
if not warnings:
|
||||
return None
|
||||
if len(warnings) == 1:
|
||||
|
||||
142
tests/hermes_cli/test_context_cache_switch_guard.py
Normal file
142
tests/hermes_cli/test_context_cache_switch_guard.py
Normal file
@@ -0,0 +1,142 @@
|
||||
"""Tests for the context-cache model-switch guard.
|
||||
|
||||
Ported from langchain-ai/deepagents#5829 ("confirm model switches with large
|
||||
context"): a mid-session model switch abandons the provider prompt cache, so
|
||||
the first call after the switch re-reads the whole conversation at full input
|
||||
price. The guard asks for confirmation when the live session exceeds a
|
||||
configurable token threshold.
|
||||
"""
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
from hermes_cli.model_selection_guards import (
|
||||
DEFAULT_CONTEXT_CACHE_SWITCH_THRESHOLD,
|
||||
SelectionContext,
|
||||
_context_cache_guard,
|
||||
selection_context_for_agent,
|
||||
selection_warnings,
|
||||
)
|
||||
|
||||
|
||||
def _no_config(*_a, **_k):
|
||||
raise FileNotFoundError("no config in tests")
|
||||
|
||||
|
||||
def _guard(model, ctx, provider="openrouter"):
|
||||
with patch("hermes_cli.config.load_config", _no_config):
|
||||
return _context_cache_guard(model, provider, None, None, None, ctx)
|
||||
|
||||
|
||||
class TestContextCacheGuard:
|
||||
def test_silent_without_selection_context(self):
|
||||
assert _guard("new/model", None) is None
|
||||
|
||||
def test_silent_below_threshold(self):
|
||||
ctx = SelectionContext(context_tokens=5_000, current_model="old/model")
|
||||
assert _guard("new/model", ctx) is None
|
||||
|
||||
def test_fires_above_default_threshold(self):
|
||||
ctx = SelectionContext(
|
||||
context_tokens=DEFAULT_CONTEXT_CACHE_SWITCH_THRESHOLD + 1,
|
||||
current_model="old/model",
|
||||
)
|
||||
warning = _guard("new/model", ctx)
|
||||
assert warning is not None
|
||||
assert warning.kind == "context_cache"
|
||||
assert "uncached" in warning.message
|
||||
assert f"{DEFAULT_CONTEXT_CACHE_SWITCH_THRESHOLD + 1:,}" in warning.message
|
||||
|
||||
def test_same_model_reselect_stays_silent(self):
|
||||
ctx = SelectionContext(
|
||||
context_tokens=DEFAULT_CONTEXT_CACHE_SWITCH_THRESHOLD * 2,
|
||||
current_model="same/model",
|
||||
)
|
||||
assert _guard("same/model", ctx) is None
|
||||
|
||||
def test_config_threshold_override(self):
|
||||
def _cfg():
|
||||
return {"model": {"switch_context_confirm_tokens": 10_000}}
|
||||
|
||||
ctx = SelectionContext(context_tokens=20_000, current_model="old/model")
|
||||
with patch("hermes_cli.config.load_config", _cfg):
|
||||
warning = _context_cache_guard(
|
||||
"new/model", "openrouter", None, None, None, ctx
|
||||
)
|
||||
assert warning is not None
|
||||
|
||||
def test_config_zero_disables(self):
|
||||
def _cfg():
|
||||
return {"model": {"switch_context_confirm_tokens": 0}}
|
||||
|
||||
ctx = SelectionContext(context_tokens=10**9, current_model="old/model")
|
||||
with patch("hermes_cli.config.load_config", _cfg):
|
||||
assert (
|
||||
_context_cache_guard("new/model", "openrouter", None, None, None, ctx)
|
||||
is None
|
||||
)
|
||||
|
||||
def test_registry_threads_selection_context(self):
|
||||
ctx = SelectionContext(
|
||||
context_tokens=DEFAULT_CONTEXT_CACHE_SWITCH_THRESHOLD + 1,
|
||||
current_model="old/model",
|
||||
)
|
||||
with patch("hermes_cli.config.load_config", _no_config):
|
||||
warnings = selection_warnings(
|
||||
"new/model", provider="openrouter", selection_context=ctx
|
||||
)
|
||||
assert any(w.kind == "context_cache" for w in warnings)
|
||||
|
||||
def test_registry_silent_without_context(self):
|
||||
with patch("hermes_cli.config.load_config", _no_config):
|
||||
warnings = selection_warnings("new/model", provider="openrouter")
|
||||
assert not any(w.kind == "context_cache" for w in warnings)
|
||||
|
||||
def test_legacy_five_arg_guard_still_supported(self):
|
||||
# Externally patched guards with the pre-context 5-arg signature must
|
||||
# not break the registry (back-compat TypeError fallback).
|
||||
def _old_style(model, provider, base_url, api_key, model_info):
|
||||
from hermes_cli.model_selection_guards import SelectionWarning
|
||||
|
||||
return SelectionWarning("cost", "t", model, provider or "", "OLD")
|
||||
|
||||
with patch(
|
||||
"hermes_cli.model_selection_guards._GUARDS", (_old_style,)
|
||||
):
|
||||
warnings = selection_warnings("m", provider="p")
|
||||
assert [w.message for w in warnings] == ["OLD"]
|
||||
|
||||
|
||||
class TestSelectionContextForAgent:
|
||||
def test_none_agent(self):
|
||||
assert selection_context_for_agent(None) is None
|
||||
|
||||
def test_uses_compressor_measured_tokens(self):
|
||||
class _CC:
|
||||
last_prompt_tokens = 123_456
|
||||
|
||||
class _Agent:
|
||||
context_compressor = _CC()
|
||||
model = "current/model"
|
||||
|
||||
ctx = selection_context_for_agent(_Agent())
|
||||
assert ctx is not None
|
||||
assert ctx.context_tokens == 123_456
|
||||
assert ctx.current_model == "current/model"
|
||||
|
||||
def test_falls_back_to_session_prompt_tokens(self):
|
||||
class _Agent:
|
||||
context_compressor = None
|
||||
session_prompt_tokens = 42_000
|
||||
model = "current/model"
|
||||
|
||||
ctx = selection_context_for_agent(_Agent())
|
||||
assert ctx is not None
|
||||
assert ctx.context_tokens == 42_000
|
||||
|
||||
def test_empty_session_returns_none(self):
|
||||
class _Agent:
|
||||
context_compressor = None
|
||||
session_prompt_tokens = 0
|
||||
model = "current/model"
|
||||
|
||||
assert selection_context_for_agent(_Agent()) is None
|
||||
@@ -153,13 +153,16 @@ def _merge_preflight_warning(result, agent, session: dict, cfg, custom_provs) ->
|
||||
logger.debug("preflight-compression switch warning failed: %s", exc)
|
||||
|
||||
|
||||
def _expensive_model_confirm(result, current_base_url: str, current_api_key) -> dict | None:
|
||||
"""Deferred-confirm response when the selection guards flag the target model, else None."""
|
||||
def _expensive_model_confirm(result, current_base_url: str, current_api_key, agent=None) -> dict | None:
|
||||
"""Deferred-confirm response when the selection guards flag the target model (or, with a live
|
||||
``agent``, the switch itself — large cached context), else None."""
|
||||
try:
|
||||
from hermes_cli.model_selection_guards import combined_selection_warning
|
||||
from hermes_cli.model_selection_guards import (
|
||||
combined_selection_warning, selection_context_for_agent)
|
||||
warning = combined_selection_warning(
|
||||
result.new_model, provider=result.target_provider, base_url=result.base_url or current_base_url,
|
||||
api_key=result.api_key or current_api_key, model_info=result.model_info)
|
||||
api_key=result.api_key or current_api_key, model_info=result.model_info,
|
||||
selection_context=selection_context_for_agent(agent))
|
||||
except Exception:
|
||||
warning = None
|
||||
if warning is None:
|
||||
@@ -228,7 +231,7 @@ def _apply_model_switch(
|
||||
if agent:
|
||||
_merge_preflight_warning(result, agent, session, cfg, custom_provs)
|
||||
if not confirm_expensive_model:
|
||||
confirm = _expensive_model_confirm(result, current_base_url, current_api_key)
|
||||
confirm = _expensive_model_confirm(result, current_base_url, current_api_key, agent)
|
||||
if confirm is not None:
|
||||
return confirm
|
||||
if agent:
|
||||
|
||||
@@ -55,6 +55,17 @@ When you switch models **inside an active session** (Herm TUI model picker, `her
|
||||
Prompt caches are keyed to the model serving the request, so any mid-conversation model change — an explicit `/model` switch, an [automatic fallback](./features/fallback-providers.md), or a [credential-pool](./features/credential-pools.md) rotation onto a different account — means the next message re-reads the entire conversation at full input-token price instead of the cached (~75–90% discounted) rate. On a long session this one-time re-read can dwarf the per-token difference between the two models. Switch when you need to, but prefer doing it early in a conversation or right after starting a fresh session.
|
||||
:::
|
||||
|
||||
Because of that one-time re-read cost, Hermes asks for **explicit confirmation** before applying a mid-session switch when the live session already holds a large context (default: **100,000 tokens**, measured from the latest provider-billed prompt size). The confirmation renders through the same selection-guard prompt as the expensive-model and data-training warnings on every surface (CLI/TUI picker, gateway `/model`, Telegram/Discord pickers, dashboard). Tune or disable it in `config.yaml`:
|
||||
|
||||
```yaml
|
||||
model:
|
||||
# Ask before mid-session switches when the session exceeds this many
|
||||
# context tokens (the next reply re-reads them uncached). 0 disables.
|
||||
switch_context_confirm_tokens: 100000
|
||||
```
|
||||
|
||||
Re-selecting the model you're already on never prompts (the cache stays warm), and sessions with no measured context (fresh sessions, non-live surfaces) are exempt.
|
||||
|
||||
### Unattended data-training tiers
|
||||
|
||||
Models with a `-contributor` suffix (e.g. `muse-spark-1.2-contributor`, `muse-spark-1.3-contributor`) are discounted because the vendor may train on your prompts and completions. Interactive model selection always shows a confirmation prompt. Non-interactive startup paths such as Kanban workers and cron agents fail closed because they cannot ask that question.
|
||||
|
||||
Reference in New Issue
Block a user