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

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270
  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.components.tools import FunctionTool
  18. from autogen_ext.models 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].models_usage is not None
  119. assert result.messages[3].models_usage.completion_tokens == 5
  120. assert result.messages[3].models_usage.prompt_tokens == 10
  121. # Test streaming.
  122. mock._curr_index = 0 # pyright: ignore
  123. index = 0
  124. async for message in agent.run_stream(task="task"):
  125. if isinstance(message, TaskResult):
  126. assert message == result
  127. else:
  128. assert message == result.messages[index]
  129. index += 1
  130. # Test state saving and loading.
  131. state = await agent.save_state()
  132. agent2 = AssistantAgent(
  133. "tool_use_agent",
  134. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  135. tools=[_pass_function, _fail_function, FunctionTool(_echo_function, description="Echo")],
  136. )
  137. await agent2.load_state(state)
  138. state2 = await agent2.save_state()
  139. assert state == state2
  140. @pytest.mark.asyncio
  141. async def test_handoffs(monkeypatch: pytest.MonkeyPatch) -> None:
  142. handoff = Handoff(target="agent2")
  143. model = "gpt-4o-2024-05-13"
  144. chat_completions = [
  145. ChatCompletion(
  146. id="id1",
  147. choices=[
  148. Choice(
  149. finish_reason="tool_calls",
  150. index=0,
  151. message=ChatCompletionMessage(
  152. content=None,
  153. tool_calls=[
  154. ChatCompletionMessageToolCall(
  155. id="1",
  156. type="function",
  157. function=Function(
  158. name=handoff.name,
  159. arguments=json.dumps({}),
  160. ),
  161. )
  162. ],
  163. role="assistant",
  164. ),
  165. )
  166. ],
  167. created=0,
  168. model=model,
  169. object="chat.completion",
  170. usage=CompletionUsage(prompt_tokens=42, completion_tokens=43, total_tokens=85),
  171. ),
  172. ]
  173. mock = _MockChatCompletion(chat_completions)
  174. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  175. tool_use_agent = AssistantAgent(
  176. "tool_use_agent",
  177. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  178. tools=[_pass_function, _fail_function, FunctionTool(_echo_function, description="Echo")],
  179. handoffs=[handoff],
  180. )
  181. assert HandoffMessage in tool_use_agent.produced_message_types
  182. result = await tool_use_agent.run(task="task")
  183. assert len(result.messages) == 4
  184. assert isinstance(result.messages[0], TextMessage)
  185. assert result.messages[0].models_usage is None
  186. assert isinstance(result.messages[1], ToolCallMessage)
  187. assert result.messages[1].models_usage is not None
  188. assert result.messages[1].models_usage.completion_tokens == 43
  189. assert result.messages[1].models_usage.prompt_tokens == 42
  190. assert isinstance(result.messages[2], ToolCallResultMessage)
  191. assert result.messages[2].models_usage is None
  192. assert isinstance(result.messages[3], HandoffMessage)
  193. assert result.messages[3].content == handoff.message
  194. assert result.messages[3].target == handoff.target
  195. assert result.messages[3].models_usage is None
  196. # Test streaming.
  197. mock._curr_index = 0 # pyright: ignore
  198. index = 0
  199. async for message in tool_use_agent.run_stream(task="task"):
  200. if isinstance(message, TaskResult):
  201. assert message == result
  202. else:
  203. assert message == result.messages[index]
  204. index += 1
  205. @pytest.mark.asyncio
  206. async def test_multi_modal_task(monkeypatch: pytest.MonkeyPatch) -> None:
  207. model = "gpt-4o-2024-05-13"
  208. chat_completions = [
  209. ChatCompletion(
  210. id="id2",
  211. choices=[
  212. Choice(finish_reason="stop", index=0, message=ChatCompletionMessage(content="Hello", role="assistant"))
  213. ],
  214. created=0,
  215. model=model,
  216. object="chat.completion",
  217. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  218. ),
  219. ]
  220. mock = _MockChatCompletion(chat_completions)
  221. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  222. agent = AssistantAgent(name="assistant", model_client=OpenAIChatCompletionClient(model=model, api_key=""))
  223. # Generate a random base64 image.
  224. img_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGP4//8/AAX+Av4N70a4AAAAAElFTkSuQmCC"
  225. result = await agent.run(task=MultiModalMessage(source="user", content=["Test", Image.from_base64(img_base64)]))
  226. assert len(result.messages) == 2
  227. @pytest.mark.asyncio
  228. async def test_invalid_model_capabilities() -> None:
  229. model = "random-model"
  230. model_client = OpenAIChatCompletionClient(
  231. model=model, api_key="", model_capabilities={"vision": False, "function_calling": False, "json_output": False}
  232. )
  233. with pytest.raises(ValueError):
  234. agent = AssistantAgent(
  235. name="assistant",
  236. model_client=model_client,
  237. tools=[_pass_function, _fail_function, FunctionTool(_echo_function, description="Echo")],
  238. )
  239. with pytest.raises(ValueError):
  240. agent = AssistantAgent(name="assistant", model_client=model_client, handoffs=["agent2"])
  241. with pytest.raises(ValueError):
  242. agent = AssistantAgent(name="assistant", model_client=model_client)
  243. # Generate a random base64 image.
  244. img_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGP4//8/AAX+Av4N70a4AAAAAElFTkSuQmCC"
  245. await agent.run(task=MultiModalMessage(source="user", content=["Test", Image.from_base64(img_base64)]))