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

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933934935936937938939940941942943944945946947948949950951952953954955956957958959960961962963964965966967968969970971972973974975976977978979980981982983984985986987988989990991992993994995996997998999100010011002100310041005100610071008100910101011101210131014101510161017101810191020102110221023102410251026102710281029103010311032103310341035
  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, FunctionCall
  35. from autogen_core.components.models import FunctionExecutionResult
  36. from autogen_core.components.tools import FunctionTool
  37. from autogen_ext.code_executors.local import LocalCommandLineCodeExecutor
  38. from autogen_ext.models import OpenAIChatCompletionClient, 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. mock = _MockChatCompletion(chat_completions)
  280. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  281. tool = FunctionTool(_pass_function, name="pass", description="pass function")
  282. tool_use_agent = AssistantAgent(
  283. "tool_use_agent",
  284. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  285. tools=[tool],
  286. )
  287. echo_agent = _EchoAgent("echo_agent", description="echo agent")
  288. termination = TextMentionTermination("TERMINATE")
  289. team = RoundRobinGroupChat(participants=[tool_use_agent, echo_agent], termination_condition=termination)
  290. result = await team.run(
  291. task="Write a program that prints 'Hello, world!'",
  292. )
  293. assert len(result.messages) == 6
  294. assert isinstance(result.messages[0], TextMessage) # task
  295. assert isinstance(result.messages[1], ToolCallMessage) # tool call
  296. assert isinstance(result.messages[2], ToolCallResultMessage) # tool call result
  297. assert isinstance(result.messages[3], TextMessage) # tool use agent response
  298. assert isinstance(result.messages[4], TextMessage) # echo agent response
  299. assert isinstance(result.messages[5], TextMessage) # tool use agent response
  300. assert result.stop_reason is not None and result.stop_reason == "Text 'TERMINATE' mentioned"
  301. context = tool_use_agent._model_context # pyright: ignore
  302. assert context[0].content == "Write a program that prints 'Hello, world!'"
  303. assert isinstance(context[1].content, list)
  304. assert isinstance(context[1].content[0], FunctionCall)
  305. assert context[1].content[0].name == "pass"
  306. assert context[1].content[0].arguments == json.dumps({"input": "pass"})
  307. assert isinstance(context[2].content, list)
  308. assert isinstance(context[2].content[0], FunctionExecutionResult)
  309. assert context[2].content[0].content == "pass"
  310. assert context[2].content[0].call_id == "1"
  311. assert context[3].content == "Hello"
  312. # Test streaming.
  313. tool_use_agent._model_context.clear() # pyright: ignore
  314. mock.reset()
  315. index = 0
  316. await team.reset()
  317. async for message in team.run_stream(
  318. task="Write a program that prints 'Hello, world!'",
  319. ):
  320. if isinstance(message, TaskResult):
  321. assert message == result
  322. else:
  323. assert message == result.messages[index]
  324. index += 1
  325. # Test Console.
  326. tool_use_agent._model_context.clear() # pyright: ignore
  327. mock.reset()
  328. index = 0
  329. await team.reset()
  330. result2 = await Console(team.run_stream(task="Write a program that prints 'Hello, world!'"))
  331. assert result2 == result
  332. @pytest.mark.asyncio
  333. async def test_round_robin_group_chat_with_resume_and_reset() -> None:
  334. agent_1 = _EchoAgent("agent_1", description="echo agent 1")
  335. agent_2 = _EchoAgent("agent_2", description="echo agent 2")
  336. agent_3 = _EchoAgent("agent_3", description="echo agent 3")
  337. agent_4 = _EchoAgent("agent_4", description="echo agent 4")
  338. termination = MaxMessageTermination(3)
  339. team = RoundRobinGroupChat(participants=[agent_1, agent_2, agent_3, agent_4], termination_condition=termination)
  340. result = await team.run(
  341. task="Write a program that prints 'Hello, world!'",
  342. )
  343. assert len(result.messages) == 3
  344. assert result.messages[1].source == "agent_1"
  345. assert result.messages[2].source == "agent_2"
  346. assert result.stop_reason is not None
  347. # Resume.
  348. result = await team.run()
  349. assert len(result.messages) == 3
  350. assert result.messages[0].source == "agent_3"
  351. assert result.messages[1].source == "agent_4"
  352. assert result.messages[2].source == "agent_1"
  353. assert result.stop_reason is not None
  354. # Reset.
  355. await team.reset()
  356. result = await team.run(task="Write a program that prints 'Hello, world!'")
  357. assert len(result.messages) == 3
  358. assert result.messages[1].source == "agent_1"
  359. assert result.messages[2].source == "agent_2"
  360. assert result.stop_reason is not None
  361. @pytest.mark.asyncio
  362. async def test_round_robin_group_chat_max_turn() -> None:
  363. agent_1 = _EchoAgent("agent_1", description="echo agent 1")
  364. agent_2 = _EchoAgent("agent_2", description="echo agent 2")
  365. agent_3 = _EchoAgent("agent_3", description="echo agent 3")
  366. agent_4 = _EchoAgent("agent_4", description="echo agent 4")
  367. team = RoundRobinGroupChat(participants=[agent_1, agent_2, agent_3, agent_4], max_turns=3)
  368. result = await team.run(
  369. task="Write a program that prints 'Hello, world!'",
  370. )
  371. assert len(result.messages) == 4
  372. assert result.messages[1].source == "agent_1"
  373. assert result.messages[2].source == "agent_2"
  374. assert result.messages[3].source == "agent_3"
  375. assert result.stop_reason is not None
  376. # Resume.
  377. result = await team.run()
  378. assert len(result.messages) == 3
  379. assert result.messages[0].source == "agent_4"
  380. assert result.messages[1].source == "agent_1"
  381. assert result.messages[2].source == "agent_2"
  382. assert result.stop_reason is not None
  383. # Reset.
  384. await team.reset()
  385. result = await team.run(task="Write a program that prints 'Hello, world!'")
  386. assert len(result.messages) == 4
  387. assert result.messages[1].source == "agent_1"
  388. assert result.messages[2].source == "agent_2"
  389. assert result.messages[3].source == "agent_3"
  390. assert result.stop_reason is not None
  391. @pytest.mark.asyncio
  392. async def test_round_robin_group_chat_cancellation() -> None:
  393. agent_1 = _EchoAgent("agent_1", description="echo agent 1")
  394. agent_2 = _EchoAgent("agent_2", description="echo agent 2")
  395. agent_3 = _EchoAgent("agent_3", description="echo agent 3")
  396. agent_4 = _EchoAgent("agent_4", description="echo agent 4")
  397. # Set max_turns to a large number to avoid stopping due to max_turns before cancellation.
  398. team = RoundRobinGroupChat(participants=[agent_1, agent_2, agent_3, agent_4], max_turns=1000)
  399. cancellation_token = CancellationToken()
  400. run_task = asyncio.create_task(
  401. team.run(
  402. task="Write a program that prints 'Hello, world!'",
  403. cancellation_token=cancellation_token,
  404. )
  405. )
  406. await asyncio.sleep(0.1)
  407. # Cancel the task.
  408. cancellation_token.cancel()
  409. with pytest.raises(asyncio.CancelledError):
  410. await run_task
  411. # Total messages produced so far.
  412. total_messages = agent_1.total_messages + agent_2.total_messages + agent_3.total_messages + agent_4.total_messages
  413. # Still can run again and finish the task.
  414. result = await team.run()
  415. assert len(result.messages) + total_messages == 1000
  416. @pytest.mark.asyncio
  417. async def test_selector_group_chat(monkeypatch: pytest.MonkeyPatch) -> None:
  418. model = "gpt-4o-2024-05-13"
  419. chat_completions = [
  420. ChatCompletion(
  421. id="id2",
  422. choices=[
  423. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent3", role="assistant"))
  424. ],
  425. created=0,
  426. model=model,
  427. object="chat.completion",
  428. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  429. ),
  430. ChatCompletion(
  431. id="id2",
  432. choices=[
  433. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent2", role="assistant"))
  434. ],
  435. created=0,
  436. model=model,
  437. object="chat.completion",
  438. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  439. ),
  440. ChatCompletion(
  441. id="id2",
  442. choices=[
  443. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent1", role="assistant"))
  444. ],
  445. created=0,
  446. model=model,
  447. object="chat.completion",
  448. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  449. ),
  450. ChatCompletion(
  451. id="id2",
  452. choices=[
  453. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent2", role="assistant"))
  454. ],
  455. created=0,
  456. model=model,
  457. object="chat.completion",
  458. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  459. ),
  460. ChatCompletion(
  461. id="id2",
  462. choices=[
  463. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent1", role="assistant"))
  464. ],
  465. created=0,
  466. model=model,
  467. object="chat.completion",
  468. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  469. ),
  470. ]
  471. mock = _MockChatCompletion(chat_completions)
  472. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  473. agent1 = _StopAgent("agent1", description="echo agent 1", stop_at=2)
  474. agent2 = _EchoAgent("agent2", description="echo agent 2")
  475. agent3 = _EchoAgent("agent3", description="echo agent 3")
  476. termination = TextMentionTermination("TERMINATE")
  477. team = SelectorGroupChat(
  478. participants=[agent1, agent2, agent3],
  479. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  480. termination_condition=termination,
  481. )
  482. result = await team.run(
  483. task="Write a program that prints 'Hello, world!'",
  484. )
  485. assert len(result.messages) == 6
  486. assert result.messages[0].content == "Write a program that prints 'Hello, world!'"
  487. assert result.messages[1].source == "agent3"
  488. assert result.messages[2].source == "agent2"
  489. assert result.messages[3].source == "agent1"
  490. assert result.messages[4].source == "agent2"
  491. assert result.messages[5].source == "agent1"
  492. assert result.stop_reason is not None and result.stop_reason == "Text 'TERMINATE' mentioned"
  493. # Test streaming.
  494. mock.reset()
  495. agent1._count = 0 # pyright: ignore
  496. index = 0
  497. await team.reset()
  498. async for message in team.run_stream(
  499. task="Write a program that prints 'Hello, world!'",
  500. ):
  501. if isinstance(message, TaskResult):
  502. assert message == result
  503. else:
  504. assert message == result.messages[index]
  505. index += 1
  506. # Test Console.
  507. mock.reset()
  508. agent1._count = 0 # pyright: ignore
  509. index = 0
  510. await team.reset()
  511. result2 = await Console(team.run_stream(task="Write a program that prints 'Hello, world!'"))
  512. assert result2 == result
  513. @pytest.mark.asyncio
  514. async def test_selector_group_chat_state() -> None:
  515. model_client = ReplayChatCompletionClient(
  516. ["agent1", "No facts", "agent2", "No plan", "agent1", "print('Hello, world!')", "agent2", "TERMINATE"],
  517. )
  518. agent1 = AssistantAgent("agent1", model_client=model_client)
  519. agent2 = AssistantAgent("agent2", model_client=model_client)
  520. termination = TextMentionTermination("TERMINATE")
  521. team1 = SelectorGroupChat(
  522. participants=[agent1, agent2], termination_condition=termination, model_client=model_client
  523. )
  524. await team1.run(task="Write a program that prints 'Hello, world!'")
  525. state = await team1.save_state()
  526. agent3 = AssistantAgent("agent1", model_client=model_client)
  527. agent4 = AssistantAgent("agent2", model_client=model_client)
  528. team2 = SelectorGroupChat(
  529. participants=[agent3, agent4], termination_condition=termination, model_client=model_client
  530. )
  531. await team2.load_state(state)
  532. state2 = await team2.save_state()
  533. assert state == state2
  534. assert agent3._model_context == agent1._model_context # pyright: ignore
  535. assert agent4._model_context == agent2._model_context # pyright: ignore
  536. manager_1 = await team1._runtime.try_get_underlying_agent_instance( # pyright: ignore
  537. AgentId("group_chat_manager", team1._team_id), # pyright: ignore
  538. SelectorGroupChatManager, # pyright: ignore
  539. ) # pyright: ignore
  540. manager_2 = await team2._runtime.try_get_underlying_agent_instance( # pyright: ignore
  541. AgentId("group_chat_manager", team2._team_id), # pyright: ignore
  542. SelectorGroupChatManager, # pyright: ignore
  543. ) # pyright: ignore
  544. assert manager_1._message_thread == manager_2._message_thread # pyright: ignore
  545. assert manager_1._previous_speaker == manager_2._previous_speaker # pyright: ignore
  546. @pytest.mark.asyncio
  547. async def test_selector_group_chat_two_speakers(monkeypatch: pytest.MonkeyPatch) -> None:
  548. model = "gpt-4o-2024-05-13"
  549. chat_completions = [
  550. ChatCompletion(
  551. id="id2",
  552. choices=[
  553. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent2", role="assistant"))
  554. ],
  555. created=0,
  556. model=model,
  557. object="chat.completion",
  558. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  559. ),
  560. ]
  561. mock = _MockChatCompletion(chat_completions)
  562. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  563. agent1 = _StopAgent("agent1", description="echo agent 1", stop_at=2)
  564. agent2 = _EchoAgent("agent2", description="echo agent 2")
  565. termination = TextMentionTermination("TERMINATE")
  566. team = SelectorGroupChat(
  567. participants=[agent1, agent2],
  568. termination_condition=termination,
  569. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  570. )
  571. result = await team.run(
  572. task="Write a program that prints 'Hello, world!'",
  573. )
  574. assert len(result.messages) == 5
  575. assert result.messages[0].content == "Write a program that prints 'Hello, world!'"
  576. assert result.messages[1].source == "agent2"
  577. assert result.messages[2].source == "agent1"
  578. assert result.messages[3].source == "agent2"
  579. assert result.messages[4].source == "agent1"
  580. # only one chat completion was called
  581. assert mock._curr_index == 1 # pyright: ignore
  582. assert result.stop_reason is not None and result.stop_reason == "Text 'TERMINATE' mentioned"
  583. # Test streaming.
  584. mock.reset()
  585. agent1._count = 0 # pyright: ignore
  586. index = 0
  587. await team.reset()
  588. async for message in team.run_stream(task="Write a program that prints 'Hello, world!'"):
  589. if isinstance(message, TaskResult):
  590. assert message == result
  591. else:
  592. assert message == result.messages[index]
  593. index += 1
  594. # Test Console.
  595. mock.reset()
  596. agent1._count = 0 # pyright: ignore
  597. index = 0
  598. await team.reset()
  599. result2 = await Console(team.run_stream(task="Write a program that prints 'Hello, world!'"))
  600. assert result2 == result
  601. @pytest.mark.asyncio
  602. async def test_selector_group_chat_two_speakers_allow_repeated(monkeypatch: pytest.MonkeyPatch) -> None:
  603. model = "gpt-4o-2024-05-13"
  604. chat_completions = [
  605. ChatCompletion(
  606. id="id2",
  607. choices=[
  608. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent2", role="assistant"))
  609. ],
  610. created=0,
  611. model=model,
  612. object="chat.completion",
  613. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  614. ),
  615. ChatCompletion(
  616. id="id2",
  617. choices=[
  618. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent2", role="assistant"))
  619. ],
  620. created=0,
  621. model=model,
  622. object="chat.completion",
  623. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  624. ),
  625. ChatCompletion(
  626. id="id2",
  627. choices=[
  628. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent1", role="assistant"))
  629. ],
  630. created=0,
  631. model=model,
  632. object="chat.completion",
  633. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  634. ),
  635. ]
  636. mock = _MockChatCompletion(chat_completions)
  637. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  638. agent1 = _StopAgent("agent1", description="echo agent 1", stop_at=1)
  639. agent2 = _EchoAgent("agent2", description="echo agent 2")
  640. termination = TextMentionTermination("TERMINATE")
  641. team = SelectorGroupChat(
  642. participants=[agent1, agent2],
  643. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  644. termination_condition=termination,
  645. allow_repeated_speaker=True,
  646. )
  647. result = await team.run(task="Write a program that prints 'Hello, world!'")
  648. assert len(result.messages) == 4
  649. assert result.messages[0].content == "Write a program that prints 'Hello, world!'"
  650. assert result.messages[1].source == "agent2"
  651. assert result.messages[2].source == "agent2"
  652. assert result.messages[3].source == "agent1"
  653. assert result.stop_reason is not None and result.stop_reason == "Text 'TERMINATE' mentioned"
  654. # Test streaming.
  655. mock.reset()
  656. index = 0
  657. await team.reset()
  658. async for message in team.run_stream(task="Write a program that prints 'Hello, world!'"):
  659. if isinstance(message, TaskResult):
  660. assert message == result
  661. else:
  662. assert message == result.messages[index]
  663. index += 1
  664. # Test Console.
  665. mock.reset()
  666. index = 0
  667. await team.reset()
  668. result2 = await Console(team.run_stream(task="Write a program that prints 'Hello, world!'"))
  669. assert result2 == result
  670. @pytest.mark.asyncio
  671. async def test_selector_group_chat_custom_selector(monkeypatch: pytest.MonkeyPatch) -> None:
  672. model = "gpt-4o-2024-05-13"
  673. chat_completions = [
  674. ChatCompletion(
  675. id="id2",
  676. choices=[
  677. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="agent3", role="assistant"))
  678. ],
  679. created=0,
  680. model=model,
  681. object="chat.completion",
  682. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  683. ),
  684. ]
  685. mock = _MockChatCompletion(chat_completions)
  686. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  687. agent1 = _EchoAgent("agent1", description="echo agent 1")
  688. agent2 = _EchoAgent("agent2", description="echo agent 2")
  689. agent3 = _EchoAgent("agent3", description="echo agent 3")
  690. agent4 = _EchoAgent("agent4", description="echo agent 4")
  691. def _select_agent(messages: Sequence[AgentMessage]) -> str | None:
  692. if len(messages) == 0:
  693. return "agent1"
  694. elif messages[-1].source == "agent1":
  695. return "agent2"
  696. elif messages[-1].source == "agent2":
  697. return None
  698. elif messages[-1].source == "agent3":
  699. return "agent4"
  700. else:
  701. return "agent1"
  702. termination = MaxMessageTermination(6)
  703. team = SelectorGroupChat(
  704. participants=[agent1, agent2, agent3, agent4],
  705. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  706. selector_func=_select_agent,
  707. termination_condition=termination,
  708. )
  709. result = await team.run(task="task")
  710. assert len(result.messages) == 6
  711. assert result.messages[1].source == "agent1"
  712. assert result.messages[2].source == "agent2"
  713. assert result.messages[3].source == "agent3"
  714. assert result.messages[4].source == "agent4"
  715. assert result.messages[5].source == "agent1"
  716. assert (
  717. result.stop_reason is not None
  718. and result.stop_reason == "Maximum number of messages 6 reached, current message count: 6"
  719. )
  720. class _HandOffAgent(BaseChatAgent):
  721. def __init__(self, name: str, description: str, next_agent: str) -> None:
  722. super().__init__(name, description)
  723. self._next_agent = next_agent
  724. @property
  725. def produced_message_types(self) -> List[type[ChatMessage]]:
  726. return [HandoffMessage]
  727. async def on_messages(self, messages: Sequence[ChatMessage], cancellation_token: CancellationToken) -> Response:
  728. return Response(
  729. chat_message=HandoffMessage(
  730. content=f"Transferred to {self._next_agent}.", target=self._next_agent, source=self.name
  731. )
  732. )
  733. async def on_reset(self, cancellation_token: CancellationToken) -> None:
  734. pass
  735. @pytest.mark.asyncio
  736. async def test_swarm_handoff() -> None:
  737. first_agent = _HandOffAgent("first_agent", description="first agent", next_agent="second_agent")
  738. second_agent = _HandOffAgent("second_agent", description="second agent", next_agent="third_agent")
  739. third_agent = _HandOffAgent("third_agent", description="third agent", next_agent="first_agent")
  740. termination = MaxMessageTermination(6)
  741. team = Swarm([second_agent, first_agent, third_agent], termination_condition=termination)
  742. result = await team.run(task="task")
  743. assert len(result.messages) == 6
  744. assert result.messages[0].content == "task"
  745. assert result.messages[1].content == "Transferred to third_agent."
  746. assert result.messages[2].content == "Transferred to first_agent."
  747. assert result.messages[3].content == "Transferred to second_agent."
  748. assert result.messages[4].content == "Transferred to third_agent."
  749. assert result.messages[5].content == "Transferred to first_agent."
  750. assert (
  751. result.stop_reason is not None
  752. and result.stop_reason == "Maximum number of messages 6 reached, current message count: 6"
  753. )
  754. # Test streaming.
  755. index = 0
  756. await team.reset()
  757. stream = team.run_stream(task="task")
  758. async for message in stream:
  759. if isinstance(message, TaskResult):
  760. assert message == result
  761. else:
  762. assert message == result.messages[index]
  763. index += 1
  764. # Test save and load.
  765. state = await team.save_state()
  766. first_agent2 = _HandOffAgent("first_agent", description="first agent", next_agent="second_agent")
  767. second_agent2 = _HandOffAgent("second_agent", description="second agent", next_agent="third_agent")
  768. third_agent2 = _HandOffAgent("third_agent", description="third agent", next_agent="first_agent")
  769. team2 = Swarm([second_agent2, first_agent2, third_agent2], termination_condition=termination)
  770. await team2.load_state(state)
  771. state2 = await team2.save_state()
  772. assert state == state2
  773. manager_1 = await team._runtime.try_get_underlying_agent_instance( # pyright: ignore
  774. AgentId("group_chat_manager", team._team_id), # pyright: ignore
  775. SwarmGroupChatManager, # pyright: ignore
  776. ) # pyright: ignore
  777. manager_2 = await team2._runtime.try_get_underlying_agent_instance( # pyright: ignore
  778. AgentId("group_chat_manager", team2._team_id), # pyright: ignore
  779. SwarmGroupChatManager, # pyright: ignore
  780. ) # pyright: ignore
  781. assert manager_1._message_thread == manager_2._message_thread # pyright: ignore
  782. assert manager_1._current_speaker == manager_2._current_speaker # pyright: ignore
  783. @pytest.mark.asyncio
  784. async def test_swarm_handoff_using_tool_calls(monkeypatch: pytest.MonkeyPatch) -> None:
  785. model = "gpt-4o-2024-05-13"
  786. chat_completions = [
  787. ChatCompletion(
  788. id="id1",
  789. choices=[
  790. Choice(
  791. finish_reason="tool_calls",
  792. index=0,
  793. message=ChatCompletionMessage(
  794. content=None,
  795. tool_calls=[
  796. ChatCompletionMessageToolCall(
  797. id="1",
  798. type="function",
  799. function=Function(
  800. name="handoff_to_agent2",
  801. arguments=json.dumps({}),
  802. ),
  803. )
  804. ],
  805. role="assistant",
  806. ),
  807. )
  808. ],
  809. created=0,
  810. model=model,
  811. object="chat.completion",
  812. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  813. ),
  814. ChatCompletion(
  815. id="id2",
  816. choices=[
  817. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="Hello", role="assistant"))
  818. ],
  819. created=0,
  820. model=model,
  821. object="chat.completion",
  822. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  823. ),
  824. ChatCompletion(
  825. id="id2",
  826. choices=[
  827. Choice(
  828. finish_reason="stop", index=0, message=ChatCompletionMessage(content="TERMINATE", role="assistant")
  829. )
  830. ],
  831. created=0,
  832. model=model,
  833. object="chat.completion",
  834. usage=CompletionUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
  835. ),
  836. ]
  837. mock = _MockChatCompletion(chat_completions)
  838. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  839. agent1 = AssistantAgent(
  840. "agent1",
  841. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  842. handoffs=[Handoff(target="agent2", name="handoff_to_agent2", message="handoff to agent2")],
  843. )
  844. agent2 = _HandOffAgent("agent2", description="agent 2", next_agent="agent1")
  845. termination = TextMentionTermination("TERMINATE")
  846. team = Swarm([agent1, agent2], termination_condition=termination)
  847. result = await team.run(task="task")
  848. assert len(result.messages) == 7
  849. assert result.messages[0].content == "task"
  850. assert isinstance(result.messages[1], ToolCallMessage)
  851. assert isinstance(result.messages[2], ToolCallResultMessage)
  852. assert result.messages[3].content == "handoff to agent2"
  853. assert result.messages[4].content == "Transferred to agent1."
  854. assert result.messages[5].content == "Hello"
  855. assert result.messages[6].content == "TERMINATE"
  856. assert result.stop_reason is not None and result.stop_reason == "Text 'TERMINATE' mentioned"
  857. # Test streaming.
  858. agent1._model_context.clear() # pyright: ignore
  859. mock.reset()
  860. index = 0
  861. await team.reset()
  862. stream = team.run_stream(task="task")
  863. async for message in stream:
  864. if isinstance(message, TaskResult):
  865. assert message == result
  866. else:
  867. assert message == result.messages[index]
  868. index += 1
  869. # Test Console
  870. agent1._model_context.clear() # pyright: ignore
  871. mock.reset()
  872. index = 0
  873. await team.reset()
  874. result2 = await Console(team.run_stream(task="task"))
  875. assert result2 == result
  876. @pytest.mark.asyncio
  877. async def test_swarm_pause_and_resume() -> None:
  878. first_agent = _HandOffAgent("first_agent", description="first agent", next_agent="second_agent")
  879. second_agent = _HandOffAgent("second_agent", description="second agent", next_agent="third_agent")
  880. third_agent = _HandOffAgent("third_agent", description="third agent", next_agent="first_agent")
  881. team = Swarm([second_agent, first_agent, third_agent], max_turns=1)
  882. result = await team.run(task="task")
  883. assert len(result.messages) == 2
  884. assert result.messages[0].content == "task"
  885. assert result.messages[1].content == "Transferred to third_agent."
  886. # Resume with a new task.
  887. result = await team.run(task="new task")
  888. assert len(result.messages) == 2
  889. assert result.messages[0].content == "new task"
  890. assert result.messages[1].content == "Transferred to first_agent."
  891. # Resume with the same task.
  892. result = await team.run()
  893. assert len(result.messages) == 1
  894. assert result.messages[0].content == "Transferred to second_agent."
  895. @pytest.mark.asyncio
  896. async def test_swarm_with_handoff_termination() -> None:
  897. first_agent = _HandOffAgent("first_agent", description="first agent", next_agent="second_agent")
  898. second_agent = _HandOffAgent("second_agent", description="second agent", next_agent="third_agent")
  899. third_agent = _HandOffAgent("third_agent", description="third agent", next_agent="first_agent")
  900. # Handoff to an existing agent.
  901. termination = HandoffTermination(target="third_agent")
  902. team = Swarm([second_agent, first_agent, third_agent], termination_condition=termination)
  903. # Start
  904. result = await team.run(task="task")
  905. assert len(result.messages) == 2
  906. assert result.messages[0].content == "task"
  907. assert result.messages[1].content == "Transferred to third_agent."
  908. # Resume existing.
  909. result = await team.run()
  910. assert len(result.messages) == 3
  911. assert result.messages[0].content == "Transferred to first_agent."
  912. assert result.messages[1].content == "Transferred to second_agent."
  913. assert result.messages[2].content == "Transferred to third_agent."
  914. # Resume new task.
  915. result = await team.run(task="new task")
  916. assert len(result.messages) == 4
  917. assert result.messages[0].content == "new task"
  918. assert result.messages[1].content == "Transferred to first_agent."
  919. assert result.messages[2].content == "Transferred to second_agent."
  920. assert result.messages[3].content == "Transferred to third_agent."
  921. # Handoff to a non-existing agent.
  922. third_agent = _HandOffAgent("third_agent", description="third agent", next_agent="non_existing_agent")
  923. termination = HandoffTermination(target="non_existing_agent")
  924. team = Swarm([second_agent, first_agent, third_agent], termination_condition=termination)
  925. # Start
  926. result = await team.run(task="task")
  927. assert len(result.messages) == 3
  928. assert result.messages[0].content == "task"
  929. assert result.messages[1].content == "Transferred to third_agent."
  930. assert result.messages[2].content == "Transferred to non_existing_agent."
  931. # Attempt to resume.
  932. with pytest.raises(ValueError):
  933. await team.run()
  934. # Attempt to resume with a new task.
  935. with pytest.raises(ValueError):
  936. await team.run(task="new task")
  937. # Resume with a HandoffMessage
  938. result = await team.run(task=HandoffMessage(content="Handoff to first_agent.", target="first_agent", source="user"))
  939. assert len(result.messages) == 4
  940. assert result.messages[0].content == "Handoff to first_agent."
  941. assert result.messages[1].content == "Transferred to second_agent."
  942. assert result.messages[2].content == "Transferred to third_agent."
  943. assert result.messages[3].content == "Transferred to non_existing_agent."