Skip to content

Commit 3fa209c

Browse files
committed
test: add coverage for Agent abort, resume, reset, and queue paths
Cover uncovered lines in agent/agent.py: - _MessageQueue "all" mode drain, has_items, clear - AgentState.pending_tool_calls setter (copy semantics) - Agent.reset() clearing state and queues - Agent.abort() setting the signal during active run - Agent.wait_for_idle() returning when no run and waiting for completion - Agent.prompt() with Message object and list[Message] inputs - Agent.resume() draining steering queue, error on empty assistant last Closes #47
1 parent b25975f commit 3fa209c

1 file changed

Lines changed: 247 additions & 1 deletion

File tree

tests/agent/test_agent.py

Lines changed: 247 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,8 @@
11
import asyncio
22

3-
from cubepi.agent.agent import Agent
3+
import pytest
4+
5+
from cubepi.agent.agent import Agent, _MessageQueue
46
from cubepi.agent.types import AgentTool
57
from cubepi.providers.base import (
68
AssistantMessage,
@@ -285,3 +287,247 @@ async def test_resume_processes_follow_up_messages(self):
285287
)
286288
assert has_follow_up
287289
assert isinstance(agent.state.messages[-1], AssistantMessage)
290+
291+
async def test_resume_drains_steering_queue_before_follow_up(self):
292+
"""resume() should drain the steering queue first when the last
293+
message is from the assistant."""
294+
provider = FauxProvider()
295+
provider.set_responses(
296+
[
297+
faux_assistant_message("Initial"),
298+
faux_assistant_message("Steered"),
299+
]
300+
)
301+
agent = Agent(provider=provider, model=make_model())
302+
303+
await agent.prompt("hello")
304+
agent.steer(UserMessage(content=[TextContent(text="steer-msg")]))
305+
await agent.resume()
306+
307+
has_steer = any(
308+
isinstance(m, UserMessage)
309+
and any(
310+
isinstance(c, TextContent) and c.text == "steer-msg" for c in m.content
311+
)
312+
for m in agent.state.messages
313+
)
314+
assert has_steer
315+
assert isinstance(agent.state.messages[-1], AssistantMessage)
316+
317+
async def test_resume_raises_on_assistant_last_with_empty_queues(self):
318+
"""resume() raises RuntimeError when the last message is from the
319+
assistant and both steering and follow-up queues are empty."""
320+
provider = FauxProvider()
321+
provider.set_responses([faux_assistant_message("done")])
322+
agent = Agent(provider=provider, model=make_model())
323+
324+
await agent.prompt("hello")
325+
326+
with pytest.raises(RuntimeError, match="Cannot continue from message role"):
327+
await agent.resume()
328+
329+
330+
class TestMessageQueueAllMode:
331+
def test_drain_returns_all_messages_at_once(self):
332+
q = _MessageQueue(mode="all")
333+
m1 = UserMessage(content=[TextContent(text="a")])
334+
m2 = UserMessage(content=[TextContent(text="b")])
335+
m3 = UserMessage(content=[TextContent(text="c")])
336+
337+
q.enqueue(m1)
338+
q.enqueue(m2)
339+
q.enqueue(m3)
340+
341+
drained = q.drain()
342+
assert drained == [m1, m2, m3]
343+
assert not q.has_items()
344+
345+
def test_drain_returns_empty_when_no_items(self):
346+
q = _MessageQueue(mode="all")
347+
assert q.drain() == []
348+
349+
def test_has_items_reflects_state(self):
350+
q = _MessageQueue(mode="all")
351+
assert not q.has_items()
352+
q.enqueue(UserMessage(content=[TextContent(text="x")]))
353+
assert q.has_items()
354+
355+
def test_clear_removes_all(self):
356+
q = _MessageQueue(mode="all")
357+
q.enqueue(UserMessage(content=[TextContent(text="a")]))
358+
q.enqueue(UserMessage(content=[TextContent(text="b")]))
359+
q.clear()
360+
assert not q.has_items()
361+
assert q.drain() == []
362+
363+
364+
class TestAgentStatePendingToolCalls:
365+
def test_setter_makes_a_copy(self):
366+
provider = FauxProvider()
367+
agent = Agent(provider=provider, model=make_model())
368+
369+
original = {"call-1", "call-2"}
370+
agent.state.pending_tool_calls = original
371+
372+
retrieved = agent.state.pending_tool_calls
373+
assert retrieved == {"call-1", "call-2"}
374+
# Must be a distinct set, not the same object
375+
assert retrieved is not original
376+
377+
378+
class TestAgentReset:
379+
async def test_reset_clears_state_after_prompt(self):
380+
provider = FauxProvider()
381+
provider.set_responses([faux_assistant_message("response")])
382+
agent = Agent(provider=provider, model=make_model())
383+
384+
await agent.prompt("hello")
385+
assert len(agent.state.messages) > 0
386+
387+
agent.reset()
388+
389+
assert agent.state.messages == []
390+
assert agent.state.is_streaming is False
391+
assert agent.state.streaming_message is None
392+
assert agent.state.pending_tool_calls == set()
393+
assert agent.state.error_message is None
394+
395+
async def test_reset_clears_queues(self):
396+
provider = FauxProvider()
397+
agent = Agent(provider=provider, model=make_model())
398+
399+
agent.steer(UserMessage(content=[TextContent(text="steer")]))
400+
agent.follow_up(UserMessage(content=[TextContent(text="follow")]))
401+
402+
agent.reset()
403+
404+
# After reset, queues should be empty — drain returns nothing
405+
assert agent._steering_queue.drain() == []
406+
assert agent._follow_up_queue.drain() == []
407+
408+
409+
class TestAgentAbortSignal:
410+
async def test_abort_sets_signal_during_active_run(self):
411+
barrier = asyncio.Event()
412+
provider = FauxProvider()
413+
414+
async def slow_stream(*args, **kwargs):
415+
from cubepi.providers.base import MessageStream, StreamEvent
416+
417+
ms = MessageStream()
418+
419+
async def produce():
420+
await barrier.wait()
421+
msg = faux_assistant_message("ok")
422+
ms.push(StreamEvent(type="done"))
423+
ms.set_result(msg)
424+
425+
asyncio.create_task(produce())
426+
return ms
427+
428+
provider.stream = slow_stream
429+
agent = Agent(provider=provider, model=make_model())
430+
431+
task = asyncio.create_task(agent.prompt("hello"))
432+
await asyncio.sleep(0.02)
433+
434+
assert agent._active_signal is not None
435+
assert not agent._active_signal.is_set()
436+
437+
agent.abort()
438+
assert agent._active_signal.is_set()
439+
440+
barrier.set()
441+
await task
442+
443+
444+
class TestAgentWaitForIdle:
445+
async def test_returns_immediately_when_no_active_run(self):
446+
provider = FauxProvider()
447+
agent = Agent(provider=provider, model=make_model())
448+
449+
# No active run, _active_done is None — should return immediately
450+
await agent.wait_for_idle()
451+
452+
async def test_waits_until_prompt_completes(self):
453+
barrier = asyncio.Event()
454+
provider = FauxProvider()
455+
456+
async def slow_stream(*args, **kwargs):
457+
from cubepi.providers.base import MessageStream, StreamEvent
458+
459+
ms = MessageStream()
460+
461+
async def produce():
462+
await barrier.wait()
463+
msg = faux_assistant_message("ok")
464+
ms.push(StreamEvent(type="done"))
465+
ms.set_result(msg)
466+
467+
asyncio.create_task(produce())
468+
return ms
469+
470+
provider.stream = slow_stream
471+
agent = Agent(provider=provider, model=make_model())
472+
473+
prompt_task = asyncio.create_task(agent.prompt("hello"))
474+
await asyncio.sleep(0.02)
475+
476+
idle_resolved = False
477+
478+
async def wait():
479+
nonlocal idle_resolved
480+
await agent.wait_for_idle()
481+
idle_resolved = True
482+
483+
wait_task = asyncio.create_task(wait())
484+
await asyncio.sleep(0.02)
485+
assert not idle_resolved
486+
487+
barrier.set()
488+
await prompt_task
489+
await wait_task
490+
assert idle_resolved
491+
492+
493+
class TestAgentPromptInputTypes:
494+
async def test_prompt_with_message_object(self):
495+
provider = FauxProvider()
496+
provider.set_responses([faux_assistant_message("response")])
497+
agent = Agent(provider=provider, model=make_model())
498+
499+
msg = UserMessage(content=[TextContent(text="direct message")])
500+
await agent.prompt(msg)
501+
502+
has_direct = any(
503+
isinstance(m, UserMessage)
504+
and any(
505+
isinstance(c, TextContent) and c.text == "direct message"
506+
for c in m.content
507+
)
508+
for m in agent.state.messages
509+
)
510+
assert has_direct
511+
assert isinstance(agent.state.messages[-1], AssistantMessage)
512+
513+
async def test_prompt_with_list_of_messages(self):
514+
provider = FauxProvider()
515+
provider.set_responses([faux_assistant_message("response")])
516+
agent = Agent(provider=provider, model=make_model())
517+
518+
msgs = [
519+
UserMessage(content=[TextContent(text="first")]),
520+
UserMessage(content=[TextContent(text="second")]),
521+
]
522+
await agent.prompt(msgs)
523+
524+
texts = [
525+
c.text
526+
for m in agent.state.messages
527+
if isinstance(m, UserMessage)
528+
for c in m.content
529+
if isinstance(c, TextContent)
530+
]
531+
assert "first" in texts
532+
assert "second" in texts
533+
assert isinstance(agent.state.messages[-1], AssistantMessage)

0 commit comments

Comments
 (0)