diff --git a/tests/test_tui_gateway_server.py b/tests/test_tui_gateway_server.py index 48addce8b2..d196047ea2 100644 --- a/tests/test_tui_gateway_server.py +++ b/tests/test_tui_gateway_server.py @@ -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() diff --git a/tui_gateway/server.py b/tui_gateway/server.py index ccd7bafd07..0a27a7811a 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -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( "",