fix(chat): resolve model switch lag by tracking model/provider key (#180)

* fix(chat): resolve model switch lag by tracking model/provider key for cache invalidation

* style: format set_chat_model function for improved readability

* fix(chat): improve model switching logic to prevent unnecessary cache rebuilds

* fix(model): ensure globals are restored on agent load failure to prevent model switch issues
This commit is contained in:
Xi Zhang
2026-04-24 12:36:33 +02:00
committed by GitHub
parent 92df3e1844
commit 3831198f19
4 changed files with 391 additions and 19 deletions
+52 -17
View File
@@ -46,6 +46,12 @@ SKILLS_DIR = str(Path(__file__).parent / "skills")
_config = None
_chat_model = None
# Track the (model, provider) binding of _chat_model so cache invalidates
# when config.model/provider change (e.g. via /model). Without this,
# _ensure_chat_model() returns the stale cached instance even after
# _ensure_config(new_cfg) has overwritten the active config — causing
# /model switch to lag one step (see issue #179).
_chat_model_key: tuple[str | None, str | None] | None = None
# Cache MCP tools by the effective config signature to avoid reconnecting
# to MCP servers on every `/new` when config is unchanged.
@@ -75,29 +81,58 @@ def _ensure_config(config=None):
return _config
def _ensure_chat_model():
"""Return cached chat model, creating it on first call."""
global _chat_model
if _chat_model is None:
from .llm import get_chat_model
def _replace_chat_model(instance, key: tuple[str | None, str | None]) -> None:
"""Install a new chat model and propagate the related invariants.
cfg = _ensure_config()
_chat_model = get_chat_model(model=cfg.model, provider=cfg.provider)
Single write point for ``_chat_model`` / ``_chat_model_key`` /
``_EvoScientist_agent``: both ``_ensure_chat_model`` (cache-miss
rebuild) and ``set_chat_model`` (explicit switch via ``/model``)
funnel through here so the three globals can never drift.
"""
global _chat_model, _chat_model_key, _EvoScientist_agent
_chat_model = instance
_chat_model_key = key
# The lazy default agent captured a reference to the previous
# ``_chat_model`` at build time, so it must be rebuilt on next access.
_EvoScientist_agent = None
def _ensure_chat_model():
"""Return cached chat model, rebuilding if cfg.model/provider changed.
The cache key is the current config's ``(model, provider)``. If it
differs from the key that built ``_chat_model``, rebuild — this makes
``create_cli_agent(config=temp_cfg)`` bind the freshly requested model
into the new agent without requiring callers to interleave
``set_chat_model()`` calls in any particular order.
"""
from .llm import get_chat_model
cfg = _ensure_config()
key = (cfg.model, cfg.provider)
if _chat_model is None or _chat_model_key != key:
_replace_chat_model(
get_chat_model(model=cfg.model, provider=cfg.provider),
key,
)
return _chat_model
def set_chat_model(model: str, provider: str | None = None):
"""Replace the cached chat model with a new one.
Called by ``/model`` to switch the LLM mid-session.
Returns the new chat model instance.
Called by ``/model`` to switch the LLM mid-session. No-op when the
cache already holds the requested ``(model, provider)`` — avoids
spawning a second ``get_chat_model`` instance (and its HTTP client)
under the ``/model`` flow where ``_ensure_chat_model`` has already
rebuilt ``_chat_model`` during the preceding ``_load_agent`` call.
Returns the current chat model instance.
"""
global _chat_model, _EvoScientist_agent
from .llm import get_chat_model
_chat_model = get_chat_model(model=model, provider=provider)
# Invalidate the cached default agent so it gets rebuilt with the new model.
_EvoScientist_agent = None
key = (model, provider)
if _chat_model is None or _chat_model_key != key:
_replace_chat_model(get_chat_model(model=model, provider=provider), key)
return _chat_model
@@ -168,8 +203,8 @@ def _inject_subagent_middleware(subs: list[dict]) -> None:
for sa in subs:
sa.setdefault("middleware", []).extend(
[
# Uses main agent's model for trigger — subagents currently
# share the same model, so context window matches.
# No ``model=`` — subagents share the main agent's model,
# so defer to the factory's ``_ensure_chat_model()`` fallback.
create_context_editing_middleware(),
ToolErrorHandlerMiddleware(),
ContextOverflowMapperMiddleware(),
@@ -326,7 +361,7 @@ def _get_default_middleware():
create_context_editing_middleware(model),
ContextOverflowMapperMiddleware(),
ToolErrorHandlerMiddleware(),
*create_tool_selector_middleware(),
*create_tool_selector_middleware(model=model),
create_memory_middleware(memory_dir, extraction_model=model),
]
@@ -460,7 +495,7 @@ def create_cli_agent(
create_context_editing_middleware(model),
ContextOverflowMapperMiddleware(),
ToolErrorHandlerMiddleware(),
*create_tool_selector_middleware(),
*create_tool_selector_middleware(model=model),
create_memory_middleware(_mem_dir, extraction_model=model),
]
if cfg.enable_ask_user and not cfg.auto_mode:
@@ -111,6 +111,7 @@ class ModelCommand(Command):
) -> None:
import copy
from ... import EvoScientist as _mod
from ...cli.agent import _load_agent
from ...EvoScientist import _ensure_config, set_chat_model
@@ -122,6 +123,22 @@ class ModelCommand(Command):
temp_cfg.model = model_name
temp_cfg.provider = provider
# create_cli_agent(config=temp_cfg) calls _ensure_config and
# _ensure_chat_model before finishing, so a failure further
# down (middleware build, MCP reconnect, deepagents wiring)
# would leave the session pointing at the new model. Snapshot
# the four globals those helpers write so we can restore on
# error. Best-effort: references already captured by concurrent
# readers (e.g. a channel thread mid-turn) are not retroactively
# patched, but /model is user-initiated from an idle prompt in
# practice.
snap = (
_mod._config,
_mod._chat_model,
_mod._chat_model_key,
_mod._EvoScientist_agent,
)
try:
new_agent = _load_agent(
workspace_dir=ctx.workspace_dir,
@@ -129,6 +146,12 @@ class ModelCommand(Command):
config=temp_cfg,
)
except Exception as e:
(
_mod._config,
_mod._chat_model,
_mod._chat_model_key,
_mod._EvoScientist_agent,
) = snap
ctx.ui.append_system(f"Failed to switch model: {e}", style="red")
return
+314
View File
@@ -262,6 +262,240 @@ class TestModelCommandFailure:
assert call_args[1]["style"] == "red"
@pytest.fixture
def evo_module_state():
"""Snapshot and restore ``EvoScientist.EvoScientist`` module globals.
The chat-model cache tests mutate ``_chat_model`` / ``_chat_model_key``
/ ``_config`` / ``_EvoScientist_agent`` directly. This fixture
guarantees all four are restored — even if a test body grows an early
return — so later tests in the suite see a clean module state.
"""
import EvoScientist.EvoScientist as mod
snapshot = (
mod._chat_model,
mod._chat_model_key,
mod._config,
mod._EvoScientist_agent,
)
try:
yield mod
finally:
(
mod._chat_model,
mod._chat_model_key,
mod._config,
mod._EvoScientist_agent,
) = snapshot
class TestEnsureChatModelCacheInvalidation:
"""Regression tests for issue #179: /model switch lagged by one step.
Root cause: ``_ensure_chat_model()`` returned a cached ``_chat_model``
without checking whether ``cfg.model`` / ``cfg.provider`` had changed
since the cache was populated. ``ModelCommand._apply_model`` builds
a new agent *before* ``set_chat_model`` runs, so the new agent was
bound to the *previous* cached model.
Fix: track a ``(model, provider)`` key alongside ``_chat_model`` and
rebuild on mismatch.
"""
def test_cache_rebuilds_when_config_model_changes(self, evo_module_state):
"""After cfg.model changes, _ensure_chat_model must return a new instance."""
mod = evo_module_state
cfg = SimpleNamespace(model="claude-sonnet-4-6", provider="anthropic")
m1 = MagicMock(name="model-1")
m2 = MagicMock(name="model-2")
mod._chat_model = None
mod._chat_model_key = None
mod._config = cfg
mod._EvoScientist_agent = None
with patch("EvoScientist.llm.get_chat_model", side_effect=[m1, m2]) as gm:
first = mod._ensure_chat_model()
assert first is m1
# Same config → cache hit, no rebuild.
again = mod._ensure_chat_model()
assert again is m1
assert gm.call_count == 1
# Simulate /model switch writing the new choice into cfg.
cfg.model = "minimax-m2.7"
cfg.provider = "openrouter"
second = mod._ensure_chat_model()
# Must be the NEW model instance, not the cached one.
assert second is m2
assert second is not first
assert gm.call_count == 2
# Second call used the new model name + provider.
_, kwargs = gm.call_args
assert kwargs == {"model": "minimax-m2.7", "provider": "openrouter"}
def test_cache_rebuilds_when_only_provider_changes(self, evo_module_state):
"""Same model name, different provider (openrouter vs anthropic) must rebuild."""
mod = evo_module_state
cfg = SimpleNamespace(model="claude-sonnet-4-6", provider="anthropic")
m1 = MagicMock(name="anthropic-model")
m2 = MagicMock(name="openrouter-model")
mod._chat_model = None
mod._chat_model_key = None
mod._config = cfg
mod._EvoScientist_agent = None
with patch("EvoScientist.llm.get_chat_model", side_effect=[m1, m2]) as gm:
assert mod._ensure_chat_model() is m1
cfg.provider = "openrouter"
assert mod._ensure_chat_model() is m2
assert gm.call_count == 2
def test_set_chat_model_updates_key(self, evo_module_state):
"""set_chat_model must keep _chat_model_key in sync to avoid
an extra rebuild on the very next _ensure_chat_model() call."""
mod = evo_module_state
cfg = SimpleNamespace(model="claude-sonnet-4-6", provider="anthropic")
m_set = MagicMock(name="explicit-set-model")
mod._chat_model = None
mod._chat_model_key = None
mod._config = cfg
mod._EvoScientist_agent = None
with patch("EvoScientist.llm.get_chat_model", return_value=m_set) as gm:
mod.set_chat_model("minimax-m2.7", provider="openrouter")
assert mod._chat_model is m_set
assert mod._chat_model_key == ("minimax-m2.7", "openrouter")
# Align cfg to what set_chat_model was called with.
cfg.model = "minimax-m2.7"
cfg.provider = "openrouter"
# Now _ensure_chat_model should NOT rebuild (key matches cfg).
assert mod._ensure_chat_model() is m_set
assert gm.call_count == 1
def test_set_chat_model_is_no_op_when_key_already_matches(self, evo_module_state):
"""set_chat_model must NOT rebuild when the cache already holds the
requested (model, provider).
Under the /model flow, ``_load_agent`` already rebuilt ``_chat_model``
via ``_ensure_chat_model`` before ``set_chat_model`` is reached, so
the subsequent set should be idempotent — reusing the same Python
instance (and thus the same underlying HTTP client) that ``ctx.agent``
is already bound to.
"""
mod = evo_module_state
existing = MagicMock(name="existing-model")
mod._chat_model = existing
mod._chat_model_key = ("minimax-m2.7", "openrouter")
mod._config = SimpleNamespace(model="minimax-m2.7", provider="openrouter")
mod._EvoScientist_agent = None
with patch("EvoScientist.llm.get_chat_model") as gm:
returned = mod.set_chat_model("minimax-m2.7", provider="openrouter")
# Fast path: returned the EXISTING instance, no rebuild.
assert returned is existing
assert mod._chat_model is existing
gm.assert_not_called()
class TestApplyModelIntegration:
"""End-to-end regression for #179: `_apply_model` must produce an agent
bound to the NEW model, not a stale cached one.
Exercises the real chain:
``_apply_model → _load_agent → _ensure_config → _ensure_chat_model``
Only ``_load_agent`` is replaced by a minimal fake that mirrors the
exact two globals ``create_cli_agent`` mutates (``_ensure_config``
+ ``_ensure_chat_model``) — so the bug path is fully exercised
without having to spin up deepagents, MCP tools, middleware, and
subagent YAML. ``get_chat_model`` returns a distinct sentinel per
``(model, provider)`` pair so we can assert on identity.
"""
def test_new_agent_is_bound_to_newly_selected_model(self, evo_module_state):
from EvoScientist.commands.implementation.model import ModelCommand
from EvoScientist.config.settings import EvoScientistConfig
mod = evo_module_state
sentinels: dict[tuple[str, str | None], MagicMock] = {}
def _fake_get_chat_model(model, provider=None):
key = (model, provider)
if key not in sentinels:
sentinels[key] = MagicMock(name=f"chat_model[{model}|{provider}]")
sentinels[key]._bound_model = model
sentinels[key]._bound_provider = provider
return sentinels[key]
def _fake_load_agent(
workspace_dir=None, checkpointer=None, config=None, *, on_mcp_progress=None
):
# Replicates the two create_cli_agent side effects that
# reveal the bug: ``_ensure_config`` writes the new cfg,
# then ``_ensure_chat_model`` must rebuild to match it.
mod._ensure_config(config)
agent = MagicMock(name="fake-agent")
agent._bound_model = mod._ensure_chat_model()
return agent
cfg = EvoScientistConfig(model="claude-sonnet-4-6", provider="anthropic")
ctx = MagicMock()
ctx.ui = MagicMock()
ctx.ui.supports_interactive = True
ctx.workspace_dir = "/tmp/test_integration"
ctx.checkpointer = None
# Prime: _chat_model already holds the OLD (default) model —
# this is the state that caused the off-by-one in production.
mod._config = cfg
mod._chat_model = _fake_get_chat_model("claude-sonnet-4-6", "anthropic")
mod._chat_model_key = ("claude-sonnet-4-6", "anthropic")
mod._EvoScientist_agent = None
old_model = mod._chat_model
with (
patch(
"EvoScientist.llm.get_chat_model",
side_effect=_fake_get_chat_model,
),
patch(
"EvoScientist.cli.agent._load_agent",
side_effect=_fake_load_agent,
),
):
cmd = ModelCommand()
_run(cmd._apply_model(ctx, "minimax-m2.7", "openrouter"))
# The agent produced by _apply_model must be bound to the
# NEWLY requested model, not the previously cached one.
assert ctx.agent._bound_model is not old_model
assert ctx.agent._bound_model._bound_model == "minimax-m2.7"
assert ctx.agent._bound_model._bound_provider == "openrouter"
# Global state reflects the switch end-to-end.
assert mod._chat_model_key == ("minimax-m2.7", "openrouter")
assert mod._chat_model is sentinels[("minimax-m2.7", "openrouter")]
assert cfg.model == "minimax-m2.7"
assert cfg.provider == "openrouter"
# User-visible success message.
msg = ctx.ui.append_system.call_args[0][0]
assert "minimax-m2.7" in msg
assert "openrouter" in msg
class TestModelCommandLoadAgentFailure:
"""Verify the transactional ordering: when ``_load_agent`` raises,
nothing downstream (``set_chat_model``, ``cfg`` mutation,
@@ -319,3 +553,83 @@ class TestModelCommandLoadAgentFailure:
call_args = ui.append_system.call_args
assert "Failed to switch model" in call_args[0][0]
assert call_args[1]["style"] == "red"
class TestApplyModelLoadAgentFailureTransactional:
"""Regression: if ``create_cli_agent`` raises AFTER partially mutating
module globals via ``_ensure_config`` + ``_ensure_chat_model``,
``_apply_model`` must roll those back so the session stays on the
original model.
Complements :class:`TestModelCommandLoadAgentFailure`, which tests
the early-failure path where ``_load_agent`` never reaches
``create_cli_agent`` and no globals get mutated.
"""
def test_globals_restored_after_create_cli_agent_partial_mutation(
self, evo_module_state
):
from EvoScientist.commands.implementation.model import ModelCommand
from EvoScientist.config.settings import EvoScientistConfig
mod = evo_module_state
sentinels: dict[tuple[str, str | None], MagicMock] = {}
def _fake_get_chat_model(model, provider=None):
key = (model, provider)
sentinels.setdefault(key, MagicMock(name=f"chat_model[{model}|{provider}]"))
return sentinels[key]
def _fake_load_agent(
workspace_dir=None,
checkpointer=None,
config=None,
*,
on_mcp_progress=None,
):
# Mimic ``create_cli_agent``: mutate globals via
# ``_ensure_config`` + ``_ensure_chat_model``, then raise
# (as if middleware construction or deepagents wiring failed).
mod._ensure_config(config)
mod._ensure_chat_model()
raise RuntimeError("middleware build failed")
cfg = EvoScientistConfig(model="claude-sonnet-4-6", provider="anthropic")
old_model = _fake_get_chat_model("claude-sonnet-4-6", "anthropic")
old_agent = MagicMock(name="old-default-agent")
mod._config = cfg
mod._chat_model = old_model
mod._chat_model_key = ("claude-sonnet-4-6", "anthropic")
mod._EvoScientist_agent = old_agent
ctx = MagicMock()
ctx.ui = MagicMock()
ctx.ui.supports_interactive = True
ctx.workspace_dir = "/tmp/test_rollback"
ctx.checkpointer = None
with (
patch(
"EvoScientist.llm.get_chat_model",
side_effect=_fake_get_chat_model,
),
patch(
"EvoScientist.cli.agent._load_agent",
side_effect=_fake_load_agent,
),
):
cmd = ModelCommand()
_run(cmd._apply_model(ctx, "minimax-m2.7", "openrouter"))
# All four globals restored to their pre-call state.
assert mod._config is cfg
assert mod._chat_model is old_model
assert mod._chat_model_key == ("claude-sonnet-4-6", "anthropic")
assert mod._EvoScientist_agent is old_agent
# The ``cfg`` object itself was not mutated.
assert cfg.model == "claude-sonnet-4-6"
assert cfg.provider == "anthropic"
# User sees an error message.
msg = ctx.ui.append_system.call_args[0][0]
assert "Failed to switch model" in msg
+2 -2
View File
@@ -152,7 +152,7 @@ def test_tracker_captures_tools():
@patch(
"EvoScientist.middleware.create_tool_selector_middleware",
side_effect=lambda: _patched_create(),
side_effect=lambda *a, **kw: _patched_create(),
)
@patch("EvoScientist.EvoScientist._ensure_chat_model")
@patch("EvoScientist.EvoScientist._ensure_config")
@@ -186,7 +186,7 @@ def test_subagent_no_tool_selector(mock_model):
@patch(
"EvoScientist.middleware.create_tool_selector_middleware",
side_effect=lambda: _patched_create(),
side_effect=lambda *a, **kw: _patched_create(),
)
@patch("EvoScientist.EvoScientist._ensure_chat_model")
@patch("EvoScientist.EvoScientist._ensure_config")