fix(tui): recover model switch after failed resume
This commit is contained in:
@@ -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
@@ -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(
|
||||
"",
|
||||
|
||||
Reference in New Issue
Block a user