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

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