From 635841d2108511408e744eb9894a11f3911a4fe8 Mon Sep 17 00:00:00 2001 From: Tranquil-Flow <66773372+Tranquil-Flow@users.noreply.github.com> Date: Sat, 27 Jun 2026 04:07:44 -0700 Subject: [PATCH] fix(agent): reload credential pool on switch_model provider change (#52727) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit switch_model() swapped model/provider/base_url/api_key but never refreshed agent._credential_pool, which stays bound to the original provider. recover_with_credential_pool() then sees a pool.provider != agent.provider mismatch and short-circuits — so a 429/401 on the new provider gets no rotation and falls through to fallback instead. Reload load_pool(new_provider) inside switch_model when the provider changes (or the pool is missing). The reload is inside the protected swap block and the pool is added to the rollback snapshot, so a failed client rebuild restores the original pool. Fixes #16678, #52727. --- agent/agent_runtime_helpers.py | 25 ++ .../test_switch_model_pool_reload_52727.py | 218 ++++++++++++++++++ 2 files changed, 243 insertions(+) create mode 100644 tests/run_agent/test_switch_model_pool_reload_52727.py diff --git a/agent/agent_runtime_helpers.py b/agent/agent_runtime_helpers.py index 3473eb0b5..1f7a3b0b6 100644 --- a/agent/agent_runtime_helpers.py +++ b/agent/agent_runtime_helpers.py @@ -1499,6 +1499,10 @@ def switch_model(agent, new_model, new_provider, api_key='', base_url='', api_mo # _client_kwargs is a dict — snapshot a shallow copy so mutating the # live dict doesn't poison the rollback target. _snapshot["_client_kwargs"] = dict(getattr(agent, "_client_kwargs", {}) or {}) + # Snapshot the credential pool reference so a failed client rebuild can + # restore the original pool (issue #52727: pool reload is part of this + # switch and must be reversible on rollback). + _snapshot["_credential_pool"] = getattr(agent, "_credential_pool", _MISSING) try: # Clear the per-config context_length override so the new model's @@ -1523,6 +1527,27 @@ def switch_model(agent, new_model, new_provider, api_key='', base_url='', api_mo if api_key: agent.api_key = api_key + # ── Reload credential pool for the new provider (issue #52727) ── + # Without this, ``recover_with_credential_pool`` sees a + # ``pool.provider != agent.provider`` mismatch and short-circuits, + # leaving the new provider with no rotation/recovery on 401/429 and + # burning the original pool's entries. Only reload when the provider + # actually changed (or the pool was missing) — re-selecting the same + # provider must not churn the pool reference. A reload failure is + # logged + swallowed: the switch itself must still complete. + old_norm = (old_provider or "").strip().lower() + new_norm = (new_provider or "").strip().lower() + if old_norm != new_norm or getattr(agent, "_credential_pool", None) is None: + try: + from agent.credential_pool import load_pool + agent._credential_pool = load_pool(new_provider) + except Exception as _pool_exc: # noqa: BLE001 + logger.warning( + "switch_model: credential pool reload failed for %s (%s); " + "continuing without pool rotation this turn", + new_provider, _pool_exc, + ) + # ── Build new client ── if api_mode == "anthropic_messages": from agent.anthropic_adapter import ( diff --git a/tests/run_agent/test_switch_model_pool_reload_52727.py b/tests/run_agent/test_switch_model_pool_reload_52727.py new file mode 100644 index 000000000..a1dba807f --- /dev/null +++ b/tests/run_agent/test_switch_model_pool_reload_52727.py @@ -0,0 +1,218 @@ +"""Regression tests for #52727: switch_model() must reload the credential pool +when the provider changes. + +When the desktop model picker swaps providers mid-session, ``switch_model`` +mutates ``agent.model``/``agent.provider``/``agent.base_url``/``agent.api_key`` +but never refreshes ``agent._credential_pool``. The pool stays bound to the +ORIGINAL provider. ``recover_with_credential_pool`` then sees a +``pool.provider != agent.provider`` mismatch and skips rotation entirely — +the 401 burns the whole retry cycle with no recovery, and the original +provider's pool entry gets marked ``STATUS_EXHAUSTED`` and persisted in +``auth.json`` (issue #52727). + +The fix reloads the pool via ``load_pool(new_provider)`` inside ``switch_model`` +whenever the provider changes (or the pool is missing). This keeps the +defensive mismatch guard in ``recover_with_credential_pool`` intact while +making it impossible for a legitimate same-call switch to trip the guard. +""" + +from unittest.mock import MagicMock, patch + +import pytest + +from agent.agent_runtime_helpers import switch_model + + +# --------------------------------------------------------------------------- +# helpers +# --------------------------------------------------------------------------- + + +def _make_agent(current_provider, current_model, current_pool): + """Bare agent object with the minimum attributes switch_model touches. + + Uses ``MagicMock`` so ``_anthropic_prompt_cache_policy(...)`` and + ``_ensure_lmstudio_runtime_loaded()`` work without real implementations. + """ + agent = MagicMock(name=f"Agent[{current_provider}]") + agent.provider = current_provider + agent.model = current_model + agent.base_url = f"https://{current_provider}.example/v1" + agent.api_key = f"{current_provider}-key" + agent.api_mode = "chat_completions" + agent.client = MagicMock(name="Client") + agent._client_kwargs = { + "api_key": "***", + "base_url": f"https://{current_provider}.example/v1", + } + agent._anthropic_client = None + agent._anthropic_api_key = "" + agent._anthropic_base_url = None + agent._is_anthropic_oauth = False + agent._config_context_length = None + agent._transport_cache = {} + agent._cached_system_prompt = "cached-system-prompt" + agent.context_compressor = None + agent._use_prompt_caching = False + agent._use_native_cache_layout = False + agent._primary_runtime = {} + agent._fallback_activated = False + agent._fallback_index = 0 + agent._fallback_chain = [] + agent._fallback_model = None + agent._credential_pool = current_pool + # Real-ish instance methods that switch_model calls + agent._anthropic_prompt_cache_policy = MagicMock(return_value=(False, False)) + agent._ensure_lmstudio_runtime_loaded = MagicMock() + return agent + + +def _make_pool(provider): + pool = MagicMock(name=f"Pool[{provider}]") + pool.provider = provider + return pool + + +# --------------------------------------------------------------------------- +# the fix +# --------------------------------------------------------------------------- + + +class TestSwitchModelReloadsCredentialPool: + """Issue #52727: switch_model must refresh _credential_pool on provider change.""" + + def test_switch_to_different_provider_reloads_pool(self): + """opencode-go -> groq must replace the agent's pool with a groq pool.""" + old_pool = _make_pool("opencode-go") + new_pool = _make_pool("groq") + agent = _make_agent("opencode-go", "qwen-coder", old_pool) + + with patch( + "agent.credential_pool.load_pool", + return_value=new_pool, + ) as load_pool_mock: + switch_model( + agent, + new_model="llama-3.3-70b", + new_provider="groq", + api_key="groq-key-new", + base_url="https://api.groq.com/openai/v1", + api_mode="chat_completions", + ) + + # The agent's pool must now point at the NEW provider's pool. + assert agent._credential_pool is new_pool, ( + f"agent._credential_pool was not reloaded on provider switch " + f"(still references {old_pool.provider})" + ) + assert agent._credential_pool.provider == "groq" + assert agent._credential_pool is not old_pool + # load_pool MUST have been called with the new provider. + load_pool_mock.assert_called_once_with("groq") + + def test_switch_to_same_provider_does_not_reload_pool(self): + """Re-selecting the current provider must NOT churn the pool reference.""" + existing_pool = _make_pool("opencode-go") + agent = _make_agent("opencode-go", "qwen-coder", existing_pool) + + load_pool_mock = MagicMock(name="load_pool") + + with patch("agent.credential_pool.load_pool", load_pool_mock): + switch_model( + agent, + new_model="qwen-coder", + new_provider="opencode-go", # SAME provider + api_key="opencode-go-key-new", + base_url="https://opencode.example/v1", + api_mode="chat_completions", + ) + + # Pool must remain the same object — no churn for same-provider switch. + assert agent._credential_pool is existing_pool + load_pool_mock.assert_not_called() + + def test_switch_creates_pool_when_agent_had_none(self): + """An agent without a pool that switches providers must acquire one.""" + new_pool = _make_pool("groq") + agent = _make_agent("opencode-go", "qwen-coder", None) + + with patch("agent.credential_pool.load_pool", return_value=new_pool): + switch_model( + agent, + new_model="llama-3.3-70b", + new_provider="groq", + api_key="groq-key-new", + base_url="https://api.groq.com/openai/v1", + api_mode="chat_completions", + ) + + assert agent._credential_pool is new_pool + assert agent._credential_pool.provider == "groq" + + def test_recover_pool_mismatch_guard_no_longer_trips_after_switch(self): + """End-to-end: after a provider switch, recover_with_credential_pool + must not skip rotation due to a provider mismatch. + + Before the fix: pool.provider=='opencode-go', agent.provider=='groq' + → mismatch guard fires → recovery skipped → 401 burns the cycle. + After the fix: switch_model reloaded the pool to groq, so the guard + is a no-op and recovery proceeds. + """ + from agent.agent_runtime_helpers import recover_with_credential_pool + from agent.error_classifier import FailoverReason + + old_pool = _make_pool("opencode-go") + new_pool = _make_pool("groq") + new_pool.mark_exhausted_and_rotate.return_value = None + agent = _make_agent("opencode-go", "qwen-coder", old_pool) + + with patch("agent.credential_pool.load_pool", return_value=new_pool): + switch_model( + agent, + new_model="llama-3.3-70b", + new_provider="groq", + api_key="groq-key-new", + base_url="https://api.groq.com/openai/v1", + api_mode="chat_completions", + ) + + # After the switch, the pool's provider matches the agent's provider. + # A 429 on groq should now reach pool.mark_exhausted_and_rotate. + recover_with_credential_pool( + agent, + status_code=429, + has_retried_429=False, + classified_reason=FailoverReason.rate_limit, + ) + + # The guard would have returned (False, has_retried_429) early + # without touching the pool. After the fix, the pool is consulted. + assert new_pool.current.called, ( + "pool.current() was never called — mismatch guard short-circuited" + ) + + def test_pool_reload_failure_does_not_block_switch(self): + """If load_pool raises (e.g. corrupt auth.json), switch_model must + still complete — the pool will simply be missing for this turn, and + the next turn can re-attempt. Crashing the whole switch is worse + than a transient pool gap.""" + agent = _make_agent("opencode-go", "qwen-coder", _make_pool("opencode-go")) + + with patch( + "agent.credential_pool.load_pool", + side_effect=RuntimeError("simulated corrupt auth.json"), + ): + # Should NOT raise — pool reload failure is logged+swallowed. + switch_model( + agent, + new_model="llama-3.3-70b", + new_provider="groq", + api_key="groq-key-new", + base_url="https://api.groq.com/openai/v1", + api_mode="chat_completions", + ) + + # The switch itself completed (provider/model updated) even though + # the pool reload failed. + assert agent.provider == "groq" + assert agent.model == "llama-3.3-70b" \ No newline at end of file