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:
215
tests/tui_gateway/test_deferred_model_switch_confirm.py
Normal file
215
tests/tui_gateway/test_deferred_model_switch_confirm.py
Normal 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}
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user