From 9d4ef04ed00055414c13fcf33925d85790221a3f Mon Sep 17 00:00:00 2001 From: SmokeDev Date: Wed, 5 Aug 2026 21:50:52 -0700 Subject: [PATCH] fix(delegation): bind steering to session generation --- tests/test_tui_gateway_server.py | 52 ++++ tests/tools/test_subagent_steer.py | 443 ++++++++++++++++++++++++++++- tools/delegate_tool.py | 57 +++- tui_gateway/methods_session.py | 17 +- tui_gateway/server.py | 39 +++ 5 files changed, 585 insertions(+), 23 deletions(-) diff --git a/tests/test_tui_gateway_server.py b/tests/test_tui_gateway_server.py index e864f185ce9ee..019af044609ba 100644 --- a/tests/test_tui_gateway_server.py +++ b/tests/test_tui_gateway_server.py @@ -4748,6 +4748,58 @@ def _configure_immediate_prompt_run( monkeypatch.setattr(server, "_get_db", lambda: None) +def test_run_prompt_submit_binds_exact_steer_authority_and_resets_contextvars( + monkeypatch, tmp_path +): + """The turn thread commissions children with this session generation only.""" + from tools.delegate_tool import _capture_gateway_steer_authority + from tui_gateway.transport import ( + bind_transport, + current_transport, + reset_transport, + ) + + class _Transport: + def write(self, _obj): + return True + + def close(self): + return None + + observed = {} + owner_transport = _Transport() + previous_transport = _Transport() + previous_record = {"session_key": "previous-generation"} + + class _CapturingAgent(_RecordingAgent): + def run_conversation(self, prompt, **kwargs): + authority = _capture_gateway_steer_authority("sid-owner") + observed["transport"] = authority[0] + observed["record"] = authority[1] + return super().run_conversation(prompt, **kwargs) + + _configure_immediate_prompt_run(monkeypatch, tmp_path) + session = _session( + session_key="session-owner", + agent=_CapturingAgent([]), + running=True, + transport=owner_transport, + ) + server._sessions["sid-owner"] = session + transport_token = bind_transport(previous_transport) + record_token = server._current_runtime_session_record.set(previous_record) + try: + server._run_prompt_submit("rid-owner", "sid-owner", session, "commission") + + assert observed == {"transport": owner_transport, "record": session} + assert current_transport() is previous_transport + assert server._current_runtime_session_record.get() is previous_record + finally: + server._current_runtime_session_record.reset(record_token) + reset_transport(transport_token) + server._sessions.pop("sid-owner", None) + + class _RecordingAgent: model = "test-model" provider = "test-provider" diff --git a/tests/tools/test_subagent_steer.py b/tests/tools/test_subagent_steer.py index 5a284f45b4e4b..f72c6b9a3273c 100644 --- a/tests/tools/test_subagent_steer.py +++ b/tests/tools/test_subagent_steer.py @@ -31,7 +31,14 @@ class _StubAgent: return self.accept -def _with_registered(sid: str, agent, *, owner_session_id: str | None = None) -> None: +def _with_registered( + sid: str, + agent, + *, + owner_session_id: str | None = None, + owner_transport=None, + owner_session_record=None, +) -> None: _register_subagent( { "subagent_id": sid, @@ -41,6 +48,8 @@ def _with_registered(sid: str, agent, *, owner_session_id: str | None = None) -> "status": "running", "agent": agent, "owner_session_id": owner_session_id, + "owner_transport": owner_transport, + "owner_session_record": owner_session_record, } ) @@ -106,7 +115,6 @@ def test_stale_agent_teardown_cannot_unregister_recycled_id(): steer_subagent( "sid-recycled-teardown", "replacement remains live", - owner_session_id="new-owner", ) is True ) @@ -120,7 +128,15 @@ def test_status_snapshot_never_leaks_owner_or_lifecycle_metadata(): from tools.delegate_tool import list_active_subagents agent = _StubAgent() - _with_registered("sid-private-metadata", agent, owner_session_id="private-owner") + owner_transport = object() + owner_session_record = {"session_key": "private-owner"} + _with_registered( + "sid-private-metadata", + agent, + owner_session_id="private-owner", + owner_transport=owner_transport, + owner_session_record=owner_session_record, + ) try: snapshot = next( item @@ -130,8 +146,12 @@ def test_status_snapshot_never_leaks_owner_or_lifecycle_metadata(): assert snapshot["status"] == "running" assert "agent" not in snapshot assert "owner_session_id" not in snapshot + assert "owner_transport" not in snapshot + assert "owner_session_record" not in snapshot assert "accepting_steer" not in snapshot assert "private-owner" not in repr(snapshot) + assert all(value is not owner_transport for value in snapshot.values()) + assert all(value is not owner_session_record for value in snapshot.values()) finally: _unregister_subagent("sid-private-metadata", agent=agent) @@ -326,15 +346,31 @@ class TestMissedSteerRetention: class TestSubagentSteerRPC: """subagent.steer gateway RPC — the programmatic caller beside subagent.interrupt.""" - def _call(self, params: dict) -> dict: + class _Transport: + def __init__(self) -> None: + self.frames: list[dict] = [] + + def write(self, obj: dict) -> bool: + self.frames.append(obj) + return True + + def close(self) -> None: + return None + + def _call(self, params: dict, *, transport=None, session_record=None) -> dict: import tui_gateway.server as srv session_id = params.get("session_id") if session_id: - srv._sessions[session_id] = {"session_key": session_id, "history": []} + srv._sessions[session_id] = session_record or { + "session_key": session_id, + "history": [], + "transport": transport, + } try: - return srv.handle_request( - {"id": 1, "method": "subagent.steer", "params": params} + return srv.dispatch( + {"id": 1, "method": "subagent.steer", "params": params}, + transport=transport, ) finally: if session_id: @@ -349,15 +385,29 @@ class TestSubagentSteerRPC: assert envelope["error"]["code"] == 4002 def test_live_child_queues_and_receives_text(self): + owner_transport = self._Transport() + owner_record = { + "session_key": "owner-session", + "history": [], + "transport": owner_transport, + } agent = _StubAgent() - _with_registered("sid-rpc-2", agent, owner_session_id="owner-session") + _with_registered( + "sid-rpc-2", + agent, + owner_session_id="owner-session", + owner_transport=owner_transport, + owner_session_record=owner_record, + ) try: envelope = self._call( { "session_id": "owner-session", "subagent_id": "sid-rpc-2", "text": "check the edge cases", - } + }, + transport=owner_transport, + session_record=owner_record, ) assert envelope["result"] == { "status": "queued", @@ -368,11 +418,17 @@ class TestSubagentSteerRPC: finally: _unregister_subagent("sid-rpc-2") - def test_run_single_child_binds_owner_from_task_local_ui_session(self): + def test_run_single_child_binds_exact_runtime_owner_artifacts(self): from gateway.session_context import clear_session_vars, set_session_vars from tools.delegate_tool import _run_single_child observed: dict[str, bool] = {} + owner_transport = self._Transport() + owner_session_record = { + "session_key": "durable-parent", + "history": [], + "transport": owner_transport, + } child = MagicMock() child._subagent_id = "sid-context-owner" child._delegate_depth = 1 @@ -384,11 +440,15 @@ class TestSubagentSteerRPC: child._subagent_id, "owned steer", owner_session_id="ui-owner", + owner_transport=owner_transport, + owner_session_record=owner_session_record, ) observed["foreign"] = steer_subagent( child._subagent_id, "foreign steer", - owner_session_id="ui-foreign", + owner_session_id="ui-owner", + owner_transport=self._Transport(), + owner_session_record=owner_session_record, ) return { "final_response": "done", @@ -405,7 +465,14 @@ class TestSubagentSteerRPC: ui_session_id="ui-owner", ) try: - _run_single_child(0, "owner binding", child=child, parent_agent=MagicMock()) + _run_single_child( + 0, + "owner binding", + child=child, + parent_agent=MagicMock(), + owner_transport=owner_transport, + owner_session_record=owner_session_record, + ) finally: clear_session_vars(tokens) @@ -438,16 +505,362 @@ class TestSubagentSteerRPC: finally: _unregister_subagent("sid-rpc-foreign") - def test_record_without_owner_cannot_be_steered_by_rpc(self): + def test_foreign_transport_with_correct_session_id_is_denied(self): + owner_transport = self._Transport() + foreign_transport = self._Transport() + owner_record = { + "session_key": "owner-session", + "history": [], + "transport": owner_transport, + } agent = _StubAgent() - _with_registered("sid-rpc-owner-missing", agent) + _with_registered( + "sid-rpc-foreign-transport", + agent, + owner_session_id="owner-session", + owner_transport=owner_transport, + owner_session_record=owner_record, + ) + try: + envelope = self._call( + { + "session_id": "owner-session", + "subagent_id": "sid-rpc-foreign-transport", + "text": "stolen identifier", + }, + transport=foreign_transport, + session_record=owner_record, + ) + assert envelope["result"]["status"] == "rejected" + assert agent.steered == [] + finally: + _unregister_subagent("sid-rpc-foreign-transport") + + def test_recycled_session_record_with_same_id_is_denied(self): + owner_transport = self._Transport() + original_record = { + "session_key": "owner-session", + "history": [], + "transport": owner_transport, + } + recycled_record = { + "session_key": "owner-session", + "history": [], + "transport": owner_transport, + } + agent = _StubAgent() + _with_registered( + "sid-rpc-recycled-session", + agent, + owner_session_id="owner-session", + owner_transport=owner_transport, + owner_session_record=original_record, + ) + try: + envelope = self._call( + { + "session_id": "owner-session", + "subagent_id": "sid-rpc-recycled-session", + "text": "new generation", + }, + transport=owner_transport, + session_record=recycled_record, + ) + assert envelope["result"]["status"] == "rejected" + assert agent.steered == [] + finally: + _unregister_subagent("sid-rpc-recycled-session") + + def test_server_resolves_exact_runtime_authority_from_dispatch_context(self): + import tui_gateway.server as srv + + owner_transport = self._Transport() + owner_record = { + "session_key": "owner-session", + "history": [], + "transport": owner_transport, + } + srv._sessions["owner-session"] = owner_record + + def capture(rid, _params): + authority = srv._current_session_steer_authority("owner-session") + return srv._ok( + rid, + { + "transport_matches": authority[0] is owner_transport, + "record_matches": authority[1] is owner_record, + }, + ) + + srv._methods["test.capture-steer-authority"] = capture + try: + envelope = srv.dispatch( + { + "id": 1, + "method": "test.capture-steer-authority", + "params": { + "session_id": "owner-session", + "owner_transport": "spoof", + "owner_session_record": "spoof", + }, + }, + transport=owner_transport, + ) + finally: + srv._methods.pop("test.capture-steer-authority", None) + srv._sessions.pop("owner-session", None) + + assert envelope["result"] == { + "transport_matches": True, + "record_matches": True, + } + + def test_delegate_capture_uses_dispatch_runtime_artifacts(self): + import tui_gateway.server as srv + from tools import delegate_tool + + owner_transport = self._Transport() + owner_record = { + "session_key": "owner-session", + "history": [], + "transport": owner_transport, + } + srv._sessions["owner-session"] = owner_record + + def capture(rid, _params): + transport, record = delegate_tool._capture_gateway_steer_authority( + "owner-session" + ) + return srv._ok( + rid, + { + "transport_matches": transport is owner_transport, + "record_matches": record is owner_record, + }, + ) + + srv._methods["test.capture-delegate-authority"] = capture + try: + envelope = srv.dispatch( + { + "id": 1, + "method": "test.capture-delegate-authority", + "params": {"session_id": "owner-session"}, + }, + transport=owner_transport, + ) + finally: + srv._methods.pop("test.capture-delegate-authority", None) + srv._sessions.pop("owner-session", None) + + assert envelope["result"] == { + "transport_matches": True, + "record_matches": True, + } + + def test_commissioning_context_rejects_recycled_runtime_session_record(self): + import tui_gateway.server as srv + from tools import delegate_tool + from tui_gateway.transport import bind_transport, reset_transport + + owner_transport = self._Transport() + original_record = { + "session_key": "owner-session", + "history": [], + "transport": owner_transport, + } + recycled_record = { + "session_key": "owner-session", + "history": [], + "transport": owner_transport, + } + srv._sessions["owner-session"] = recycled_record + transport_token = bind_transport(owner_transport) + record_token = srv._current_runtime_session_record.set(original_record) + try: + assert delegate_tool._capture_gateway_steer_authority("owner-session") == ( + None, + None, + ) + finally: + srv._current_runtime_session_record.reset(record_token) + reset_transport(transport_token) + srv._sessions.pop("owner-session", None) + + def test_rpc_params_cannot_spoof_runtime_artifacts(self): + owner_transport = self._Transport() + owner_record = { + "session_key": "owner-session", + "history": [], + "transport": owner_transport, + } + agent = _StubAgent() + _with_registered( + "sid-rpc-param-spoof", + agent, + owner_session_id="owner-session", + owner_transport=owner_transport, + owner_session_record=owner_record, + ) + try: + envelope = self._call( + { + "session_id": "owner-session", + "subagent_id": "sid-rpc-param-spoof", + "text": "ignore serialized capabilities", + "owner_transport": self._Transport(), + "owner_session_record": {"session_key": "owner-session"}, + "owner_token": "forged", + }, + transport=owner_transport, + session_record=owner_record, + ) + assert envelope["result"]["status"] == "queued" + assert agent.steered == ["ignore serialized capabilities"] + finally: + _unregister_subagent("sid-rpc-param-spoof") + + def test_session_transport_rebinding_does_not_transfer_ownership(self): + original_transport = self._Transport() + rebound_transport = self._Transport() + owner_record = { + "session_key": "owner-session", + "history": [], + "transport": original_transport, + } + agent = _StubAgent() + _with_registered( + "sid-rpc-rebound", + agent, + owner_session_id="owner-session", + owner_transport=original_transport, + owner_session_record=owner_record, + ) + owner_record["transport"] = rebound_transport + try: + for transport in (original_transport, rebound_transport): + envelope = self._call( + { + "session_id": "owner-session", + "subagent_id": "sid-rpc-rebound", + "text": "rebound authority", + }, + transport=transport, + session_record=owner_record, + ) + assert envelope["result"]["status"] == "rejected" + assert agent.steered == [] + finally: + _unregister_subagent("sid-rpc-rebound") + + def test_concurrent_sessions_cannot_cross_steer(self): + transports = [self._Transport(), self._Transport()] + records = [ + {"session_key": f"session-{i}", "history": [], "transport": transports[i]} + for i in range(2) + ] + agents = [_StubAgent(), _StubAgent()] + for i in range(2): + _with_registered( + f"sid-concurrent-{i}", + agents[i], + owner_session_id=f"session-{i}", + owner_transport=transports[i], + owner_session_record=records[i], + ) + + barrier = threading.Barrier(2) + results: list[str] = [] + + def cross_call(caller: int) -> None: + barrier.wait(5) + envelope = self._call( + { + "session_id": f"session-{caller}", + "subagent_id": f"sid-concurrent-{1 - caller}", + "text": f"cross-{caller}", + }, + transport=transports[caller], + session_record=records[caller], + ) + results.append(envelope["result"]["status"]) + + threads = [threading.Thread(target=cross_call, args=(i,)) for i in range(2)] + try: + for thread in threads: + thread.start() + for thread in threads: + thread.join(5) + assert all(not thread.is_alive() for thread in threads) + assert sorted(results) == ["rejected", "rejected"] + assert agents[0].steered == [] + assert agents[1].steered == [] + finally: + for i in range(2): + _unregister_subagent(f"sid-concurrent-{i}") + + def test_owner_still_works_after_unrelated_dispatch_queries(self): + import tui_gateway.server as srv + + owner_transport = self._Transport() + owner_record = { + "session_key": "owner-session", + "history": [], + "transport": owner_transport, + } + agent = _StubAgent() + _with_registered( + "sid-rpc-after-query", + agent, + owner_session_id="owner-session", + owner_transport=owner_transport, + owner_session_record=owner_record, + ) + srv._methods["test.unrelated-query"] = lambda rid, _params: srv._ok( + rid, {"ok": True} + ) + try: + assert srv.dispatch( + {"id": 8, "method": "test.unrelated-query", "params": {}}, + transport=self._Transport(), + )["result"] == {"ok": True} + envelope = self._call( + { + "session_id": "owner-session", + "subagent_id": "sid-rpc-after-query", + "text": "still mine", + }, + transport=owner_transport, + session_record=owner_record, + ) + assert envelope["result"]["status"] == "queued" + assert agent.steered == ["still mine"] + finally: + srv._methods.pop("test.unrelated-query", None) + _unregister_subagent("sid-rpc-after-query") + + def test_record_missing_runtime_artifacts_cannot_be_steered_by_rpc(self): + agent = _StubAgent() + owner_transport = self._Transport() + owner_record = { + "session_key": "claiming-session", + "history": [], + "transport": owner_transport, + } + _with_registered( + "sid-rpc-owner-missing", + agent, + owner_session_id="claiming-session", + ) try: envelope = self._call( { "session_id": "claiming-session", "subagent_id": "sid-rpc-owner-missing", "text": "ambiguous authority", - } + }, + transport=owner_transport, + session_record=owner_record, ) assert envelope["result"]["status"] == "rejected" assert agent.steered == [] diff --git a/tools/delegate_tool.py b/tools/delegate_tool.py index 672d2f3be25a6..5a7b3eb88085d 100644 --- a/tools/delegate_tool.py +++ b/tools/delegate_tool.py @@ -237,6 +237,8 @@ def steer_subagent( text: str, *, owner_session_id: Optional[str] = None, + owner_transport: Any = None, + owner_session_record: Any = None, ) -> bool: """Queue steering text into a single running subagent without stopping it. @@ -259,8 +261,15 @@ def steer_subagent( record = _active_subagents.get(subagent_id) if not record or not record.get("accepting_steer", False): return False - if owner_session_id is not None and record.get("owner_session_id") != owner_session_id: - return False + if owner_session_id is not None: + if ( + record.get("owner_session_id") != owner_session_id + or owner_transport is None + or record.get("owner_transport") is not owner_transport + or owner_session_record is None + or record.get("owner_session_record") is not owner_session_record + ): + return False agent = record.get("agent") if agent is None: return False @@ -271,6 +280,24 @@ def steer_subagent( return False +def _capture_gateway_steer_authority( + owner_session_id: Optional[str], +) -> tuple[Any, Any]: + """Capture exact request transport + live session generation, if any. + + This is intentionally an in-process bridge, not a serializable capability. + Non-gateway hosts (including the CLI helper path) receive ``(None, None)``. + """ + if not owner_session_id: + return None, None + try: + from tui_gateway.server import _current_session_steer_authority + + return _current_session_steer_authority(owner_session_id) + except Exception: + return None, None + + def list_active_subagents() -> List[Dict[str, Any]]: """Snapshot of the currently running subagent tree. @@ -282,7 +309,14 @@ def list_active_subagents() -> List[Dict[str, Any]]: { k: v for k, v in r.items() - if k not in {"agent", "owner_session_id", "accepting_steer"} + if k + not in { + "agent", + "owner_session_id", + "owner_transport", + "owner_session_record", + "accepting_steer", + } } for r in _active_subagents.values() ] @@ -2045,6 +2079,8 @@ def _run_single_child( parent_agent=None, *, owner_session_id: Optional[str] = None, + owner_transport: Any = None, + owner_session_record: Any = None, **_kwargs, ) -> Dict[str, Any]: """ @@ -2186,6 +2222,12 @@ def _run_single_child( owner_session_id = get_session_env("HERMES_UI_SESSION_ID", "") or None except Exception: owner_session_id = None + if owner_session_id and ( + owner_transport is None or owner_session_record is None + ): + owner_transport, owner_session_record = ( + _capture_gateway_steer_authority(owner_session_id) + ) _raw_depth = getattr(child, "_delegate_depth", 1) _tui_depth = max(0, _raw_depth - 1) if isinstance(_raw_depth, int) else 0 _parent_sid = getattr(child, "_parent_subagent_id", None) @@ -2207,6 +2249,8 @@ def _run_single_child( # Immutable live gateway/TUI session that commissioned this # child. Empty outside those hosts; RPC authority fails closed. "owner_session_id": owner_session_id, + "owner_transport": owner_transport, + "owner_session_record": owner_session_record, } ) @@ -3090,6 +3134,9 @@ def delegate_task( _origin_ui_session_id = get_session_env("HERMES_UI_SESSION_ID", "") except Exception: _origin_ui_session_id = "" + _origin_owner_transport, _origin_owner_session_record = ( + _capture_gateway_steer_authority(_origin_ui_session_id) + ) # Build all child agents on the main thread (thread-safe construction). # _build_child_preserving_parent_tools saves/restores the parent's @@ -3154,6 +3201,8 @@ def delegate_task( child, parent_agent, owner_session_id=_origin_ui_session_id or None, + owner_transport=_origin_owner_transport, + owner_session_record=_origin_owner_session_record, ) results.append(result) else: @@ -3177,6 +3226,8 @@ def delegate_task( child=child, parent_agent=parent_agent, owner_session_id=_origin_ui_session_id or None, + owner_transport=_origin_owner_transport, + owner_session_record=_origin_owner_session_record, ) futures[future] = i diff --git a/tui_gateway/methods_session.py b/tui_gateway/methods_session.py index d8bf95d2b903d..8a04660557a7c 100644 --- a/tui_gateway/methods_session.py +++ b/tui_gateway/methods_session.py @@ -2954,15 +2954,22 @@ def _(rid, params: dict) -> dict: text = (params.get("text") or "").strip() if not text: return _err(rid, 4002, "text is required") - _session, err = _sess_nowait(params, rid) + _invoking_session, err = _sess_nowait(params, rid) if err: return err invoking_session_id = str(params.get("session_id") or "").strip() - queued = steer_subagent( - subagent_id, - text, - owner_session_id=invoking_session_id, + invoking_transport, invoking_session = _current_session_steer_authority( + invoking_session_id ) + queued = False + if invoking_transport is not None and invoking_session is not None: + queued = steer_subagent( + subagent_id, + text, + owner_session_id=invoking_session_id, + owner_transport=invoking_transport, + owner_session_record=invoking_session, + ) return _ok( rid, { diff --git a/tui_gateway/server.py b/tui_gateway/server.py index 87568f7f9094f..f9e1981e1ac66 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -295,6 +295,12 @@ _pool = concurrent.futures.ThreadPoolExecutor( ) atexit.register(lambda: _pool.shutdown(wait=False, cancel_futures=True)) +# Exact in-memory session generation executing on the current turn thread. +# Unlike a public session id, this object identity cannot be supplied by RPC. +_current_runtime_session_record: contextvars.ContextVar[dict | None] = ( + contextvars.ContextVar("hermes_gateway_runtime_session_record", default=None) +) + # Reserve real stdout for JSON-RPC only; redirect Python's stdout to stderr # so stray print() from libraries/tools becomes harmless gateway.stderr instead # of corrupting the JSON protocol. @@ -1900,6 +1906,31 @@ def handle_request(req: dict) -> dict | None: return fn(rid, params) +def _current_session_steer_authority( + session_id: str, +) -> tuple[Transport | None, dict | None]: + """Resolve unforgeable steering authority for this exact RPC context. + + The public session id is only a lookup hint. Authority is the identity of + both the request's ContextVar-bound transport and the live in-memory + session record currently stored under that id. Session transport rebinding, + removal, or id reuse therefore invalidates an earlier generation. + """ + transport = current_transport() + if transport is None or not session_id: + return None, None + expected_session = _current_runtime_session_record.get() + with _sessions_lock: + session = _sessions.get(session_id) + if ( + session is None + or (expected_session is not None and session is not expected_session) + or session.get("transport") is not transport + ): + return None, None + return transport, session + + def dispatch(req: dict, transport: Optional[Transport] = None) -> dict | None: """Route inbound RPCs — long handlers to the pool, everything else inline. @@ -9386,6 +9417,12 @@ def _run_prompt_submit( _emit("message.start", sid) def run(): + # The conversation runs on a fresh thread, so ContextVars from the RPC + # dispatcher do not follow automatically. Rebind the exact transport + # stored on this session generation before any tool can commission a + # child; delegate_task then captures it as non-serializable authority. + transport_token = bind_transport(session.get("transport")) + runtime_session_token = _current_runtime_session_record.set(session) approval_token = None session_tokens = [] home_token = None # per-turn HERMES_HOME override for a resumed remote profile @@ -10089,6 +10126,8 @@ def _run_prompt_submit( if secret_token is not None: reset_secret_scope(secret_token) _clear_session_context(session_tokens) + _current_runtime_session_record.reset(runtime_session_token) + reset_transport(transport_token) # Clear the per-turn interim callback so a stale closure from # this turn can't fire during a later turn on the same agent. agent.interim_assistant_callback = None