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

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