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

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