Skip to content

Commit cffaa34

Browse files
danielmillerpclaude
andcommitted
test(tutorials): poll task state instead of a fixed 1s sleep
The async base tutorials sent the agent's reply, slept 1s, then asserted the task state held 3 messages. The agent persists the turn to state only after it sends the reply, so a slow model call left the read seeing the previous turn and failed with "assert 1 == 3". Add wait_for_state_messages to test_utils, which polls the state until it reaches the expected count or a 30s timeout, and use it in 010_multiturn and 020_streaming. A state that never updates still fails with its real length. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
1 parent 85b0342 commit cffaa34

5 files changed

Lines changed: 123 additions & 22 deletions

File tree

‎examples/tutorials/10_async/00_base/010_multiturn/tests/test_agent.py‎

Lines changed: 3 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@
2424
import pytest_asyncio
2525
from test_utils.async_utils import (
2626
stream_agent_response,
27+
wait_for_state_messages,
2728
send_event_and_poll_yielding,
2829
)
2930

@@ -124,11 +125,7 @@ async def test_send_event_and_poll(self, client: AsyncAgentex, agent_id: str):
124125
assert user_message_found, "User message not found"
125126
assert agent_response_found, "Agent response not found"
126127

127-
await asyncio.sleep(1) # wait for state to be updated
128-
states = await client.states.list(agent_id=agent_id, task_id=task.id)
129-
assert len(states) == 1
130-
state = states[0].state
131-
messages = state.get("messages", [])
128+
messages = await wait_for_state_messages(client, agent_id, task.id, expected_count=3)
132129

133130
assert isinstance(messages, list)
134131
assert len(messages) == 3
@@ -207,11 +204,7 @@ async def stream_messages() -> None:
207204
assert agent_response_found, "Agent response not found in stream"
208205

209206
# Verify the state has been updated
210-
await asyncio.sleep(1) # wait for state to be updated
211-
states = await client.states.list(agent_id=agent_id, task_id=task.id)
212-
assert len(states) == 1
213-
state = states[0].state
214-
messages = state.get("messages", [])
207+
messages = await wait_for_state_messages(client, agent_id, task.id, expected_count=3)
215208

216209
assert isinstance(messages, list)
217210
assert len(messages) == 3

‎examples/tutorials/10_async/00_base/020_streaming/tests/test_agent.py‎

Lines changed: 3 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@
2424
import pytest_asyncio
2525
from test_utils.async_utils import (
2626
stream_agent_response,
27+
wait_for_state_messages,
2728
send_event_and_poll_yielding,
2829
)
2930

@@ -122,11 +123,7 @@ async def test_send_event_and_poll(self, client: AsyncAgentex, agent_id: str):
122123
assert agent_response_found, "Agent response not found"
123124

124125
# assert the state has been updated
125-
await asyncio.sleep(1) # wait for state to be updated
126-
states = await client.states.list(agent_id=agent_id, task_id=task.id)
127-
assert len(states) == 1
128-
state = states[0].state
129-
messages = state.get("messages", [])
126+
messages = await wait_for_state_messages(client, agent_id, task.id, expected_count=3)
130127

131128
assert isinstance(messages, list)
132129
assert len(messages) == 3
@@ -205,11 +202,7 @@ async def stream_messages() -> None:
205202
assert delta_messages_found, "Delta messages not found in stream (streaming response expected)"
206203

207204
# Verify the state has been updated
208-
await asyncio.sleep(1) # wait for state to be updated
209-
states = await client.states.list(agent_id=agent_id, task_id=task.id)
210-
assert len(states) == 1
211-
state: dict[str, object] = states[0].state
212-
messages = state.get("messages", [])
205+
messages = await wait_for_state_messages(client, agent_id, task.id, expected_count=3)
213206

214207
assert isinstance(messages, list)
215208
assert len(messages) == 3

‎examples/tutorials/test_utils/async_utils.py‎

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -156,6 +156,42 @@ async def poll_messages(
156156
await asyncio.sleep(sleep_interval)
157157

158158

159+
async def wait_for_state_messages(
160+
client: AsyncAgentex,
161+
agent_id: str,
162+
task_id: str,
163+
expected_count: int,
164+
timeout: float = 30,
165+
sleep_interval: float = 0.5,
166+
) -> list:
167+
"""
168+
Poll the task state until its "messages" list reaches expected_count.
169+
170+
Agents typically send their reply before they persist the turn to task state,
171+
so reading the state right after the reply arrives can see the previous turn.
172+
173+
Args:
174+
client: AgentEx client instance
175+
agent_id: The agent ID
176+
task_id: The task ID
177+
expected_count: Number of state messages to wait for
178+
timeout: Maximum seconds to poll (default: 30)
179+
sleep_interval: Seconds to sleep between polls (default: 0.5)
180+
181+
Returns:
182+
The state's "messages" list from the last poll. Callers assert on it, so a
183+
state that never reaches expected_count still fails with the real length.
184+
"""
185+
deadline = time.monotonic() + timeout
186+
while True:
187+
states = await client.states.list(agent_id=agent_id, task_id=task_id)
188+
assert len(states) == 1, f"Expected exactly one state, got {len(states)}"
189+
messages = states[0].state.get("messages", [])
190+
if len(messages) >= expected_count or time.monotonic() >= deadline:
191+
return messages
192+
await asyncio.sleep(sleep_interval)
193+
194+
159195
async def send_event_and_stream(
160196
client: AsyncAgentex,
161197
agent_id: str,
Lines changed: 79 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,79 @@
1+
"""Tests for wait_for_state_messages in the tutorials' test_utils.
2+
3+
The tutorials' `test_utils` package shares its name with `tests/test_utils`, so the
4+
module is loaded from its file path instead of imported by name.
5+
"""
6+
7+
from __future__ import annotations
8+
9+
import time
10+
import importlib.util
11+
from types import ModuleType, SimpleNamespace
12+
from pathlib import Path
13+
14+
import pytest
15+
16+
_ASYNC_UTILS = Path(__file__).resolve().parents[2] / "examples" / "tutorials" / "test_utils" / "async_utils.py"
17+
18+
19+
def _load_async_utils() -> ModuleType:
20+
spec = importlib.util.spec_from_file_location("tutorial_async_utils", _ASYNC_UTILS)
21+
assert spec is not None and spec.loader is not None
22+
module = importlib.util.module_from_spec(spec)
23+
spec.loader.exec_module(module)
24+
return module
25+
26+
27+
wait_for_state_messages = _load_async_utils().wait_for_state_messages
28+
29+
30+
class _FakeStates:
31+
"""Returns one state whose message count grows from 1 to 3 once `write_after` seconds pass."""
32+
33+
def __init__(self, write_after: float | None):
34+
self._start = time.monotonic()
35+
self._write_after = write_after
36+
self.calls = 0
37+
38+
async def list(self, agent_id: str, task_id: str) -> list[SimpleNamespace]:
39+
del agent_id, task_id
40+
self.calls += 1
41+
written = self._write_after is not None and time.monotonic() - self._start >= self._write_after
42+
return [SimpleNamespace(state={"messages": [{"role": "system"}] * (3 if written else 1)})]
43+
44+
45+
def _client(write_after: float | None) -> SimpleNamespace:
46+
return SimpleNamespace(states=_FakeStates(write_after))
47+
48+
49+
@pytest.mark.asyncio
50+
async def test_waits_for_a_state_write_that_lands_after_the_reply() -> None:
51+
client = _client(write_after=0.3)
52+
53+
messages = await wait_for_state_messages(client, "agent", "task", expected_count=3, timeout=5, sleep_interval=0.05)
54+
55+
assert len(messages) == 3
56+
assert client.states.calls > 1
57+
58+
59+
@pytest.mark.asyncio
60+
async def test_returns_the_real_state_at_the_deadline_when_it_never_updates() -> None:
61+
client = _client(write_after=None)
62+
start = time.monotonic()
63+
64+
messages = await wait_for_state_messages(client, "agent", "task", expected_count=3, timeout=0.3, sleep_interval=0.05)
65+
66+
# It kept polling until the deadline instead of returning early. No upper bound:
67+
# a busy worker may wake late, and returning at all already proves it stopped.
68+
assert len(messages) == 1
69+
assert time.monotonic() - start >= 0.3
70+
71+
72+
@pytest.mark.asyncio
73+
async def test_returns_on_the_first_poll_when_the_state_is_already_written() -> None:
74+
client = _client(write_after=0)
75+
76+
messages = await wait_for_state_messages(client, "agent", "task", expected_count=3)
77+
78+
assert len(messages) == 3
79+
assert client.states.calls == 1

‎uv.lock‎

Lines changed: 2 additions & 2 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

0 commit comments

Comments
 (0)