diff --git a/tests/test_tui_gateway_server.py b/tests/test_tui_gateway_server.py index c55798fa04..c9e65bc002 100644 --- a/tests/test_tui_gateway_server.py +++ b/tests/test_tui_gateway_server.py @@ -9272,34 +9272,49 @@ def test_config_set_model_explicit_provider_skips_broken_default_init(monkeypatc @pytest.mark.parametrize( - "provider_flag", [" --provider custom:new-provider", ""] + ("provider_flag", "failure_text"), + [ + (" --provider custom:new-provider", "Unknown provider 'removed-provider'"), + ("", ""), + ], ) def test_config_set_model_recovers_failed_profile_resume_after_build_completes( - monkeypatch, tmp_path, provider_flag + monkeypatch, tmp_path, provider_flag, failure_text ): - """Recovery waits for the failed generation and uses the session profile. + """Recovery waits for the real failed build and uses its owning profile. - Exercise the real model-switch and deferred-build boundaries: only their - provider/network and agent-construction leaves are replaced. + Both the failed and replacement generations cross the real deferred-build + boundary. Provider resolution is the only model-switch leaf replaced. """ from agent.secret_scope import current_secret_scope from hermes_constants import get_hermes_home + launch_url = "https://launch.example/v1" + profile_url = "https://profile.example/v1" launch_home = tmp_path / "launch" profile_home = tmp_path / "profiles" / "work" launch_home.mkdir() profile_home.mkdir(parents=True) (launch_home / "config.yaml").write_text( - "model:\n default: launch/model\n provider: launch-provider\n", + "model:\n" + " default: launch/model\n" + " provider: custom:new-provider\n" + "providers:\n" + " new-provider:\n" + f" base_url: {launch_url}\n" + " key_env: LAUNCH_API_KEY\n", encoding="utf-8", ) + (launch_home / ".env").write_text( + "LAUNCH_API_KEY=launch-secret\n", encoding="utf-8" + ) (profile_home / "config.yaml").write_text( "model:\n" " default: old/model\n" " provider: custom:new-provider\n" "providers:\n" " new-provider:\n" - " base_url: https://profile.example/v1\n" + f" base_url: {profile_url}\n" " key_env: PROFILE_API_KEY\n", encoding="utf-8", ) @@ -9330,8 +9345,7 @@ def test_config_set_model_recovers_failed_profile_resume_after_build_completes( old_override = {"model": "old/model", "provider": "removed-provider"} session = _session( agent_ready=old_ready, - agent_build_started=True, - agent_error="Unknown provider 'removed-provider'", + agent_error=None, model_override=old_override, resume_runtime_overrides={ "model_override": old_override, @@ -9342,53 +9356,86 @@ def test_config_set_model_recovers_failed_profile_resume_after_build_completes( ) session["agent"] = None server._sessions["sid"] = session + old_finally_entered = threading.Event() + release_old_finally = threading.Event() switch_called = threading.Event() - seen = {"switch": None, "build": None, "persist": 0} - - result = types.SimpleNamespace( - success=True, - new_model="new/model", - target_provider="custom:new-provider", - api_key="profile-secret", - base_url="https://profile.example/v1", - api_mode="chat_completions", - warning_message="", - model_info=None, - error_message="", - ) + seen = {"switch": None, "build": None, "persisted": []} + make_calls = 0 def fake_switch_model(**kwargs): + provider = kwargs["user_providers"]["new-provider"] + secrets = dict(current_secret_scope() or {}) + api_key = secrets[provider["key_env"]] seen["switch"] = { "home": get_hermes_home(), - "secrets": dict(current_secret_scope() or {}), - "providers": kwargs["user_providers"], + "secrets": secrets, + "base_url": provider["base_url"], + "api_key": api_key, "current_provider": kwargs["current_provider"], "current_api_key": kwargs["current_api_key"], } switch_called.set() - return result + return types.SimpleNamespace( + success=True, + new_model="new/model", + target_provider="custom:new-provider", + api_key=api_key, + base_url=provider["base_url"], + api_mode="chat_completions", + warning_message="", + model_info=None, + error_message="", + ) class FakeDb: - def __init__(self, **_kwargs): + def __init__(self, *_args, **_kwargs): pass + def get_session(self, _key): + return {"model_config": {}} + + def update_session_meta(self, key, model_config, model): + seen["persisted"].append( + { + "key": key, + "model": model, + "config": json.loads(model_config), + } + ) + def close(self): pass def fake_make_agent(_sid, _key, **kwargs): + nonlocal make_calls + make_calls += 1 + if make_calls == 1: + raise RuntimeError(failure_text) + override = kwargs["model_override"] seen["build"] = { "home": get_hermes_home(), "secrets": dict(current_secret_scope() or {}), "overrides": kwargs, } return types.SimpleNamespace( - model="new/model", - provider="custom:new-provider", - base_url="https://profile.example/v1", - api_mode="chat_completions", + model=override["model"], + provider="custom", + base_url=override["base_url"], + api_key=override["api_key"], + api_mode=override["api_mode"], + reasoning_config=kwargs.get("reasoning_config_override"), + service_tier=None, _session_db=kwargs.get("session_db"), ) + real_transfer = server._transfer_db_to_agent + + def barrier_transfer(agent, db): + if agent is None and not old_finally_entered.is_set(): + old_finally_entered.set() + assert release_old_finally.wait(timeout=10) + return real_transfer(agent, db) + monkeypatch.setattr("hermes_cli.model_switch.switch_model", fake_switch_model) monkeypatch.setattr( "hermes_cli.model_selection_guards.combined_selection_warning", @@ -9396,6 +9443,7 @@ def test_config_set_model_recovers_failed_profile_resume_after_build_completes( ) monkeypatch.setattr("hermes_state.SessionDB", FakeDb) monkeypatch.setattr(server, "_make_agent", fake_make_agent) + monkeypatch.setattr(server, "_transfer_db_to_agent", barrier_transfer) monkeypatch.setattr( "tui_gateway.entry.ensure_mcp_discovery_started", lambda: None ) @@ -9406,17 +9454,12 @@ def test_config_set_model_recovers_failed_profile_resume_after_build_completes( monkeypatch.setattr(server, "_probe_config_health", lambda *_args: None) monkeypatch.setattr(server, "_schedule_mcp_late_refresh", lambda *a, **k: None) monkeypatch.setattr(server, "_emit", lambda *a, **k: None) - monkeypatch.setattr( - server, - "_persist_live_session_runtime", - lambda _session: seen.__setitem__("persist", seen["persist"] + 1), - ) - response = [] + response = {} def run_request(): - response.append( - server.handle_request( + try: + response["value"] = server.handle_request( { "id": "1", "method": "config.set", @@ -9427,30 +9470,35 @@ def test_config_set_model_recovers_failed_profile_resume_after_build_completes( }, } ) - ) + except BaseException as exc: + response["error"] = exc request_thread = threading.Thread(target=run_request) + old_build_thread = None try: + server._start_agent_build("sid", session) + old_build_thread = session["_agent_build_thread"] + assert old_finally_entered.wait(timeout=10) + assert session["agent_error"] == failure_text + assert not old_ready.is_set() + request_thread.start() assert old_ready.wait_entered.wait(timeout=2), ( "model recovery did not wait for the failed build generation" ) assert not switch_called.is_set() - old_ready.set() + release_old_finally.set() request_thread.join(timeout=10) assert not request_thread.is_alive() - assert response[0]["result"]["value"] == "new/model" - assert seen["persist"] == 1 + assert "error" not in response + assert response["value"]["result"]["value"] == "new/model" + assert make_calls == 2 assert seen["switch"] == { "home": profile_home, "secrets": {"PROFILE_API_KEY": "profile-secret"}, - "providers": { - "new-provider": { - "base_url": "https://profile.example/v1", - "key_env": "PROFILE_API_KEY", - } - }, + "base_url": profile_url, + "api_key": "profile-secret", "current_provider": ( "custom:new-provider" if provider_flag else "custom" ), @@ -9466,9 +9514,30 @@ def test_config_set_model_recovers_failed_profile_resume_after_build_completes( assert overrides["reasoning_config_override"] == reasoning assert session["agent_error"] is None assert session["agent"].model == "new/model" + assert session["agent"].base_url == profile_url + assert session["agent"].api_key == "profile-secret" + assert seen["persisted"] == [ + { + "key": "session-key", + "model": "new/model", + "config": { + "model": "new/model", + "provider": "custom:new-provider", + "base_url": profile_url, + "api_mode": "chat_completions", + "reasoning_config": reasoning, + }, + } + ] finally: + release_old_finally.set() old_ready.set() request_thread.join(timeout=10) + if old_build_thread is not None: + old_build_thread.join(timeout=10) + new_build_thread = session.get("_agent_build_thread") + if new_build_thread is not None: + new_build_thread.join(timeout=10) server._sessions.pop("sid", None) diff --git a/tui_gateway/server.py b/tui_gateway/server.py index 444384bddf..6f0f6ee2ef 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -6153,6 +6153,39 @@ def _session_profile_runtime_scope(session: dict): reset_hermes_home_override(home_token) +def _restart_completed_failed_agent_build( + sid: str, session: dict, failed_ready: threading.Event | None +) -> bool: + """Replace one completed failed build generation and start its retry.""" + if failed_ready is None: + return False + build_lock = session.setdefault("agent_build_lock", threading.Lock()) + with build_lock: + if ( + session.get("agent") is not None + or session.get("agent_error") is None + or session.get("agent_ready") is not failed_ready + or not failed_ready.is_set() + ): + return False + 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(sid, session) + return True + + def _apply_model_switch( sid: str, session: dict, @@ -13755,21 +13788,27 @@ def _(rid, params: dict) -> dict: ) parsed_flags = parse_model_switch_args(value) explicit_provider = parsed_flags.explicit_provider - failed_agent_init = session.get("agent") is None and bool( - session.get("agent_error") + failed_agent_init = ( + session.get("agent") is None + and session.get("agent_error") is not None ) failed_ready = session.get("agent_ready") if failed_agent_init else None - if ( - failed_agent_init - and failed_ready is not None - and not failed_ready.wait(timeout=30.0) - ): - return _err(rid, 5032, "agent initialization timed out") + if failed_agent_init: + if failed_ready is None: + return _err( + rid, + 5032, + session.get("agent_error") + or "agent initialization failed", + ) + if not failed_ready.wait(timeout=30.0): + return _err(rid, 5032, "agent initialization timed out") failed_agent_init = ( failed_agent_init and session.get("agent") is None - and bool(session.get("agent_error")) + and session.get("agent_error") is not None and session.get("agent_ready") is failed_ready + and failed_ready.is_set() ) if ( session.get("agent") is None @@ -13794,42 +13833,16 @@ def _(rid, params: dict) -> dict: parsed_flags=parsed_flags, ) if failed_agent_init and not result.get("confirm_required"): - restart_build = False - build_lock = session.setdefault( - "agent_build_lock", threading.Lock() + _restart_completed_failed_agent_build( + params.get("session_id", ""), session, failed_ready ) - with build_lock: - if ( - session.get("agent") is None - and session.get("agent_error") - and session.get("agent_ready") is failed_ready - and (failed_ready is None or failed_ready.is_set()) - ): - 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) - restart_build = True - if restart_build: - _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) + with _session_profile_runtime_scope(session): + _persist_live_session_runtime(session) else: result = _apply_model_switch( "",