mirror of
https://github.com/usestrix/strix.git
synced 2026-08-21 02:45:31 +02:00
fix: restore pre-v1 lifecycle resilience: mailbox delivery, uniform revival, waiting timeout, broader retries
This commit is contained in:
+78
-31
@@ -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))
|
||||
|
||||
+86
-17
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
+98
-1
@@ -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"
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user