fix(tui): harden failed build recovery handoff

This commit is contained in:
Inimitable Mind
2026-08-22 17:07:11 -07:00
committed by Teknium
parent 258d0283b0
commit 4e860c093a
2 changed files with 169 additions and 87 deletions
+117 -48
View File
@@ -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)
+52 -39
View File
@@ -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(
"",