loop.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300
  1. from __future__ import annotations
  2. import json
  3. from collections.abc import AsyncIterator, Callable, Iterator
  4. from typing import TYPE_CHECKING
  5. from supply_agent.llm.client import LLMClient
  6. from supply_agent.tools.registry import ToolRegistry
  7. from supply_agent.types import (
  8. AgentEvent,
  9. AgentEventType,
  10. AgentResult,
  11. Message,
  12. Role,
  13. )
  14. if TYPE_CHECKING:
  15. from supply_agent.logging.logger import AgentLogger
  16. class AgentLoop:
  17. """
  18. ReAct-style agent loop: Reason → Act (tool call) → Observe → Repeat.
  19. Implements the standard tool-calling pattern used by modern agent frameworks.
  20. """
  21. def __init__(
  22. self,
  23. llm: LLMClient,
  24. tools: ToolRegistry,
  25. system_message: Message,
  26. messages: list[Message],
  27. max_iterations: int = 20,
  28. temperature: float | None = None,
  29. logger: AgentLogger | None = None,
  30. active_skills: list[str] | None = None,
  31. system_message_builder: Callable[[], Message] | None = None,
  32. ) -> None:
  33. self.llm = llm
  34. self.tools = tools
  35. self.system_message = system_message
  36. self.system_message_builder = system_message_builder
  37. self.messages = messages
  38. self.max_iterations = max_iterations
  39. self.temperature = temperature
  40. self.logger = logger
  41. self.active_skills = active_skills or []
  42. self.tool_calls_made = 0
  43. def _all_messages(self) -> list[Message]:
  44. return [self.system_message, *self.messages]
  45. def _on_skill_loaded(self, arguments: str) -> None:
  46. """Refresh system message after a skill is loaded."""
  47. try:
  48. args = json.loads(arguments)
  49. skill_name = args.get("name", "")
  50. if skill_name and skill_name not in self.active_skills:
  51. self.active_skills.append(skill_name)
  52. except json.JSONDecodeError:
  53. pass
  54. if self.system_message_builder:
  55. self.system_message = self.system_message_builder()
  56. def run(self) -> AgentResult:
  57. iterations = 0
  58. while iterations < self.max_iterations:
  59. iterations += 1
  60. response = self.llm.chat(
  61. self._all_messages(),
  62. tools=self.tools.definitions or None,
  63. temperature=self.temperature,
  64. iteration=iterations,
  65. )
  66. self.messages.append(response)
  67. if not response.tool_calls:
  68. return self._build_result(response.content or "", iterations)
  69. for tc in response.tool_calls:
  70. self.tool_calls_made += 1
  71. result = self.tools.execute(tc.id, tc.name, tc.arguments)
  72. if self.logger:
  73. self.logger.log_tool_call(
  74. iterations,
  75. tc.name,
  76. tc.arguments,
  77. result.content,
  78. result.is_error,
  79. tool_call_id=tc.id,
  80. )
  81. if tc.name == "load_skill" and not result.is_error:
  82. self._on_skill_loaded(tc.arguments)
  83. if self.logger:
  84. self.logger.log_skill_loaded(iterations, tc.arguments)
  85. self.messages.append(
  86. Message(
  87. role=Role.TOOL,
  88. content=result.content,
  89. tool_call_id=result.tool_call_id,
  90. name=result.name,
  91. )
  92. )
  93. self.messages.append(
  94. Message(
  95. role=Role.USER,
  96. content="Maximum iterations reached. Please provide your best answer now.",
  97. )
  98. )
  99. final = self.llm.chat(
  100. self._all_messages(),
  101. temperature=self.temperature,
  102. iteration=iterations + 1,
  103. )
  104. self.messages.append(final)
  105. return self._build_result(final.content or "", iterations)
  106. async def arun(self) -> AgentResult:
  107. iterations = 0
  108. while iterations < self.max_iterations:
  109. iterations += 1
  110. response = await self.llm.achat(
  111. self._all_messages(),
  112. tools=self.tools.definitions or None,
  113. temperature=self.temperature,
  114. iteration=iterations,
  115. )
  116. self.messages.append(response)
  117. if not response.tool_calls:
  118. return self._build_result(response.content or "", iterations)
  119. for tc in response.tool_calls:
  120. self.tool_calls_made += 1
  121. result = await self.tools.aexecute(tc.id, tc.name, tc.arguments)
  122. if self.logger:
  123. self.logger.log_tool_call(
  124. iterations,
  125. tc.name,
  126. tc.arguments,
  127. result.content,
  128. result.is_error,
  129. tool_call_id=tc.id,
  130. )
  131. if tc.name == "load_skill" and not result.is_error:
  132. self._on_skill_loaded(tc.arguments)
  133. if self.logger:
  134. self.logger.log_skill_loaded(iterations, tc.arguments)
  135. self.messages.append(
  136. Message(
  137. role=Role.TOOL,
  138. content=result.content,
  139. tool_call_id=result.tool_call_id,
  140. name=result.name,
  141. )
  142. )
  143. self.messages.append(
  144. Message(
  145. role=Role.USER,
  146. content="Maximum iterations reached. Please provide your best answer now.",
  147. )
  148. )
  149. final = await self.llm.achat(
  150. self._all_messages(),
  151. temperature=self.temperature,
  152. iteration=iterations + 1,
  153. )
  154. self.messages.append(final)
  155. return self._build_result(final.content or "", iterations)
  156. def stream(self) -> Iterator[AgentEvent]:
  157. iterations = 0
  158. while iterations < self.max_iterations:
  159. iterations += 1
  160. yield AgentEvent(
  161. type=AgentEventType.THINKING,
  162. data={"iteration": iterations},
  163. )
  164. response = self.llm.chat(
  165. self._all_messages(),
  166. tools=self.tools.definitions or None,
  167. temperature=self.temperature,
  168. iteration=iterations,
  169. )
  170. self.messages.append(response)
  171. if not response.tool_calls:
  172. yield AgentEvent(
  173. type=AgentEventType.MESSAGE,
  174. data={"content": response.content or ""},
  175. )
  176. yield AgentEvent(
  177. type=AgentEventType.DONE,
  178. data=self._build_result(response.content or "", iterations).model_dump(),
  179. )
  180. return
  181. for tc in response.tool_calls:
  182. self.tool_calls_made += 1
  183. yield AgentEvent(
  184. type=AgentEventType.TOOL_CALL,
  185. data={"name": tc.name, "arguments": tc.arguments, "id": tc.id},
  186. )
  187. result = self.tools.execute(tc.id, tc.name, tc.arguments)
  188. if self.logger:
  189. self.logger.log_tool_call(
  190. iterations,
  191. tc.name,
  192. tc.arguments,
  193. result.content,
  194. result.is_error,
  195. tool_call_id=tc.id,
  196. )
  197. yield AgentEvent(
  198. type=AgentEventType.TOOL_RESULT,
  199. data={"name": result.name, "content": result.content, "is_error": result.is_error},
  200. )
  201. self.messages.append(
  202. Message(
  203. role=Role.TOOL,
  204. content=result.content,
  205. tool_call_id=result.tool_call_id,
  206. name=result.name,
  207. )
  208. )
  209. yield AgentEvent(type=AgentEventType.DONE, data={"content": "Max iterations reached"})
  210. async def astream(self) -> AsyncIterator[AgentEvent]:
  211. iterations = 0
  212. while iterations < self.max_iterations:
  213. iterations += 1
  214. yield AgentEvent(
  215. type=AgentEventType.THINKING,
  216. data={"iteration": iterations},
  217. )
  218. response = await self.llm.achat(
  219. self._all_messages(),
  220. tools=self.tools.definitions or None,
  221. temperature=self.temperature,
  222. iteration=iterations,
  223. )
  224. self.messages.append(response)
  225. if not response.tool_calls:
  226. yield AgentEvent(
  227. type=AgentEventType.MESSAGE,
  228. data={"content": response.content or ""},
  229. )
  230. yield AgentEvent(
  231. type=AgentEventType.DONE,
  232. data=self._build_result(response.content or "", iterations).model_dump(),
  233. )
  234. return
  235. for tc in response.tool_calls:
  236. self.tool_calls_made += 1
  237. yield AgentEvent(
  238. type=AgentEventType.TOOL_CALL,
  239. data={"name": tc.name, "arguments": tc.arguments, "id": tc.id},
  240. )
  241. result = await self.tools.aexecute(tc.id, tc.name, tc.arguments)
  242. if self.logger:
  243. self.logger.log_tool_call(
  244. iterations,
  245. tc.name,
  246. tc.arguments,
  247. result.content,
  248. result.is_error,
  249. tool_call_id=tc.id,
  250. )
  251. yield AgentEvent(
  252. type=AgentEventType.TOOL_RESULT,
  253. data={"name": result.name, "content": result.content, "is_error": result.is_error},
  254. )
  255. self.messages.append(
  256. Message(
  257. role=Role.TOOL,
  258. content=result.content,
  259. tool_call_id=result.tool_call_id,
  260. name=result.name,
  261. )
  262. )
  263. yield AgentEvent(type=AgentEventType.DONE, data={"content": "Max iterations reached"})
  264. def _build_result(self, content: str, iterations: int) -> AgentResult:
  265. return AgentResult(
  266. content=content,
  267. messages=self.messages,
  268. iterations=iterations,
  269. tool_calls_made=self.tool_calls_made,
  270. skills_used=list(self.active_skills),
  271. )