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_group_chat.py 50 kB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933934935936937938939940941942943944945946947948949950951952953954955956957958959960961962963964965966967968969970971972973974975976977978979980981982983984985986987988989990991992993994995996997998999100010011002100310041005100610071008100910101011101210131014101510161017101810191020102110221023102410251026102710281029103010311032103310341035103610371038103910401041104210431044104510461047104810491050105110521053105410551056105710581059106010611062106310641065106610671068106910701071107210731074107510761077107810791080108110821083108410851086108710881089109010911092109310941095109610971098109911001101110211031104110511061107110811091110111111121113111411151116111711181119112011211122112311241125112611271128112911301131113211331134113511361137113811391140114111421143114411451146114711481149115011511152115311541155115611571158115911601161116211631164116511661167116811691170117111721173117411751176117711781179118011811182118311841185118611871188118911901191119211931194119511961197119811991200120112021203120412051206120712081209121012111212121312141215121612171218121912201221
  1. import asyncio
  2. import json
  3. import logging
  4. import tempfile
  5. from typing import Any, AsyncGenerator, List, Sequence
  6. import pytest
  7. from autogen_agentchat import EVENT_LOGGER_NAME
  8. from autogen_agentchat.agents import (
  9. AssistantAgent,
  10. BaseChatAgent,
  11. CodeExecutorAgent,
  12. )
  13. from autogen_agentchat.base import Handoff, Response, TaskResult
  14. from autogen_agentchat.conditions import HandoffTermination, MaxMessageTermination, TextMentionTermination
  15. from autogen_agentchat.messages import (
  16. AgentEvent,
  17. ChatMessage,
  18. HandoffMessage,
  19. MultiModalMessage,
  20. StopMessage,
  21. TextMessage,
  22. ToolCallExecutionEvent,
  23. ToolCallRequestEvent,
  24. ToolCallSummaryMessage,
  25. )
  26. from autogen_agentchat.teams import (
  27. RoundRobinGroupChat,
  28. SelectorGroupChat,
  29. Swarm,
  30. )
  31. from autogen_agentchat.teams._group_chat._round_robin_group_chat import RoundRobinGroupChatManager
  32. from autogen_agentchat.teams._group_chat._selector_group_chat import SelectorGroupChatManager
  33. from autogen_agentchat.teams._group_chat._swarm_group_chat import SwarmGroupChatManager
  34. from autogen_agentchat.ui import Console
  35. from autogen_core import AgentId, CancellationToken, FunctionCall
  36. from autogen_core.models import (
  37. AssistantMessage,
  38. FunctionExecutionResult,
  39. FunctionExecutionResultMessage,
  40. LLMMessage,
  41. UserMessage,
  42. )
  43. from autogen_core.tools import FunctionTool
  44. from autogen_ext.code_executors.local import LocalCommandLineCodeExecutor
  45. from autogen_ext.models.openai import OpenAIChatCompletionClient
  46. from autogen_ext.models.replay import ReplayChatCompletionClient
  47. from openai.resources.chat.completions import AsyncCompletions
  48. from openai.types.chat.chat_completion import ChatCompletion, Choice
  49. from openai.types.chat.chat_completion_chunk import ChatCompletionChunk
  50. from openai.types.chat.chat_completion_message import ChatCompletionMessage
  51. from openai.types.chat.chat_completion_message_tool_call import ChatCompletionMessageToolCall, Function
  52. from openai.types.completion_usage import CompletionUsage
  53. from utils import FileLogHandler
  54. logger = logging.getLogger(EVENT_LOGGER_NAME)
  55. logger.setLevel(logging.DEBUG)
  56. logger.addHandler(FileLogHandler("test_group_chat.log"))
  57. class _MockChatCompletion:
  58. def __init__(self, chat_completions: List[ChatCompletion]) -> None:
  59. self._saved_chat_completions = chat_completions
  60. self._curr_index = 0
  61. async def mock_create(
  62. self, *args: Any, **kwargs: Any
  63. ) -> ChatCompletion | AsyncGenerator[ChatCompletionChunk, None]:
  64. await asyncio.sleep(0.1)
  65. completion = self._saved_chat_completions[self._curr_index]
  66. self._curr_index += 1
  67. return completion
  68. def reset(self) -> None:
  69. self._curr_index = 0
  70. class _EchoAgent(BaseChatAgent):
  71. def __init__(self, name: str, description: str) -> None:
  72. super().__init__(name, description)
  73. self._last_message: str | None = None
  74. self._total_messages = 0
  75. @property
  76. def produced_message_types(self) -> Sequence[type[ChatMessage]]:
  77. return (TextMessage,)
  78. @property
  79. def total_messages(self) -> int:
  80. return self._total_messages
  81. async def on_messages(self, messages: Sequence[ChatMessage], cancellation_token: CancellationToken) -> Response:
  82. if len(messages) > 0:
  83. assert isinstance(messages[0], TextMessage)
  84. self._last_message = messages[0].content
  85. self._total_messages += 1
  86. return Response(chat_message=TextMessage(content=messages[0].content, source=self.name))
  87. else:
  88. assert self._last_message is not None
  89. self._total_messages += 1
  90. return Response(chat_message=TextMessage(content=self._last_message, source=self.name))
  91. async def on_reset(self, cancellation_token: CancellationToken) -> None:
  92. self._last_message = None
  93. class _StopAgent(_EchoAgent):
  94. def __init__(self, name: str, description: str, *, stop_at: int = 1) -> None:
  95. super().__init__(name, description)
  96. self._count = 0
  97. self._stop_at = stop_at
  98. @property
  99. def produced_message_types(self) -> Sequence[type[ChatMessage]]:
  100. return (TextMessage, StopMessage)
  101. async def on_messages(self, messages: Sequence[ChatMessage], cancellation_token: CancellationToken) -> Response:
  102. self._count += 1
  103. if self._count < self._stop_at:
  104. return await super().on_messages(messages, cancellation_token)
  105. return Response(chat_message=StopMessage(content="TERMINATE", source=self.name))
  106. def _pass_function(input: str) -> str:
  107. return "pass"
  108. @pytest.mark.asyncio
  109. async def test_round_robin_group_chat(monkeypatch: pytest.MonkeyPatch) -> None:
  110. model = "gpt-4o-2024-05-13"
  111. chat_completions = [
  112. ChatCompletion(
  113. id="id1",
  114. choices=[
  115. Choice(
  116. finish_reason="stop",
  117. index=0,
  118. message=ChatCompletionMessage(
  119. content="""Here is the program\n ```python\nprint("Hello, world!")\n```""",
  120. role="assistant",
  121. ),
  122. )
  123. ],
  124. created=0,
  125. model=model,
  126. object="chat.completion",
  127. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  128. ),
  129. ChatCompletion(
  130. id="id2",
  131. choices=[
  132. Choice(
  133. finish_reason="stop",
  134. index=0,
  135. message=ChatCompletionMessage(content="TERMINATE", role="assistant"),
  136. )
  137. ],
  138. created=0,
  139. model=model,
  140. object="chat.completion",
  141. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  142. ),
  143. ]
  144. mock = _MockChatCompletion(chat_completions)
  145. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  146. with tempfile.TemporaryDirectory() as temp_dir:
  147. code_executor_agent = CodeExecutorAgent(
  148. "code_executor", code_executor=LocalCommandLineCodeExecutor(work_dir=temp_dir)
  149. )
  150. coding_assistant_agent = AssistantAgent(
  151. "coding_assistant", model_client=OpenAIChatCompletionClient(model=model, api_key="")
  152. )
  153. termination = TextMentionTermination("TERMINATE")
  154. team = RoundRobinGroupChat(
  155. participants=[coding_assistant_agent, code_executor_agent], termination_condition=termination
  156. )
  157. result = await team.run(
  158. task="Write a program that prints 'Hello, world!'",
  159. )
  160. expected_messages = [
  161. "Write a program that prints 'Hello, world!'",
  162. 'Here is the program\n ```python\nprint("Hello, world!")\n```',
  163. "Hello, world!",
  164. "TERMINATE",
  165. ]
  166. # Normalize the messages to remove \r\n and any leading/trailing whitespace.
  167. normalized_messages = [
  168. msg.content.replace("\r\n", "\n").rstrip("\n") if isinstance(msg.content, str) else msg.content
  169. for msg in result.messages
  170. ]
  171. # Assert that all expected messages are in the collected messages
  172. assert normalized_messages == expected_messages
  173. assert result.stop_reason is not None and result.stop_reason == "Text 'TERMINATE' mentioned"
  174. # Test streaming.
  175. mock.reset()
  176. index = 0
  177. await team.reset()
  178. async for message in team.run_stream(
  179. task="Write a program that prints 'Hello, world!'",
  180. ):
  181. if isinstance(message, TaskResult):
  182. assert message == result
  183. else:
  184. assert message == result.messages[index]
  185. index += 1
  186. # Test message input.
  187. # Text message.
  188. mock.reset()
  189. index = 0
  190. await team.reset()
  191. result_2 = await team.run(
  192. task=TextMessage(content="Write a program that prints 'Hello, world!'", source="user")
  193. )
  194. assert result == result_2
  195. # Test multi-modal message.
  196. mock.reset()
  197. index = 0
  198. await team.reset()
  199. result_2 = await team.run(
  200. task=MultiModalMessage(content=["Write a program that prints 'Hello, world!'"], source="user")
  201. )
  202. assert result.messages[0].content == result_2.messages[0].content[0]
  203. assert result.messages[1:] == result_2.messages[1:]
  204. @pytest.mark.asyncio
  205. async def test_round_robin_group_chat_state() -> None:
  206. model_client = ReplayChatCompletionClient(
  207. ["No facts", "No plan", "print('Hello, world!')", "TERMINATE"],
  208. )
  209. agent1 = AssistantAgent("agent1", model_client=model_client)
  210. agent2 = AssistantAgent("agent2", model_client=model_client)
  211. termination = TextMentionTermination("TERMINATE")
  212. team1 = RoundRobinGroupChat(participants=[agent1, agent2], termination_condition=termination)
  213. await team1.run(task="Write a program that prints 'Hello, world!'")
  214. state = await team1.save_state()
  215. agent3 = AssistantAgent("agent1", model_client=model_client)
  216. agent4 = AssistantAgent("agent2", model_client=model_client)
  217. team2 = RoundRobinGroupChat(participants=[agent3, agent4], termination_condition=termination)
  218. await team2.load_state(state)
  219. state2 = await team2.save_state()
  220. assert state == state2
  221. agent1_model_ctx_messages = await agent1._model_context.get_messages() # pyright: ignore
  222. agent2_model_ctx_messages = await agent2._model_context.get_messages() # pyright: ignore
  223. agent3_model_ctx_messages = await agent3._model_context.get_messages() # pyright: ignore
  224. agent4_model_ctx_messages = await agent4._model_context.get_messages() # pyright: ignore
  225. assert agent3_model_ctx_messages == agent1_model_ctx_messages
  226. assert agent4_model_ctx_messages == agent2_model_ctx_messages
  227. manager_1 = await team1._runtime.try_get_underlying_agent_instance( # pyright: ignore
  228. AgentId("group_chat_manager", team1._team_id), # pyright: ignore
  229. RoundRobinGroupChatManager, # pyright: ignore
  230. ) # pyright: ignore
  231. manager_2 = await team2._runtime.try_get_underlying_agent_instance( # pyright: ignore
  232. AgentId("group_chat_manager", team2._team_id), # pyright: ignore
  233. RoundRobinGroupChatManager, # pyright: ignore
  234. ) # pyright: ignore
  235. assert manager_1._current_turn == manager_2._current_turn # pyright: ignore
  236. assert manager_1._message_thread == manager_2._message_thread # pyright: ignore
  237. @pytest.mark.asyncio
  238. async def test_round_robin_group_chat_with_tools(monkeypatch: pytest.MonkeyPatch) -> None:
  239. model = "gpt-4o-2024-05-13"
  240. chat_completions = [
  241. ChatCompletion(
  242. id="id1",
  243. choices=[
  244. Choice(
  245. finish_reason="tool_calls",
  246. index=0,
  247. message=ChatCompletionMessage(
  248. content=None,
  249. tool_calls=[
  250. ChatCompletionMessageToolCall(
  251. id="1",
  252. type="function",
  253. function=Function(
  254. name="pass",
  255. arguments=json.dumps({"input": "pass"}),
  256. ),
  257. )
  258. ],
  259. role="assistant",
  260. ),
  261. )
  262. ],
  263. created=0,
  264. model=model,
  265. object="chat.completion",
  266. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  267. ),
  268. ChatCompletion(
  269. id="id2",
  270. choices=[
  271. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="Hello", role="assistant"))
  272. ],
  273. created=0,
  274. model=model,
  275. object="chat.completion",
  276. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  277. ),
  278. ChatCompletion(
  279. id="id2",
  280. choices=[
  281. Choice(
  282. finish_reason="stop", index=0, message=ChatCompletionMessage(content="TERMINATE", role="assistant")
  283. )
  284. ],
  285. created=0,
  286. model=model,
  287. object="chat.completion",
  288. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  289. ),
  290. ]
  291. # Test with repeat tool calls once
  292. mock = _MockChatCompletion(chat_completions)
  293. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  294. tool = FunctionTool(_pass_function, name="pass", description="pass function")
  295. tool_use_agent = AssistantAgent(
  296. "tool_use_agent",
  297. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  298. tools=[tool],
  299. )
  300. echo_agent = _EchoAgent("echo_agent", description="echo agent")
  301. termination = TextMentionTermination("TERMINATE")
  302. team = RoundRobinGroupChat(participants=[tool_use_agent, echo_agent], termination_condition=termination)
  303. result = await team.run(
  304. task="Write a program that prints 'Hello, world!'",
  305. )
  306. assert len(result.messages) == 8
  307. assert isinstance(result.messages[0], TextMessage) # task
  308. assert isinstance(result.messages[1], ToolCallRequestEvent) # tool call
  309. assert isinstance(result.messages[2], ToolCallExecutionEvent) # tool call result
  310. assert isinstance(result.messages[3], ToolCallSummaryMessage) # tool use agent response
  311. assert result.messages[3].content == "pass" # ensure the tool call was executed
  312. assert isinstance(result.messages[4], TextMessage) # echo agent response
  313. assert isinstance(result.messages[5], TextMessage) # tool use agent response
  314. assert isinstance(result.messages[6], TextMessage) # echo agent response
  315. assert isinstance(result.messages[7], TextMessage) # tool use agent response, that has TERMINATE
  316. assert result.messages[7].content == "TERMINATE"
  317. assert result.stop_reason is not None and result.stop_reason == "Text 'TERMINATE' mentioned"
  318. # Test streaming.
  319. await tool_use_agent._model_context.clear() # pyright: ignore
  320. mock.reset()
  321. index = 0
  322. await team.reset()
  323. async for message in team.run_stream(
  324. task="Write a program that prints 'Hello, world!'",
  325. ):
  326. if isinstance(message, TaskResult):
  327. assert message == result
  328. else:
  329. assert message == result.messages[index]
  330. index += 1
  331. # Test Console.
  332. await tool_use_agent._model_context.clear() # pyright: ignore
  333. mock.reset()
  334. index = 0
  335. await team.reset()
  336. result2 = await Console(team.run_stream(task="Write a program that prints 'Hello, world!'"))
  337. assert result2 == result
  338. @pytest.mark.asyncio
  339. async def test_round_robin_group_chat_with_resume_and_reset() -> None:
  340. agent_1 = _EchoAgent("agent_1", description="echo agent 1")
  341. agent_2 = _EchoAgent("agent_2", description="echo agent 2")
  342. agent_3 = _EchoAgent("agent_3", description="echo agent 3")
  343. agent_4 = _EchoAgent("agent_4", description="echo agent 4")
  344. termination = MaxMessageTermination(3)
  345. team = RoundRobinGroupChat(participants=[agent_1, agent_2, agent_3, agent_4], termination_condition=termination)
  346. result = await team.run(
  347. task="Write a program that prints 'Hello, world!'",
  348. )
  349. assert len(result.messages) == 3
  350. assert result.messages[1].source == "agent_1"
  351. assert result.messages[2].source == "agent_2"
  352. assert result.stop_reason is not None
  353. # Resume.
  354. result = await team.run()
  355. assert len(result.messages) == 3
  356. assert result.messages[0].source == "agent_3"
  357. assert result.messages[1].source == "agent_4"
  358. assert result.messages[2].source == "agent_1"
  359. assert result.stop_reason is not None
  360. # Reset.
  361. await team.reset()
  362. result = await team.run(task="Write a program that prints 'Hello, world!'")
  363. assert len(result.messages) == 3
  364. assert result.messages[1].source == "agent_1"
  365. assert result.messages[2].source == "agent_2"
  366. assert result.stop_reason is not None
  367. @pytest.mark.asyncio
  368. async def test_round_robin_group_chat_max_turn() -> None:
  369. agent_1 = _EchoAgent("agent_1", description="echo agent 1")
  370. agent_2 = _EchoAgent("agent_2", description="echo agent 2")
  371. agent_3 = _EchoAgent("agent_3", description="echo agent 3")
  372. agent_4 = _EchoAgent("agent_4", description="echo agent 4")
  373. team = RoundRobinGroupChat(participants=[agent_1, agent_2, agent_3, agent_4], max_turns=3)
  374. result = await team.run(
  375. task="Write a program that prints 'Hello, world!'",
  376. )
  377. assert len(result.messages) == 4
  378. assert result.messages[1].source == "agent_1"
  379. assert result.messages[2].source == "agent_2"
  380. assert result.messages[3].source == "agent_3"
  381. assert result.stop_reason is not None
  382. # Resume.
  383. result = await team.run()
  384. assert len(result.messages) == 3
  385. assert result.messages[0].source == "agent_4"
  386. assert result.messages[1].source == "agent_1"
  387. assert result.messages[2].source == "agent_2"
  388. assert result.stop_reason is not None
  389. # Reset.
  390. await team.reset()
  391. result = await team.run(task="Write a program that prints 'Hello, world!'")
  392. assert len(result.messages) == 4
  393. assert result.messages[1].source == "agent_1"
  394. assert result.messages[2].source == "agent_2"
  395. assert result.messages[3].source == "agent_3"
  396. assert result.stop_reason is not None
  397. @pytest.mark.asyncio
  398. async def test_round_robin_group_chat_cancellation() -> None:
  399. agent_1 = _EchoAgent("agent_1", description="echo agent 1")
  400. agent_2 = _EchoAgent("agent_2", description="echo agent 2")
  401. agent_3 = _EchoAgent("agent_3", description="echo agent 3")
  402. agent_4 = _EchoAgent("agent_4", description="echo agent 4")
  403. # Set max_turns to a large number to avoid stopping due to max_turns before cancellation.
  404. team = RoundRobinGroupChat(participants=[agent_1, agent_2, agent_3, agent_4], max_turns=1000)
  405. cancellation_token = CancellationToken()
  406. run_task = asyncio.create_task(
  407. team.run(
  408. task="Write a program that prints 'Hello, world!'",
  409. cancellation_token=cancellation_token,
  410. )
  411. )
  412. await asyncio.sleep(0.1)
  413. # Cancel the task.
  414. cancellation_token.cancel()
  415. with pytest.raises(asyncio.CancelledError):
  416. await run_task
  417. # Total messages produced so far.
  418. total_messages = agent_1.total_messages + agent_2.total_messages + agent_3.total_messages + agent_4.total_messages
  419. # Still can run again and finish the task.
  420. result = await team.run()
  421. assert len(result.messages) + total_messages == 1000
  422. @pytest.mark.asyncio
  423. async def test_selector_group_chat(monkeypatch: pytest.MonkeyPatch) -> None:
  424. model = "gpt-4o-2024-05-13"
  425. chat_completions = [
  426. ChatCompletion(
  427. id="id2",
  428. choices=[
  429. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent3", role="assistant"))
  430. ],
  431. created=0,
  432. model=model,
  433. object="chat.completion",
  434. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  435. ),
  436. ChatCompletion(
  437. id="id2",
  438. choices=[
  439. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent2", role="assistant"))
  440. ],
  441. created=0,
  442. model=model,
  443. object="chat.completion",
  444. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  445. ),
  446. ChatCompletion(
  447. id="id2",
  448. choices=[
  449. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent1", role="assistant"))
  450. ],
  451. created=0,
  452. model=model,
  453. object="chat.completion",
  454. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  455. ),
  456. ChatCompletion(
  457. id="id2",
  458. choices=[
  459. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent2", role="assistant"))
  460. ],
  461. created=0,
  462. model=model,
  463. object="chat.completion",
  464. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  465. ),
  466. ChatCompletion(
  467. id="id2",
  468. choices=[
  469. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent1", role="assistant"))
  470. ],
  471. created=0,
  472. model=model,
  473. object="chat.completion",
  474. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  475. ),
  476. ]
  477. mock = _MockChatCompletion(chat_completions)
  478. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  479. agent1 = _StopAgent("agent1", description="echo agent 1", stop_at=2)
  480. agent2 = _EchoAgent("agent2", description="echo agent 2")
  481. agent3 = _EchoAgent("agent3", description="echo agent 3")
  482. termination = TextMentionTermination("TERMINATE")
  483. team = SelectorGroupChat(
  484. participants=[agent1, agent2, agent3],
  485. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  486. termination_condition=termination,
  487. )
  488. result = await team.run(
  489. task="Write a program that prints 'Hello, world!'",
  490. )
  491. assert len(result.messages) == 6
  492. assert result.messages[0].content == "Write a program that prints 'Hello, world!'"
  493. assert result.messages[1].source == "agent3"
  494. assert result.messages[2].source == "agent2"
  495. assert result.messages[3].source == "agent1"
  496. assert result.messages[4].source == "agent2"
  497. assert result.messages[5].source == "agent1"
  498. assert result.stop_reason is not None and result.stop_reason == "Text 'TERMINATE' mentioned"
  499. # Test streaming.
  500. mock.reset()
  501. agent1._count = 0 # pyright: ignore
  502. index = 0
  503. await team.reset()
  504. async for message in team.run_stream(
  505. task="Write a program that prints 'Hello, world!'",
  506. ):
  507. if isinstance(message, TaskResult):
  508. assert message == result
  509. else:
  510. assert message == result.messages[index]
  511. index += 1
  512. # Test Console.
  513. mock.reset()
  514. agent1._count = 0 # pyright: ignore
  515. index = 0
  516. await team.reset()
  517. result2 = await Console(team.run_stream(task="Write a program that prints 'Hello, world!'"))
  518. assert result2 == result
  519. @pytest.mark.asyncio
  520. async def test_selector_group_chat_state() -> None:
  521. model_client = ReplayChatCompletionClient(
  522. ["agent1", "No facts", "agent2", "No plan", "agent1", "print('Hello, world!')", "agent2", "TERMINATE"],
  523. )
  524. agent1 = AssistantAgent("agent1", model_client=model_client)
  525. agent2 = AssistantAgent("agent2", model_client=model_client)
  526. termination = TextMentionTermination("TERMINATE")
  527. team1 = SelectorGroupChat(
  528. participants=[agent1, agent2], termination_condition=termination, model_client=model_client
  529. )
  530. await team1.run(task="Write a program that prints 'Hello, world!'")
  531. state = await team1.save_state()
  532. agent3 = AssistantAgent("agent1", model_client=model_client)
  533. agent4 = AssistantAgent("agent2", model_client=model_client)
  534. team2 = SelectorGroupChat(
  535. participants=[agent3, agent4], termination_condition=termination, model_client=model_client
  536. )
  537. await team2.load_state(state)
  538. state2 = await team2.save_state()
  539. assert state == state2
  540. agent1_model_ctx_messages = await agent1._model_context.get_messages() # pyright: ignore
  541. agent2_model_ctx_messages = await agent2._model_context.get_messages() # pyright: ignore
  542. agent3_model_ctx_messages = await agent3._model_context.get_messages() # pyright: ignore
  543. agent4_model_ctx_messages = await agent4._model_context.get_messages() # pyright: ignore
  544. assert agent3_model_ctx_messages == agent1_model_ctx_messages
  545. assert agent4_model_ctx_messages == agent2_model_ctx_messages
  546. manager_1 = await team1._runtime.try_get_underlying_agent_instance( # pyright: ignore
  547. AgentId("group_chat_manager", team1._team_id), # pyright: ignore
  548. SelectorGroupChatManager, # pyright: ignore
  549. ) # pyright: ignore
  550. manager_2 = await team2._runtime.try_get_underlying_agent_instance( # pyright: ignore
  551. AgentId("group_chat_manager", team2._team_id), # pyright: ignore
  552. SelectorGroupChatManager, # pyright: ignore
  553. ) # pyright: ignore
  554. assert manager_1._message_thread == manager_2._message_thread # pyright: ignore
  555. assert manager_1._previous_speaker == manager_2._previous_speaker # pyright: ignore
  556. @pytest.mark.asyncio
  557. async def test_selector_group_chat_two_speakers(monkeypatch: pytest.MonkeyPatch) -> None:
  558. model = "gpt-4o-2024-05-13"
  559. chat_completions = [
  560. ChatCompletion(
  561. id="id2",
  562. choices=[
  563. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent2", role="assistant"))
  564. ],
  565. created=0,
  566. model=model,
  567. object="chat.completion",
  568. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  569. ),
  570. ]
  571. mock = _MockChatCompletion(chat_completions)
  572. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  573. agent1 = _StopAgent("agent1", description="echo agent 1", stop_at=2)
  574. agent2 = _EchoAgent("agent2", description="echo agent 2")
  575. termination = TextMentionTermination("TERMINATE")
  576. team = SelectorGroupChat(
  577. participants=[agent1, agent2],
  578. termination_condition=termination,
  579. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  580. )
  581. result = await team.run(
  582. task="Write a program that prints 'Hello, world!'",
  583. )
  584. assert len(result.messages) == 5
  585. assert result.messages[0].content == "Write a program that prints 'Hello, world!'"
  586. assert result.messages[1].source == "agent2"
  587. assert result.messages[2].source == "agent1"
  588. assert result.messages[3].source == "agent2"
  589. assert result.messages[4].source == "agent1"
  590. # only one chat completion was called
  591. assert mock._curr_index == 1 # pyright: ignore
  592. assert result.stop_reason is not None and result.stop_reason == "Text 'TERMINATE' mentioned"
  593. # Test streaming.
  594. mock.reset()
  595. agent1._count = 0 # pyright: ignore
  596. index = 0
  597. await team.reset()
  598. async for message in team.run_stream(task="Write a program that prints 'Hello, world!'"):
  599. if isinstance(message, TaskResult):
  600. assert message == result
  601. else:
  602. assert message == result.messages[index]
  603. index += 1
  604. # Test Console.
  605. mock.reset()
  606. agent1._count = 0 # pyright: ignore
  607. index = 0
  608. await team.reset()
  609. result2 = await Console(team.run_stream(task="Write a program that prints 'Hello, world!'"))
  610. assert result2 == result
  611. @pytest.mark.asyncio
  612. async def test_selector_group_chat_two_speakers_allow_repeated(monkeypatch: pytest.MonkeyPatch) -> None:
  613. model = "gpt-4o-2024-05-13"
  614. chat_completions = [
  615. ChatCompletion(
  616. id="id2",
  617. choices=[
  618. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent2", role="assistant"))
  619. ],
  620. created=0,
  621. model=model,
  622. object="chat.completion",
  623. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  624. ),
  625. ChatCompletion(
  626. id="id2",
  627. choices=[
  628. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent2", role="assistant"))
  629. ],
  630. created=0,
  631. model=model,
  632. object="chat.completion",
  633. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  634. ),
  635. ChatCompletion(
  636. id="id2",
  637. choices=[
  638. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent1", role="assistant"))
  639. ],
  640. created=0,
  641. model=model,
  642. object="chat.completion",
  643. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  644. ),
  645. ]
  646. mock = _MockChatCompletion(chat_completions)
  647. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  648. agent1 = _StopAgent("agent1", description="echo agent 1", stop_at=1)
  649. agent2 = _EchoAgent("agent2", description="echo agent 2")
  650. termination = TextMentionTermination("TERMINATE")
  651. team = SelectorGroupChat(
  652. participants=[agent1, agent2],
  653. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  654. termination_condition=termination,
  655. allow_repeated_speaker=True,
  656. )
  657. result = await team.run(task="Write a program that prints 'Hello, world!'")
  658. assert len(result.messages) == 4
  659. assert result.messages[0].content == "Write a program that prints 'Hello, world!'"
  660. assert result.messages[1].source == "agent2"
  661. assert result.messages[2].source == "agent2"
  662. assert result.messages[3].source == "agent1"
  663. assert result.stop_reason is not None and result.stop_reason == "Text 'TERMINATE' mentioned"
  664. # Test streaming.
  665. mock.reset()
  666. index = 0
  667. await team.reset()
  668. async for message in team.run_stream(task="Write a program that prints 'Hello, world!'"):
  669. if isinstance(message, TaskResult):
  670. assert message == result
  671. else:
  672. assert message == result.messages[index]
  673. index += 1
  674. # Test Console.
  675. mock.reset()
  676. index = 0
  677. await team.reset()
  678. result2 = await Console(team.run_stream(task="Write a program that prints 'Hello, world!'"))
  679. assert result2 == result
  680. @pytest.mark.asyncio
  681. async def test_selector_group_chat_custom_selector(monkeypatch: pytest.MonkeyPatch) -> None:
  682. model = "gpt-4o-2024-05-13"
  683. chat_completions = [
  684. ChatCompletion(
  685. id="id2",
  686. choices=[
  687. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent3", role="assistant"))
  688. ],
  689. created=0,
  690. model=model,
  691. object="chat.completion",
  692. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  693. ),
  694. ]
  695. mock = _MockChatCompletion(chat_completions)
  696. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  697. agent1 = _EchoAgent("agent1", description="echo agent 1")
  698. agent2 = _EchoAgent("agent2", description="echo agent 2")
  699. agent3 = _EchoAgent("agent3", description="echo agent 3")
  700. agent4 = _EchoAgent("agent4", description="echo agent 4")
  701. def _select_agent(messages: Sequence[AgentEvent | ChatMessage]) -> str | None:
  702. if len(messages) == 0:
  703. return "agent1"
  704. elif messages[-1].source == "agent1":
  705. return "agent2"
  706. elif messages[-1].source == "agent2":
  707. return None
  708. elif messages[-1].source == "agent3":
  709. return "agent4"
  710. else:
  711. return "agent1"
  712. termination = MaxMessageTermination(6)
  713. team = SelectorGroupChat(
  714. participants=[agent1, agent2, agent3, agent4],
  715. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  716. selector_func=_select_agent,
  717. termination_condition=termination,
  718. )
  719. result = await team.run(task="task")
  720. assert len(result.messages) == 6
  721. assert result.messages[1].source == "agent1"
  722. assert result.messages[2].source == "agent2"
  723. assert result.messages[3].source == "agent3"
  724. assert result.messages[4].source == "agent4"
  725. assert result.messages[5].source == "agent1"
  726. assert (
  727. result.stop_reason is not None
  728. and result.stop_reason == "Maximum number of messages 6 reached, current message count: 6"
  729. )
  730. class _HandOffAgent(BaseChatAgent):
  731. def __init__(self, name: str, description: str, next_agent: str) -> None:
  732. super().__init__(name, description)
  733. self._next_agent = next_agent
  734. @property
  735. def produced_message_types(self) -> Sequence[type[ChatMessage]]:
  736. return (HandoffMessage,)
  737. async def on_messages(self, messages: Sequence[ChatMessage], cancellation_token: CancellationToken) -> Response:
  738. return Response(
  739. chat_message=HandoffMessage(
  740. content=f"Transferred to {self._next_agent}.", target=self._next_agent, source=self.name
  741. )
  742. )
  743. async def on_reset(self, cancellation_token: CancellationToken) -> None:
  744. pass
  745. @pytest.mark.asyncio
  746. async def test_swarm_handoff() -> None:
  747. first_agent = _HandOffAgent("first_agent", description="first agent", next_agent="second_agent")
  748. second_agent = _HandOffAgent("second_agent", description="second agent", next_agent="third_agent")
  749. third_agent = _HandOffAgent("third_agent", description="third agent", next_agent="first_agent")
  750. termination = MaxMessageTermination(6)
  751. team = Swarm([second_agent, first_agent, third_agent], termination_condition=termination)
  752. result = await team.run(task="task")
  753. assert len(result.messages) == 6
  754. assert result.messages[0].content == "task"
  755. assert result.messages[1].content == "Transferred to third_agent."
  756. assert result.messages[2].content == "Transferred to first_agent."
  757. assert result.messages[3].content == "Transferred to second_agent."
  758. assert result.messages[4].content == "Transferred to third_agent."
  759. assert result.messages[5].content == "Transferred to first_agent."
  760. assert (
  761. result.stop_reason is not None
  762. and result.stop_reason == "Maximum number of messages 6 reached, current message count: 6"
  763. )
  764. # Test streaming.
  765. index = 0
  766. await team.reset()
  767. stream = team.run_stream(task="task")
  768. async for message in stream:
  769. if isinstance(message, TaskResult):
  770. assert message == result
  771. else:
  772. assert message == result.messages[index]
  773. index += 1
  774. # Test save and load.
  775. state = await team.save_state()
  776. first_agent2 = _HandOffAgent("first_agent", description="first agent", next_agent="second_agent")
  777. second_agent2 = _HandOffAgent("second_agent", description="second agent", next_agent="third_agent")
  778. third_agent2 = _HandOffAgent("third_agent", description="third agent", next_agent="first_agent")
  779. team2 = Swarm([second_agent2, first_agent2, third_agent2], termination_condition=termination)
  780. await team2.load_state(state)
  781. state2 = await team2.save_state()
  782. assert state == state2
  783. manager_1 = await team._runtime.try_get_underlying_agent_instance( # pyright: ignore
  784. AgentId("group_chat_manager", team._team_id), # pyright: ignore
  785. SwarmGroupChatManager, # pyright: ignore
  786. ) # pyright: ignore
  787. manager_2 = await team2._runtime.try_get_underlying_agent_instance( # pyright: ignore
  788. AgentId("group_chat_manager", team2._team_id), # pyright: ignore
  789. SwarmGroupChatManager, # pyright: ignore
  790. ) # pyright: ignore
  791. assert manager_1._message_thread == manager_2._message_thread # pyright: ignore
  792. assert manager_1._current_speaker == manager_2._current_speaker # pyright: ignore
  793. @pytest.mark.asyncio
  794. async def test_swarm_handoff_using_tool_calls(monkeypatch: pytest.MonkeyPatch) -> None:
  795. model = "gpt-4o-2024-05-13"
  796. chat_completions = [
  797. ChatCompletion(
  798. id="id1",
  799. choices=[
  800. Choice(
  801. finish_reason="tool_calls",
  802. index=0,
  803. message=ChatCompletionMessage(
  804. content=None,
  805. tool_calls=[
  806. ChatCompletionMessageToolCall(
  807. id="1",
  808. type="function",
  809. function=Function(
  810. name="handoff_to_agent2",
  811. arguments=json.dumps({}),
  812. ),
  813. )
  814. ],
  815. role="assistant",
  816. ),
  817. )
  818. ],
  819. created=0,
  820. model=model,
  821. object="chat.completion",
  822. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  823. ),
  824. ChatCompletion(
  825. id="id2",
  826. choices=[
  827. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="Hello", role="assistant"))
  828. ],
  829. created=0,
  830. model=model,
  831. object="chat.completion",
  832. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  833. ),
  834. ChatCompletion(
  835. id="id2",
  836. choices=[
  837. Choice(
  838. finish_reason="stop", index=0, message=ChatCompletionMessage(content="TERMINATE", role="assistant")
  839. )
  840. ],
  841. created=0,
  842. model=model,
  843. object="chat.completion",
  844. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  845. ),
  846. ]
  847. mock = _MockChatCompletion(chat_completions)
  848. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  849. agent1 = AssistantAgent(
  850. "agent1",
  851. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  852. handoffs=[Handoff(target="agent2", name="handoff_to_agent2", message="handoff to agent2")],
  853. )
  854. agent2 = _HandOffAgent("agent2", description="agent 2", next_agent="agent1")
  855. termination = TextMentionTermination("TERMINATE")
  856. team = Swarm([agent1, agent2], termination_condition=termination)
  857. result = await team.run(task="task")
  858. assert len(result.messages) == 7
  859. assert result.messages[0].content == "task"
  860. assert isinstance(result.messages[1], ToolCallRequestEvent)
  861. assert isinstance(result.messages[2], ToolCallExecutionEvent)
  862. assert result.messages[3].content == "handoff to agent2"
  863. assert result.messages[4].content == "Transferred to agent1."
  864. assert result.messages[5].content == "Hello"
  865. assert result.messages[6].content == "TERMINATE"
  866. assert result.stop_reason is not None and result.stop_reason == "Text 'TERMINATE' mentioned"
  867. # Test streaming.
  868. await agent1._model_context.clear() # pyright: ignore
  869. mock.reset()
  870. index = 0
  871. await team.reset()
  872. stream = team.run_stream(task="task")
  873. async for message in stream:
  874. if isinstance(message, TaskResult):
  875. assert message == result
  876. else:
  877. assert message == result.messages[index]
  878. index += 1
  879. # Test Console
  880. await agent1._model_context.clear() # pyright: ignore
  881. mock.reset()
  882. index = 0
  883. await team.reset()
  884. result2 = await Console(team.run_stream(task="task"))
  885. assert result2 == result
  886. @pytest.mark.asyncio
  887. async def test_swarm_pause_and_resume() -> None:
  888. first_agent = _HandOffAgent("first_agent", description="first agent", next_agent="second_agent")
  889. second_agent = _HandOffAgent("second_agent", description="second agent", next_agent="third_agent")
  890. third_agent = _HandOffAgent("third_agent", description="third agent", next_agent="first_agent")
  891. team = Swarm([second_agent, first_agent, third_agent], max_turns=1)
  892. result = await team.run(task="task")
  893. assert len(result.messages) == 2
  894. assert result.messages[0].content == "task"
  895. assert result.messages[1].content == "Transferred to third_agent."
  896. # Resume with a new task.
  897. result = await team.run(task="new task")
  898. assert len(result.messages) == 2
  899. assert result.messages[0].content == "new task"
  900. assert result.messages[1].content == "Transferred to first_agent."
  901. # Resume with the same task.
  902. result = await team.run()
  903. assert len(result.messages) == 1
  904. assert result.messages[0].content == "Transferred to second_agent."
  905. @pytest.mark.asyncio
  906. async def test_swarm_with_parallel_tool_calls(monkeypatch: pytest.MonkeyPatch) -> None:
  907. model = "gpt-4o-2024-05-13"
  908. chat_completions = [
  909. ChatCompletion(
  910. id="id1",
  911. choices=[
  912. Choice(
  913. finish_reason="tool_calls",
  914. index=0,
  915. message=ChatCompletionMessage(
  916. content=None,
  917. tool_calls=[
  918. ChatCompletionMessageToolCall(
  919. id="1",
  920. type="function",
  921. function=Function(
  922. name="tool1",
  923. arguments=json.dumps({}),
  924. ),
  925. ),
  926. ChatCompletionMessageToolCall(
  927. id="2",
  928. type="function",
  929. function=Function(
  930. name="tool2",
  931. arguments=json.dumps({}),
  932. ),
  933. ),
  934. ChatCompletionMessageToolCall(
  935. id="3",
  936. type="function",
  937. function=Function(
  938. name="handoff_to_agent2",
  939. arguments=json.dumps({}),
  940. ),
  941. ),
  942. ],
  943. role="assistant",
  944. ),
  945. )
  946. ],
  947. created=0,
  948. model=model,
  949. object="chat.completion",
  950. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  951. ),
  952. ChatCompletion(
  953. id="id2",
  954. choices=[
  955. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="Hello", role="assistant"))
  956. ],
  957. created=0,
  958. model=model,
  959. object="chat.completion",
  960. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  961. ),
  962. ChatCompletion(
  963. id="id2",
  964. choices=[
  965. Choice(
  966. finish_reason="stop", index=0, message=ChatCompletionMessage(content="TERMINATE", role="assistant")
  967. )
  968. ],
  969. created=0,
  970. model=model,
  971. object="chat.completion",
  972. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  973. ),
  974. ]
  975. mock = _MockChatCompletion(chat_completions)
  976. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  977. expected_handoff_context: List[LLMMessage] = [
  978. AssistantMessage(
  979. source="agent1",
  980. content=[
  981. FunctionCall(id="1", name="tool1", arguments="{}"),
  982. FunctionCall(id="2", name="tool2", arguments="{}"),
  983. ],
  984. ),
  985. FunctionExecutionResultMessage(
  986. content=[
  987. FunctionExecutionResult(content="tool1", call_id="1"),
  988. FunctionExecutionResult(content="tool2", call_id="2"),
  989. ]
  990. ),
  991. ]
  992. def tool1() -> str:
  993. return "tool1"
  994. def tool2() -> str:
  995. return "tool2"
  996. agent1 = AssistantAgent(
  997. "agent1",
  998. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  999. handoffs=[Handoff(target="agent2", name="handoff_to_agent2", message="handoff to agent2")],
  1000. tools=[tool1, tool2],
  1001. )
  1002. agent2 = AssistantAgent(
  1003. "agent2",
  1004. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  1005. )
  1006. termination = TextMentionTermination("TERMINATE")
  1007. team = Swarm([agent1, agent2], termination_condition=termination)
  1008. result = await team.run(task="task")
  1009. assert len(result.messages) == 6
  1010. assert result.messages[0] == TextMessage(content="task", source="user")
  1011. assert isinstance(result.messages[1], ToolCallRequestEvent)
  1012. assert isinstance(result.messages[2], ToolCallExecutionEvent)
  1013. assert result.messages[3] == HandoffMessage(
  1014. content="handoff to agent2",
  1015. target="agent2",
  1016. source="agent1",
  1017. context=expected_handoff_context,
  1018. )
  1019. assert result.messages[4].content == "Hello"
  1020. assert result.messages[4].source == "agent2"
  1021. assert result.messages[5].content == "TERMINATE"
  1022. assert result.messages[5].source == "agent2"
  1023. # Verify the tool calls are in agent2's context.
  1024. agent2_model_ctx_messages = await agent2._model_context.get_messages() # pyright: ignore
  1025. assert agent2_model_ctx_messages[0] == UserMessage(content="task", source="user")
  1026. assert agent2_model_ctx_messages[1] == expected_handoff_context[0]
  1027. assert agent2_model_ctx_messages[2] == expected_handoff_context[1]
  1028. @pytest.mark.asyncio
  1029. async def test_swarm_with_handoff_termination() -> None:
  1030. first_agent = _HandOffAgent("first_agent", description="first agent", next_agent="second_agent")
  1031. second_agent = _HandOffAgent("second_agent", description="second agent", next_agent="third_agent")
  1032. third_agent = _HandOffAgent("third_agent", description="third agent", next_agent="first_agent")
  1033. # Handoff to an existing agent.
  1034. termination = HandoffTermination(target="third_agent")
  1035. team = Swarm([second_agent, first_agent, third_agent], termination_condition=termination)
  1036. # Start
  1037. result = await team.run(task="task")
  1038. assert len(result.messages) == 2
  1039. assert result.messages[0].content == "task"
  1040. assert result.messages[1].content == "Transferred to third_agent."
  1041. # Resume existing.
  1042. result = await team.run()
  1043. assert len(result.messages) == 3
  1044. assert result.messages[0].content == "Transferred to first_agent."
  1045. assert result.messages[1].content == "Transferred to second_agent."
  1046. assert result.messages[2].content == "Transferred to third_agent."
  1047. # Resume new task.
  1048. result = await team.run(task="new task")
  1049. assert len(result.messages) == 4
  1050. assert result.messages[0].content == "new task"
  1051. assert result.messages[1].content == "Transferred to first_agent."
  1052. assert result.messages[2].content == "Transferred to second_agent."
  1053. assert result.messages[3].content == "Transferred to third_agent."
  1054. # Handoff to a non-existing agent.
  1055. third_agent = _HandOffAgent("third_agent", description="third agent", next_agent="non_existing_agent")
  1056. termination = HandoffTermination(target="non_existing_agent")
  1057. team = Swarm([second_agent, first_agent, third_agent], termination_condition=termination)
  1058. # Start
  1059. result = await team.run(task="task")
  1060. assert len(result.messages) == 3
  1061. assert result.messages[0].content == "task"
  1062. assert result.messages[1].content == "Transferred to third_agent."
  1063. assert result.messages[2].content == "Transferred to non_existing_agent."
  1064. # Attempt to resume.
  1065. with pytest.raises(ValueError):
  1066. await team.run()
  1067. # Attempt to resume with a new task.
  1068. with pytest.raises(ValueError):
  1069. await team.run(task="new task")
  1070. # Resume with a HandoffMessage
  1071. result = await team.run(task=HandoffMessage(content="Handoff to first_agent.", target="first_agent", source="user"))
  1072. assert len(result.messages) == 4
  1073. assert result.messages[0].content == "Handoff to first_agent."
  1074. assert result.messages[1].content == "Transferred to second_agent."
  1075. assert result.messages[2].content == "Transferred to third_agent."
  1076. assert result.messages[3].content == "Transferred to non_existing_agent."
  1077. @pytest.mark.asyncio
  1078. async def test_round_robin_group_chat_with_message_list() -> None:
  1079. # Create a simple team with echo agents
  1080. agent1 = _EchoAgent("Agent1", "First agent")
  1081. agent2 = _EchoAgent("Agent2", "Second agent")
  1082. termination = MaxMessageTermination(4) # Stop after 4 messages
  1083. team = RoundRobinGroupChat([agent1, agent2], termination_condition=termination)
  1084. # Create a list of messages
  1085. messages: List[ChatMessage] = [
  1086. TextMessage(content="Message 1", source="user"),
  1087. TextMessage(content="Message 2", source="user"),
  1088. TextMessage(content="Message 3", source="user"),
  1089. ]
  1090. # Run the team with the message list
  1091. result = await team.run(task=messages)
  1092. # Verify the messages were processed in order
  1093. assert len(result.messages) == 4 # Initial messages + echo until termination
  1094. assert result.messages[0].content == "Message 1" # First message
  1095. assert result.messages[1].content == "Message 2" # Second message
  1096. assert result.messages[2].content == "Message 3" # Third message
  1097. assert result.messages[3].content == "Message 1" # Echo from first agent
  1098. assert result.stop_reason == "Maximum number of messages 4 reached, current message count: 4"
  1099. # Test with streaming
  1100. await team.reset()
  1101. index = 0
  1102. async for message in team.run_stream(task=messages):
  1103. if isinstance(message, TaskResult):
  1104. assert message == result
  1105. else:
  1106. assert message == result.messages[index]
  1107. index += 1
  1108. # Test with invalid message list
  1109. with pytest.raises(ValueError, match="All messages in task list must be valid ChatMessage types"):
  1110. await team.run(task=["not a message"]) # type: ignore[list-item, arg-type] # intentionally testing invalid input
  1111. # Test with empty message list
  1112. with pytest.raises(ValueError, match="Task list cannot be empty"):
  1113. await team.run(task=[])