fix(tui): recover model switch after failed resume

This commit is contained in:
Inimitable Mind
2026-08-22 13:44:35 -07:00
committed by Teknium
parent d48adfa789
commit dbb9854591
2 changed files with 109 additions and 1 deletions
+77
View File
@@ -9271,6 +9271,83 @@ def test_config_set_model_explicit_provider_skips_broken_default_init(monkeypatc
server._sessions.pop("sid", None)
@pytest.mark.parametrize(
"provider_flag", [" --provider custom:new-provider", ""]
)
def test_config_set_model_recovers_failed_deferred_resume(monkeypatch, provider_flag):
old_ready = threading.Event()
old_ready.set()
reasoning = {"effort": "high"}
old_override = {"model": "old/model", "provider": "removed-provider"}
session = _session(
agent_ready=old_ready,
agent_build_started=True,
agent_error="Unknown provider 'removed-provider'",
model_override=old_override,
resume_runtime_overrides={
"model_override": old_override,
"provider_override": "removed-provider",
"reasoning_config_override": reasoning,
},
)
session["agent"] = None
server._sessions["sid"] = session
seen = {"build": 0, "persist": 0}
def fake_apply_model_switch(_sid, current, _value, **_kwargs):
current["model_override"] = {
"model": "new/model",
"provider": "custom:new-provider",
}
return {
"value": "new/model",
"warning": "",
"confirm_required": False,
"scope": "session",
}
def fake_start_agent_build(sid, current):
seen["build"] += 1
assert sid == "sid"
assert current["agent_error"] is None
assert current["agent_ready"] is not old_ready
assert current.get("agent_build_started") is None
overrides = current["resume_runtime_overrides"]
assert overrides["model_override"] == current["model_override"]
assert overrides["provider_override"] == "custom:new-provider"
assert overrides["reasoning_config_override"] == reasoning
current["agent"] = types.SimpleNamespace(model="new/model")
current["agent_ready"].set()
monkeypatch.setattr(server, "_apply_model_switch", fake_apply_model_switch)
monkeypatch.setattr(server, "_start_agent_build", fake_start_agent_build)
monkeypatch.setattr(
server,
"_persist_live_session_runtime",
lambda _session: seen.__setitem__("persist", seen["persist"] + 1),
)
try:
resp = server.handle_request(
{
"id": "1",
"method": "config.set",
"params": {
"session_id": "sid",
"key": "model",
"value": f"new/model{provider_flag}",
},
}
)
assert resp["result"]["value"] == "new/model"
assert seen == {"build": 1, "persist": 1}
assert session["agent_error"] is None
assert session["agent"].model == "new/model"
finally:
server._sessions.pop("sid", None)
def test_config_set_model_explicit_provider_surfaces_selected_provider_errors(monkeypatch):
seen = {"build": 0, "wait": 0}
session = _session()
+32 -1
View File
@@ -13739,7 +13739,14 @@ def _(rid, params: dict) -> dict:
)
parsed_flags = parse_model_switch_args(value)
explicit_provider = parsed_flags.explicit_provider
if session.get("agent") is None and not explicit_provider.strip():
failed_agent_init = session.get("agent") is None and bool(
session.get("agent_error")
)
if (
session.get("agent") is None
and not explicit_provider.strip()
and not failed_agent_init
):
session_id = params.get("session_id", "")
_start_agent_build(session_id, session)
init_err = _wait_agent(session, rid)
@@ -13756,6 +13763,30 @@ def _(rid, params: dict) -> dict:
),
parsed_flags=parsed_flags,
)
if failed_agent_init and not result.get("confirm_required"):
model_override = session.get("model_override")
resume_overrides = session.get("resume_runtime_overrides")
if isinstance(model_override, dict) and isinstance(
resume_overrides, dict
):
resume_overrides = dict(resume_overrides)
resume_overrides["model_override"] = model_override
if provider := model_override.get("provider"):
resume_overrides["provider_override"] = provider
else:
resume_overrides.pop("provider_override", None)
session["resume_runtime_overrides"] = resume_overrides
session["agent_error"] = None
session["agent_ready"] = threading.Event()
session.pop("agent_build_started", None)
session.pop("_agent_build_thread", None)
_start_agent_build(params.get("session_id", ""), session)
init_err = _wait_agent(session, rid)
if init_err:
return init_err
if session.get("agent") is None:
return _err(rid, 5032, "agent initialization failed")
_persist_live_session_runtime(session)
else:
result = _apply_model_switch(
"",