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

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