fix(tui): harden failed build recovery handoff
This commit is contained in:
@@ -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
@@ -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(
|
||||
"",
|
||||
|
||||
Reference in New Issue
Block a user