loop.py 10 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280
  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, tc.name, tc.arguments, result.content, result.is_error
  75. )
  76. if tc.name == "load_skill" and not result.is_error:
  77. self._on_skill_loaded(tc.arguments)
  78. if self.logger:
  79. self.logger.log_skill_loaded(iterations, tc.arguments)
  80. self.messages.append(
  81. Message(
  82. role=Role.TOOL,
  83. content=result.content,
  84. tool_call_id=result.tool_call_id,
  85. name=result.name,
  86. )
  87. )
  88. self.messages.append(
  89. Message(
  90. role=Role.USER,
  91. content="Maximum iterations reached. Please provide your best answer now.",
  92. )
  93. )
  94. final = self.llm.chat(
  95. self._all_messages(),
  96. temperature=self.temperature,
  97. iteration=iterations + 1,
  98. )
  99. self.messages.append(final)
  100. return self._build_result(final.content or "", iterations)
  101. async def arun(self) -> AgentResult:
  102. iterations = 0
  103. while iterations < self.max_iterations:
  104. iterations += 1
  105. response = await self.llm.achat(
  106. self._all_messages(),
  107. tools=self.tools.definitions or None,
  108. temperature=self.temperature,
  109. iteration=iterations,
  110. )
  111. self.messages.append(response)
  112. if not response.tool_calls:
  113. return self._build_result(response.content or "", iterations)
  114. for tc in response.tool_calls:
  115. self.tool_calls_made += 1
  116. result = await self.tools.aexecute(tc.id, tc.name, tc.arguments)
  117. if self.logger:
  118. self.logger.log_tool_call(
  119. iterations, tc.name, tc.arguments, result.content, result.is_error
  120. )
  121. if tc.name == "load_skill" and not result.is_error:
  122. self._on_skill_loaded(tc.arguments)
  123. if self.logger:
  124. self.logger.log_skill_loaded(iterations, tc.arguments)
  125. self.messages.append(
  126. Message(
  127. role=Role.TOOL,
  128. content=result.content,
  129. tool_call_id=result.tool_call_id,
  130. name=result.name,
  131. )
  132. )
  133. self.messages.append(
  134. Message(
  135. role=Role.USER,
  136. content="Maximum iterations reached. Please provide your best answer now.",
  137. )
  138. )
  139. final = await self.llm.achat(
  140. self._all_messages(),
  141. temperature=self.temperature,
  142. iteration=iterations + 1,
  143. )
  144. self.messages.append(final)
  145. return self._build_result(final.content or "", iterations)
  146. def stream(self) -> Iterator[AgentEvent]:
  147. iterations = 0
  148. while iterations < self.max_iterations:
  149. iterations += 1
  150. yield AgentEvent(
  151. type=AgentEventType.THINKING,
  152. data={"iteration": iterations},
  153. )
  154. response = self.llm.chat(
  155. self._all_messages(),
  156. tools=self.tools.definitions or None,
  157. temperature=self.temperature,
  158. iteration=iterations,
  159. )
  160. self.messages.append(response)
  161. if not response.tool_calls:
  162. yield AgentEvent(
  163. type=AgentEventType.MESSAGE,
  164. data={"content": response.content or ""},
  165. )
  166. yield AgentEvent(
  167. type=AgentEventType.DONE,
  168. data=self._build_result(response.content or "", iterations).model_dump(),
  169. )
  170. return
  171. for tc in response.tool_calls:
  172. self.tool_calls_made += 1
  173. yield AgentEvent(
  174. type=AgentEventType.TOOL_CALL,
  175. data={"name": tc.name, "arguments": tc.arguments},
  176. )
  177. result = self.tools.execute(tc.id, tc.name, tc.arguments)
  178. if self.logger:
  179. self.logger.log_tool_call(
  180. iterations, tc.name, tc.arguments, result.content, result.is_error
  181. )
  182. yield AgentEvent(
  183. type=AgentEventType.TOOL_RESULT,
  184. data={"name": result.name, "content": result.content, "is_error": result.is_error},
  185. )
  186. self.messages.append(
  187. Message(
  188. role=Role.TOOL,
  189. content=result.content,
  190. tool_call_id=result.tool_call_id,
  191. name=result.name,
  192. )
  193. )
  194. yield AgentEvent(type=AgentEventType.DONE, data={"content": "Max iterations reached"})
  195. async def astream(self) -> AsyncIterator[AgentEvent]:
  196. iterations = 0
  197. while iterations < self.max_iterations:
  198. iterations += 1
  199. yield AgentEvent(
  200. type=AgentEventType.THINKING,
  201. data={"iteration": iterations},
  202. )
  203. response = await self.llm.achat(
  204. self._all_messages(),
  205. tools=self.tools.definitions or None,
  206. temperature=self.temperature,
  207. iteration=iterations,
  208. )
  209. self.messages.append(response)
  210. if not response.tool_calls:
  211. yield AgentEvent(
  212. type=AgentEventType.MESSAGE,
  213. data={"content": response.content or ""},
  214. )
  215. yield AgentEvent(
  216. type=AgentEventType.DONE,
  217. data=self._build_result(response.content or "", iterations).model_dump(),
  218. )
  219. return
  220. for tc in response.tool_calls:
  221. self.tool_calls_made += 1
  222. yield AgentEvent(
  223. type=AgentEventType.TOOL_CALL,
  224. data={"name": tc.name, "arguments": tc.arguments},
  225. )
  226. result = await self.tools.aexecute(tc.id, tc.name, tc.arguments)
  227. if self.logger:
  228. self.logger.log_tool_call(
  229. iterations, tc.name, tc.arguments, result.content, result.is_error
  230. )
  231. yield AgentEvent(
  232. type=AgentEventType.TOOL_RESULT,
  233. data={"name": result.name, "content": result.content, "is_error": result.is_error},
  234. )
  235. self.messages.append(
  236. Message(
  237. role=Role.TOOL,
  238. content=result.content,
  239. tool_call_id=result.tool_call_id,
  240. name=result.name,
  241. )
  242. )
  243. yield AgentEvent(type=AgentEventType.DONE, data={"content": "Max iterations reached"})
  244. def _build_result(self, content: str, iterations: int) -> AgentResult:
  245. return AgentResult(
  246. content=content,
  247. messages=self.messages,
  248. iterations=iterations,
  249. tool_calls_made=self.tool_calls_made,
  250. skills_used=list(self.active_skills),
  251. )