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

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842
  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. MemoryQueryEvent,
  13. ModelClientStreamingChunkEvent,
  14. MultiModalMessage,
  15. TextMessage,
  16. ToolCallExecutionEvent,
  17. ToolCallRequestEvent,
  18. ToolCallSummaryMessage,
  19. )
  20. from autogen_core import FunctionCall, Image
  21. from autogen_core.memory import ListMemory, Memory, MemoryContent, MemoryMimeType, MemoryQueryResult
  22. from autogen_core.model_context import BufferedChatCompletionContext
  23. from autogen_core.models import CreateResult, FunctionExecutionResult, LLMMessage, RequestUsage
  24. from autogen_core.models._model_client import ModelFamily
  25. from autogen_core.tools import FunctionTool
  26. from autogen_ext.models.openai import OpenAIChatCompletionClient
  27. from autogen_ext.models.replay import ReplayChatCompletionClient
  28. from openai.resources.chat.completions import AsyncCompletions
  29. from openai.types.chat.chat_completion import ChatCompletion, Choice
  30. from openai.types.chat.chat_completion_chunk import ChatCompletionChunk
  31. from openai.types.chat.chat_completion_message import ChatCompletionMessage
  32. from openai.types.chat.chat_completion_message_tool_call import (
  33. ChatCompletionMessageToolCall,
  34. Function,
  35. )
  36. from openai.types.completion_usage import CompletionUsage
  37. from utils import FileLogHandler
  38. logger = logging.getLogger(EVENT_LOGGER_NAME)
  39. logger.setLevel(logging.DEBUG)
  40. logger.addHandler(FileLogHandler("test_assistant_agent.log"))
  41. class _MockChatCompletion:
  42. def __init__(self, chat_completions: List[ChatCompletion]) -> None:
  43. self._saved_chat_completions = chat_completions
  44. self.curr_index = 0
  45. self.calls: List[List[LLMMessage]] = []
  46. async def mock_create(
  47. self, *args: Any, **kwargs: Any
  48. ) -> ChatCompletion | AsyncGenerator[ChatCompletionChunk, None]:
  49. self.calls.append(kwargs["messages"]) # Save the call
  50. await asyncio.sleep(0.1)
  51. completion = self._saved_chat_completions[self.curr_index]
  52. self.curr_index += 1
  53. return completion
  54. def _pass_function(input: str) -> str:
  55. return "pass"
  56. async def _fail_function(input: str) -> str:
  57. return "fail"
  58. async def _echo_function(input: str) -> str:
  59. return input
  60. @pytest.mark.asyncio
  61. async def test_run_with_tools(monkeypatch: pytest.MonkeyPatch) -> None:
  62. model = "gpt-4o-2024-05-13"
  63. chat_completions = [
  64. ChatCompletion(
  65. id="id1",
  66. choices=[
  67. Choice(
  68. finish_reason="tool_calls",
  69. index=0,
  70. message=ChatCompletionMessage(
  71. content=None,
  72. tool_calls=[
  73. ChatCompletionMessageToolCall(
  74. id="1",
  75. type="function",
  76. function=Function(
  77. name="_pass_function",
  78. arguments=json.dumps({"input": "task"}),
  79. ),
  80. )
  81. ],
  82. role="assistant",
  83. ),
  84. )
  85. ],
  86. created=0,
  87. model=model,
  88. object="chat.completion",
  89. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  90. ),
  91. ChatCompletion(
  92. id="id2",
  93. choices=[
  94. Choice(
  95. finish_reason="stop",
  96. index=0,
  97. message=ChatCompletionMessage(content="pass", role="assistant"),
  98. )
  99. ],
  100. created=0,
  101. model=model,
  102. object="chat.completion",
  103. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  104. ),
  105. ChatCompletion(
  106. id="id2",
  107. choices=[
  108. Choice(
  109. finish_reason="stop",
  110. index=0,
  111. message=ChatCompletionMessage(content="TERMINATE", role="assistant"),
  112. )
  113. ],
  114. created=0,
  115. model=model,
  116. object="chat.completion",
  117. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  118. ),
  119. ]
  120. mock = _MockChatCompletion(chat_completions)
  121. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  122. agent = AssistantAgent(
  123. "tool_use_agent",
  124. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  125. tools=[
  126. _pass_function,
  127. _fail_function,
  128. FunctionTool(_echo_function, description="Echo"),
  129. ],
  130. )
  131. result = await agent.run(task="task")
  132. assert len(result.messages) == 4
  133. assert isinstance(result.messages[0], TextMessage)
  134. assert result.messages[0].models_usage is None
  135. assert isinstance(result.messages[1], ToolCallRequestEvent)
  136. assert result.messages[1].models_usage is not None
  137. assert result.messages[1].models_usage.completion_tokens == 5
  138. assert result.messages[1].models_usage.prompt_tokens == 10
  139. assert isinstance(result.messages[2], ToolCallExecutionEvent)
  140. assert result.messages[2].models_usage is None
  141. assert isinstance(result.messages[3], ToolCallSummaryMessage)
  142. assert result.messages[3].content == "pass"
  143. assert result.messages[3].models_usage is None
  144. # Test streaming.
  145. mock.curr_index = 0 # Reset the mock
  146. index = 0
  147. async for message in agent.run_stream(task="task"):
  148. if isinstance(message, TaskResult):
  149. assert message == result
  150. else:
  151. assert message == result.messages[index]
  152. index += 1
  153. # Test state saving and loading.
  154. state = await agent.save_state()
  155. agent2 = AssistantAgent(
  156. "tool_use_agent",
  157. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  158. tools=[_pass_function, _fail_function, FunctionTool(_echo_function, description="Echo")],
  159. )
  160. await agent2.load_state(state)
  161. state2 = await agent2.save_state()
  162. assert state == state2
  163. @pytest.mark.asyncio
  164. async def test_run_with_tools_and_reflection(monkeypatch: pytest.MonkeyPatch) -> None:
  165. model = "gpt-4o-2024-05-13"
  166. chat_completions = [
  167. ChatCompletion(
  168. id="id1",
  169. choices=[
  170. Choice(
  171. finish_reason="tool_calls",
  172. index=0,
  173. message=ChatCompletionMessage(
  174. content=None,
  175. tool_calls=[
  176. ChatCompletionMessageToolCall(
  177. id="1",
  178. type="function",
  179. function=Function(
  180. name="_pass_function",
  181. arguments=json.dumps({"input": "task"}),
  182. ),
  183. )
  184. ],
  185. role="assistant",
  186. ),
  187. )
  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(finish_reason="stop", index=0, message=ChatCompletionMessage(content="Hello", role="assistant"))
  198. ],
  199. created=0,
  200. model=model,
  201. object="chat.completion",
  202. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  203. ),
  204. ChatCompletion(
  205. id="id2",
  206. choices=[
  207. Choice(
  208. finish_reason="stop", index=0, message=ChatCompletionMessage(content="TERMINATE", role="assistant")
  209. )
  210. ],
  211. created=0,
  212. model=model,
  213. object="chat.completion",
  214. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  215. ),
  216. ]
  217. mock = _MockChatCompletion(chat_completions)
  218. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  219. agent = AssistantAgent(
  220. "tool_use_agent",
  221. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  222. tools=[_pass_function, _fail_function, FunctionTool(_echo_function, description="Echo")],
  223. reflect_on_tool_use=True,
  224. )
  225. result = await agent.run(task="task")
  226. assert len(result.messages) == 4
  227. assert isinstance(result.messages[0], TextMessage)
  228. assert result.messages[0].models_usage is None
  229. assert isinstance(result.messages[1], ToolCallRequestEvent)
  230. assert result.messages[1].models_usage is not None
  231. assert result.messages[1].models_usage.completion_tokens == 5
  232. assert result.messages[1].models_usage.prompt_tokens == 10
  233. assert isinstance(result.messages[2], ToolCallExecutionEvent)
  234. assert result.messages[2].models_usage is None
  235. assert isinstance(result.messages[3], TextMessage)
  236. assert result.messages[3].content == "Hello"
  237. assert result.messages[3].models_usage is not None
  238. assert result.messages[3].models_usage.completion_tokens == 5
  239. assert result.messages[3].models_usage.prompt_tokens == 10
  240. # Test streaming.
  241. mock.curr_index = 0 # pyright: ignore
  242. index = 0
  243. async for message in agent.run_stream(task="task"):
  244. if isinstance(message, TaskResult):
  245. assert message == result
  246. else:
  247. assert message == result.messages[index]
  248. index += 1
  249. # Test state saving and loading.
  250. state = await agent.save_state()
  251. agent2 = AssistantAgent(
  252. "tool_use_agent",
  253. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  254. tools=[
  255. _pass_function,
  256. _fail_function,
  257. FunctionTool(_echo_function, description="Echo"),
  258. ],
  259. )
  260. await agent2.load_state(state)
  261. state2 = await agent2.save_state()
  262. assert state == state2
  263. @pytest.mark.asyncio
  264. async def test_run_with_parallel_tools(monkeypatch: pytest.MonkeyPatch) -> None:
  265. model = "gpt-4o-2024-05-13"
  266. chat_completions = [
  267. ChatCompletion(
  268. id="id1",
  269. choices=[
  270. Choice(
  271. finish_reason="tool_calls",
  272. index=0,
  273. message=ChatCompletionMessage(
  274. content=None,
  275. tool_calls=[
  276. ChatCompletionMessageToolCall(
  277. id="1",
  278. type="function",
  279. function=Function(
  280. name="_pass_function",
  281. arguments=json.dumps({"input": "task1"}),
  282. ),
  283. ),
  284. ChatCompletionMessageToolCall(
  285. id="2",
  286. type="function",
  287. function=Function(
  288. name="_pass_function",
  289. arguments=json.dumps({"input": "task2"}),
  290. ),
  291. ),
  292. ChatCompletionMessageToolCall(
  293. id="3",
  294. type="function",
  295. function=Function(
  296. name="_echo_function",
  297. arguments=json.dumps({"input": "task3"}),
  298. ),
  299. ),
  300. ],
  301. role="assistant",
  302. ),
  303. )
  304. ],
  305. created=0,
  306. model=model,
  307. object="chat.completion",
  308. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  309. ),
  310. ChatCompletion(
  311. id="id2",
  312. choices=[
  313. Choice(
  314. finish_reason="stop",
  315. index=0,
  316. message=ChatCompletionMessage(content="pass", role="assistant"),
  317. )
  318. ],
  319. created=0,
  320. model=model,
  321. object="chat.completion",
  322. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  323. ),
  324. ChatCompletion(
  325. id="id2",
  326. choices=[
  327. Choice(
  328. finish_reason="stop",
  329. index=0,
  330. message=ChatCompletionMessage(content="TERMINATE", role="assistant"),
  331. )
  332. ],
  333. created=0,
  334. model=model,
  335. object="chat.completion",
  336. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  337. ),
  338. ]
  339. mock = _MockChatCompletion(chat_completions)
  340. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  341. agent = AssistantAgent(
  342. "tool_use_agent",
  343. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  344. tools=[
  345. _pass_function,
  346. _fail_function,
  347. FunctionTool(_echo_function, description="Echo"),
  348. ],
  349. )
  350. result = await agent.run(task="task")
  351. assert len(result.messages) == 4
  352. assert isinstance(result.messages[0], TextMessage)
  353. assert result.messages[0].models_usage is None
  354. assert isinstance(result.messages[1], ToolCallRequestEvent)
  355. assert result.messages[1].content == [
  356. FunctionCall(id="1", arguments=r'{"input": "task1"}', name="_pass_function"),
  357. FunctionCall(id="2", arguments=r'{"input": "task2"}', name="_pass_function"),
  358. FunctionCall(id="3", arguments=r'{"input": "task3"}', name="_echo_function"),
  359. ]
  360. assert result.messages[1].models_usage is not None
  361. assert result.messages[1].models_usage.completion_tokens == 5
  362. assert result.messages[1].models_usage.prompt_tokens == 10
  363. assert isinstance(result.messages[2], ToolCallExecutionEvent)
  364. expected_content = [
  365. FunctionExecutionResult(call_id="1", content="pass"),
  366. FunctionExecutionResult(call_id="2", content="pass"),
  367. FunctionExecutionResult(call_id="3", content="task3"),
  368. ]
  369. for expected in expected_content:
  370. assert expected in result.messages[2].content
  371. assert result.messages[2].models_usage is None
  372. assert isinstance(result.messages[3], ToolCallSummaryMessage)
  373. assert result.messages[3].content == "pass\npass\ntask3"
  374. assert result.messages[3].models_usage is None
  375. # Test streaming.
  376. mock.curr_index = 0 # Reset the mock
  377. index = 0
  378. async for message in agent.run_stream(task="task"):
  379. if isinstance(message, TaskResult):
  380. assert message == result
  381. else:
  382. assert message == result.messages[index]
  383. index += 1
  384. # Test state saving and loading.
  385. state = await agent.save_state()
  386. agent2 = AssistantAgent(
  387. "tool_use_agent",
  388. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  389. tools=[_pass_function, _fail_function, FunctionTool(_echo_function, description="Echo")],
  390. )
  391. await agent2.load_state(state)
  392. state2 = await agent2.save_state()
  393. assert state == state2
  394. @pytest.mark.asyncio
  395. async def test_handoffs(monkeypatch: pytest.MonkeyPatch) -> None:
  396. handoff = Handoff(target="agent2")
  397. model = "gpt-4o-2024-05-13"
  398. chat_completions = [
  399. ChatCompletion(
  400. id="id1",
  401. choices=[
  402. Choice(
  403. finish_reason="tool_calls",
  404. index=0,
  405. message=ChatCompletionMessage(
  406. content=None,
  407. tool_calls=[
  408. ChatCompletionMessageToolCall(
  409. id="1",
  410. type="function",
  411. function=Function(
  412. name=handoff.name,
  413. arguments=json.dumps({}),
  414. ),
  415. )
  416. ],
  417. role="assistant",
  418. ),
  419. )
  420. ],
  421. created=0,
  422. model=model,
  423. object="chat.completion",
  424. usage=CompletionUsage(prompt_tokens=42, completion_tokens=43, total_tokens=85),
  425. ),
  426. ]
  427. mock = _MockChatCompletion(chat_completions)
  428. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  429. tool_use_agent = AssistantAgent(
  430. "tool_use_agent",
  431. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  432. tools=[
  433. _pass_function,
  434. _fail_function,
  435. FunctionTool(_echo_function, description="Echo"),
  436. ],
  437. handoffs=[handoff],
  438. )
  439. assert HandoffMessage in tool_use_agent.produced_message_types
  440. result = await tool_use_agent.run(task="task")
  441. assert len(result.messages) == 4
  442. assert isinstance(result.messages[0], TextMessage)
  443. assert result.messages[0].models_usage is None
  444. assert isinstance(result.messages[1], ToolCallRequestEvent)
  445. assert result.messages[1].models_usage is not None
  446. assert result.messages[1].models_usage.completion_tokens == 43
  447. assert result.messages[1].models_usage.prompt_tokens == 42
  448. assert isinstance(result.messages[2], ToolCallExecutionEvent)
  449. assert result.messages[2].models_usage is None
  450. assert isinstance(result.messages[3], HandoffMessage)
  451. assert result.messages[3].content == handoff.message
  452. assert result.messages[3].target == handoff.target
  453. assert result.messages[3].models_usage is None
  454. # Test streaming.
  455. mock.curr_index = 0 # pyright: ignore
  456. index = 0
  457. async for message in tool_use_agent.run_stream(task="task"):
  458. if isinstance(message, TaskResult):
  459. assert message == result
  460. else:
  461. assert message == result.messages[index]
  462. index += 1
  463. @pytest.mark.asyncio
  464. async def test_multi_modal_task(monkeypatch: pytest.MonkeyPatch) -> None:
  465. model = "gpt-4o-2024-05-13"
  466. chat_completions = [
  467. ChatCompletion(
  468. id="id2",
  469. choices=[
  470. Choice(
  471. finish_reason="stop",
  472. index=0,
  473. message=ChatCompletionMessage(content="Hello", role="assistant"),
  474. )
  475. ],
  476. created=0,
  477. model=model,
  478. object="chat.completion",
  479. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  480. ),
  481. ]
  482. mock = _MockChatCompletion(chat_completions)
  483. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  484. agent = AssistantAgent(
  485. name="assistant",
  486. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  487. )
  488. # Generate a random base64 image.
  489. img_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGP4//8/AAX+Av4N70a4AAAAAElFTkSuQmCC"
  490. result = await agent.run(task=MultiModalMessage(source="user", content=["Test", Image.from_base64(img_base64)]))
  491. assert len(result.messages) == 2
  492. @pytest.mark.asyncio
  493. async def test_invalid_model_capabilities() -> None:
  494. model = "random-model"
  495. model_client = OpenAIChatCompletionClient(
  496. model=model,
  497. api_key="",
  498. model_info={"vision": False, "function_calling": False, "json_output": False, "family": ModelFamily.UNKNOWN},
  499. )
  500. with pytest.raises(ValueError):
  501. agent = AssistantAgent(
  502. name="assistant",
  503. model_client=model_client,
  504. tools=[
  505. _pass_function,
  506. _fail_function,
  507. FunctionTool(_echo_function, description="Echo"),
  508. ],
  509. )
  510. with pytest.raises(ValueError):
  511. agent = AssistantAgent(name="assistant", model_client=model_client, handoffs=["agent2"])
  512. with pytest.raises(ValueError):
  513. agent = AssistantAgent(name="assistant", model_client=model_client)
  514. # Generate a random base64 image.
  515. img_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGP4//8/AAX+Av4N70a4AAAAAElFTkSuQmCC"
  516. await agent.run(task=MultiModalMessage(source="user", content=["Test", Image.from_base64(img_base64)]))
  517. @pytest.mark.asyncio
  518. async def test_list_chat_messages(monkeypatch: pytest.MonkeyPatch) -> None:
  519. model = "gpt-4o-2024-05-13"
  520. chat_completions = [
  521. ChatCompletion(
  522. id="id1",
  523. choices=[
  524. Choice(
  525. finish_reason="stop",
  526. index=0,
  527. message=ChatCompletionMessage(content="Response to message 1", role="assistant"),
  528. )
  529. ],
  530. created=0,
  531. model=model,
  532. object="chat.completion",
  533. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=15),
  534. ),
  535. ]
  536. mock = _MockChatCompletion(chat_completions)
  537. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  538. agent = AssistantAgent(
  539. "test_agent",
  540. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  541. )
  542. # Create a list of chat messages
  543. messages: List[ChatMessage] = [
  544. TextMessage(content="Message 1", source="user"),
  545. TextMessage(content="Message 2", source="user"),
  546. ]
  547. # Test run method with list of messages
  548. result = await agent.run(task=messages)
  549. assert len(result.messages) == 3 # 2 input messages + 1 response message
  550. assert isinstance(result.messages[0], TextMessage)
  551. assert result.messages[0].content == "Message 1"
  552. assert result.messages[0].source == "user"
  553. assert isinstance(result.messages[1], TextMessage)
  554. assert result.messages[1].content == "Message 2"
  555. assert result.messages[1].source == "user"
  556. assert isinstance(result.messages[2], TextMessage)
  557. assert result.messages[2].content == "Response to message 1"
  558. assert result.messages[2].source == "test_agent"
  559. assert result.messages[2].models_usage is not None
  560. assert result.messages[2].models_usage.completion_tokens == 5
  561. assert result.messages[2].models_usage.prompt_tokens == 10
  562. # Test run_stream method with list of messages
  563. mock.curr_index = 0 # Reset mock index using public attribute
  564. index = 0
  565. async for message in agent.run_stream(task=messages):
  566. if isinstance(message, TaskResult):
  567. assert message == result
  568. else:
  569. assert message == result.messages[index]
  570. index += 1
  571. @pytest.mark.asyncio
  572. async def test_model_context(monkeypatch: pytest.MonkeyPatch) -> None:
  573. model = "gpt-4o-2024-05-13"
  574. chat_completions = [
  575. ChatCompletion(
  576. id="id1",
  577. choices=[
  578. Choice(
  579. finish_reason="stop",
  580. index=0,
  581. message=ChatCompletionMessage(content="Response to message 3", role="assistant"),
  582. )
  583. ],
  584. created=0,
  585. model=model,
  586. object="chat.completion",
  587. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=15),
  588. ),
  589. ]
  590. mock = _MockChatCompletion(chat_completions)
  591. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  592. model_context = BufferedChatCompletionContext(buffer_size=2)
  593. agent = AssistantAgent(
  594. "test_agent",
  595. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  596. model_context=model_context,
  597. )
  598. messages = [
  599. TextMessage(content="Message 1", source="user"),
  600. TextMessage(content="Message 2", source="user"),
  601. TextMessage(content="Message 3", source="user"),
  602. ]
  603. await agent.run(task=messages)
  604. # Check if the mock client is called with only the last two messages.
  605. assert len(mock.calls) == 1
  606. # 2 message from the context + 1 system message
  607. assert len(mock.calls[0]) == 3
  608. @pytest.mark.asyncio
  609. async def test_run_with_memory(monkeypatch: pytest.MonkeyPatch) -> None:
  610. model = "gpt-4o-2024-05-13"
  611. chat_completions = [
  612. ChatCompletion(
  613. id="id1",
  614. choices=[
  615. Choice(
  616. finish_reason="stop",
  617. index=0,
  618. message=ChatCompletionMessage(content="Hello", role="assistant"),
  619. )
  620. ],
  621. created=0,
  622. model=model,
  623. object="chat.completion",
  624. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=0),
  625. ),
  626. ]
  627. b64_image_str = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGP4//8/AAX+Av4N70a4AAAAAElFTkSuQmCC"
  628. mock = _MockChatCompletion(chat_completions)
  629. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  630. # Test basic memory properties and empty context
  631. memory = ListMemory(name="test_memory")
  632. assert memory.name == "test_memory"
  633. empty_context = BufferedChatCompletionContext(buffer_size=2)
  634. empty_results = await memory.update_context(empty_context)
  635. assert len(empty_results.memories.results) == 0
  636. # Test various content types
  637. memory = ListMemory()
  638. await memory.add(MemoryContent(content="text content", mime_type=MemoryMimeType.TEXT))
  639. await memory.add(MemoryContent(content={"key": "value"}, mime_type=MemoryMimeType.JSON))
  640. await memory.add(MemoryContent(content=Image.from_base64(b64_image_str), mime_type=MemoryMimeType.IMAGE))
  641. # Test query functionality
  642. query_result = await memory.query(MemoryContent(content="", mime_type=MemoryMimeType.TEXT))
  643. assert isinstance(query_result, MemoryQueryResult)
  644. # Should have all three memories we added
  645. assert len(query_result.results) == 3
  646. # Test clear and cleanup
  647. await memory.clear()
  648. empty_query = await memory.query(MemoryContent(content="", mime_type=MemoryMimeType.TEXT))
  649. assert len(empty_query.results) == 0
  650. await memory.close() # Should not raise
  651. # Test invalid memory type
  652. with pytest.raises(TypeError):
  653. AssistantAgent(
  654. "test_agent",
  655. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  656. memory="invalid", # type: ignore
  657. )
  658. # Test with agent
  659. memory2 = ListMemory()
  660. await memory2.add(MemoryContent(content="test instruction", mime_type=MemoryMimeType.TEXT))
  661. agent = AssistantAgent(
  662. "test_agent", model_client=OpenAIChatCompletionClient(model=model, api_key=""), memory=[memory2]
  663. )
  664. result = await agent.run(task="test task")
  665. assert len(result.messages) > 0
  666. memory_event = next((msg for msg in result.messages if isinstance(msg, MemoryQueryEvent)), None)
  667. assert memory_event is not None
  668. assert len(memory_event.content) > 0
  669. assert isinstance(memory_event.content[0], MemoryContent)
  670. # Test memory protocol
  671. class BadMemory:
  672. pass
  673. assert not isinstance(BadMemory(), Memory)
  674. assert isinstance(ListMemory(), Memory)
  675. @pytest.mark.asyncio
  676. async def test_assistant_agent_declarative(monkeypatch: pytest.MonkeyPatch) -> None:
  677. model = "gpt-4o-2024-05-13"
  678. chat_completions = [
  679. ChatCompletion(
  680. id="id1",
  681. choices=[
  682. Choice(
  683. finish_reason="stop",
  684. index=0,
  685. message=ChatCompletionMessage(content="Response to message 3", role="assistant"),
  686. )
  687. ],
  688. created=0,
  689. model=model,
  690. object="chat.completion",
  691. usage=CompletionUsage(prompt_tokens=10, completion_tokens=5, total_tokens=15),
  692. ),
  693. ]
  694. mock = _MockChatCompletion(chat_completions)
  695. monkeypatch.setattr(AsyncCompletions, "create", mock.mock_create)
  696. model_context = BufferedChatCompletionContext(buffer_size=2)
  697. agent = AssistantAgent(
  698. "test_agent",
  699. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  700. model_context=model_context,
  701. )
  702. agent_config = agent.dump_component()
  703. assert agent_config.provider == "autogen_agentchat.agents.AssistantAgent"
  704. agent2 = AssistantAgent.load_component(agent_config)
  705. assert agent2.name == agent.name
  706. agent3 = AssistantAgent(
  707. "test_agent",
  708. model_client=OpenAIChatCompletionClient(model=model, api_key=""),
  709. model_context=model_context,
  710. tools=[
  711. _pass_function,
  712. _fail_function,
  713. FunctionTool(_echo_function, description="Echo"),
  714. ],
  715. )
  716. agent3_config = agent3.dump_component()
  717. assert agent3_config.provider == "autogen_agentchat.agents.AssistantAgent"
  718. @pytest.mark.asyncio
  719. async def test_model_client_stream() -> None:
  720. mock_client = ReplayChatCompletionClient(
  721. [
  722. "Response to message 3",
  723. ]
  724. )
  725. agent = AssistantAgent(
  726. "test_agent",
  727. model_client=mock_client,
  728. model_client_stream=True,
  729. )
  730. chunks: List[str] = []
  731. async for message in agent.run_stream(task="task"):
  732. if isinstance(message, TaskResult):
  733. assert message.messages[-1].content == "Response to message 3"
  734. elif isinstance(message, ModelClientStreamingChunkEvent):
  735. chunks.append(message.content)
  736. assert "".join(chunks) == "Response to message 3"
  737. @pytest.mark.asyncio
  738. async def test_model_client_stream_with_tool_calls() -> None:
  739. mock_client = ReplayChatCompletionClient(
  740. [
  741. CreateResult(
  742. content=[
  743. FunctionCall(id="1", name="_pass_function", arguments=r'{"input": "task"}'),
  744. FunctionCall(id="3", name="_echo_function", arguments=r'{"input": "task"}'),
  745. ],
  746. finish_reason="function_calls",
  747. usage=RequestUsage(prompt_tokens=10, completion_tokens=5),
  748. cached=False,
  749. ),
  750. "Example response 2 to task",
  751. ]
  752. )
  753. mock_client._model_info["function_calling"] = True # pyright: ignore
  754. agent = AssistantAgent(
  755. "test_agent",
  756. model_client=mock_client,
  757. model_client_stream=True,
  758. reflect_on_tool_use=True,
  759. tools=[_pass_function, _echo_function],
  760. )
  761. chunks: List[str] = []
  762. async for message in agent.run_stream(task="task"):
  763. if isinstance(message, TaskResult):
  764. assert message.messages[-1].content == "Example response 2 to task"
  765. assert message.messages[1].content == [
  766. FunctionCall(id="1", name="_pass_function", arguments=r'{"input": "task"}'),
  767. FunctionCall(id="3", name="_echo_function", arguments=r'{"input": "task"}'),
  768. ]
  769. assert message.messages[2].content == [
  770. FunctionExecutionResult(call_id="1", content="pass"),
  771. FunctionExecutionResult(call_id="3", content="task"),
  772. ]
  773. elif isinstance(message, ModelClientStreamingChunkEvent):
  774. chunks.append(message.content)
  775. assert "".join(chunks) == "Example response 2 to task"