fix(runtime): isolate nested invocation ownership

This commit is contained in:
Alexandre Teixeira
2026-10-01 02:11:53 +01:00
parent 4c122de880
commit aedec7d005
6 changed files with 691 additions and 64 deletions
+112 -38
View File
@@ -24,6 +24,10 @@ itself wasn't confident about.
from __future__ import annotations
import asyncio
from contextlib import aclosing
from contextvars import ContextVar
from copy import deepcopy
from functools import wraps
import logging
import re
from typing import Any, Dict, List, Optional, Tuple
@@ -32,6 +36,58 @@ from urllib.parse import urlparse
logger = logging.getLogger(__name__)
_TAKEOVER: ContextVar[dict | None] = ContextVar('teacher_takeover_request', default=None)
def request_teacher_takeover(**parameters):
"""Record a finished student's handoff; execution waits for its gate to close."""
from src.agent_runtime.journal import current_journal
request = _TAKEOVER.get()
if request is not None:
journal = current_journal()
request.update(parameters, parent_run_id=journal.run_id if journal is not None else None)
def with_teacher_takeover(func):
"""Orchestrate gated invocations and own the single outer stream terminator."""
@wraps(func)
async def wrapped(*args, **kwargs):
request = {}
token = _TAKEOVER.set(request)
done = False
failed = False
try:
async with aclosing(func(*args, **kwargs)) as stream:
async for chunk in stream:
if chunk.strip() == 'data: [DONE]':
done = True
continue
failed |= chunk.startswith('event: error')
yield chunk
# The parent generator and gate have both unwound. Child control
# events now belong only to the child, never to the parent gate.
if request and not failed:
try:
async with aclosing(run_teacher_inline(**request)) as stream:
async for chunk in stream:
if chunk.strip() == 'data: [DONE]':
continue
failed |= chunk.startswith('event: error')
yield chunk
except Exception as exc:
logger.warning('teacher escalation hook failed: %s', exc, exc_info=True)
if not failed:
import json
yield 'data: ' + json.dumps({'type': 'escalation_failed', 'reason': str(exc),
'teacher': True}) + '\n\n'
if done and not failed:
yield 'data: [DONE]\n\n'
finally:
_TAKEOVER.reset(token)
return wrapped
# Hosts considered SOTA / paid APIs — if the student's endpoint URL
# hits one of these, the loop is OFF (the user is already paying for
# a top-tier model; no need to escalate).
@@ -524,6 +580,11 @@ async def run_teacher_inline(
tool_policy: Any = None,
active_document: Any = None,
active_email: Optional[Dict[str, str]] = None,
turn_contract=None,
parent_run_id: Optional[str] = None,
external_untrusted_context_seen: bool = False,
client_runtime_context: Optional[Dict[str, Any]] = None,
plan_mode: bool = False,
):
"""Async generator. Yields SSE event strings.
@@ -606,7 +667,7 @@ async def run_teacher_inline(
# user/assistant/tool history so the teacher sees what the student
# tried. The appended note leads with the user request text so RAG
# tool selection picks the right tools for the teacher's turn.
history = [m for m in student_messages if m.get("role") != "system"]
history = deepcopy([m for m in student_messages if m.get("role") != "system"])
note_content = (
f"{user_request or '(no user request captured)'}\n\n"
"[teacher-takeover] The previous attempt by the student model "
@@ -623,8 +684,9 @@ async def run_teacher_inline(
captured_tool_events: List[Dict[str, Any]] = []
captured_text_parts: List[str] = []
captured_metrics: Dict[str, Any] = {}
captured_decision: Dict[str, Any] = {}
async for evt_str in stream_agent_loop(
async with aclosing(stream_agent_loop(
endpoint_url=teacher_url,
model=teacher_model,
messages=teacher_messages,
@@ -632,51 +694,63 @@ async def run_teacher_inline(
owner=owner,
session_id=session_id,
workspace=workspace,
disabled_tools=disabled_tools,
disabled_tools=set(disabled_tools) if disabled_tools is not None else None,
tool_policy=tool_policy,
active_document=active_document,
active_email=active_email,
turn_contract=turn_contract,
_parent_run_id=parent_run_id,
external_untrusted_context_seen=external_untrusted_context_seen,
client_runtime_context=deepcopy(client_runtime_context),
plan_mode=plan_mode,
_is_teacher_run=True,
):
# Swallow teacher's own [DONE] — outer loop emits the real one
if "[DONE]" in evt_str:
continue
if evt_str.startswith("data: "):
try:
payload = json.loads(evt_str[6:].strip())
except Exception:
)) as stream:
async for evt_str in stream:
# Swallow teacher's own [DONE] — outer loop emits the real one
if evt_str.strip() == 'data: [DONE]':
continue
if evt_str.startswith('event: error'):
yield evt_str
continue
if isinstance(payload, dict):
payload["teacher"] = True
typ = payload.get("type")
if typ == "metrics" and isinstance(payload.get("data"), dict):
# The outer chat route persists only the last metrics
# payload. Keep a copy so any approval produced after the
# recursive teacher run's metrics remains reloadable.
captured_metrics = dict(payload["data"])
if typ == "tool_output":
captured_tool_event = {
"tool": payload.get("tool"),
"command": payload.get("command"),
"output": payload.get("output"),
"exit_code": payload.get("exit_code"),
}
if isinstance(payload.get("ask_user"), dict):
captured_tool_event["ask_user"] = payload["ask_user"]
captured_tool_events.append(captured_tool_event)
if "delta" in payload and isinstance(payload["delta"], str):
if payload.get("thinking"):
continue
captured_text_parts.append(payload["delta"])
yield 'data: ' + json.dumps(payload) + '\n\n'
continue
yield evt_str
return
if evt_str.startswith("data: "):
try:
payload = json.loads(evt_str[6:].strip())
except Exception:
yield evt_str
continue
if isinstance(payload, dict):
payload["teacher"] = True
typ = payload.get("type")
if typ == 'completion_decision' and isinstance(payload.get('data'), dict):
captured_decision = payload['data']
if typ == "metrics" and isinstance(payload.get("data"), dict):
# The outer chat route persists only the last metrics
# payload. Keep a copy so any approval produced after the
# recursive teacher run's metrics remains reloadable.
captured_metrics = dict(payload["data"])
if typ == "tool_output":
captured_tool_event = {
"tool": payload.get("tool"),
"command": payload.get("command"),
"output": payload.get("output"),
"exit_code": payload.get("exit_code"),
}
if isinstance(payload.get("ask_user"), dict):
captured_tool_event["ask_user"] = payload["ask_user"]
captured_tool_events.append(captured_tool_event)
if "delta" in payload and isinstance(payload["delta"], str):
if payload.get("thinking"):
continue
captured_text_parts.append(payload["delta"])
yield 'data: ' + json.dumps(payload) + '\n\n'
continue
yield evt_str
# A takeover that paused for a question or exact action has not completed
# yet. Its server-owned approval card is already in the live/persisted tool
# events; do not evaluate the partial trace or distill it into a skill.
if any(event.get("ask_user") for event in captured_tool_events):
if (any(event.get("ask_user") for event in captured_tool_events)
or (captured_decision and not captured_decision.get('can_complete', False))):
return
teacher_text = "".join(captured_text_parts).strip()