Compare commits

...
8 changed files with 384 additions and 3 deletions
+1
View File
@@ -238,6 +238,7 @@ ignore = [
"tests/test_disable_streaming.py" = ["N802"]
"tests/test_tool_call_ids.py" = ["N802"]
"tests/test_tool_call_limits.py" = ["N802", "SLF001"]
"tests/test_stream_idle_timeout.py" = ["N802", "SLF001"]
"tests/test_unknown_tool_recovery.py" = ["N802"]
"tests/test_report_pdf.py" = ["S105", "S106"]
# Stdlib HTTP handler overrides (do_GET/do_POST) and lazy imports that avoid a
+1
View File
@@ -28,6 +28,7 @@ INTER-AGENT MESSAGES:
- Messages from other agents arrive prefixed with a header like `[Message from agent <name> | type=... | priority=...]`. Treat them as internal context — never repeat them verbatim in your own output.
- Treat agent identity / inherited-context preambles as internal metadata; do not echo them in outputs or tool calls.
- Minimize inter-agent messaging: only message when essential for coordination or assistance; avoid routine status updates; batch non-urgent information; prefer parent/child completion flows and shared artifacts over messaging
- wait_for_agents blocks and resumes you automatically, so it is never a poll you repeat: issue exactly ONE wait, then stop and react to what it returns. Never write out a wait/check loop (wait → view_agent_graph → wait → ...) ahead of time — those extra calls only strand you and are collapsed anyway
{% if interactive %}
INTERACTIVE BEHAVIOR:
+54 -3
View File
@@ -2,11 +2,13 @@
from __future__ import annotations
import asyncio
import contextlib
import inspect
import logging
import os
import time
from collections.abc import AsyncGenerator
from typing import TYPE_CHECKING, Any, cast
from agents import (
@@ -252,11 +254,23 @@ class _TurnGuardModel(Model):
Tool-call volume: a degenerate response can queue hundreds of calls that
the run loop then honours one by one. Only the first
``LLM_MAX_TOOL_CALLS_PER_TURN`` calls of a response are kept.
Stalled streams: a turn that emits a few tokens and then goes silent is
not covered by the request timeout, which resets on any byte (keepalives
included). ``LLM_STREAM_IDLE_TIMEOUT`` bounds the gap between events so the
turn fails instead of hanging, and the existing retry path replays it.
"""
def __init__(self, inner: Model, *, max_tool_calls_per_turn: int = 0) -> None:
def __init__(
self,
inner: Model,
*,
max_tool_calls_per_turn: int = 0,
stream_idle_timeout: float = 0.0,
) -> None:
self._inner = inner
self._max_tool_calls_per_turn = max_tool_calls_per_turn
self._stream_idle_timeout = stream_idle_timeout
def _limiter(self) -> TurnToolCallLimiter:
return TurnToolCallLimiter(self._max_tool_calls_per_turn)
@@ -337,13 +351,41 @@ class _TurnGuardModel(Model):
conversation_id=conversation_id,
prompt=prompt,
)
async for event in stream:
async for event in _with_idle_timeout(stream, self._stream_idle_timeout):
guarded = _guard_event(event, rewriter, limiter)
if guarded is not None:
yield guarded
self._log_dropped(limiter)
async def _aclose(stream: AsyncIterator[TResponseStreamEvent]) -> None:
if isinstance(stream, AsyncGenerator):
with contextlib.suppress(Exception):
await stream.aclose()
async def _with_idle_timeout(
stream: AsyncIterator[TResponseStreamEvent], timeout: float
) -> AsyncIterator[TResponseStreamEvent]:
if timeout <= 0:
async for event in stream:
yield event
return
iterator = stream.__aiter__()
while True:
try:
event = await asyncio.wait_for(iterator.__anext__(), timeout)
except StopAsyncIteration:
return
except TimeoutError:
await _aclose(stream)
message = f"model stream produced no event for {timeout:.0f}s"
logger.warning("%s; abandoning the turn", message)
raise TimeoutError(message) from None
yield event
def _guard_event(
event: TResponseStreamEvent, rewriter: TurnCallIdRewriter, limiter: TurnToolCallLimiter
) -> TResponseStreamEvent | None:
@@ -429,6 +471,7 @@ class StrixProvider(MultiProvider):
def get_model(self, model_name: str | None) -> Model:
llm = load_settings().llm
slug = codex.subscription_model(model_name)
idle_timeout = float(llm.stream_idle_timeout)
if slug:
# The ChatGPT subscription backend is always streamed; it has no
# non-streaming mode to fall back to, so LLM_DISABLE_STREAMING
@@ -442,7 +485,15 @@ class StrixProvider(MultiProvider):
model = super().get_model(model_name)
if llm.disable_streaming:
model = _NonStreamingModel(model)
return _TurnGuardModel(model, max_tool_calls_per_turn=llm.max_tool_calls_per_turn)
# The wrapper emits its single event only once the whole request
# is done, so an idle gap is meaningless here; the request
# timeout bounds it instead.
idle_timeout = 0.0
return _TurnGuardModel(
model,
max_tool_calls_per_turn=llm.max_tool_calls_per_turn,
stream_idle_timeout=idle_timeout,
)
DEFAULT_MODEL_RETRY = ModelRetrySettings(
+1
View File
@@ -57,6 +57,7 @@ class LlmSettings(BaseSettings):
alias="LLM_DISABLE_STREAMING",
)
timeout: int = Field(default=300, alias="LLM_TIMEOUT")
stream_idle_timeout: int = Field(default=300, ge=0, alias="LLM_STREAM_IDLE_TIMEOUT")
max_tool_calls_per_turn: int = Field(
default=32,
ge=0,
+3
View File
@@ -20,6 +20,8 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__)
LLM_TURN_KEY = "llm_turn"
_STAGE_LABELS: tuple[str, ...] = ("NOTICE", "URGENT", "CRITICAL")
_TURN_WARN_BANDS: tuple[float, ...] = (0.70, 0.85, 0.95)
_ROOT_BUDGET_WARN_BANDS: tuple[float, ...] = (0.70, 0.85, 0.95)
@@ -144,6 +146,7 @@ class ReportUsageHooks(RunHooks[dict[str, Any]]):
system_prompt: str | None, # noqa: ARG002
input_items: list[TResponseInputItem],
) -> None:
context.context[LLM_TURN_KEY] = int(context.context.get(LLM_TURN_KEY, 0)) + 1
try:
self._maybe_warn_turns(context, input_items)
self._maybe_warn_budget(context, input_items)
+25
View File
@@ -14,6 +14,7 @@ from agents import RunContextWrapper, function_tool
from strix.core.agents import Status, coordinator_from_context
from strix.core.execution import notify_parent_on_terminal
from strix.core.hooks import LLM_TURN_KEY
from strix.skills import validate_requested_skills
@@ -224,6 +225,7 @@ _WAIT_DEFAULT_TIMEOUT_S = 300
# ``timeout_seconds`` the model asks for. One second of headroom lets the
# tool's own timeout fire first and return a clean result.
_WAIT_HARD_CEILING_S = _WAIT_DEFAULT_TIMEOUT_S + 1
_WAITED_TURN_KEY = "waited_llm_turn"
@function_tool(timeout=_WAIT_HARD_CEILING_S)
@@ -239,6 +241,11 @@ async def wait_for_agents( # noqa: PLR0911
completion reports. You resume the instant any message arrives, so
size ``timeout_seconds`` to the work you're awaiting.
**Issue exactly one wait, then stop and react to what it returns.**
This call blocks and resumes on its own; it is not a poll you repeat.
Do not write out a wait/check loop ahead of time — a second wait in
the same turn returns immediately without waiting.
**This tool is only for waiting on other agents.** Two things it is
NOT for:
@@ -290,6 +297,24 @@ async def wait_for_agents( # noqa: PLR0911
default=str,
)
turn = inner.get(LLM_TURN_KEY)
if turn is not None and inner.get(_WAITED_TURN_KEY) == turn:
return json.dumps(
{
"success": True,
"wait_outcome": "already_waited",
"reason": reason,
"note": (
"You already waited in this turn. A single wait_for_agents blocks and "
"resumes on its own, so queueing more waits only strands you — issue one "
"wait, then react to what it returns."
),
},
ensure_ascii=False,
default=str,
)
inner[_WAITED_TURN_KEY] = turn
async with coordinator._lock:
stopped = coordinator.statuses.get(me) == "stopped"
if stopped:
+173
View File
@@ -0,0 +1,173 @@
"""Tests for the model-stream idle watchdog.
A turn that streams a few tokens and then goes silent is not covered by the
request timeout: the read timeout resets on every byte, keepalives included.
The watchdog bounds the gap between events so the turn fails and can be
retried instead of parking the agent forever.
"""
from __future__ import annotations
import asyncio
import json
import threading
import time
from http.server import BaseHTTPRequestHandler, HTTPServer
from typing import TYPE_CHECKING, Any
import pytest
from agents.model_settings import ModelSettings
from agents.models.interface import Model, ModelTracing
from agents.models.openai_chatcompletions import OpenAIChatCompletionsModel
from openai import AsyncOpenAI
from strix.config import loader
from strix.config.loader import load_settings
from strix.config.models import StrixProvider, _TurnGuardModel, _with_idle_timeout
if TYPE_CHECKING:
from collections.abc import AsyncIterator, Iterator
_STALL_SECONDS = 30.0
def _chunk(text: str) -> bytes:
payload = {
"id": "chatcmpl-1",
"object": "chat.completion.chunk",
"created": 0,
"model": "gw-model",
"choices": [{"index": 0, "delta": {"content": text}, "finish_reason": None}],
}
return b"data: " + json.dumps(payload).encode() + b"\n\n"
class _StallingHandler(BaseHTTPRequestHandler):
"""Streams a couple of tokens, then stops producing anything."""
stop = threading.Event()
def log_message(self, *args: Any) -> None:
pass
def do_POST(self) -> None:
length = int(self.headers.get("Content-Length", 0))
self.rfile.read(length)
self.send_response(200)
self.send_header("Content-Type", "text/event-stream")
self.end_headers()
self.wfile.write(_chunk("Now"))
self.wfile.write(_chunk(" spawning"))
self.wfile.flush()
self.stop.wait(_STALL_SECONDS)
@pytest.fixture
def stalling_gateway() -> Iterator[str]:
_StallingHandler.stop.clear()
server = HTTPServer(("127.0.0.1", 0), _StallingHandler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield f"http://127.0.0.1:{server.server_address[1]}/v1"
finally:
_StallingHandler.stop.set()
server.shutdown()
server.server_close()
def _stream(base_url: str, *, idle_timeout: float) -> AsyncIterator[Any]:
client = AsyncOpenAI(api_key="tok", base_url=base_url, max_retries=0, timeout=_STALL_SECONDS)
inner: Model = OpenAIChatCompletionsModel(model="gw-model", openai_client=client)
guarded = _TurnGuardModel(inner, stream_idle_timeout=idle_timeout)
return guarded.stream_response(
None,
"go",
ModelSettings(),
[],
None,
[],
ModelTracing.DISABLED,
previous_response_id=None,
conversation_id=None,
prompt=None,
)
async def _drain(base_url: str, *, idle_timeout: float) -> list[Any]:
return [event async for event in _stream(base_url, idle_timeout=idle_timeout)]
@pytest.mark.asyncio
async def test_stalled_stream_hangs_without_the_watchdog(stalling_gateway: str) -> None:
# Repro: tokens arrive, then nothing. Un-watched, the turn just sits there;
# the request timeout is far away and would reset on any keepalive byte.
with pytest.raises(TimeoutError):
await asyncio.wait_for(_drain(stalling_gateway, idle_timeout=0), timeout=2)
@pytest.mark.asyncio
async def test_stalled_stream_is_abandoned_by_the_watchdog(stalling_gateway: str) -> None:
started = time.monotonic()
with pytest.raises(TimeoutError, match="produced no event"):
await _drain(stalling_gateway, idle_timeout=1)
assert time.monotonic() - started < _STALL_SECONDS
@pytest.mark.asyncio
async def test_events_keep_flowing_while_the_stream_is_alive() -> None:
async def _live() -> AsyncIterator[Any]:
for i in range(5):
await asyncio.sleep(0.05)
yield f"event-{i}"
seen: list[Any] = [event async for event in _with_idle_timeout(_live(), 1.0)]
assert seen == [f"event-{i}" for i in range(5)]
@pytest.fixture
def _reset_settings(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
for key in ("STRIX_LLM", "LLM_DISABLE_STREAMING", "LLM_STREAM_IDLE_TIMEOUT"):
monkeypatch.delenv(key, raising=False)
monkeypatch.setattr(loader, "_cached", None)
monkeypatch.setattr(loader, "_override", None)
yield
class _DummyModel(Model):
async def get_response(self, *args: Any, **kwargs: Any) -> Any:
raise NotImplementedError
def stream_response(self, *args: Any, **kwargs: Any) -> Any:
raise NotImplementedError
def test_idle_timeout_is_configurable(
monkeypatch: pytest.MonkeyPatch, _reset_settings: None
) -> None:
monkeypatch.setattr("strix.config.models.MultiProvider.get_model", lambda *_: _DummyModel())
monkeypatch.setenv("LLM_STREAM_IDLE_TIMEOUT", "45")
load_settings()
model = StrixProvider().get_model("openai/gpt-4o-mini")
assert isinstance(model, _TurnGuardModel)
assert model._stream_idle_timeout == 45
def test_idle_timeout_is_off_without_streaming(
monkeypatch: pytest.MonkeyPatch, _reset_settings: None
) -> None:
# LLM_DISABLE_STREAMING turns the whole request into one event, so an idle
# gap would just be the request duration — the request timeout bounds that.
monkeypatch.setattr("strix.config.models.MultiProvider.get_model", lambda *_: _DummyModel())
monkeypatch.setenv("LLM_STREAM_IDLE_TIMEOUT", "45")
monkeypatch.setenv("LLM_DISABLE_STREAMING", "true")
load_settings()
model = StrixProvider().get_model("openai/gpt-4o-mini")
assert isinstance(model, _TurnGuardModel)
assert model._stream_idle_timeout == 0
+126
View File
@@ -0,0 +1,126 @@
"""Tests for collapsing repeated waits queued inside one model turn.
An orchestrator that writes out its whole poll loop ahead of time queues
many ``wait_for_agents`` calls in a single response. Each one parks for its
full timeout, so the agent stops reacting for hours while its children run
unsupervised. Only the first wait of a turn parks; the rest return at once.
"""
from __future__ import annotations
import asyncio
import json
import time
from typing import TYPE_CHECKING, Any, cast
import pytest
from agents import RunContextWrapper
from agents.tool_context import ToolContext
from strix.core.agents import AgentCoordinator
from strix.core.hooks import LLM_TURN_KEY, ReportUsageHooks
from strix.tools.agents_graph.tools import wait_for_agents
if TYPE_CHECKING:
from collections.abc import Iterator
_WAIT_SECONDS = 2
@pytest.fixture
def _fast_wait(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
# The real ceiling is 300s per wait; the shape of the bug is the same.
monkeypatch.setattr(
"strix.tools.agents_graph.tools._WAIT_DEFAULT_TIMEOUT_S", _WAIT_SECONDS, raising=True
)
yield
async def _context() -> dict[str, Any]:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
return {"agent_id": "root", "coordinator": coordinator}
async def _wait(inner: dict[str, Any]) -> dict[str, Any]:
ctx = ToolContext(
context=inner,
tool_name="wait_for_agents",
tool_call_id="call-1",
tool_arguments="{}",
)
raw: str = await wait_for_agents.on_invoke_tool(
ctx, json.dumps({"reason": "waiting for wave 1", "timeout_seconds": _WAIT_SECONDS})
)
return cast("dict[str, Any]", json.loads(raw))
@pytest.mark.asyncio
async def test_waits_queued_in_one_turn_each_park_without_the_guard(_fast_wait: None) -> None:
# Repro: no turn marker in context (as before the fix) — every queued wait
# parks for its full timeout, so N waits cost N x timeout.
inner = await _context()
started = time.monotonic()
outcomes = [(await _wait(inner))["wait_outcome"] for _ in range(3)]
elapsed = time.monotonic() - started
assert outcomes == ["timeout", "timeout", "timeout"]
assert elapsed >= 3 * _WAIT_SECONDS
@pytest.mark.asyncio
async def test_repeated_waits_in_one_turn_are_collapsed(_fast_wait: None) -> None:
inner = await _context()
inner[LLM_TURN_KEY] = 1
started = time.monotonic()
outcomes = [(await _wait(inner))["wait_outcome"] for _ in range(3)]
elapsed = time.monotonic() - started
assert outcomes == ["timeout", "already_waited", "already_waited"]
assert elapsed < 2 * _WAIT_SECONDS
@pytest.mark.asyncio
async def test_a_wait_in_the_next_turn_still_parks(_fast_wait: None) -> None:
inner = await _context()
inner[LLM_TURN_KEY] = 1
assert (await _wait(inner))["wait_outcome"] == "timeout"
assert (await _wait(inner))["wait_outcome"] == "already_waited"
inner[LLM_TURN_KEY] = 2
assert (await _wait(inner))["wait_outcome"] == "timeout"
@pytest.mark.asyncio
async def test_each_model_turn_bumps_the_turn_marker() -> None:
hooks = ReportUsageHooks(model="gw-model")
context: RunContextWrapper[dict[str, Any]] = RunContextWrapper(context={})
agent = cast("Any", None)
await hooks.on_llm_start(context, agent, None, [])
await hooks.on_llm_start(context, agent, None, [])
assert context.context[LLM_TURN_KEY] == 2
@pytest.mark.asyncio
async def test_a_collapsed_wait_still_reports_arriving_messages(_fast_wait: None) -> None:
inner = await _context()
inner[LLM_TURN_KEY] = 1
coordinator = cast("AgentCoordinator", inner["coordinator"])
async def _send() -> None:
await asyncio.sleep(0.1)
await coordinator.send("root", {"type": "information", "content": "child done"})
task = asyncio.create_task(_send())
first = await _wait(inner)
await task
assert first["wait_outcome"] == "message_arrived"
assert (await _wait(inner))["wait_outcome"] == "already_waited"