graph.py 6.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164
  1. """One complete business round as a deterministic DAG.
  2. The project does not depend on LangGraph, so this module preserves the same boundary
  3. without adding a second Agent runtime: one invocation is one acyclic business round.
  4. """
  5. from __future__ import annotations
  6. from typing import Protocol
  7. from find_agent_v2.context import build_node_slots, render_node_context
  8. from find_agent_v2.observability import InputSlot, NullObserver
  9. from find_agent_v2.prompts import EVALUATOR_PROMPT, EVIDENCE_PROMPT, PLANNER_PROMPT, SEARCH_PROMPT
  10. from find_agent_v2.service import FindAgentV2Service
  11. from find_agent_v2.state import FindAgentState, NodeRun
  12. from find_agent_v2.tools import EVALUATION_TOOLS, EVIDENCE_TOOLS, SEARCH_TOOLS, ToolFn
  13. class NodeRunner(Protocol):
  14. async def run_node(
  15. self,
  16. *,
  17. node: str,
  18. round_index: int,
  19. system_prompt: str,
  20. user_content: str,
  21. tools: tuple[ToolFn, ...] = (),
  22. max_iterations: int = 12,
  23. slots: tuple[InputSlot, ...] = (),
  24. ) -> NodeRun: ...
  25. class FindAgentRoundGraph:
  26. """planner -> search -> evidence -> evaluator -> END."""
  27. def __init__(
  28. self, *, service: FindAgentV2Service, runner: NodeRunner,
  29. observer=None,
  30. ) -> None:
  31. self.service = service
  32. self.runner = runner
  33. self.observer = observer or NullObserver()
  34. def _full_state(self, state: FindAgentState, *, pending_only: bool = False):
  35. return self.service.get_full_state(
  36. state.run_id, pending_only=pending_only,
  37. )
  38. def _context(self, state: FindAgentState, *, pending_only: bool = False) -> str:
  39. return render_node_context(
  40. user_input=state.user_input,
  41. full_state=self._full_state(state, pending_only=pending_only),
  42. round_index=state.round_index,
  43. plan=state.plan,
  44. )
  45. def _slots(
  46. self, state: FindAgentState, *, pending_only: bool = False,
  47. ) -> tuple[InputSlot, ...]:
  48. return build_node_slots(
  49. user_input=state.user_input,
  50. full_state=self._full_state(state, pending_only=pending_only),
  51. round_index=state.round_index,
  52. plan=state.plan,
  53. )
  54. async def invoke(self, state: FindAgentState) -> FindAgentState:
  55. with self.observer.round(round_index=state.round_index) as round_observation:
  56. state = await self._invoke_nodes(state)
  57. round_observation.set_output({
  58. "状态快照": state.snapshot.__dict__ if state.snapshot else {},
  59. "本轮计划": state.plan,
  60. }, ok=True)
  61. return state
  62. async def _invoke_nodes(self, state: FindAgentState) -> FindAgentState:
  63. state.phase = "planning"
  64. plan = await self.runner.run_node(
  65. node="planner",
  66. round_index=state.round_index,
  67. system_prompt=PLANNER_PROMPT,
  68. user_content=self._context(state),
  69. tools=(),
  70. max_iterations=2,
  71. slots=self._slots(state),
  72. )
  73. state.node_runs.append(plan)
  74. state.plan = plan.content.strip()
  75. self.service.update_round(
  76. state.run_id, state.round_index, phase="searching", plan=state.plan,
  77. )
  78. state.phase = "searching"
  79. search = await self.runner.run_node(
  80. node="search",
  81. round_index=state.round_index,
  82. system_prompt=SEARCH_PROMPT,
  83. user_content=self._context(state),
  84. tools=SEARCH_TOOLS,
  85. max_iterations=10,
  86. slots=self._slots(state),
  87. )
  88. state.node_runs.append(search)
  89. after_search = self.service.snapshot(state.run_id)
  90. if after_search.pending_count:
  91. state.phase = "evidence"
  92. self.service.update_round(state.run_id, state.round_index, phase="evidence")
  93. evidence = await self.runner.run_node(
  94. node="evidence",
  95. round_index=state.round_index,
  96. system_prompt=EVIDENCE_PROMPT,
  97. user_content=self._context(state, pending_only=True),
  98. tools=EVIDENCE_TOOLS,
  99. max_iterations=12,
  100. slots=self._slots(state, pending_only=True),
  101. )
  102. state.node_runs.append(evidence)
  103. state.phase = "evaluating"
  104. self.service.update_round(state.run_id, state.round_index, phase="evaluating")
  105. # A model may intentionally keep one tool payload small (for example,
  106. # evaluate ten candidates at a time). One evaluator invocation is
  107. # therefore not proof that the queue has been drained. Re-enter the
  108. # node with a fresh DB snapshot until every candidate has a terminal
  109. # bucket, while failing fast if an invocation makes no progress.
  110. stagnant_attempts = 0
  111. for _ in range(64):
  112. before_evaluation = self.service.snapshot(state.run_id)
  113. if not before_evaluation.pending_count:
  114. break
  115. evaluation = await self.runner.run_node(
  116. node="evaluator",
  117. round_index=state.round_index,
  118. system_prompt=EVALUATOR_PROMPT,
  119. user_content=self._context(state, pending_only=True),
  120. tools=EVALUATION_TOOLS,
  121. max_iterations=12,
  122. slots=self._slots(state, pending_only=True),
  123. )
  124. state.node_runs.append(evaluation)
  125. after_evaluation = self.service.snapshot(state.run_id)
  126. if after_evaluation.pending_count >= before_evaluation.pending_count:
  127. stagnant_attempts += 1
  128. if stagnant_attempts >= 3:
  129. raise RuntimeError(
  130. "评估节点连续 3 次未消费 pending_evaluation 候选:"
  131. f"remaining={after_evaluation.pending_count}"
  132. )
  133. else:
  134. stagnant_attempts = 0
  135. else:
  136. raise RuntimeError("评估批次超过安全上限 64")
  137. state.snapshot = self.service.snapshot(state.run_id)
  138. state.phase = "done"
  139. self.service.update_round(
  140. state.run_id,
  141. state.round_index,
  142. phase="done",
  143. status="done",
  144. snapshot=state.snapshot,
  145. )
  146. return state