runtime.py 5.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162
  1. """Thin ReAct host used by each deterministic workflow node."""
  2. from __future__ import annotations
  3. from collections.abc import Iterable
  4. from typing import Any
  5. from find_agent_v2.observability import InputSlot, ObagentObserver
  6. from find_agent_v2.state import NodeRun
  7. from find_agent_v2.tools import ToolFn, build_tool_registry
  8. from supply_agent import Agent
  9. from supply_agent.config import Settings
  10. class _ObagentEventCollector:
  11. """Collect ReAct events in memory for obagent; never writes project log artifacts."""
  12. def __init__(self) -> None:
  13. self.events: list[dict[str, Any]] = []
  14. def start_run(self, *_args, **_kwargs) -> str:
  15. return ""
  16. def log_llm_input(self, iteration, model, messages, tools, temperature) -> None:
  17. self.events.append({
  18. "type": "llm_input",
  19. "iteration": iteration,
  20. "model": model,
  21. "temperature": temperature,
  22. "messages": [message.model_dump(mode="json") for message in messages],
  23. "tools": [item.model_dump(mode="json") for item in (tools or [])],
  24. })
  25. def log_llm_output(
  26. self, iteration, response, raw_response=None, *, model=None, provider="openrouter",
  27. ) -> None:
  28. usage = getattr(raw_response, "usage", None)
  29. if hasattr(usage, "model_dump"):
  30. usage = usage.model_dump()
  31. self.events.append({
  32. "type": "llm_output",
  33. "iteration": iteration,
  34. "model": model,
  35. "provider": provider,
  36. "response": response.model_dump(mode="json"),
  37. "usage": usage,
  38. })
  39. def log_tool_call(
  40. self, iteration, name, arguments, result, is_error=False, *, tool_call_id=None,
  41. ) -> None:
  42. self.events.append({
  43. "type": "tool_call",
  44. "iteration": iteration,
  45. "name": name,
  46. "arguments": arguments,
  47. "result": result,
  48. "is_error": bool(is_error),
  49. "tool_call_id": tool_call_id,
  50. })
  51. def log_skill_loaded(self, *_args, **_kwargs) -> None:
  52. return None
  53. def end_run(self, *_args, **_kwargs) -> None:
  54. return None
  55. class FindAgentNodeHost:
  56. """Build a fresh, physically capability-limited Agent for every node."""
  57. def __init__(
  58. self,
  59. *,
  60. settings: Settings | None = None,
  61. models_by_role: dict[str, str] | None = None,
  62. default_model: str = "google/gemini-3-flash-preview",
  63. observer: ObagentObserver | None = None,
  64. ) -> None:
  65. self.settings = settings
  66. self.models_by_role = dict(models_by_role or {})
  67. self.default_model = default_model
  68. self.observer = observer or ObagentObserver()
  69. async def run_node(
  70. self,
  71. *,
  72. node: str,
  73. round_index: int,
  74. system_prompt: str,
  75. user_content: str,
  76. tools: Iterable[ToolFn] = (),
  77. max_iterations: int = 12,
  78. slots: tuple[InputSlot, ...] = (),
  79. ) -> NodeRun:
  80. tool_functions = tuple(tools)
  81. model = self.models_by_role.get(node, self.default_model)
  82. event_collector = _ObagentEventCollector()
  83. agent = Agent(
  84. settings=self.settings,
  85. name=f"find_agent_v2.{node}",
  86. model=model,
  87. system_prompt=system_prompt,
  88. tools=build_tool_registry(tool_functions),
  89. max_iterations=max_iterations,
  90. temperature=0.2,
  91. # v2 must not write the project's JSONL/log/OSS visualization artifacts.
  92. logger=event_collector,
  93. )
  94. # Stage capabilities are exact. The generic load_skill tool is not part of this workflow.
  95. agent.tools.unregister("load_skill")
  96. try:
  97. with self.observer.node(node=node) as observation:
  98. try:
  99. actual_user_content = observation.declare(
  100. fallback=user_content,
  101. system_prompt=system_prompt,
  102. slots=slots,
  103. tools=tool_functions,
  104. model=model,
  105. )
  106. result = await agent.arun_core(actual_user_content)
  107. messages = [message.model_dump(mode="json") for message in result.messages]
  108. process = {
  109. "messages": messages,
  110. "events": event_collector.events,
  111. "iterations": result.iterations,
  112. "tool_calls_made": result.tool_calls_made,
  113. }
  114. observation.record_react(output=process, ok=True)
  115. observation.set_output(
  116. {"agent输出": result.content, **process}, ok=True,
  117. )
  118. return NodeRun.from_agent_result(node, round_index, result)
  119. except Exception as exc:
  120. error = {"error": f"{type(exc).__name__}: {exc}"}
  121. observation.record_react(output=error, ok=False)
  122. observation.set_output(error, ok=False)
  123. raise
  124. finally:
  125. client = getattr(agent.llm, "_async_client", None)
  126. if client is not None:
  127. await client.close()
  128. def normalize_models(
  129. *,
  130. model: str | None = None,
  131. planning: str | None = None,
  132. search: str | None = None,
  133. evidence: str | None = None,
  134. evaluation: str | None = None,
  135. report: str | None = None,
  136. ) -> dict[str, str]:
  137. base = model or "google/gemini-3-flash-preview"
  138. return {
  139. "planner": planning or base,
  140. "search": search or base,
  141. "evidence": evidence or base,
  142. "evaluator": evaluation or base,
  143. "report": report or base,
  144. }