runtime.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302
  1. """Independent LangChain ReAct host for find_agent_v2.
  2. The host owns model construction, retries, usage accounting, tool adaptation and
  3. controlled nested-agent delegation. It never imports the legacy find-agent package.
  4. """
  5. from __future__ import annotations
  6. import asyncio
  7. import json
  8. import os
  9. from collections.abc import Iterable
  10. from typing import Any
  11. from langchain.agents import create_agent
  12. from langchain.agents.middleware import ModelFallbackMiddleware, ModelRetryMiddleware, ToolRetryMiddleware
  13. from langchain_core.messages import AIMessage, BaseMessage, ToolMessage
  14. from langchain_core.tools import StructuredTool
  15. from langchain_openai import ChatOpenAI
  16. from pydantic import BaseModel, Field
  17. from find_agent_v2.observability import InputSlot, ObagentObserver
  18. from find_agent_v2.state import NodeRun
  19. from find_agent_v2.tools import ToolFn
  20. from supply_agent.config import Settings, get_settings
  21. class DelegateRequest(BaseModel):
  22. task: str = Field(description="交给子 Agent 的完整、独立任务")
  23. class DelegateArgs(BaseModel):
  24. requests: list[DelegateRequest] = Field(
  25. min_length=1, max_length=8,
  26. description="可并发执行的子 Agent 任务;有依赖的任务不要放在同一批",
  27. )
  28. def _tool_name(fn: ToolFn) -> str:
  29. return str(getattr(fn, "_tool_name", fn.__name__))
  30. def _as_langchain_tool(fn: ToolFn) -> StructuredTool:
  31. kwargs = {
  32. "name": _tool_name(fn),
  33. "description": str(getattr(fn, "_tool_description", fn.__doc__ or "")),
  34. }
  35. if asyncio.iscoroutinefunction(fn):
  36. return StructuredTool.from_function(coroutine=fn, **kwargs)
  37. return StructuredTool.from_function(func=fn, **kwargs)
  38. def _message_dict(message: BaseMessage) -> dict[str, Any]:
  39. data = message.model_dump(mode="json")
  40. data["type"] = message.type
  41. return data
  42. def _usage(messages: list[BaseMessage]) -> dict[str, int | float]:
  43. totals: dict[str, int | float] = {
  44. "input_tokens": 0, "output_tokens": 0, "total_tokens": 0, "cost": 0.0,
  45. }
  46. for message in messages:
  47. usage = getattr(message, "usage_metadata", None) or {}
  48. totals["input_tokens"] += int(usage.get("input_tokens") or 0)
  49. totals["output_tokens"] += int(usage.get("output_tokens") or 0)
  50. totals["total_tokens"] += int(usage.get("total_tokens") or 0)
  51. response = getattr(message, "response_metadata", None) or {}
  52. cost = response.get("cost") or (response.get("usage") or {}).get("cost")
  53. if cost is not None:
  54. totals["cost"] += float(cost)
  55. totals["cost"] = round(float(totals["cost"]), 8)
  56. return totals
  57. def _events(messages: list[BaseMessage]) -> list[dict[str, Any]]:
  58. events: list[dict[str, Any]] = []
  59. for message in messages:
  60. if isinstance(message, AIMessage):
  61. events.append({
  62. "type": "llm_output",
  63. "content": message.content,
  64. "tool_calls": message.tool_calls,
  65. "usage": message.usage_metadata,
  66. })
  67. elif isinstance(message, ToolMessage):
  68. events.append({
  69. "type": "tool_call",
  70. "name": message.name,
  71. "tool_call_id": message.tool_call_id,
  72. "result": message.content,
  73. "status": message.status,
  74. })
  75. return events
  76. class FindAgentNodeHost:
  77. """Create a fresh LangChain Agent for every workflow or delegated node."""
  78. def __init__(
  79. self,
  80. *,
  81. settings: Settings | None = None,
  82. models_by_role: dict[str, str] | None = None,
  83. default_model: str = "google/gemini-3-flash-preview",
  84. observer: ObagentObserver | None = None,
  85. ) -> None:
  86. self.settings = settings or get_settings()
  87. self.models_by_role = dict(models_by_role or {})
  88. self.default_model = default_model
  89. self.observer = observer or ObagentObserver()
  90. self.usage = {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0, "cost": 0.0}
  91. def reset_usage(self) -> None:
  92. self.usage = {
  93. "input_tokens": 0,
  94. "output_tokens": 0,
  95. "total_tokens": 0,
  96. "cost": 0.0,
  97. }
  98. def _model(self, role: str) -> ChatOpenAI:
  99. model_name = self.models_by_role.get(role, self.default_model)
  100. return ChatOpenAI(
  101. model=model_name,
  102. api_key=self.settings.openrouter_api_key,
  103. base_url=self.settings.openrouter_base_url,
  104. timeout=self.settings.openrouter_timeout_seconds,
  105. temperature=0.2,
  106. max_retries=0, # retries are observable LangChain middleware below
  107. default_headers={
  108. "HTTP-Referer": self.settings.openrouter_site_url,
  109. "X-Title": self.settings.openrouter_site_name,
  110. },
  111. )
  112. def _middleware(self, role: str):
  113. items: list[Any] = [
  114. ModelRetryMiddleware(max_retries=2, on_failure="error"),
  115. ToolRetryMiddleware(max_retries=2, on_failure="continue"),
  116. ]
  117. fallback = os.getenv("FIND_AGENT_V2_FALLBACK_MODEL", "").strip()
  118. if fallback and fallback != self.models_by_role.get(role, self.default_model):
  119. items.insert(0, ModelFallbackMiddleware(self._model_for_name(fallback)))
  120. return items
  121. def _model_for_name(self, model_name: str) -> ChatOpenAI:
  122. return ChatOpenAI(
  123. model=model_name,
  124. api_key=self.settings.openrouter_api_key,
  125. base_url=self.settings.openrouter_base_url,
  126. timeout=self.settings.openrouter_timeout_seconds,
  127. temperature=0.2,
  128. max_retries=0,
  129. )
  130. def _delegate_tool(
  131. self,
  132. *,
  133. parent_node: str,
  134. round_index: int,
  135. system_prompt: str,
  136. tools: tuple[ToolFn, ...],
  137. max_iterations: int,
  138. slots: tuple[InputSlot, ...],
  139. ) -> StructuredTool:
  140. async def delegate_agents_v2(requests: list[DelegateRequest]) -> str:
  141. """并发委派多个相互独立的任务给同阶段子 Agent。"""
  142. normalized = [
  143. item if isinstance(item, DelegateRequest) else DelegateRequest.model_validate(item)
  144. for item in requests
  145. ]
  146. async def run_one(index: int, request: DelegateRequest) -> dict[str, Any]:
  147. result = await self.run_node(
  148. node=parent_node,
  149. round_index=round_index,
  150. system_prompt=system_prompt,
  151. user_content=request.task,
  152. tools=tools,
  153. max_iterations=max_iterations,
  154. slots=slots,
  155. branch_key=f"delegate-{index}",
  156. allow_delegation=False,
  157. )
  158. return {
  159. "branch": index,
  160. "content": result.content,
  161. "iterations": result.iterations,
  162. "tool_calls_made": result.tool_calls_made,
  163. }
  164. results = await asyncio.gather(*(
  165. run_one(index, request) for index, request in enumerate(normalized, start=1)
  166. ), return_exceptions=True)
  167. output = []
  168. for index, result in enumerate(results, start=1):
  169. if isinstance(result, BaseException):
  170. output.append({"branch": index, "error": f"{type(result).__name__}: {result}"})
  171. else:
  172. output.append(result)
  173. return json.dumps({"delegated": output}, ensure_ascii=False)
  174. return StructuredTool.from_function(
  175. coroutine=delegate_agents_v2,
  176. name="delegate_agents_v2",
  177. description="将多个无依赖任务并发委派给拥有相同阶段工具权限的子 Agent。",
  178. args_schema=DelegateArgs,
  179. )
  180. async def run_node(
  181. self,
  182. *,
  183. node: str,
  184. round_index: int,
  185. system_prompt: str,
  186. user_content: str,
  187. tools: Iterable[ToolFn] = (),
  188. max_iterations: int = 12,
  189. slots: tuple[InputSlot, ...] = (),
  190. branch_key: str = "",
  191. allow_delegation: bool = True,
  192. ) -> NodeRun:
  193. tool_functions = tuple(tools)
  194. model_name = self.models_by_role.get(node, self.default_model)
  195. langchain_tools = [_as_langchain_tool(fn) for fn in tool_functions]
  196. if allow_delegation and node in {"search", "evidence", "evaluator"}:
  197. langchain_tools.append(self._delegate_tool(
  198. parent_node=node,
  199. round_index=round_index,
  200. system_prompt=system_prompt,
  201. tools=tool_functions,
  202. max_iterations=max_iterations,
  203. slots=slots,
  204. ))
  205. agent = create_agent(
  206. model=self._model(node),
  207. tools=langchain_tools,
  208. system_prompt=system_prompt,
  209. middleware=self._middleware(node),
  210. name=f"find_agent_v2_{node}",
  211. )
  212. with self.observer.node(node=node, branch_key=branch_key) as observation:
  213. actual_user_content = observation.declare(
  214. fallback=user_content,
  215. system_prompt=system_prompt,
  216. slots=slots,
  217. tools=tuple(langchain_tools),
  218. model=model_name,
  219. refs=({"delegate": "same-stage-worker"} if allow_delegation else None),
  220. )
  221. result = await agent.ainvoke(
  222. {"messages": [{"role": "user", "content": actual_user_content}]},
  223. # create_agent counts model and tool nodes separately, and middleware
  224. # retries also consume supersteps. Keep this budget distinct from
  225. # the business-level no-progress guard in graph.py.
  226. config={"recursion_limit": max(64, max_iterations * 6)},
  227. )
  228. messages: list[BaseMessage] = list(result.get("messages") or [])
  229. usage = _usage(messages)
  230. for key in ("input_tokens", "output_tokens", "total_tokens"):
  231. self.usage[key] = int(self.usage[key]) + int(usage[key])
  232. self.usage["cost"] = round(float(self.usage["cost"]) + float(usage["cost"]), 8)
  233. process = {
  234. "messages": [_message_dict(message) for message in messages],
  235. "events": _events(messages),
  236. "usage": usage,
  237. "iterations": sum(isinstance(message, AIMessage) for message in messages),
  238. "tool_calls_made": sum(
  239. len(message.tool_calls) for message in messages if isinstance(message, AIMessage)
  240. ),
  241. }
  242. last_ai = next((message for message in reversed(messages) if isinstance(message, AIMessage)), None)
  243. content = str(last_ai.content if last_ai is not None else "")
  244. observation.record_react(output=process, ok=True)
  245. observation.set_output({"agent输出": content, **process}, ok=True)
  246. return NodeRun(
  247. node=node,
  248. round_index=round_index,
  249. content=content,
  250. iterations=int(process["iterations"]),
  251. tool_calls_made=int(process["tool_calls_made"]),
  252. )
  253. def normalize_models(
  254. *,
  255. model: str | None = None,
  256. planning: str | None = None,
  257. search: str | None = None,
  258. evidence: str | None = None,
  259. evaluation: str | None = None,
  260. report: str | None = None,
  261. ) -> dict[str, str]:
  262. base = model or "google/gemini-3-flash-preview"
  263. return {
  264. "supervisor": planning or base,
  265. "search": search or base,
  266. "evidence": evidence or base,
  267. "evaluator": evaluation or base,
  268. "report": report or base,
  269. }