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:
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user