|
1 | 1 | import asyncio |
2 | 2 |
|
3 | | -from cubepi.agent.agent import Agent |
| 3 | +import pytest |
| 4 | + |
| 5 | +from cubepi.agent.agent import Agent, _MessageQueue |
4 | 6 | from cubepi.agent.types import AgentTool |
5 | 7 | from cubepi.providers.base import ( |
6 | 8 | AssistantMessage, |
@@ -285,3 +287,247 @@ async def test_resume_processes_follow_up_messages(self): |
285 | 287 | ) |
286 | 288 | assert has_follow_up |
287 | 289 | 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