fix(delegation): bind steering to session generation

This commit is contained in:
SmokeDev 2026-08-05 21:50:52 -07:00 committed by Teknium
parent a94ebf5f5e
commit 9d4ef04ed0
5 changed files with 585 additions and 23 deletions

View File

@ -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"

View File

@ -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 == []

View File

@ -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

View File

@ -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,
{

View File

@ -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