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_assistant_agent.py 17 kB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469
  1. import asyncio
  2. import json
  3. import logging
  4. from typing import Any, AsyncGenerator, List
  5. import pytest
  6. from autogen_agentchat import EVENT_LOGGER_NAME
  7. from autogen_agentchat.agents import AssistantAgent
  8. from autogen_agentchat.base import Handoff, TaskResult
  9. from autogen_agentchat.messages import (
  10. ChatMessage,
  11. HandoffMessage,
  12. MultiModalMessage,
  13. TextMessage,
  14. ToolCallMessage,
  15. ToolCallResultMessage,
  16. )
  17. from autogen_core import Image
  18. from autogen_core.tools import FunctionTool
  19. from autogen_ext.models.openai import OpenAIChatCompletionClient
  20. from openai.resources.chat.completions import AsyncCompletions
  21. from openai.types.chat.chat_completion import ChatCompletion, Choice
  22. from openai.types.chat.chat_completion_chunk import ChatCompletionChunk
  23. from openai.types.chat.chat_completion_message import ChatCompletionMessage
  24. from openai.types.chat.chat_completion_message_tool_call import (
  25. ChatCompletionMessageToolCall,
  26. Function,
  27. )
  28. from openai.types.completion_usage import CompletionUsage
  29. from utils import FileLogHandler
  30. logger = logging.getLogger(EVENT_LOGGER_NAME)
  31. logger.setLevel(logging.DEBUG)
  32. logger.addHandler(FileLogHandler("test_assistant_agent.log"))
  33. class _MockChatCompletion:
  34. def __init__(self, chat_completions: List[ChatCompletion]) -> None:
  35. self._saved_chat_completions = chat_completions
  36. self.curr_index = 0
  37. async def mock_create(
  38. self, *args: Any, **kwargs: Any
  39. ) -> ChatCompletion | AsyncGenerator[ChatCompletionChunk, None]:
  40. await asyncio.sleep(0.1)
  41. completion = self._saved_chat_completions[self.curr_index]
  42. self.curr_index += 1
  43. return completion
  44. def _pass_function(input: str) -> str:
  45. return "pass"
  46. async def _fail_function(input: str) -> str:
  47. return "fail"
  48. async def _echo_function(input: str) -> str:
  49. return input
  50. @pytest.mark.asyncio
  51. async def test_run_with_tools(monkeypatch: pytest.MonkeyPatch) -> None:
  52. model = "gpt-4o-2024-05-13"
  53. chat_completions = [
  54. ChatCompletion(
  55. id="id1",
  56. choices=[
  57. Choice(
  58. finish_reason="tool_calls",
  59. index=0,
  60. message=ChatCompletionMessage(
  61. content=None,
  62. tool_calls=[
  63. ChatCompletionMessageToolCall(
  64. id="1",
  65. type="function",
  66. function=Function(
  67. name="_pass_function",
  68. arguments=json.dumps({"input": "task"}),
  69. ),
  70. )
  71. ],
  72. role="assistant",
  73. ),
  74. )
  75. ],
  76. created=0,
  77. model=model,
  78. object="chat.completion",
  79. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  80. ),
  81. ChatCompletion(
  82. id="id2",
  83. choices=[
  84. Choice(
  85. finish_reason="stop",
  86. index=0,
  87. message=ChatCompletionMessage(content="pass", role="assistant"),
  88. )
  89. ],
  90. created=0,
  91. model=model,
  92. object="chat.completion",
  93. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  94. ),
  95. ChatCompletion(
  96. id="id2",
  97. choices=[
  98. Choice(
  99. finish_reason="stop",
  100. index=0,
  101. message=ChatCompletionMessage(content="TERMINATE", role="assistant"),
  102. )
  103. ],
  104. created=0,
  105. model=model,
  106. object="chat.completion",
  107. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  108. ),
  109. ]
  110. mock = _MockChatCompletion(chat_completions)
  111. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  112. agent = AssistantAgent(
  113. "tool_use_agent",
  114. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  115. tools=[
  116. _pass_function,
  117. _fail_function,
  118. FunctionTool(_echo_function, description="Echo"),
  119. ],
  120. )
  121. result = await agent.run(task="task")
  122. assert len(result.messages) == 4
  123. assert isinstance(result.messages[0], TextMessage)
  124. assert result.messages[0].models_usage is None
  125. assert isinstance(result.messages[1], ToolCallMessage)
  126. assert result.messages[1].models_usage is not None
  127. assert result.messages[1].models_usage.completion_tokens == 5
  128. assert result.messages[1].models_usage.prompt_tokens == 10
  129. assert isinstance(result.messages[2], ToolCallResultMessage)
  130. assert result.messages[2].models_usage is None
  131. assert isinstance(result.messages[3], TextMessage)
  132. assert result.messages[3].content == "pass"
  133. assert result.messages[3].models_usage is None
  134. # Test streaming.
  135. mock.curr_index = 0 # Reset the mock
  136. index = 0
  137. async for message in agent.run_stream(task="task"):
  138. if isinstance(message, TaskResult):
  139. assert message == result
  140. else:
  141. assert message == result.messages[index]
  142. index += 1
  143. # Test state saving and loading.
  144. state = await agent.save_state()
  145. agent2 = AssistantAgent(
  146. "tool_use_agent",
  147. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  148. tools=[_pass_function, _fail_function, FunctionTool(_echo_function, description="Echo")],
  149. )
  150. await agent2.load_state(state)
  151. state2 = await agent2.save_state()
  152. assert state == state2
  153. @pytest.mark.asyncio
  154. async def test_run_with_tools_and_reflection(monkeypatch: pytest.MonkeyPatch) -> None:
  155. model = "gpt-4o-2024-05-13"
  156. chat_completions = [
  157. ChatCompletion(
  158. id="id1",
  159. choices=[
  160. Choice(
  161. finish_reason="tool_calls",
  162. index=0,
  163. message=ChatCompletionMessage(
  164. content=None,
  165. tool_calls=[
  166. ChatCompletionMessageToolCall(
  167. id="1",
  168. type="function",
  169. function=Function(
  170. name="_pass_function",
  171. arguments=json.dumps({"input": "task"}),
  172. ),
  173. )
  174. ],
  175. role="assistant",
  176. ),
  177. )
  178. ],
  179. created=0,
  180. model=model,
  181. object="chat.completion",
  182. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  183. ),
  184. ChatCompletion(
  185. id="id2",
  186. choices=[
  187. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="Hello", role="assistant"))
  188. ],
  189. created=0,
  190. model=model,
  191. object="chat.completion",
  192. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  193. ),
  194. ChatCompletion(
  195. id="id2",
  196. choices=[
  197. Choice(
  198. finish_reason="stop", index=0, message=ChatCompletionMessage(content="TERMINATE", role="assistant")
  199. )
  200. ],
  201. created=0,
  202. model=model,
  203. object="chat.completion",
  204. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  205. ),
  206. ]
  207. mock = _MockChatCompletion(chat_completions)
  208. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  209. agent = AssistantAgent(
  210. "tool_use_agent",
  211. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  212. tools=[_pass_function, _fail_function, FunctionTool(_echo_function, description="Echo")],
  213. reflect_on_tool_use=True,
  214. )
  215. result = await agent.run(task="task")
  216. assert len(result.messages) == 4
  217. assert isinstance(result.messages[0], TextMessage)
  218. assert result.messages[0].models_usage is None
  219. assert isinstance(result.messages[1], ToolCallMessage)
  220. assert result.messages[1].models_usage is not None
  221. assert result.messages[1].models_usage.completion_tokens == 5
  222. assert result.messages[1].models_usage.prompt_tokens == 10
  223. assert isinstance(result.messages[2], ToolCallResultMessage)
  224. assert result.messages[2].models_usage is None
  225. assert isinstance(result.messages[3], TextMessage)
  226. assert result.messages[3].content == "Hello"
  227. assert result.messages[3].models_usage is not None
  228. assert result.messages[3].models_usage.completion_tokens == 5
  229. assert result.messages[3].models_usage.prompt_tokens == 10
  230. # Test streaming.
  231. mock.curr_index = 0 # pyright: ignore
  232. index = 0
  233. async for message in agent.run_stream(task="task"):
  234. if isinstance(message, TaskResult):
  235. assert message == result
  236. else:
  237. assert message == result.messages[index]
  238. index += 1
  239. # Test state saving and loading.
  240. state = await agent.save_state()
  241. agent2 = AssistantAgent(
  242. "tool_use_agent",
  243. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  244. tools=[
  245. _pass_function,
  246. _fail_function,
  247. FunctionTool(_echo_function, description="Echo"),
  248. ],
  249. )
  250. await agent2.load_state(state)
  251. state2 = await agent2.save_state()
  252. assert state == state2
  253. @pytest.mark.asyncio
  254. async def test_handoffs(monkeypatch: pytest.MonkeyPatch) -> None:
  255. handoff = Handoff(target="agent2")
  256. model = "gpt-4o-2024-05-13"
  257. chat_completions = [
  258. ChatCompletion(
  259. id="id1",
  260. choices=[
  261. Choice(
  262. finish_reason="tool_calls",
  263. index=0,
  264. message=ChatCompletionMessage(
  265. content=None,
  266. tool_calls=[
  267. ChatCompletionMessageToolCall(
  268. id="1",
  269. type="function",
  270. function=Function(
  271. name=handoff.name,
  272. arguments=json.dumps({}),
  273. ),
  274. )
  275. ],
  276. role="assistant",
  277. ),
  278. )
  279. ],
  280. created=0,
  281. model=model,
  282. object="chat.completion",
  283. usage=CompletionUsage(prompt_tokens=42, completion_tokens=43, total_tokens=85),
  284. ),
  285. ]
  286. mock = _MockChatCompletion(chat_completions)
  287. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  288. tool_use_agent = AssistantAgent(
  289. "tool_use_agent",
  290. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  291. tools=[
  292. _pass_function,
  293. _fail_function,
  294. FunctionTool(_echo_function, description="Echo"),
  295. ],
  296. handoffs=[handoff],
  297. )
  298. assert HandoffMessage in tool_use_agent.produced_message_types
  299. result = await tool_use_agent.run(task="task")
  300. assert len(result.messages) == 4
  301. assert isinstance(result.messages[0], TextMessage)
  302. assert result.messages[0].models_usage is None
  303. assert isinstance(result.messages[1], ToolCallMessage)
  304. assert result.messages[1].models_usage is not None
  305. assert result.messages[1].models_usage.completion_tokens == 43
  306. assert result.messages[1].models_usage.prompt_tokens == 42
  307. assert isinstance(result.messages[2], ToolCallResultMessage)
  308. assert result.messages[2].models_usage is None
  309. assert isinstance(result.messages[3], HandoffMessage)
  310. assert result.messages[3].content == handoff.message
  311. assert result.messages[3].target == handoff.target
  312. assert result.messages[3].models_usage is None
  313. # Test streaming.
  314. mock.curr_index = 0 # pyright: ignore
  315. index = 0
  316. async for message in tool_use_agent.run_stream(task="task"):
  317. if isinstance(message, TaskResult):
  318. assert message == result
  319. else:
  320. assert message == result.messages[index]
  321. index += 1
  322. @pytest.mark.asyncio
  323. async def test_multi_modal_task(monkeypatch: pytest.MonkeyPatch) -> None:
  324. model = "gpt-4o-2024-05-13"
  325. chat_completions = [
  326. ChatCompletion(
  327. id="id2",
  328. choices=[
  329. Choice(
  330. finish_reason="stop",
  331. index=0,
  332. message=ChatCompletionMessage(content="Hello", role="assistant"),
  333. )
  334. ],
  335. created=0,
  336. model=model,
  337. object="chat.completion",
  338. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  339. ),
  340. ]
  341. mock = _MockChatCompletion(chat_completions)
  342. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  343. agent = AssistantAgent(
  344. name="assistant",
  345. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  346. )
  347. # Generate a random base64 image.
  348. img_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGP4//8/AAX+Av4N70a4AAAAAElFTkSuQmCC"
  349. result = await agent.run(task=MultiModalMessage(source="user", content=["Test", Image.from_base64(img_base64)]))
  350. assert len(result.messages) == 2
  351. @pytest.mark.asyncio
  352. async def test_invalid_model_capabilities() -> None:
  353. model = "random-model"
  354. model_client = OpenAIChatCompletionClient(
  355. model=model,
  356. api_key="",
  357. model_capabilities={
  358. "vision": False,
  359. "function_calling": False,
  360. "json_output": False,
  361. },
  362. )
  363. with pytest.raises(ValueError):
  364. agent = AssistantAgent(
  365. name="assistant",
  366. model_client=model_client,
  367. tools=[
  368. _pass_function,
  369. _fail_function,
  370. FunctionTool(_echo_function, description="Echo"),
  371. ],
  372. )
  373. with pytest.raises(ValueError):
  374. agent = AssistantAgent(name="assistant", model_client=model_client, handoffs=["agent2"])
  375. with pytest.raises(ValueError):
  376. agent = AssistantAgent(name="assistant", model_client=model_client)
  377. # Generate a random base64 image.
  378. img_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGP4//8/AAX+Av4N70a4AAAAAElFTkSuQmCC"
  379. await agent.run(task=MultiModalMessage(source="user", content=["Test", Image.from_base64(img_base64)]))
  380. @pytest.mark.asyncio
  381. async def test_list_chat_messages(monkeypatch: pytest.MonkeyPatch) -> None:
  382. model = "gpt-4o-2024-05-13"
  383. chat_completions = [
  384. ChatCompletion(
  385. id="id1",
  386. choices=[
  387. Choice(
  388. finish_reason="stop",
  389. index=0,
  390. message=ChatCompletionMessage(content="Response to message 1", role="assistant"),
  391. )
  392. ],
  393. created=0,
  394. model=model,
  395. object="chat.completion",
  396. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=15),
  397. ),
  398. ]
  399. mock = _MockChatCompletion(chat_completions)
  400. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  401. agent = AssistantAgent(
  402. "test_agent",
  403. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  404. )
  405. # Create a list of chat messages
  406. messages: List[ChatMessage] = [
  407. TextMessage(content="Message 1", source="user"),
  408. TextMessage(content="Message 2", source="user"),
  409. ]
  410. # Test run method with list of messages
  411. result = await agent.run(task=messages)
  412. assert len(result.messages) == 3 # 2 input messages + 1 response message
  413. assert isinstance(result.messages[0], TextMessage)
  414. assert result.messages[0].content == "Message 1"
  415. assert result.messages[0].source == "user"
  416. assert isinstance(result.messages[1], TextMessage)
  417. assert result.messages[1].content == "Message 2"
  418. assert result.messages[1].source == "user"
  419. assert isinstance(result.messages[2], TextMessage)
  420. assert result.messages[2].content == "Response to message 1"
  421. assert result.messages[2].source == "test_agent"
  422. assert result.messages[2].models_usage is not None
  423. assert result.messages[2].models_usage.completion_tokens == 5
  424. assert result.messages[2].models_usage.prompt_tokens == 10
  425. # Test run_stream method with list of messages
  426. mock.curr_index = 0 # Reset mock index using public attribute
  427. index = 0
  428. async for message in agent.run_stream(task=messages):
  429. if isinstance(message, TaskResult):
  430. assert message == result
  431. else:
  432. assert message == result.messages[index]
  433. index += 1