loop.py 23 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601
  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, MalformedFunctionCallError
  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. ToolCall,
  14. ToolDefinition,
  15. ToolResult,
  16. )
  17. if TYPE_CHECKING:
  18. from supply_agent.logging.logger import AgentLogger
  19. class AgentLoop:
  20. """
  21. ReAct-style agent loop: Reason → Act (tool call) → Observe → Repeat.
  22. Implements the standard tool-calling pattern used by modern agent frameworks.
  23. """
  24. def __init__(
  25. self,
  26. llm: LLMClient,
  27. tools: ToolRegistry,
  28. system_message: Message,
  29. messages: list[Message],
  30. max_iterations: int = 20,
  31. temperature: float | None = None,
  32. logger: AgentLogger | None = None,
  33. active_skills: list[str] | None = None,
  34. system_message_builder: Callable[[], Message] | None = None,
  35. completion_guard: Callable[[list[Message]], str | None] | None = None,
  36. completion_guard_blocks: bool = False,
  37. tool_call_budgets: dict[str, tuple[set[str], int]] | None = None,
  38. tool_repeat_requires_change: dict[str, set[str]] | None = None,
  39. ) -> None:
  40. self.llm = llm
  41. self.tools = tools
  42. self.system_message = system_message
  43. self.system_message_builder = system_message_builder
  44. self.messages = messages
  45. self.max_iterations = max_iterations
  46. self.temperature = temperature
  47. self.logger = logger
  48. self.active_skills = active_skills or []
  49. self.completion_guard = completion_guard
  50. self.completion_guard_blocks = completion_guard_blocks
  51. self.tool_call_budgets = tool_call_budgets or {}
  52. self.tool_repeat_requires_change = tool_repeat_requires_change or {}
  53. self._tool_budget_counts = {
  54. group: 0 for group in self.tool_call_budgets
  55. }
  56. self.tool_calls_made = 0
  57. def _all_messages(self) -> list[Message]:
  58. return [self.system_message, *self.messages]
  59. def _on_skill_loaded(self, arguments: str) -> None:
  60. """Refresh system message after a skill is loaded."""
  61. try:
  62. args = json.loads(arguments)
  63. skill_name = args.get("name", "")
  64. if skill_name and skill_name not in self.active_skills:
  65. self.active_skills.append(skill_name)
  66. except json.JSONDecodeError:
  67. pass
  68. if self.system_message_builder:
  69. self.system_message = self.system_message_builder()
  70. def _completion_block_reason(self) -> str | None:
  71. if self.completion_guard is None:
  72. return None
  73. return self.completion_guard(self.messages)
  74. def _should_block_completion(self, reason: str | None) -> bool:
  75. return bool(reason and self.completion_guard_blocks)
  76. def _completion_response_content(self, content: str) -> str:
  77. reason = self._completion_block_reason()
  78. if reason and not self.completion_guard_blocks:
  79. return f"{content}\n\n---\n完成提示(未阻断结束):{reason}"
  80. return content
  81. def _append_completion_feedback(self, reason: str) -> None:
  82. self.messages.append(
  83. Message(
  84. role=Role.USER,
  85. content=(
  86. "完成守卫拒绝当前最终回答:"
  87. f"{reason}。请继续调用必要工具修复这些问题;"
  88. "只有守卫条件全部满足后才能输出最终答案。"
  89. ),
  90. )
  91. )
  92. def _append_malformed_tool_feedback(self) -> None:
  93. self.messages.append(
  94. Message(
  95. role=Role.USER,
  96. content=(
  97. "上一步连续生成了无效工具参数。请继续任务,下一步只调用一个"
  98. "最必要的工具,严格使用其 schema,不得添加未定义字段。"
  99. ),
  100. )
  101. )
  102. def _consume_tool_budget(self, tool_name: str) -> str | None:
  103. matching = [
  104. (group, limit)
  105. for group, (tool_names, limit) in self.tool_call_budgets.items()
  106. if tool_name in tool_names
  107. ]
  108. for group, limit in matching:
  109. used = self._tool_budget_counts[group]
  110. if used >= limit:
  111. return (
  112. f"工具预算已耗尽: group={group}, limit={limit}。"
  113. "请使用已有结果完成筛选、持久化和审计,不要继续扩大搜索。"
  114. )
  115. for group, _ in matching:
  116. self._tool_budget_counts[group] += 1
  117. return None
  118. def _repeat_policy_error(self, tool_name: str) -> str | None:
  119. dependencies = self.tool_repeat_requires_change.get(tool_name)
  120. if not dependencies:
  121. return None
  122. last_call_index = -1
  123. last_change_index = -1
  124. for index, message in enumerate(self.messages):
  125. if message.role != Role.TOOL or not message.name:
  126. continue
  127. try:
  128. payload = json.loads(message.content or "")
  129. except (TypeError, json.JSONDecodeError):
  130. continue
  131. if isinstance(payload, dict) and payload.get("error"):
  132. continue
  133. if message.name == tool_name:
  134. last_call_index = index
  135. if message.name in dependencies:
  136. last_change_index = index
  137. if last_call_index >= 0 and last_change_index < last_call_index:
  138. return (
  139. f"工具 {tool_name} 在状态未变化时不得重复调用。"
  140. f"必须先成功执行以下任一状态变更工具: {sorted(dependencies)}"
  141. )
  142. return None
  143. def _available_tool_definitions(self) -> list[ToolDefinition] | None:
  144. exhausted_tools: set[str] = set()
  145. for group, (tool_names, limit) in self.tool_call_budgets.items():
  146. if self._tool_budget_counts[group] >= limit:
  147. exhausted_tools.update(tool_names)
  148. definitions = [
  149. definition
  150. for definition in self.tools.definitions
  151. if definition.name not in exhausted_tools
  152. and self._repeat_policy_error(definition.name) is None
  153. ]
  154. return definitions or None
  155. def _refund_tool_budget_for_input_error(
  156. self,
  157. tool_name: str,
  158. result: ToolResult,
  159. ) -> None:
  160. try:
  161. payload = json.loads(result.content)
  162. except (TypeError, json.JSONDecodeError):
  163. return
  164. if not isinstance(payload, dict) or payload.get("input_error") is not True:
  165. return
  166. for group, (tool_names, _) in self.tool_call_budgets.items():
  167. if tool_name in tool_names and self._tool_budget_counts[group] > 0:
  168. self._tool_budget_counts[group] -= 1
  169. @staticmethod
  170. def _budget_error_result(tool_call: ToolCall, error: str) -> ToolResult:
  171. return ToolResult(
  172. tool_call_id=tool_call.id,
  173. name=tool_call.name,
  174. content=json.dumps(
  175. {
  176. "error": error,
  177. "budget_exhausted": True,
  178. "title": "工具预算拒绝调用",
  179. },
  180. ensure_ascii=False,
  181. ),
  182. is_error=True,
  183. )
  184. @staticmethod
  185. def _policy_error_result(tool_call: ToolCall, error: str) -> ToolResult:
  186. return ToolResult(
  187. tool_call_id=tool_call.id,
  188. name=tool_call.name,
  189. content=json.dumps(
  190. {
  191. "error": error,
  192. "policy_rejected": True,
  193. "title": "工具状态策略拒绝调用",
  194. },
  195. ensure_ascii=False,
  196. ),
  197. is_error=True,
  198. )
  199. def _max_iterations_result(self, iterations: int) -> AgentResult:
  200. reason = self._completion_block_reason()
  201. content = "Max iterations reached"
  202. if reason:
  203. content = f"Max iterations reached before completion: {reason}"
  204. return self._build_result(content, iterations)
  205. def _iterations_exhausted_result(self, iterations: int) -> AgentResult:
  206. if self.completion_guard is not None and self.completion_guard_blocks:
  207. return self._max_iterations_result(iterations)
  208. last_assistant = next(
  209. (
  210. message
  211. for message in reversed(self.messages)
  212. if message.role == Role.ASSISTANT and not message.tool_calls
  213. ),
  214. None,
  215. )
  216. content = (
  217. (last_assistant.content or "").strip()
  218. if last_assistant and (last_assistant.content or "").strip()
  219. else "Max iterations reached"
  220. )
  221. return self._build_result(
  222. self._completion_response_content(content),
  223. iterations,
  224. )
  225. def _uses_blocking_completion_guard(self) -> bool:
  226. return self.completion_guard is not None and self.completion_guard_blocks
  227. def run(self) -> AgentResult:
  228. iterations = 0
  229. while iterations < self.max_iterations:
  230. iterations += 1
  231. try:
  232. response = self.llm.chat(
  233. self._all_messages(),
  234. tools=self._available_tool_definitions(),
  235. temperature=self.temperature,
  236. iteration=iterations,
  237. )
  238. except MalformedFunctionCallError:
  239. self._append_malformed_tool_feedback()
  240. continue
  241. self.messages.append(response)
  242. if not response.tool_calls:
  243. reason = self._completion_block_reason()
  244. if self._should_block_completion(reason):
  245. self._append_completion_feedback(reason)
  246. continue
  247. return self._build_result(
  248. self._completion_response_content(response.content or ""),
  249. iterations,
  250. )
  251. for tc in response.tool_calls:
  252. self.tool_calls_made += 1
  253. policy_error = self._repeat_policy_error(tc.name)
  254. budget_error = (
  255. None if policy_error else self._consume_tool_budget(tc.name)
  256. )
  257. result = (
  258. self._policy_error_result(tc, policy_error)
  259. if policy_error
  260. else (
  261. self._budget_error_result(tc, budget_error)
  262. if budget_error
  263. else self.tools.execute(tc.id, tc.name, tc.arguments)
  264. )
  265. )
  266. if not policy_error and not budget_error:
  267. self._refund_tool_budget_for_input_error(tc.name, result)
  268. if self.logger:
  269. self.logger.log_tool_call(
  270. iterations,
  271. tc.name,
  272. tc.arguments,
  273. result.content,
  274. result.is_error,
  275. tool_call_id=tc.id,
  276. )
  277. if tc.name == "load_skill" and not result.is_error:
  278. self._on_skill_loaded(tc.arguments)
  279. if self.logger:
  280. self.logger.log_skill_loaded(iterations, tc.arguments)
  281. self.messages.append(
  282. Message(
  283. role=Role.TOOL,
  284. content=result.content,
  285. tool_call_id=result.tool_call_id,
  286. name=result.name,
  287. )
  288. )
  289. if self._uses_blocking_completion_guard():
  290. return self._max_iterations_result(iterations)
  291. self.messages.append(
  292. Message(
  293. role=Role.USER,
  294. content="Maximum iterations reached. Please provide your best answer now.",
  295. )
  296. )
  297. final = self.llm.chat(
  298. self._all_messages(),
  299. temperature=self.temperature,
  300. iteration=iterations + 1,
  301. )
  302. self.messages.append(final)
  303. return self._build_result(final.content or "", iterations)
  304. async def arun(self) -> AgentResult:
  305. iterations = 0
  306. while iterations < self.max_iterations:
  307. iterations += 1
  308. try:
  309. response = await self.llm.achat(
  310. self._all_messages(),
  311. tools=self._available_tool_definitions(),
  312. temperature=self.temperature,
  313. iteration=iterations,
  314. )
  315. except MalformedFunctionCallError:
  316. self._append_malformed_tool_feedback()
  317. continue
  318. self.messages.append(response)
  319. if not response.tool_calls:
  320. reason = self._completion_block_reason()
  321. if self._should_block_completion(reason):
  322. self._append_completion_feedback(reason)
  323. continue
  324. return self._build_result(
  325. self._completion_response_content(response.content or ""),
  326. iterations,
  327. )
  328. for tc in response.tool_calls:
  329. self.tool_calls_made += 1
  330. policy_error = self._repeat_policy_error(tc.name)
  331. budget_error = (
  332. None if policy_error else self._consume_tool_budget(tc.name)
  333. )
  334. result = (
  335. self._policy_error_result(tc, policy_error)
  336. if policy_error
  337. else (
  338. self._budget_error_result(tc, budget_error)
  339. if budget_error
  340. else await self.tools.aexecute(
  341. tc.id, tc.name, tc.arguments
  342. )
  343. )
  344. )
  345. if not policy_error and not budget_error:
  346. self._refund_tool_budget_for_input_error(tc.name, result)
  347. if self.logger:
  348. self.logger.log_tool_call(
  349. iterations,
  350. tc.name,
  351. tc.arguments,
  352. result.content,
  353. result.is_error,
  354. tool_call_id=tc.id,
  355. )
  356. if tc.name == "load_skill" and not result.is_error:
  357. self._on_skill_loaded(tc.arguments)
  358. if self.logger:
  359. self.logger.log_skill_loaded(iterations, tc.arguments)
  360. self.messages.append(
  361. Message(
  362. role=Role.TOOL,
  363. content=result.content,
  364. tool_call_id=result.tool_call_id,
  365. name=result.name,
  366. )
  367. )
  368. if self._uses_blocking_completion_guard():
  369. return self._max_iterations_result(iterations)
  370. self.messages.append(
  371. Message(
  372. role=Role.USER,
  373. content="Maximum iterations reached. Please provide your best answer now.",
  374. )
  375. )
  376. final = await self.llm.achat(
  377. self._all_messages(),
  378. temperature=self.temperature,
  379. iteration=iterations + 1,
  380. )
  381. self.messages.append(final)
  382. return self._build_result(final.content or "", iterations)
  383. def stream(self) -> Iterator[AgentEvent]:
  384. iterations = 0
  385. while iterations < self.max_iterations:
  386. iterations += 1
  387. yield AgentEvent(
  388. type=AgentEventType.THINKING,
  389. data={"iteration": iterations},
  390. )
  391. try:
  392. response = self.llm.chat(
  393. self._all_messages(),
  394. tools=self._available_tool_definitions(),
  395. temperature=self.temperature,
  396. iteration=iterations,
  397. )
  398. except MalformedFunctionCallError:
  399. self._append_malformed_tool_feedback()
  400. continue
  401. self.messages.append(response)
  402. if not response.tool_calls:
  403. reason = self._completion_block_reason()
  404. if self._should_block_completion(reason):
  405. self._append_completion_feedback(reason)
  406. continue
  407. final_content = self._completion_response_content(
  408. response.content or ""
  409. )
  410. yield AgentEvent(
  411. type=AgentEventType.MESSAGE,
  412. data={"content": final_content},
  413. )
  414. yield AgentEvent(
  415. type=AgentEventType.DONE,
  416. data=self._build_result(final_content, iterations).model_dump(),
  417. )
  418. return
  419. for tc in response.tool_calls:
  420. self.tool_calls_made += 1
  421. yield AgentEvent(
  422. type=AgentEventType.TOOL_CALL,
  423. data={"name": tc.name, "arguments": tc.arguments, "id": tc.id},
  424. )
  425. policy_error = self._repeat_policy_error(tc.name)
  426. budget_error = (
  427. None if policy_error else self._consume_tool_budget(tc.name)
  428. )
  429. result = (
  430. self._policy_error_result(tc, policy_error)
  431. if policy_error
  432. else (
  433. self._budget_error_result(tc, budget_error)
  434. if budget_error
  435. else self.tools.execute(tc.id, tc.name, tc.arguments)
  436. )
  437. )
  438. if not policy_error and not budget_error:
  439. self._refund_tool_budget_for_input_error(tc.name, result)
  440. if self.logger:
  441. self.logger.log_tool_call(
  442. iterations,
  443. tc.name,
  444. tc.arguments,
  445. result.content,
  446. result.is_error,
  447. tool_call_id=tc.id,
  448. )
  449. yield AgentEvent(
  450. type=AgentEventType.TOOL_RESULT,
  451. data={"name": result.name, "content": result.content, "is_error": result.is_error},
  452. )
  453. self.messages.append(
  454. Message(
  455. role=Role.TOOL,
  456. content=result.content,
  457. tool_call_id=result.tool_call_id,
  458. name=result.name,
  459. )
  460. )
  461. yield AgentEvent(
  462. type=AgentEventType.DONE,
  463. data=self._iterations_exhausted_result(iterations).model_dump(),
  464. )
  465. async def astream(self) -> AsyncIterator[AgentEvent]:
  466. iterations = 0
  467. while iterations < self.max_iterations:
  468. iterations += 1
  469. yield AgentEvent(
  470. type=AgentEventType.THINKING,
  471. data={"iteration": iterations},
  472. )
  473. try:
  474. response = await self.llm.achat(
  475. self._all_messages(),
  476. tools=self._available_tool_definitions(),
  477. temperature=self.temperature,
  478. iteration=iterations,
  479. )
  480. except MalformedFunctionCallError:
  481. self._append_malformed_tool_feedback()
  482. continue
  483. self.messages.append(response)
  484. if not response.tool_calls:
  485. reason = self._completion_block_reason()
  486. if self._should_block_completion(reason):
  487. self._append_completion_feedback(reason)
  488. continue
  489. final_content = self._completion_response_content(
  490. response.content or ""
  491. )
  492. yield AgentEvent(
  493. type=AgentEventType.MESSAGE,
  494. data={"content": final_content},
  495. )
  496. yield AgentEvent(
  497. type=AgentEventType.DONE,
  498. data=self._build_result(final_content, iterations).model_dump(),
  499. )
  500. return
  501. for tc in response.tool_calls:
  502. self.tool_calls_made += 1
  503. yield AgentEvent(
  504. type=AgentEventType.TOOL_CALL,
  505. data={"name": tc.name, "arguments": tc.arguments, "id": tc.id},
  506. )
  507. policy_error = self._repeat_policy_error(tc.name)
  508. budget_error = (
  509. None if policy_error else self._consume_tool_budget(tc.name)
  510. )
  511. result = (
  512. self._policy_error_result(tc, policy_error)
  513. if policy_error
  514. else (
  515. self._budget_error_result(tc, budget_error)
  516. if budget_error
  517. else await self.tools.aexecute(
  518. tc.id, tc.name, tc.arguments
  519. )
  520. )
  521. )
  522. if not policy_error and not budget_error:
  523. self._refund_tool_budget_for_input_error(tc.name, result)
  524. if self.logger:
  525. self.logger.log_tool_call(
  526. iterations,
  527. tc.name,
  528. tc.arguments,
  529. result.content,
  530. result.is_error,
  531. tool_call_id=tc.id,
  532. )
  533. yield AgentEvent(
  534. type=AgentEventType.TOOL_RESULT,
  535. data={"name": result.name, "content": result.content, "is_error": result.is_error},
  536. )
  537. self.messages.append(
  538. Message(
  539. role=Role.TOOL,
  540. content=result.content,
  541. tool_call_id=result.tool_call_id,
  542. name=result.name,
  543. )
  544. )
  545. yield AgentEvent(
  546. type=AgentEventType.DONE,
  547. data=self._iterations_exhausted_result(iterations).model_dump(),
  548. )
  549. def _build_result(self, content: str, iterations: int) -> AgentResult:
  550. return AgentResult(
  551. content=content,
  552. messages=self.messages,
  553. iterations=iterations,
  554. tool_calls_made=self.tool_calls_made,
  555. skills_used=list(self.active_skills),
  556. )