You can not select more than 25 topics Topics must start with a chinese character,a letter or number, can include dashes ('-') and can be up to 35 characters long.

test_runtime.py 2.7 kB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172
  1. import pytest
  2. from agnext.application import SingleThreadedAgentRuntime
  3. from agnext.core import AgentId, AgentInstantiationContext
  4. from test_utils import CascadingAgent, CascadingMessageType, LoopbackAgent, MessageType, NoopAgent
  5. @pytest.mark.asyncio
  6. async def test_agent_names_must_be_unique() -> None:
  7. runtime = SingleThreadedAgentRuntime()
  8. def agent_factory() -> NoopAgent:
  9. id = AgentInstantiationContext.current_agent_id()
  10. assert id == AgentId("name1", "default")
  11. agent = NoopAgent()
  12. assert agent.id == id
  13. return agent
  14. agent1 = await runtime.register_and_get("name1", agent_factory)
  15. assert agent1 == AgentId("name1", "default")
  16. with pytest.raises(ValueError):
  17. _agent1 = await runtime.register_and_get("name1", NoopAgent)
  18. _agent1 = await runtime.register_and_get("name3", NoopAgent)
  19. @pytest.mark.asyncio
  20. async def test_register_receives_publish() -> None:
  21. runtime = SingleThreadedAgentRuntime()
  22. await runtime.register("name", LoopbackAgent)
  23. run_context = runtime.start()
  24. await runtime.publish_message(MessageType(), namespace="default")
  25. await run_context.stop_when_idle()
  26. # Agent in default namespace should have received the message
  27. long_running_agent = await runtime.try_get_underlying_agent_instance(await runtime.get("name"), type=LoopbackAgent)
  28. assert long_running_agent.num_calls == 1
  29. # Agent in other namespace should not have received the message
  30. other_long_running_agent: LoopbackAgent = await runtime.try_get_underlying_agent_instance(await runtime.get("name", namespace="other"), type=LoopbackAgent)
  31. assert other_long_running_agent.num_calls == 0
  32. @pytest.mark.asyncio
  33. async def test_register_receives_publish_cascade() -> None:
  34. runtime = SingleThreadedAgentRuntime()
  35. num_agents = 5
  36. num_initial_messages = 5
  37. max_rounds = 5
  38. total_num_calls_expected = 0
  39. for i in range(0, max_rounds):
  40. total_num_calls_expected += num_initial_messages * ((num_agents - 1) ** i)
  41. # Register agents
  42. for i in range(num_agents):
  43. await runtime.register(f"name{i}", lambda: CascadingAgent(max_rounds))
  44. run_context = runtime.start()
  45. # Publish messages
  46. for _ in range(num_initial_messages):
  47. await runtime.publish_message(CascadingMessageType(round=1), namespace="default")
  48. # Process until idle.
  49. await run_context.stop_when_idle()
  50. # Check that each agent received the correct number of messages.
  51. for i in range(num_agents):
  52. agent = await runtime.try_get_underlying_agent_instance(await runtime.get(f"name{i}"), CascadingAgent)
  53. assert agent.num_calls == total_num_calls_expected