From f2e1fe782875e4e59350662e9efc68600e193406 Mon Sep 17 00:00:00 2001 From: Josh Stevenson Date: Sun, 4 Oct 2026 01:35:15 -0700 Subject: [PATCH] fix(gateway): preserve model intent across turn admission --- .../test_model_intent_admission_order.py | 535 ++++++++++++++++++ tui_gateway/server.py | 261 ++++++--- 2 files changed, 721 insertions(+), 75 deletions(-) create mode 100644 tests/tui_gateway/test_model_intent_admission_order.py diff --git a/tests/tui_gateway/test_model_intent_admission_order.py b/tests/tui_gateway/test_model_intent_admission_order.py new file mode 100644 index 000000000000..70e3c2139c67 --- /dev/null +++ b/tests/tui_gateway/test_model_intent_admission_order.py @@ -0,0 +1,535 @@ +"""Model intent ordering through real gateway admission and inert runtimes.""" +import copy +import errno +import threading +from concurrent.futures import ThreadPoolExecutor +from types import SimpleNamespace + +import pytest + +from hermes_cli import model_switch as ms +from hermes_state import SessionDB +from tests.hermes_cli.test_model_switch_route_contract import routes # noqa: F401 +from tests.hermes_cli.test_model_switch_runtime_boundary import agent, assert_route # noqa: F401 +from tests.tui_gateway.test_failed_turn_retention import turn_env # noqa: F401 +from tests.tui_gateway.test_model_once_turn_runtime_owner import live_turn # noqa: F401 +from tests.tui_gateway.test_model_switch_target_session_contract import gateway # noqa: F401 +from tests.tui_gateway.test_config_dispatch_responsiveness import Transport +from tui_gateway import server +from tui_gateway.transport import StdioTransport + +REAL_THREAD = threading.Thread +REAL_CONFIG_SYNC = server._sync_agent_model_with_config +REAL_RESTART = server._restart_slash_worker + + +@pytest.fixture +def intent_turn(live_turn, routes, monkeypatch): + session, events = live_turn + session['running'] = False + session['transport'] = Transport() + routes.config['providers']['endpoint-c'] = { + 'name': 'Endpoint C', 'base_url': 'http://c.invalid/v1', + 'api_key': 'synthetic-key-c', 'default_model': 'model-c'} + routes.write() + db = SessionDB(routes.home / 'state.db') + db.create_session(session['session_key'], source='tui') + session['agent']._session_db = db + monkeypatch.setattr(server, '_db', db) + monkeypatch.setattr(server, '_session_uses_compute_host', lambda *a, **kw: False) + monkeypatch.setattr(server, '_voice_mode_enabled', lambda: False) + monkeypatch.setattr(server, '_ensure_active_session_slot', lambda *a: None) + monkeypatch.setattr(server, '_sync_bot_capabilities', lambda *a: None) + yield session, events + db.close() + + +def select(model, provider, *, once=False, confirmed=True): + raw = f'{model} --provider {provider} ' + ('--once' if once else '--session') + return {'id': 'pick-' + model, 'method': 'config.set', 'params': { + 'session_id': 'selected', 'key': 'model', 'value': raw, + 'confirm_expensive_model': confirmed}} + + +def run_turn(session, monkeypatch, rid='offline-turn'): + observations = [] + value = session['agent'] + + def offline(*a, **kw): + info = server._session_info(value, session) + observations.append((value.model, value.provider, info['model'], info['provider'])) + return {'final_response': 'offline', 'completed': False, 'api_calls': 0} + + monkeypatch.setattr(value, 'run_conversation', offline) + session['running'] = True + server._run_prompt_submit(rid, 'selected', session, 'offline input') + session['_run_thread'].join(5) + assert not session['_run_thread'].is_alive() + return observations + + +@pytest.mark.parametrize('pause', ['admitted', 'executing', 'steered']) +def test_slow_idle_selection_cannot_repoint_admitted_send(intent_turn, monkeypatch, pause): + session, _ = intent_turn + value = session['agent'] + transport = session['transport'] + entered, resolve_release = threading.Event(), threading.Event() + turn_entered, turn_release = threading.Event(), threading.Event() + real_switch = ms.switch_model + samples = {} + monkeypatch.setattr(server.threading, 'Thread', REAL_THREAD) + + def slow(*a, **kw): + entered.set() + assert resolve_release.wait(5) + return real_switch(*a, **kw) + + def ready(*a, **kw): + if pause != 'executing': + turn_entered.set() + assert turn_release.wait(5) + return None + + def offline(*a, **kw): + samples['start'] = (value.model, value.provider, value.client) + if pause == 'executing': + turn_entered.set() + assert turn_release.wait(5) + samples['end'] = (value.model, value.provider, value.client) + return {'final_response': 'offline', 'completed': False, 'api_calls': 0} + + monkeypatch.setattr(ms, 'switch_model', slow) + monkeypatch.setattr(server, '_wait_agent_for_prompt', ready) + monkeypatch.setattr(value, 'run_conversation', offline) + original = (value.model, value.provider, value.client) + with ThreadPoolExecutor(max_workers=1) as pool: + monkeypatch.setattr(server, '_pool', pool) + assert server.dispatch(select('model-b', 'endpoint-b'), transport) is None + try: + assert entered.wait(5) + ack = server.dispatch({'id': 'send-a', 'method': 'prompt.submit', 'params': { + 'session_id': 'selected', 'text': 'offline input'}}, transport) + assert not ack.get('error'), ack + assert ack['result']['status'] == 'streaming' + assert turn_entered.wait(5) + assert session['running'] + resolve_release.set() + assert transport.done.wait(5) + if pause == 'steered': + admitted_snapshot = session['inflight_turn'] + correction = server.dispatch({'id': 'steer-a', 'method': 'session.steer', + 'params': {'session_id': 'selected', 'text': 'offline correction'}}, transport) + assert correction['result']['status'] == 'queued' + assert session['inflight_turn'] is not admitted_snapshot + assert session['inflight_turn']['started_at'] is admitted_snapshot['started_at'] + finally: + resolve_release.set() + turn_release.set() + thread = session.get('_run_thread') + while thread is not None: + thread.join(5) + assert not thread.is_alive() + successor = session.get('_run_thread') + if successor is thread: + break + thread = successor + assert not session['_run_thread'].is_alive() + assert samples == {'start': original, 'end': original} + result = transport.frames[0]['result'] + assert result['deferred'] is True and result['confirm_required'] is False + assert session['pending_model_switch']['display_model'] == 'model-b' + assert_route(value, 'http://a.invalid/v1', 'synthetic-key-a') + # The intent becomes eligible only after this admission has settled. + monkeypatch.setattr(ms, 'switch_model', real_switch) + assert run_turn(session, monkeypatch, 'next-b') == [ + ('model-b', 'endpoint-b', 'model-b', 'endpoint-b')] + assert_route(value, 'http://b.invalid/v1', 'synthetic-key-b') + + +def test_live_swap_and_once_publication_share_admission_lock(intent_turn, monkeypatch): + session, _ = intent_turn + value = session['agent'] + real_switch = value.switch_model + lock_observations = [] + + def switch(*a, **kw): + acquired = session['history_lock'].acquire(blocking=False) + if acquired: + session['history_lock'].release() + lock_observations.append(acquired) + return real_switch(*a, **kw) + + monkeypatch.setattr(value, 'switch_model', switch) + response = server.handle_request(select('model-b', 'endpoint-b', once=True)) + assert not response.get('error'), response + assert lock_observations == [False] + snapshot, runtime = server._consume_one_turn_model_runtime(session, value) + assert snapshot['model'] == 'model-a' and runtime['active'] + assert_route(value, 'http://b.invalid/v1', 'synthetic-key-b') + + +def test_new_idle_choice_retires_older_deferred_choice(intent_turn, monkeypatch, routes): + session, _ = intent_turn + other_pending = {'raw': 'other --session', 'display_model': 'other'} + server._sessions['other'] = {'pending_model_switch': other_pending} + before = (routes.home / 'config.yaml').read_bytes() + session['running'] = True + queued = server.handle_request(select('model-b', 'endpoint-b')) + assert queued['result']['deferred'] is True + session['running'] = False + chosen = server.handle_request(select('model-c', 'endpoint-c')) + assert not chosen.get('error'), chosen + assert chosen['result']['value'] == 'model-c' + assert 'pending_model_switch' not in session + info = server._session_info(session['agent'], session) + assert (info['model'], info['provider']) == ('model-c', 'endpoint-c') + assert run_turn(session, monkeypatch) == [('model-c', 'endpoint-c', 'model-c', 'endpoint-c')] + assert_route(session['agent'], 'http://c.invalid/v1', 'synthetic-key-c') + assert server._sessions['other']['pending_model_switch'] is other_pending + assert (routes.home / 'config.yaml').read_bytes() == before + + +@pytest.mark.parametrize('failure', ['disabled', 'client', 'unconfirmed']) +def test_unsuccessful_idle_choice_preserves_deferred_intent(intent_turn, routes, monkeypatch, failure): + session, _ = intent_turn + value = session['agent'] + session['running'] = True + assert server.handle_request(select('model-b', 'endpoint-b'))['result']['deferred'] + pending = session['pending_model_switch'] + session['running'] = False + old_client = value.client + if failure == 'disabled': + routes.config['providers']['endpoint-c']['enabled'] = False + routes.write() + elif failure == 'client': + create = value._create_openai_client + + def fail(kwargs, **kw): + if kwargs['base_url'] == 'http://c.invalid/v1': + raise RuntimeError('offline client construction failed') + return create(kwargs, **kw) + + monkeypatch.setattr(value, '_create_openai_client', fail) + else: + monkeypatch.setattr('hermes_cli.model_selection_guards.combined_selection_warning', + lambda *a, **kw: SimpleNamespace(message='offline consent required')) + response = server.handle_request(select('model-c', 'endpoint-c', confirmed=failure != 'unconfirmed')) + if failure == 'unconfirmed': + assert response['result']['confirm_required'] is True + else: + assert response['error']['code'] == 5001 + assert session['pending_model_switch'] is pending + assert value.client is old_client + assert_route(value, 'http://a.invalid/v1', 'synthetic-key-a') + assert session['history_lock'].acquire(blocking=False) + session['history_lock'].release() + assert run_turn(session, monkeypatch) == [('model-b', 'endpoint-b', 'model-b', 'endpoint-b')] + + +def test_replacing_unused_once_returns_to_original_runtime(intent_turn, monkeypatch, routes): + session, _ = intent_turn + pin = copy.deepcopy(session['model_override']) + before = (routes.home / 'config.yaml').read_bytes() + assert not server.handle_request(select('model-b', 'endpoint-b', once=True)).get('error') + original_restore = session['one_turn_model_restore'] + original_lease = session['_one_turn_model_runtime'] + assert not server.handle_request(select('model-c', 'endpoint-c', once=True)).get('error') + assert session['one_turn_model_restore'] is original_restore + assert session['_one_turn_model_runtime'] is not original_lease + assert run_turn(session, monkeypatch) == [('model-c', 'endpoint-c', 'model-c', 'endpoint-c')] + value = session['agent'] + assert (value.model, value.provider) == ('model-a', 'endpoint-a') + assert_route(value, 'http://a.invalid/v1', 'synthetic-key-a') + assert session['model_override'] == pin + info = server._session_info(value, session) + assert (info['model'], info['provider']) == ('model-a', 'endpoint-a') + assert not session.get('_one_turn_model_runtime') + assert (routes.home / 'config.yaml').read_bytes() == before + + +def test_durable_choice_supersedes_unused_once_baseline(intent_turn, monkeypatch): + session, _ = intent_turn + assert not server.handle_request(select('model-b', 'endpoint-b', once=True)).get('error') + assert not server.handle_request(select('model-c', 'endpoint-c')).get('error') + assert not session.get('one_turn_model_restore') and not session.get('_one_turn_model_runtime') + assert run_turn(session, monkeypatch) == [('model-c', 'endpoint-c', 'model-c', 'endpoint-c')] + assert session['model_override']['model'] == 'model-c' + assert_route(session['agent'], 'http://c.invalid/v1', 'synthetic-key-c') + + +@pytest.mark.parametrize('newer_choice', [False, True]) +def test_publication_error_does_not_discard_or_resurrect_old_intent(intent_turn, monkeypatch, newer_choice): + session, _ = intent_turn + session['running'] = True + assert server.handle_request(select('model-b', 'endpoint-b'))['result']['deferred'] + pending = session['pending_model_switch'] + session['running'] = False + collect = server._emit + + class FailingStream: + def write(self, text): + raise OSError(errno.EIO, 'offline event publication failed') + + failed_transport = StdioTransport(lambda: FailingStream(), threading.Lock()) + + def publish(name, sid, payload=None): + if name == 'session.info' and payload['model'] == 'model-c': + if newer_choice: + server._mirror_slash_side_effects('selected', session, + '/model model-d --provider endpoint-c --session') + return failed_transport.write(server._event_frame(name, sid, payload)) + return collect(name, sid, payload) + + monkeypatch.setattr(server, '_emit', publish) + response = server.handle_request(select('model-c', 'endpoint-c')) + monkeypatch.setattr(server, '_emit', collect) + assert response['error']['code'] == 5001 + assert session['history_lock'].acquire(blocking=False) + session['history_lock'].release() + if newer_choice: + assert not session.get('pending_model_switch') + assert session['model_override']['model'] == 'model-d' + assert run_turn(session, monkeypatch) == [('model-d', 'endpoint-c', 'model-d', 'endpoint-c')] + else: + assert session['pending_model_switch'] is pending + # The suffix error does not claim to roll back the already working C + # client. B remains the previously acknowledged next-turn intent. + assert session['agent'].model == 'model-c' + assert run_turn(session, monkeypatch) == [('model-b', 'endpoint-b', 'model-b', 'endpoint-b')] + + +def test_internal_one_shot_restore_keeps_user_deferred_pick(intent_turn, monkeypatch): + session, _ = intent_turn + restore = {'override': session['model_override'], 'model': 'model-a', 'provider': 'endpoint-a'} + # Use an inert custom runtime for the one-shot; its actual finally branch + # restores through the same helper used by the MoA one-shot producer. + server._apply_model_switch('selected', session, 'model-c --provider endpoint-c --session', + confirm_expensive_model=True, persist_override=False) + session['moa_one_shot_restore'] = restore + pending = [] + + def offline(*a, **kw): + assert server.handle_request(select('model-b', 'endpoint-b'))['result']['deferred'] + pending.append(session['pending_model_switch']) + return {'final_response': 'offline', 'completed': False, 'api_calls': 0} + + monkeypatch.setattr(session['agent'], 'run_conversation', offline) + session['running'] = True + server._run_prompt_submit('internal-one-shot', 'selected', session, 'offline input') + session['_run_thread'].join(5) + assert not session['_run_thread'].is_alive() + assert session['agent'].model == 'model-a' + assert session['pending_model_switch'] is pending[0] + assert run_turn(session, monkeypatch) == [('model-b', 'endpoint-b', 'model-b', 'endpoint-b')] + + +def test_publication_error_after_once_completion_keeps_old_intent(intent_turn, monkeypatch): + session, _ = intent_turn + session['running'] = True + assert server.handle_request(select('model-b', 'endpoint-b'))['result']['deferred'] + pending = session['pending_model_switch'] + session['running'] = False + collect = server._emit + + class FailingStream: + def write(self, text): + raise OSError(errno.EIO, 'offline once publication failed') + + failed_transport = StdioTransport(lambda: FailingStream(), threading.Lock()) + + def publish(name, sid, payload=None): + if name == 'session.info' and payload['model'] == 'model-c': + monkeypatch.setattr(server, '_emit', collect) + assert run_turn(session, monkeypatch, 'during-once-publication') == [ + ('model-c', 'endpoint-c', 'model-c', 'endpoint-c')] + assert session['agent'].model == 'model-a' + assert not session.get('_one_turn_model_runtime') + restored = server._session_info(session['agent'], session) + assert (restored['model'], restored['provider']) == ('model-a', 'endpoint-a') + return failed_transport.write(server._event_frame(name, sid, payload)) + return collect(name, sid, payload) + + monkeypatch.setattr(server, '_emit', publish) + response = server.handle_request(select('model-c', 'endpoint-c', once=True)) + monkeypatch.setattr(server, '_emit', collect) + assert response['error']['code'] == 5001 + assert session['pending_model_switch'] is pending + assert_route(session['agent'], 'http://a.invalid/v1', 'synthetic-key-a') + assert run_turn(session, monkeypatch) == [('model-b', 'endpoint-b', 'model-b', 'endpoint-b')] + + +@pytest.mark.parametrize('ordering', ['once_success', 'c_failure_first', 'd_failure_first']) +def test_publication_error_after_newer_once_completion_does_not_restore_old_intent(intent_turn, monkeypatch, ordering): + session, _ = intent_turn + transport = session['transport'] + session['running'] = True + assert server.handle_request(select('model-b', 'endpoint-b'))['result']['deferred'] + old_pending = session['pending_model_switch'] + session['running'] = False + entered, release = threading.Event(), threading.Event() + turn_entered, turn_release = threading.Event(), threading.Event() + turn_samples = [] + publisher = [] + newer_publisher = threading.get_ident() + collect = server._emit + monkeypatch.setattr(server.threading, 'Thread', REAL_THREAD) + + class InertWorker: + def __init__(self, block=False): + self._closed = False + self.block = block + + def close(self): + if self._closed: + return + self._closed = True + if self.block: + publisher.append(threading.get_ident()) + entered.set() + assert release.wait(5) + + class FailingStream: + def write(self, text): + raise OSError(errno.EIO, 'offline late publication failed') + + failed_transport = StdioTransport(lambda: FailingStream(), threading.Lock()) + initial_worker = InertWorker(block=True) + session['slash_worker'] = initial_worker + monkeypatch.setattr(server, '_SlashWorker', lambda *a, **kw: InertWorker()) + monkeypatch.setattr(server, '_restart_slash_worker', REAL_RESTART) + + def publish(name, sid, payload=None): + if name == 'session.info' and threading.get_ident() in publisher: + return failed_transport.write(server._event_frame(name, sid, payload)) + if (name == 'session.info' and payload['model'] == 'model-d' + and threading.get_ident() == newer_publisher and ordering != 'once_success'): + if ordering == 'c_failure_first': + release.set() + assert transport.done.wait(5) + return failed_transport.write(server._event_frame(name, sid, payload)) + return collect(name, sid, payload) + + monkeypatch.setattr(server, '_emit', publish) + with ThreadPoolExecutor(max_workers=1) as pool: + monkeypatch.setattr(server, '_pool', pool) + assert server.dispatch(select('model-c', 'endpoint-c'), transport) is None + try: + assert entered.wait(5) and initial_worker._closed + pin = session['model_override'] + if ordering == 'once_success': + server._mirror_slash_side_effects('selected', session, + '/model model-d --provider endpoint-c --once') + assert session['agent'].model == 'model-d' + assert run_turn(session, monkeypatch, 'newer-once') == [ + ('model-d', 'endpoint-c', 'model-d', 'endpoint-c')] + assert session['agent'].model == 'model-c' + assert session['model_override'] is pin + assert not session.get('_one_turn_model_runtime') + else: + output = server._mirror_slash_side_effects('selected', session, + '/model model-d --provider endpoint-c --session') + assert 'publication failed' in output + if ordering == 'd_failure_first': + # C still owns an unfinished suffix: its old B cannot + # become eligible merely because D released its claim. + server._apply_pending_model_switch('selected', session) + assert session['agent'].model == 'model-d' + value = session['agent'] + admitted_runtime = (value.model, value.provider, value.client) + + def ready(*a, **kw): + turn_entered.set() + assert turn_release.wait(5) + return None + + def offline(*a, **kw): + turn_samples.append((value.model, value.provider, value.client)) + return {'final_response': 'offline', 'completed': False, 'api_calls': 0} + + monkeypatch.setattr(server, '_wait_agent_for_prompt', ready) + monkeypatch.setattr(value, 'run_conversation', offline) + ack = server.dispatch({'id': 'send-d-between-failures', 'method': 'prompt.submit', + 'params': {'session_id': 'selected', 'text': 'offline input'}}, transport) + assert not ack.get('error'), ack + assert ack['result']['status'] == 'streaming' + assert turn_entered.wait(5) and session['running'] + release.set() + assert transport.done.wait(5) + assert old_pending['after_inflight_turn'] is session['inflight_turn'] + finally: + release.set() + turn_release.set() + thread = session.get('_run_thread') + while thread is not None: + thread.join(5) + assert not thread.is_alive() + successor = session.get('_run_thread') + if successor is thread: + break + thread = successor + assert transport.done.wait(5) + monkeypatch.setattr(server, '_emit', collect) + assert transport.frames[0]['error']['code'] == 5001 + if ordering == 'once_success': + assert not session.get('pending_model_switch') + assert run_turn(session, monkeypatch) == [('model-c', 'endpoint-c', 'model-c', 'endpoint-c')] + else: + assert session['pending_model_switch'] is old_pending + if ordering == 'd_failure_first': + assert turn_samples == [admitted_runtime] + assert_route(session['agent'], 'http://c.invalid/v1', 'synthetic-key-c') + monkeypatch.setattr(server, '_wait_agent_for_prompt', lambda *a, **kw: None) + assert run_turn(session, monkeypatch) == [('model-b', 'endpoint-b', 'model-b', 'endpoint-b')] + + +def test_config_adoption_keeps_pick_fenced_after_current_admission(intent_turn, routes, monkeypatch): + session, _ = intent_turn + session.pop('model_override') + routes.config['model'].update(default='model-c', provider='endpoint-c') + routes.write() + monkeypatch.setattr(server, '_sync_agent_model_with_config', REAL_CONFIG_SYNC) + session['running'] = True + with session['history_lock']: + server._start_inflight_turn(session, 'already admitted') + assert server.handle_request(select('model-b', 'endpoint-b'))['result']['deferred'] + pending = session['pending_model_switch'] + assert run_turn(session, monkeypatch) == [('model-c', 'endpoint-c', 'model-b', 'endpoint-b')] + assert session['pending_model_switch'] is pending + assert run_turn(session, monkeypatch) == [('model-b', 'endpoint-b', 'model-b', 'endpoint-b')] + + +@pytest.mark.parametrize('replacement', ['record', 'agent', 'transport']) +def test_delayed_idle_selection_refuses_changed_owner(intent_turn, monkeypatch, replacement): + session, events = intent_turn + value = session['agent'] + before = copy.deepcopy(session['model_override']) + entered, release = threading.Event(), threading.Event() + transport = session['transport'] + real_switch = ms.switch_model + monkeypatch.setattr(server.threading, 'Thread', REAL_THREAD) + + def slow(*a, **kw): + entered.set() + assert release.wait(5) + return real_switch(*a, **kw) + + monkeypatch.setattr(ms, 'switch_model', slow) + with ThreadPoolExecutor(max_workers=1) as pool: + monkeypatch.setattr(server, '_pool', pool) + assert server.dispatch(select('model-b', 'endpoint-b'), transport) is None + try: + assert entered.wait(5) + if replacement == 'record': + server._sessions['selected'] = {'agent': SimpleNamespace(model='replacement')} + elif replacement == 'agent': + session['agent'] = SimpleNamespace(model='replacement') + else: + session['transport'] = Transport() + finally: + release.set() + assert transport.done.wait(5) + assert transport.frames[0]['error']['code'] == 5001 + assert_route(value, 'http://a.invalid/v1', 'synthetic-key-a') + assert session['model_override'] == before + assert not session.get('_config_set_pending') + assert not events diff --git a/tui_gateway/server.py b/tui_gateway/server.py index e06205106602..0eed76de2526 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -6202,7 +6202,9 @@ def _is_pivot_marker(entry: Any) -> bool: return isinstance(entry, dict) and entry.get("display_kind") == "personality_switch" -def _append_model_switch_marker(session: dict | None, *, model: str, provider: str) -> None: +def _append_model_switch_marker( + session: dict | None, *, model: str, provider: str, history_lock_held: bool = False, +) -> None: """Record a real system-history pivot after a live model switch. Only the most recent marker is kept: each new switch first strips any @@ -6238,7 +6240,7 @@ def _replace_markers() -> None: session["history_version"] = int(session.get("history_version", 0)) + 1 lock = session.get("history_lock") - if lock is not None: + if lock is not None and not history_lock_held: with lock: _replace_markers() else: @@ -6757,6 +6759,8 @@ def _apply_model_switch( pin_session_override: bool = True, parsed_flags: Any | None = None, persist_override: bool | None = None, + defer_if_running: bool = False, + supersede_pending: bool | None = None, ) -> dict: from hermes_cli.model_switch import ( parse_model_switch_args, @@ -6797,6 +6801,13 @@ def _apply_model_switch( raise ValueError("model value required") agent = session.get("agent") + owner_transport = session.get("transport") + # Existing restore/adoption/MoA callers explicitly suppress persistence; + # the picker marks manual intent, and the pending consumer opts out. + # A direct user /model choice otherwise keeps the pinning default. + if supersede_pending is None: + supersede_pending = pin_session_override and persist_override is None + superseded_pending = session.get("pending_model_switch") if supersede_pending else None if one_turn and not agent: raise ValueError("/model --once requires a live session") if agent: @@ -6855,8 +6866,6 @@ def _apply_model_switch( if not result.success: raise ValueError(result.error_message or "model switch failed") - restore_snapshot = _snapshot_agent_model_runtime(agent) if (one_turn and agent) else None - if agent: try: from hermes_cli.context_switch_guard import merge_preflight_compression_warning @@ -6903,37 +6912,68 @@ def _apply_model_switch( "confirm_message": confirm_msg, } - if agent: - try: - agent.switch_model( - new_model=result.new_model, - new_provider=result.target_provider, - api_key=result.api_key, - base_url=result.base_url, - api_mode=result.api_mode, + model_commit_lock = session.get("history_lock") + pending_publication = None + with model_commit_lock if model_commit_lock is not None else contextlib.nullcontext(): + if defer_if_running: + if (_sessions.get(sid) is not session or session.get("agent") is not agent + or session.get("transport") is not owner_transport): + raise ValueError("session owner changed; model request not applied") + if session.get("running"): + session["pending_model_switch"] = { + "raw": raw_input, + "confirm_expensive_model": confirm_expensive_model, + "display_model": result.new_model, + "display_provider": result.target_provider, + "after_inflight_turn": session.get("inflight_turn"), + } + return { + "value": result.new_model, "warning": result.warning_message or "", + "confirm_required": False, "confirm_message": "", "deferred": True, + "scope": "once" if one_turn else ("global" if persist_global else "session"), + } + restore_snapshot = None + if one_turn and agent: + queued_restore = session.get("one_turn_model_restore") + queued_runtime = session.get("_one_turn_model_runtime") + if (queued_restore and _owns_one_turn_model_runtime(session, agent, queued_runtime) + and not queued_runtime.get("active")): + # Replacing an unused once choice changes its one eligible + # turn, not the original durable/runtime restoration target. + restore_snapshot = queued_restore + else: + restore_snapshot = _snapshot_agent_model_runtime(agent) + + if agent: + try: + agent.switch_model( + new_model=result.new_model, + new_provider=result.target_provider, + api_key=result.api_key, + base_url=result.base_url, + api_mode=result.api_mode, + ) + except Exception as exc: + # The in-place swap rolled the agent back to the old working + # model/client and re-raised. Abort the commit: do NOT restart the + # slash worker, persist runtime, append the switch marker, set a + # session model_override, or persist to config — all of which would + # otherwise leave the session pinned to a broken model and kill the + # conversation on the next turn (#50163). A failed switch is a + # no-op; surface a clean error to the client. + logger.warning("In-place model switch failed for TUI agent: %s", exc) + raise ValueError( + f"Model switch to {result.new_model} failed ({exc}); " + f"staying on {getattr(agent, 'model', current_model)}." + ) from exc + _persist_live_session_runtime(session) + _persist_live_session_system_prompt(session) + _append_model_switch_marker( + session, model=result.new_model, provider=result.target_provider, + history_lock_held=model_commit_lock is not None, ) - except Exception as exc: - # The in-place swap rolled the agent back to the old working - # model/client and re-raised. Abort the commit: do NOT restart the - # slash worker, persist runtime, append the switch marker, set a - # session model_override, or persist to config — all of which would - # otherwise leave the session pinned to a broken model and kill the - # conversation on the next turn (#50163). A failed switch is a - # no-op; surface a clean error to the client. - logger.warning("In-place model switch failed for TUI agent: %s", exc) - raise ValueError( - f"Model switch to {result.new_model} failed ({exc}); " - f"staying on {getattr(agent, 'model', current_model)}." - ) from exc - _restart_slash_worker(sid, session) - _persist_live_session_runtime(session) - _persist_live_session_system_prompt(session) - _append_model_switch_marker( - session, model=result.new_model, provider=result.target_provider - ) - # The turn consumer uses this lock: snapshot and ownership must become - # visible (or retire) together, never as a partially published lease. - with session["history_lock"]: + # The turn consumer uses this lock: snapshot and ownership must become + # visible (or retire) together, never as a partially published lease. if one_turn: session["one_turn_model_restore"] = restore_snapshot session["_one_turn_model_runtime"] = { @@ -6943,36 +6983,82 @@ def _apply_model_switch( session.pop("one_turn_model_restore", None) session.pop("_one_turn_model_runtime", None) - # Record the switch as a PER-SESSION override so a later rebuild of THIS - # session (e.g. /new via _reset_session_agent, or resume) re-derives the - # user's chosen model/provider instead of falling back to global config. - # - # We deliberately do NOT write process-global env vars (HERMES_MODEL / - # HERMES_INFERENCE_MODEL / HERMES_TUI_PROVIDER / HERMES_INFERENCE_PROVIDER) - # here. The desktop backend hosts every same-profile session in ONE process, - # so mutating os.environ on a /model switch leaked the new model/provider - # into every OTHER live session's next agent rebuild — switching the model - # in one session silently changed it in the others (the cross-session - # contamination bug). agent.switch_model() above already mutated the right - # agent in place; the override dict makes that choice survive a rebuild - # without touching shared process state. - if pin_session_override and isinstance(session, dict) and not one_turn: - session["model_override"] = { - "model": result.new_model, - "provider": result.target_provider, - "base_url": result.base_url, - "api_key": result.api_key, - "api_mode": result.api_mode, - } - if isinstance(session, dict): - session.pop("model_verified_for", None) - mirror = session.get("_metadata_mirror") - if isinstance(mirror, dict): - mirror["model_ready"] = False - if agent: - _emit("session.info", sid, _session_info(agent, session)) - if persist_global: - _persist_model_switch(result) + # Record the switch as a PER-SESSION override so a later rebuild of THIS + # session (e.g. /new via _reset_session_agent, or resume) re-derives the + # user's chosen model/provider instead of falling back to global config. + # + # We deliberately do NOT write process-global env vars (HERMES_MODEL / + # HERMES_INFERENCE_MODEL / HERMES_TUI_PROVIDER / HERMES_INFERENCE_PROVIDER) + # here. The desktop backend hosts every same-profile session in ONE process, + # so mutating os.environ on a /model switch leaked the new model/provider + # into every OTHER live session's next agent rebuild — switching the model + # in one session silently changed it in the others (the cross-session + # contamination bug). agent.switch_model() above already mutated the right + # agent in place; the override dict makes that choice survive a rebuild + # without touching shared process state. + if pin_session_override and isinstance(session, dict) and not one_turn: + session["model_override"] = { + "model": result.new_model, + "provider": result.target_provider, + "base_url": result.base_url, + "api_key": result.api_key, + "api_mode": result.api_mode, + } + if superseded_pending is not None and session.get("pending_model_switch") is superseded_pending: + # Keep the old intent in its canonical queue until a manual + # choice finishes publication. Overlapping choices share claims + # on this same old intent; a newer queued pick is a new dict. + pending_publication = superseded_pending.get("_model_switch_publications") + if pending_publication is None: + pending_publication = { + "owners": [], + "projection": {key: superseded_pending[key] for key in + ("display_model", "display_provider") if key in superseded_pending}, + } + superseded_pending["_model_switch_publications"] = pending_publication + pending_publication["owners"].append(result) + # Project the working pin/owned once runtime, including its + # natural restore, rather than the temporarily held old B. + superseded_pending.pop("display_model", None) + superseded_pending.pop("display_provider", None) + if isinstance(session, dict): + session.pop("model_verified_for", None) + mirror = session.get("_metadata_mirror") + if isinstance(mirror, dict): + mirror["model_ready"] = False + # Worker replacement can acquire the registry lock; keep it outside + # history_lock to preserve the existing registry -> history lock order. + try: + if agent: + _restart_slash_worker(sid, session) + _emit("session.info", sid, _session_info(agent, session)) + if persist_global: + _persist_model_switch(result) + except Exception: + # A raised publication/persistence suffix means config.set cannot + # acknowledge C. Retain B's already acknowledged intent, without + # undoing the working client or resurrecting B over a newer choice. + if pending_publication is not None: + with model_commit_lock if model_commit_lock is not None else contextlib.nullcontext(): + if (_sessions.get(sid) is session + and session.get("pending_model_switch") is superseded_pending + and superseded_pending.get("_model_switch_publications") is pending_publication): + pending_publication["owners"][:] = [ + owner for owner in pending_publication["owners"] if owner is not result + ] + if not pending_publication["owners"]: + superseded_pending.pop("_model_switch_publications", None) + superseded_pending.update(pending_publication["projection"]) + if session.get("running"): + superseded_pending["after_inflight_turn"] = session.get("inflight_turn") + raise + if pending_publication is not None: + with model_commit_lock if model_commit_lock is not None else contextlib.nullcontext(): + if (_sessions.get(sid) is session + and session.get("pending_model_switch") is superseded_pending): + # Any acknowledged manual choice supersedes this captured + # old intent, even if another choice's suffix is in flight. + session.pop("pending_model_switch", None) return { "value": result.new_model, "warning": result.warning_message or "", @@ -7632,7 +7718,23 @@ def _apply_pending_model_switch(sid: str, session: dict) -> None: the current model and never blocks the turn, matching ``_sync_agent_model_with_config``. """ - pending = session.pop("pending_model_switch", None) + with session["history_lock"]: + pending = session.get("pending_model_switch") + if pending and pending.get("_model_switch_publications") is not None: + return + # A pick that finished resolving AFTER Send's admission belongs to a + # later turn, even if this turn has not reached its setup yet. An + # accepted correction shallow-copies the replay dict, preserving its + # existing immutable started_at object from the same admission. + deferred_after = pending.get("after_inflight_turn") if pending else None + current_inflight = session.get("inflight_turn") + if (deferred_after is not None and ( + deferred_after is current_inflight + or (isinstance(deferred_after, dict) and isinstance(current_inflight, dict) + and deferred_after.get("started_at") is not None + and deferred_after["started_at"] is current_inflight.get("started_at")))): + return + pending = session.pop("pending_model_switch", None) if not pending or session.get("agent") is None: return try: @@ -7641,6 +7743,7 @@ def _apply_pending_model_switch(sid: str, session: dict) -> None: session, pending["raw"], confirm_expensive_model=bool(pending.get("confirm_expensive_model")), + supersede_pending=False, ) # A queued pick is a deliberate user action; honour the expensive-model # confirm by NOT applying it silently — surface the warning and drop the @@ -15114,6 +15217,8 @@ def _(rid, params: dict) -> dict: # The user gets to pick, keep typing, and send the next turn on # the new model without waiting for the swap or interrupting. if session.get("running"): + pending_agent = session.get("agent") + pending_transport = session.get("transport") parsed = parse_model_switch_args(value) try: pending_model = parsed.model_input @@ -15161,16 +15266,19 @@ def _(rid, params: dict) -> dict: "deferred": False, }, ) - session["pending_model_switch"] = { - "raw": value, - "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": pending_provider, - } + with session["history_lock"]: + if (_sessions.get(params["session_id"]) is not session + or session.get("agent") is not pending_agent + or session.get("transport") is not pending_transport): + return _err(rid, 4001, "session owner changed; request not applied") + session["pending_model_switch"] = { + "raw": value, + "confirm_expensive_model": confirmed, + # Projection names the next eligible turn's choice. + "display_model": pending_model, + "display_provider": pending_provider, + "after_inflight_turn": session.get("inflight_turn"), + } return _ok( rid, { @@ -15201,6 +15309,8 @@ def _(rid, params: dict) -> dict: params.get("confirm_expensive_model", False) ), parsed_flags=parsed_flags, + defer_if_running=True, + supersede_pending=True, ) else: result = _apply_model_switch( @@ -15220,6 +15330,7 @@ def _(rid, params: dict) -> dict: "confirm_required": result.get("confirm_required", False), "confirm_message": result.get("confirm_message", ""), "scope": result.get("scope", "session"), + **({"deferred": result["deferred"]} if "deferred" in result else {}), }, ) except Exception as e: