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

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