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 8.4 kB

1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
1 year ago
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230
  1. import logging
  2. import pytest
  3. from autogen_core import (
  4. AgentId,
  5. AgentInstantiationContext,
  6. AgentType,
  7. DefaultTopicId,
  8. SingleThreadedAgentRuntime,
  9. TopicId,
  10. TypeSubscription,
  11. try_get_known_serializers_for_type,
  12. type_subscription,
  13. )
  14. from autogen_test_utils import (
  15. CascadingAgent,
  16. CascadingMessageType,
  17. LoopbackAgent,
  18. LoopbackAgentWithDefaultSubscription,
  19. MessageType,
  20. NoopAgent,
  21. )
  22. from autogen_test_utils.telemetry_test_utils import TestExporter, get_test_tracer_provider
  23. from opentelemetry.sdk.trace import TracerProvider
  24. test_exporter = TestExporter()
  25. @pytest.fixture
  26. def tracer_provider() -> TracerProvider:
  27. test_exporter.clear()
  28. return get_test_tracer_provider(test_exporter)
  29. @pytest.mark.asyncio
  30. async def test_agent_type_must_be_unique() -> None:
  31. runtime = SingleThreadedAgentRuntime()
  32. def agent_factory() -> NoopAgent:
  33. id = AgentInstantiationContext.current_agent_id()
  34. assert id == AgentId("name1", "default")
  35. agent = NoopAgent()
  36. assert agent.id == id
  37. return agent
  38. await NoopAgent.register(runtime, "name1", agent_factory)
  39. # await runtime.register_factory(type=AgentType("name1"), agent_factory=agent_factory, expected_class=NoopAgent)
  40. with pytest.raises(ValueError):
  41. await runtime.register_factory(type=AgentType("name1"), agent_factory=agent_factory, expected_class=NoopAgent)
  42. await runtime.register_factory(type=AgentType("name2"), agent_factory=agent_factory, expected_class=NoopAgent)
  43. @pytest.mark.asyncio
  44. async def test_register_receives_publish(tracer_provider: TracerProvider) -> None:
  45. runtime = SingleThreadedAgentRuntime(tracer_provider=tracer_provider)
  46. runtime.add_message_serializer(try_get_known_serializers_for_type(MessageType))
  47. await runtime.register_factory(
  48. type=AgentType("name"), agent_factory=lambda: LoopbackAgent(), expected_class=LoopbackAgent
  49. )
  50. await runtime.add_subscription(TypeSubscription("default", "name"))
  51. runtime.start()
  52. await runtime.publish_message(MessageType(), topic_id=TopicId("default", "default"))
  53. await runtime.stop_when_idle()
  54. # Agent in default namespace should have received the message
  55. long_running_agent = await runtime.try_get_underlying_agent_instance(AgentId("name", "default"), type=LoopbackAgent)
  56. assert long_running_agent.num_calls == 1
  57. # Agent in other namespace should not have received the message
  58. other_long_running_agent: LoopbackAgent = await runtime.try_get_underlying_agent_instance(
  59. AgentId("name", key="other"), type=LoopbackAgent
  60. )
  61. assert other_long_running_agent.num_calls == 0
  62. exported_spans = test_exporter.get_exported_spans()
  63. assert len(exported_spans) == 3
  64. span_names = [span.name for span in exported_spans]
  65. assert span_names == [
  66. "autogen create default.(default)-T",
  67. "autogen process name.(default)-A",
  68. "autogen publish default.(default)-T",
  69. ]
  70. @pytest.mark.asyncio
  71. async def test_register_receives_publish_with_exception(caplog: pytest.LogCaptureFixture) -> None:
  72. runtime = SingleThreadedAgentRuntime()
  73. runtime.add_message_serializer(try_get_known_serializers_for_type(MessageType))
  74. async def agent_factory() -> LoopbackAgent:
  75. raise ValueError("test")
  76. await runtime.register_factory(type=AgentType("name"), agent_factory=agent_factory, expected_class=LoopbackAgent)
  77. await runtime.add_subscription(TypeSubscription("default", "name"))
  78. with caplog.at_level(logging.ERROR):
  79. runtime.start()
  80. await runtime.publish_message(MessageType(), topic_id=TopicId("default", "default"))
  81. await runtime.stop_when_idle()
  82. # Check if logger has the exception.
  83. assert any("Error processing publish message" in e.message for e in caplog.records)
  84. @pytest.mark.asyncio
  85. async def test_register_receives_publish_cascade() -> None:
  86. num_agents = 5
  87. num_initial_messages = 5
  88. max_rounds = 5
  89. total_num_calls_expected = 0
  90. for i in range(0, max_rounds):
  91. total_num_calls_expected += num_initial_messages * ((num_agents - 1) ** i)
  92. runtime = SingleThreadedAgentRuntime()
  93. # Register agents
  94. for i in range(num_agents):
  95. await CascadingAgent.register(runtime, f"name{i}", lambda: CascadingAgent(max_rounds))
  96. runtime.start()
  97. # Publish messages
  98. for _ in range(num_initial_messages):
  99. await runtime.publish_message(CascadingMessageType(round=1), DefaultTopicId())
  100. # Process until idle.
  101. await runtime.stop_when_idle()
  102. # Check that each agent received the correct number of messages.
  103. for i in range(num_agents):
  104. agent = await runtime.try_get_underlying_agent_instance(AgentId(f"name{i}", "default"), CascadingAgent)
  105. assert agent.num_calls == total_num_calls_expected
  106. @pytest.mark.asyncio
  107. async def test_register_factory_explicit_name() -> None:
  108. runtime = SingleThreadedAgentRuntime()
  109. await LoopbackAgent.register(runtime, "name", LoopbackAgent)
  110. await runtime.add_subscription(TypeSubscription("default", "name"))
  111. runtime.start()
  112. agent_id = AgentId("name", key="default")
  113. topic_id = TopicId("default", "default")
  114. await runtime.publish_message(MessageType(), topic_id=topic_id)
  115. await runtime.stop_when_idle()
  116. # Agent in default namespace should have received the message
  117. long_running_agent = await runtime.try_get_underlying_agent_instance(agent_id, type=LoopbackAgent)
  118. assert long_running_agent.num_calls == 1
  119. # Agent in other namespace should not have received the message
  120. other_long_running_agent: LoopbackAgent = await runtime.try_get_underlying_agent_instance(
  121. AgentId("name", key="other"), type=LoopbackAgent
  122. )
  123. assert other_long_running_agent.num_calls == 0
  124. @pytest.mark.asyncio
  125. async def test_default_subscription() -> None:
  126. runtime = SingleThreadedAgentRuntime()
  127. runtime.start()
  128. await LoopbackAgentWithDefaultSubscription.register(runtime, "name", LoopbackAgentWithDefaultSubscription)
  129. agent_id = AgentId("name", key="default")
  130. await runtime.publish_message(MessageType(), topic_id=DefaultTopicId())
  131. await runtime.stop_when_idle()
  132. long_running_agent = await runtime.try_get_underlying_agent_instance(
  133. agent_id, type=LoopbackAgentWithDefaultSubscription
  134. )
  135. assert long_running_agent.num_calls == 1
  136. other_long_running_agent = await runtime.try_get_underlying_agent_instance(
  137. AgentId("name", key="other"), type=LoopbackAgentWithDefaultSubscription
  138. )
  139. assert other_long_running_agent.num_calls == 0
  140. @pytest.mark.asyncio
  141. async def test_type_subscription() -> None:
  142. runtime = SingleThreadedAgentRuntime()
  143. runtime.start()
  144. @type_subscription(topic_type="Other")
  145. class LoopbackAgentWithSubscription(LoopbackAgent): ...
  146. await LoopbackAgentWithSubscription.register(runtime, "name", LoopbackAgentWithSubscription)
  147. agent_id = AgentId("name", key="default")
  148. await runtime.publish_message(MessageType(), topic_id=TopicId("Other", "default"))
  149. await runtime.stop_when_idle()
  150. long_running_agent = await runtime.try_get_underlying_agent_instance(agent_id, type=LoopbackAgentWithSubscription)
  151. assert long_running_agent.num_calls == 1
  152. other_long_running_agent = await runtime.try_get_underlying_agent_instance(
  153. AgentId("name", key="other"), type=LoopbackAgentWithSubscription
  154. )
  155. assert other_long_running_agent.num_calls == 0
  156. @pytest.mark.asyncio
  157. async def test_default_subscription_publish_to_other_source() -> None:
  158. runtime = SingleThreadedAgentRuntime()
  159. runtime.start()
  160. await LoopbackAgentWithDefaultSubscription.register(runtime, "name", LoopbackAgentWithDefaultSubscription)
  161. agent_id = AgentId("name", key="default")
  162. await runtime.publish_message(MessageType(), topic_id=DefaultTopicId(source="other"))
  163. await runtime.stop_when_idle()
  164. long_running_agent = await runtime.try_get_underlying_agent_instance(
  165. agent_id, type=LoopbackAgentWithDefaultSubscription
  166. )
  167. assert long_running_agent.num_calls == 0
  168. other_long_running_agent = await runtime.try_get_underlying_agent_instance(
  169. AgentId("name", key="other"), type=LoopbackAgentWithDefaultSubscription
  170. )
  171. assert other_long_running_agent.num_calls == 1