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

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