From e6c0737d874babe0b392c0a14c97498c6b941fda Mon Sep 17 00:00:00 2001 From: Ahmed Allam Date: Tue, 28 Jul 2026 00:28:07 +0000 Subject: [PATCH] fix: restore pre-v1 lifecycle resilience: mailbox delivery, uniform revival, waiting timeout, broader retries --- strix/core/agents.py | 109 +++++++++++++++++------- strix/core/execution.py | 103 ++++++++++++++++++---- strix/interface/tui/app.py | 2 +- tests/test_execution.py | 99 ++++++++++++++++++++- tests/test_execution_transient_retry.py | 21 ++++- 5 files changed, 282 insertions(+), 52 deletions(-) diff --git a/strix/core/agents.py b/strix/core/agents.py index 97d94548..41015266 100644 --- a/strix/core/agents.py +++ b/strix/core/agents.py @@ -32,6 +32,8 @@ class AgentRuntime: stream: Any | None = None interrupt_on_message: bool = False wake: asyncio.Event = field(default_factory=asyncio.Event) + mailbox: list[dict[str, Any]] = field(default_factory=list) + user_wake_required: bool = False class AgentCoordinator: @@ -179,6 +181,7 @@ class AgentCoordinator: if agent_id in self.statuses: self.statuses[agent_id] = "running" self.errors.pop(agent_id, None) + self.runtimes.setdefault(agent_id, AgentRuntime()).user_wake_required = False await self._maybe_snapshot() async def park_waiting(self, agent_id: str) -> None: @@ -196,6 +199,7 @@ class AgentCoordinator: elif status == "running": self.errors.pop(agent_id, None) runtime = self.runtimes.setdefault(agent_id, AgentRuntime()) + runtime.user_wake_required = status in {"failed", "crashed"} runtime.wake.set() logger.info("agent.status %s=%s", agent_id, status) await self._maybe_snapshot() @@ -203,49 +207,56 @@ class AgentCoordinator: async def send( self, target_agent_id: str, message: dict[str, Any], *, interrupt: bool = True ) -> bool: - """Deliver a user/peer message by appending it to the target SDK session.""" - if message.get("from") == "user" and self._budget_paused: + """Queue a user/peer message in the target's mailbox and wake it. + + Delivery never blocks: the message lands in an in-memory mailbox and + the target agent appends it to its own SDK session when it wakes + (``consume_pending``). + """ + from_user = message.get("from") == "user" + if from_user and self._budget_paused: await self.resume_from_budget_pause(exclude=target_agent_id) async with self._lock: if target_agent_id not in self.statuses: logger.debug("agent.send dropped unknown target=%s", target_agent_id) return False runtime = self.runtimes.setdefault(target_agent_id, AgentRuntime()) - session = runtime.session + runtime.mailbox.append(dict(message)) + self.pending_counts[target_agent_id] = self.pending_counts.get(target_agent_id, 0) + 1 + if from_user: + runtime.user_wake_required = False + runtime.wake.set() stream = runtime.stream interrupt_on_message = runtime.interrupt_on_message - if session is None: - logger.warning( - "agent.send dropped target=%s because its SDK session is not attached", - target_agent_id, - ) - return False - try: - async with session_write_lock(session): - await session.add_items([self._message_to_session_item(message)]) - except Exception: - logger.exception( - "agent.send failed to append to SDK session target=%s", - target_agent_id, - ) - return False - async with self._lock: - self.pending_counts[target_agent_id] = self.pending_counts.get(target_agent_id, 0) + 1 - self.runtimes.setdefault(target_agent_id, AgentRuntime()).wake.set() if stream is not None and interrupt and interrupt_on_message: stream.cancel(mode="immediate") await self._maybe_snapshot() return True - async def wait_for_message(self, agent_id: str) -> None: + async def wait_for_message(self, agent_id: str, *, timeout: float | None = None) -> bool: + """Wait until a message is ready for ``agent_id``; False on ``timeout``. + + An agent parked with an error (``failed``/``crashed``) is only released + by a user message; peer messages stay queued until then. + """ while True: async with self._lock: + runtime = self.runtimes.setdefault(agent_id, AgentRuntime()) reserve_exit = self._reserve_stopped and self.parent_of.get(agent_id) is not None - if self._budget_stopped or reserve_exit or self.pending_counts.get(agent_id, 0) > 0: - return - wake = self.runtimes.setdefault(agent_id, AgentRuntime()).wake + pending_ready = ( + self.pending_counts.get(agent_id, 0) > 0 and not runtime.user_wake_required + ) + if self._budget_stopped or reserve_exit or pending_ready: + return True + wake = runtime.wake wake.clear() - await wake.wait() + if timeout is None: + await wake.wait() + else: + try: + await asyncio.wait_for(wake.wait(), timeout) + except TimeoutError: + return False async def consume_pending( self, @@ -253,17 +264,42 @@ class AgentCoordinator: *, include_items: bool = False, ) -> tuple[int, list[Any]]: + """Drain the agent's mailbox into its own SDK session. + + Runs in the receiving agent's context, so a slow or wedged session + write can never block a sender. + """ async with self._lock: - count = self.pending_counts.get(agent_id, 0) + runtime = self.runtimes.setdefault(agent_id, AgentRuntime()) + queued = list(runtime.mailbox) + runtime.mailbox.clear() + count = max(self.pending_counts.get(agent_id, 0), len(queued)) self.pending_counts[agent_id] = 0 - session = self.runtimes.get(agent_id, AgentRuntime()).session + session = runtime.session if count <= 0: return 0, [] + items = [self._message_to_session_item(m) for m in queued] + if items: + if session is None: + logger.warning( + "agent %s has no SDK session attached; %d queued messages were not persisted", + agent_id, + len(items), + ) + else: + try: + async with session_write_lock(session): + await session.add_items(items) + except Exception: + logger.exception( + "failed to append %d queued messages to the session of %s", + len(items), + agent_id, + ) await self._maybe_snapshot() - if not include_items or session is None: + if not include_items: return count, [] - items = await session.get_items() - return count, list(items[-count:]) + return count, items async def request_stop(self, agent_id: str) -> None: async with self._lock: @@ -374,6 +410,11 @@ class AgentCoordinator: "names": dict(self.names), "metadata": {aid: dict(md) for aid, md in self.metadata.items()}, "pending_counts": dict(self.pending_counts), + "mailboxes": { + aid: [dict(m) for m in runtime.mailbox] + for aid, runtime in self.runtimes.items() + if runtime.mailbox + }, "errors": dict(self.errors), "budget_stopped": self._budget_stopped, "reserve_stopped": self._reserve_stopped, @@ -388,6 +429,12 @@ class AgentCoordinator: self.metadata = {aid: dict(md) for aid, md in snap.get("metadata", {}).items()} self.pending_counts = dict(snap.get("pending_counts", {})) self.errors = dict(snap.get("errors", {})) + mailboxes = snap.get("mailboxes", {}) + if isinstance(mailboxes, dict): + for aid, msgs in mailboxes.items(): + if isinstance(msgs, list): + runtime = self.runtimes.setdefault(aid, AgentRuntime()) + runtime.mailbox = [dict(m) for m in msgs if isinstance(m, dict)] self._budget_stopped = bool(snap.get("budget_stopped", False)) self._reserve_stopped = bool(snap.get("reserve_stopped", False)) self._budget_paused = bool(snap.get("budget_paused", False)) diff --git a/strix/core/execution.py b/strix/core/execution.py index fec727eb..747aad64 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -9,6 +9,7 @@ import uuid from collections.abc import Callable from typing import TYPE_CHECKING, Any, cast +import litellm from agents import RunConfig, Runner from agents.exceptions import AgentsException, MaxTurnsExceeded, UserError from agents.sandbox.errors import ExecTransportError @@ -16,11 +17,10 @@ from docker import errors as docker_errors # type: ignore[import-untyped, unuse from openai import ( APIConnectionError, APIError, - APIStatusError, APITimeoutError, - RateLimitError, ) +from strix.config import codex from strix.core.hooks import ( BudgetExceededError, BudgetPausedError, @@ -88,10 +88,9 @@ async def _compact_session( ) -_TRANSIENT_MODEL_STATUS_CODES = frozenset({408, 500, 502, 503, 504}) -_MAX_TRANSIENT_MODEL_RETRIES = 4 +_MAX_TRANSIENT_MODEL_RETRIES = 5 _TRANSIENT_MODEL_RETRY_BASE_DELAY_S = 2.0 -_TRANSIENT_MODEL_RETRY_MAX_DELAY_S = 30.0 +_TRANSIENT_MODEL_RETRY_MAX_DELAY_S = 90.0 def _model_error_status_code(exc: BaseException) -> int | None: @@ -100,15 +99,16 @@ def _model_error_status_code(exc: BaseException) -> int | None: def _is_transient_model_error(exc: BaseException) -> bool: - if isinstance(exc, RateLimitError): + if codex.is_content_guardrail_error(exc): return False - if isinstance(exc, APITimeoutError | APIConnectionError): + if isinstance( + exc, APITimeoutError | APIConnectionError | TimeoutError | ConnectionError | OSError + ): return True - if isinstance(exc, APIStatusError): - return exc.status_code in _TRANSIENT_MODEL_STATUS_CODES - if isinstance(exc, APIError): - return _model_error_status_code(exc) is None - return False + code = _model_error_status_code(exc) + if code is not None: + return bool(litellm._should_retry(code)) + return isinstance(exc, APIError) def _transient_model_retry_delay(attempt: int) -> float: @@ -153,7 +153,7 @@ async def run_agent_loop( if not (start_parked and interactive): if interactive: with contextlib.suppress(BudgetPausedError): - result = await _run_cycle( + result = await _run_cycle_parked( agent, coordinator, agent_id, @@ -162,7 +162,6 @@ async def run_agent_loop( context=context, max_turns=max_turns, session=session, - interactive=interactive, event_sink=event_sink, hooks=hooks, ) @@ -184,8 +183,9 @@ async def run_agent_loop( return result while True: + timeout = await _plain_waiting_timeout(coordinator, agent_id, context) try: - await coordinator.wait_for_message(agent_id) + woke = await coordinator.wait_for_message(agent_id, timeout=timeout) except asyncio.CancelledError: return result @@ -197,9 +197,21 @@ async def run_agent_loop( await coordinator.set_status(agent_id, "stopped") raise SubagentBudgetReservedError("scan reached the sub-agent budget reserve") + if not woke: + logger.info("agent %s reached its waiting timeout; auto-resuming", agent_id) + await coordinator.send( + agent_id, + { + "from": "system", + "type": "auto_resume", + "content": "Waiting timeout reached. Resuming execution.", + }, + interrupt=False, + ) + await coordinator.consume_pending(agent_id) with contextlib.suppress(BudgetPausedError): - result = await _run_cycle( + result = await _run_cycle_parked( agent, coordinator, agent_id, @@ -208,7 +220,6 @@ async def run_agent_loop( context=context, max_turns=max_turns, session=session, - interactive=interactive, event_sink=event_sink, hooks=hooks, ) @@ -431,6 +442,64 @@ async def _run_noninteractive_until_lifecycle( ) +_WAITING_AUTO_RESUME_TIMEOUT_S = 600.0 + + +async def _plain_waiting_timeout( + coordinator: AgentCoordinator, + agent_id: str, + context: dict[str, Any], +) -> float | None: + """Auto-resume timeout for a plainly-waiting subagent; None waits forever.""" + if context.get("parent_id") is None: + return None + async with coordinator._lock: + status = coordinator.statuses.get(agent_id) + has_error = agent_id in coordinator.errors + runtime = coordinator.runtimes.get(agent_id) + gated = runtime.user_wake_required if runtime is not None else False + if status == "waiting" and not has_error and not gated: + return _WAITING_AUTO_RESUME_TIMEOUT_S + return None + + +async def _run_cycle_parked( + agent: Any, + coordinator: AgentCoordinator, + agent_id: str, + *, + input_data: Any, + run_config: RunConfig, + context: dict[str, Any], + max_turns: int, + session: Session | None, + event_sink: StreamEventSink | None, + hooks: RunHooks[dict[str, Any]] | None, +) -> RunResultBase | None: + """Interactive run cycle that parks on any error instead of killing the runner.""" + try: + return await _run_cycle( + agent, + coordinator, + agent_id, + input_data=input_data, + run_config=run_config, + context=context, + max_turns=max_turns, + session=session, + interactive=True, + event_sink=event_sink, + hooks=hooks, + ) + except (BudgetExceededError, BudgetPausedError, SubagentBudgetReservedError): + raise + except Exception as exc: + logger.exception("error escaped the run cycle for %s; parking as failed", agent_id) + await coordinator.set_status(agent_id, "failed", error=str(exc) or type(exc).__name__) + await _notify_parent_on_terminal(coordinator, agent_id, "failed") + return None + + async def _run_cycle( # noqa: PLR0912, PLR0915 agent: Any, coordinator: AgentCoordinator, diff --git a/strix/interface/tui/app.py b/strix/interface/tui/app.py index 01fc5488..caf22d42 100644 --- a/strix/interface/tui/app.py +++ b/strix/interface/tui/app.py @@ -1042,7 +1042,7 @@ class StrixTUIApp(App): # type: ignore[misc] status=status, error_message=error or "", ) - if status in {"failed", "crashed", "waiting"} and error: + if error: if agent_id not in self._error_noted_agents: self._error_noted_agents.add(agent_id) self.live_view.record_agent_error(agent_id, error) diff --git a/tests/test_execution.py b/tests/test_execution.py index a16880a4..50675b41 100644 --- a/tests/test_execution.py +++ b/tests/test_execution.py @@ -5,12 +5,13 @@ from __future__ import annotations import asyncio import contextlib import json -from typing import Any +from typing import Any, cast import pytest from agents.memory import SQLiteSession from agents.tool_context import ToolContext +from strix.core import execution from strix.core.agents import AgentCoordinator from strix.core.execution import ( _notify_parent_on_terminal, @@ -495,3 +496,99 @@ async def test_terminal_notice_does_not_cancel_parent_stream(tmp_path: Any) -> N assert stream.cancelled is False assert coordinator.pending_counts.get("root", 0) > 0 session.close() + + +@pytest.mark.asyncio +async def test_send_queues_without_session_and_drains_on_consume(tmp_path: Any) -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + + assert await coordinator.send("root", {"from": "user", "content": "hello"}) is True + assert coordinator.pending_counts["root"] == 1 + + session = SQLiteSession("root", tmp_path / "agents.db") + await coordinator.attach_runtime("root", session=session) + + count, items = await coordinator.consume_pending("root", include_items=True) + assert count == 1 + assert items[0]["content"] == "hello" + stored = await session.get_items() + last = cast("dict[str, Any]", stored[-1]) + assert last["content"] == "hello" + session.close() + + +@pytest.mark.asyncio +async def test_error_parked_agent_only_released_by_user_message(tmp_path: Any) -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.register("child", "recon", parent_id="root") + session = SQLiteSession("child", tmp_path / "agents.db") + await coordinator.attach_runtime("child", session=session) + await coordinator.set_status("child", "crashed", error="boom") + + await coordinator.send("child", {"from": "root", "content": "peer nudge"}) + waiter = asyncio.create_task(coordinator.wait_for_message("child")) + await asyncio.sleep(0.05) + assert not waiter.done() + + await coordinator.send("child", {"from": "user", "content": "wake up"}) + assert await asyncio.wait_for(waiter, timeout=1.0) is True + + count, items = await coordinator.consume_pending("child", include_items=True) + assert count == 2 + assert items[0]["content"].endswith("peer nudge") + assert items[1]["content"] == "wake up" + session.close() + + +@pytest.mark.asyncio +async def test_wait_for_message_timeout_returns_false() -> None: + coordinator = AgentCoordinator() + await coordinator.register("child", "recon", parent_id="root") + + assert await coordinator.wait_for_message("child", timeout=0.05) is False + + +@pytest.mark.asyncio +async def test_snapshot_round_trip_preserves_mailboxes() -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.send("root", {"from": "user", "content": "queued"}) + + snap = await coordinator.snapshot() + restored = AgentCoordinator() + await restored.restore(snap) + + assert restored.pending_counts["root"] == 1 + assert restored.runtimes["root"].mailbox == [{"from": "user", "content": "queued"}] + + +@pytest.mark.asyncio +async def test_run_cycle_parked_parks_instead_of_raising( + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def _boom(*_args: Any, **_kwargs: Any) -> Any: + raise RuntimeError("unexpected explosion") + + monkeypatch.setattr(execution, "_run_cycle", _boom) + + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + + result = await execution._run_cycle_parked( + object(), + coordinator, + "root", + input_data=[], + run_config=None, # type: ignore[arg-type] + context={}, + max_turns=5, + session=None, + event_sink=None, + hooks=None, + ) + + assert result is None + assert coordinator.statuses["root"] == "failed" + assert coordinator.errors["root"] == "unexpected explosion" diff --git a/tests/test_execution_transient_retry.py b/tests/test_execution_transient_retry.py index 889eb96f..d81e4e70 100644 --- a/tests/test_execution_transient_retry.py +++ b/tests/test_execution_transient_retry.py @@ -15,6 +15,7 @@ from openai import ( RateLimitError, ) +from strix.config import codex from strix.core import execution from strix.core.agents import AgentCoordinator @@ -55,11 +56,27 @@ def test_server_errors_are_transient() -> None: assert execution._is_transient_model_error(_status_error(status)) is True -def test_rate_limit_is_not_retried_here() -> None: +def test_rate_limit_is_retried() -> None: rate_limited = RateLimitError( "slow down", response=httpx.Response(429, request=_request()), body=None ) - assert execution._is_transient_model_error(rate_limited) is False + assert execution._is_transient_model_error(rate_limited) is True + + +def test_dns_and_connection_errors_are_transient() -> None: + assert execution._is_transient_model_error(OSError("nodename nor servname provided")) is True + assert execution._is_transient_model_error(ConnectionError("reset")) is True + assert execution._is_transient_model_error(TimeoutError("timed out")) is True + + +def test_content_guardrail_is_not_retried() -> None: + guardrail = APIError( + "This content was flagged for possible cybersecurity risk", + _request(), + body=None, + ) + assert codex.is_content_guardrail_error(guardrail) is True + assert execution._is_transient_model_error(guardrail) is False def test_client_errors_are_not_transient() -> None: