runtime.py 15 KB

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