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

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