fix(tui_gateway): ask before queueing a guarded model picked mid-turn (#91043)

* fix(tui_gateway): ask before queueing a guarded model picked mid-turn

config.set model on a running session cannot swap the agent in place, so it
stashes the pick in session["pending_model_switch"] and applies it at the
next turn start. That branch answered confirm_required=False without ever
running the selection guards.

A client that implements the confirm round-trip was therefore told no
consent was needed and never prompted. One turn later
_apply_pending_model_switch ran the guards with the stashed (unconfirmed)
flag, saw the warning, and dropped the switch by design. The model reverted
with no confirm ever offered, because the only moment a round-trip was
possible had already passed.

Evaluate the guards before stashing, where the client still has a live
response to turn into a prompt. Nothing is queued for an unconfirmed
guarded pick, so the session is left exactly as it was and the re-send
carrying confirm_expensive_model queues it for real. The apply-time check
stays as the backstop for guards that can only decide after resolution.

The data-policy guard keys on the model id alone, which is all this branch
can see before resolution. The cost guard returns None when pricing is
unknown and its models.dev lookup is allow_network=False, so calling it
early can only under-fire and never blocks the RPC thread.

* test(tui_gateway): pin provider forwarding, name the canonical confirm field

Two review follow-ups, no behavior change.

_pending_switch_selection_warning forwards `provider=provider or None`, but
nothing asserted it: a guarded model id fires the data-policy guard on the
model alone, so the existing tests passed with `provider` dropped entirely.
Record the kwargs instead. Dropping the argument fails the first test;
removing the `or None` normalization fails the second.

The confirm responses carry `warning` and `confirm_message` with identical
text, which reads like an accident. Name which one clients should read
(`confirm_message`; `warning` is the pre-confirm-era alias that
_apply_pending_model_switch already treats as a fallback) so the two do not
drift apart later.

Both raised by @Enough1122 in review.
This commit is contained in:
Jack Lau
2026-08-26 13:38:53 -05:00
committed by GitHub
parent b519ce29ad
commit 2552579912
2 changed files with 288 additions and 6 deletions

View File

@@ -0,0 +1,215 @@
"""A model picked mid-turn must still get its selection-guard confirm step.
``config.set model`` on a *running* session cannot swap the agent in place --
the worker thread is reading ``agent.model`` / ``agent.client`` on every
iteration -- so it stashes the pick in ``session["pending_model_switch"]`` and
``_apply_pending_model_switch`` applies it at the next turn start.
That deferral used to skip the selection guards entirely: the stash branch
answered ``confirm_required: False`` without ever calling them. A client that
implements the confirm round-trip was therefore told no consent was needed, so
it never prompted. One turn later ``_apply_pending_model_switch`` ran the
guards with the stashed (unconfirmed) flag, saw the warning, and deliberately
dropped the switch -- correct on its own terms, but by then no round-trip was
possible. The user's pick silently reverted and the confirm was never offered
on this path at all.
"""
import threading
import types
import pytest
from tui_gateway import server
# A vendor-documented data-training tier. The data-policy guard keys on the
# model id alone (no base_url / api_key / model_info), which is exactly what
# the stash branch can see before resolution.
GUARDED_MODEL = "muse-spark-1.2-contributor"
UNGUARDED_MODEL = "anthropic/claude-sonnet-4.6"
def _session(**extra):
return {
"agent": types.SimpleNamespace(),
"session_key": "session-key",
"history": [],
"history_lock": threading.Lock(),
"history_version": 0,
"running": False,
"attached_images": [],
"image_counter": 0,
"cols": 80,
"slash_worker": None,
"show_reasoning": False,
"tool_progress_mode": "all",
**extra,
}
def _config_set_model(value, **extra_params):
params = {"session_id": "sid", "key": "model", "value": value}
params.update(extra_params)
return server.handle_request({"id": "1", "method": "config.set", "params": params})
@pytest.fixture
def running_session(monkeypatch):
"""A busy session whose live swap path is fatal if it is ever reached."""
def _must_not_run(*_args, **_kwargs):
raise AssertionError(
"_apply_model_switch ran on the busy path -- it would race the "
"worker thread reading agent.model / agent.client"
)
monkeypatch.setattr(server, "_apply_model_switch", _must_not_run)
server._sessions["sid"] = _session(running=True)
try:
yield server._sessions["sid"]
finally:
server._sessions.pop("sid", None)
class TestGuardedPickAsksBeforeStashing:
def test_reports_confirm_required_instead_of_deferring(self, running_session):
resp = _config_set_model(GUARDED_MODEL)
assert not resp.get("error")
result = resp["result"]
assert result["confirm_required"] is True, (
"the deferred path answered confirm_required=False without running "
"the guards, so a correct client never prompts and the pick is "
"dropped a turn later with no way to consent"
)
assert result["confirm_message"].strip()
assert result["deferred"] is False
def test_leaves_the_session_untouched(self, running_session):
_config_set_model(GUARDED_MODEL)
assert "pending_model_switch" not in running_session, (
"an unconfirmed guarded pick must not be queued -- the next turn "
"start would drop it anyway, after the pill already moved"
)
def test_confirm_message_names_the_guard(self, running_session):
message = _config_set_model(GUARDED_MODEL)["result"]["confirm_message"]
assert "CONTRIBUTOR TIER" in message
assert "train" in message.lower()
def test_reconfirming_queues_the_pick(self, running_session):
resp = _config_set_model(GUARDED_MODEL, confirm_expensive_model=True)
result = resp["result"]
assert result["deferred"] is True
assert result["confirm_required"] is False
pending = running_session["pending_model_switch"]
assert pending["raw"] == GUARDED_MODEL
assert pending["confirm_expensive_model"] is True, (
"the ack must survive into the stash or _apply_pending_model_switch "
"re-runs the guard at turn start and drops the confirmed pick"
)
class TestUnguardedPickStillDefers:
"""The queue-don't-race behaviour is the whole point of this branch."""
def test_defers_without_a_confirm_step(self, running_session):
result = _config_set_model(UNGUARDED_MODEL)["result"]
assert result["deferred"] is True
assert result["confirm_required"] is False
assert result["confirm_message"] == ""
assert result["value"] == UNGUARDED_MODEL
def test_stashes_the_pick_for_the_next_turn(self, running_session):
_config_set_model(UNGUARDED_MODEL)
pending = running_session["pending_model_switch"]
assert pending["raw"] == UNGUARDED_MODEL
assert pending["confirm_expensive_model"] is False
def test_explicit_provider_is_still_recorded_for_display(self, running_session):
_config_set_model(f"{UNGUARDED_MODEL} --provider anthropic")
pending = running_session["pending_model_switch"]
assert pending["display_provider"] == "anthropic"
class TestGuardFailureIsNotFatal:
def test_a_raising_guard_falls_back_to_deferring(self, running_session, monkeypatch):
"""A broken guard must never cost the user their model pick.
The apply-time check in ``_apply_pending_model_switch`` is still there,
so failing open here degrades to the old behaviour rather than to a
silently unguarded switch.
"""
def _boom(*_args, **_kwargs):
raise RuntimeError("guard table is broken")
monkeypatch.setattr(
"hermes_cli.model_selection_guards.combined_selection_warning", _boom
)
result = _config_set_model(GUARDED_MODEL)["result"]
assert result["deferred"] is True
assert running_session["pending_model_switch"]["raw"] == GUARDED_MODEL
class TestHelperContract:
def test_returns_none_for_an_empty_model(self):
assert server._pending_switch_selection_warning("", "") is None
def test_returns_none_when_no_guard_fires(self):
assert server._pending_switch_selection_warning(UNGUARDED_MODEL, "") is None
def test_returns_the_message_when_a_guard_fires(self):
message = server._pending_switch_selection_warning(GUARDED_MODEL, "")
assert message is not None
assert "CONTRIBUTOR TIER" in message
def test_an_explicit_provider_reaches_the_guards(self, monkeypatch):
"""Provider-keyed guards are useless if the provider is dropped here.
The docstring promises the early call can only under-fire relative to
the resolved one, and that only holds if what the caller DID say is
forwarded. Asserting on a guarded model id would pass even with
``provider`` dropped, so record the kwargs instead.
"""
seen = {}
def _fake(model, provider=None, **kwargs):
seen["model"] = model
seen["provider"] = provider
return None
import hermes_cli.model_selection_guards as guards
monkeypatch.setattr(guards, "combined_selection_warning", _fake)
server._pending_switch_selection_warning(UNGUARDED_MODEL, "openrouter")
assert seen == {"model": UNGUARDED_MODEL, "provider": "openrouter"}
def test_an_empty_provider_is_normalised_to_none(self, monkeypatch):
"""`provider or None` is load-bearing: "" is not "no provider" to a
guard that does an `is None` check, and the TUI sends "" for unset."""
seen = {}
def _fake(model, provider=None, **kwargs):
seen["provider"] = provider
return None
import hermes_cli.model_selection_guards as guards
monkeypatch.setattr(guards, "combined_selection_warning", _fake)
server._pending_switch_selection_warning(UNGUARDED_MODEL, "")
assert seen == {"provider": None}

View File

@@ -5653,6 +5653,9 @@ def _apply_model_switch(
confirm_msg = warning.message
if result.warning_message:
confirm_msg = f"{confirm_msg}\n\n{result.warning_message}"
# Same contract as the deferred branch below: confirm_message is
# canonical, warning is the pre-confirm-era alias. Identical by
# design, not by accident.
return {
"value": result.new_model,
"warning": confirm_msg,
@@ -5834,6 +5837,32 @@ def _sync_agent_model_with_config(sid: str, session: dict) -> None:
)
def _pending_switch_selection_warning(model: str, provider: str) -> str | None:
"""Selection-guard message for a model queued mid-turn, or ``None``.
Runs BEFORE the pick is stashed, while the client still has a live response
it can turn into a confirm prompt. Only pre-resolution inputs exist here --
the model id the user picked and any explicit ``--provider`` -- which is
exactly what the data-policy guard keys on. Guards that can only decide
once base_url / api_key / model_info have settled still get their chance in
``_apply_model_switch``; the cost guard returns ``None`` when pricing is
unknown, so an early call can only under-fire, never over-fire.
A misbehaving guard must never break the pick, so exceptions are swallowed
and treated as "no warning" -- the apply-time check remains the backstop.
"""
if not model:
return None
try:
from hermes_cli.model_selection_guards import combined_selection_warning
warning = combined_selection_warning(model, provider=provider or None)
except Exception:
return None
return warning.message if warning is not None else None
def _apply_pending_model_switch(sid: str, session: dict) -> None:
"""Apply a model switch queued while a turn was running.
@@ -12704,19 +12733,57 @@ def _(rid, params: dict) -> dict:
pending_model = parsed.model_input
except Exception:
pending_model = str(value)
pending_provider = (
getattr(parsed, "explicit_provider", "") or ""
).strip()
confirmed = bool(params.get("confirm_expensive_model", False))
# Run the selection guards HERE, not only at apply time.
# This branch used to answer confirm_required=False without
# consulting them, so a client that implements the confirm
# round-trip was told no consent was needed. It stashed the
# pick, and _apply_pending_model_switch -- which calls the
# guards with the stashed (unconfirmed) flag -- dropped the
# switch at the next turn start. The model reverted with no
# confirm ever offered, because the one moment a round-trip
# was possible had already passed.
if not confirmed:
pending_warning = _pending_switch_selection_warning(
pending_model, pending_provider
)
if pending_warning is not None:
# Nothing is stashed: an unconfirmed guarded pick
# leaves the session exactly as it was, and the
# client re-sends with confirm_expensive_model to
# queue it for real.
return _ok(
rid,
{
"key": key,
"value": pending_model,
# `confirm_message` is the field to read.
# `warning` carries the same text only so
# clients written before the confirm
# round-trip existed still show something;
# `_apply_pending_model_switch` already
# prefers confirm_message and falls back to
# warning. Keep them identical or drop
# `warning` -- do not let them diverge.
"warning": pending_warning,
"confirm_required": True,
"confirm_message": pending_warning,
"scope": "session",
"deferred": False,
},
)
session["pending_model_switch"] = {
"raw": value,
"confirm_expensive_model": bool(
params.get("confirm_expensive_model", False)
),
"confirm_expensive_model": confirmed,
# The resolved model/provider the next turn will run on.
# _session_info reports these while the switch is pending
# so the end-of-turn settle keeps showing the user's pick
# instead of blipping back to the still-live old model.
"display_model": pending_model,
"display_provider": (
getattr(parsed, "explicit_provider", "") or ""
).strip(),
"display_provider": pending_provider,
}
return _ok(
rid,