runtime.py 19 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465
  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 Callable, 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, PlanningDecision, SupervisorDecision
  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. ) -> 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=(InputSlot(
  155. "委派任务", request.task, "delegate_assignment",
  156. "delegate_agents_v2.requests", False,
  157. ),),
  158. branch_key=f"delegate-{index}",
  159. allow_delegation=False,
  160. )
  161. return {
  162. "branch": index,
  163. "content": result.content,
  164. "iterations": result.iterations,
  165. "tool_calls_made": result.tool_calls_made,
  166. }
  167. results = await asyncio.gather(*(
  168. run_one(index, request) for index, request in enumerate(normalized, start=1)
  169. ), return_exceptions=True)
  170. output = []
  171. for index, result in enumerate(results, start=1):
  172. if isinstance(result, BaseException):
  173. output.append({"branch": index, "error": f"{type(result).__name__}: {result}"})
  174. else:
  175. output.append(result)
  176. return json.dumps({"delegated": output}, ensure_ascii=False)
  177. return StructuredTool.from_function(
  178. coroutine=delegate_agents_v2,
  179. name="delegate_agents_v2",
  180. description="将多个无依赖任务并发委派给拥有相同阶段工具权限的子 Agent。",
  181. args_schema=DelegateArgs,
  182. )
  183. async def run_node(
  184. self,
  185. *,
  186. node: str,
  187. round_index: int,
  188. system_prompt: str,
  189. user_content: str,
  190. tools: Iterable[ToolFn] = (),
  191. max_iterations: int = 12,
  192. slots: tuple[InputSlot, ...] = (),
  193. branch_key: str = "",
  194. allow_delegation: bool = True,
  195. ) -> NodeRun:
  196. tool_functions = tuple(tools)
  197. model_name = self.models_by_role.get(node, self.default_model)
  198. langchain_tools = [_as_langchain_tool(fn) for fn in tool_functions]
  199. if allow_delegation and node in {"search", "evidence", "evaluator"}:
  200. langchain_tools.append(self._delegate_tool(
  201. parent_node=node,
  202. round_index=round_index,
  203. system_prompt=system_prompt,
  204. tools=tool_functions,
  205. max_iterations=max_iterations,
  206. ))
  207. agent = create_agent(
  208. model=self._model(node),
  209. tools=langchain_tools,
  210. system_prompt=system_prompt,
  211. middleware=self._middleware(node),
  212. name=f"find_agent_v2_{node}",
  213. )
  214. with self.observer.node(node=node, branch_key=branch_key) as observation:
  215. actual_user_content = observation.declare(
  216. fallback=user_content,
  217. system_prompt=system_prompt,
  218. slots=slots,
  219. tools=tuple(langchain_tools),
  220. model=model_name,
  221. refs=({"delegate": "same-stage-worker"} if allow_delegation else None),
  222. )
  223. result = await agent.ainvoke(
  224. {"messages": [{"role": "user", "content": actual_user_content}]},
  225. # create_agent counts model and tool nodes separately, and middleware
  226. # retries also consume supersteps. Keep this budget distinct from
  227. # the business-level no-progress guard in graph.py.
  228. config={"recursion_limit": max(64, max_iterations * 6)},
  229. )
  230. messages: list[BaseMessage] = list(result.get("messages") or [])
  231. usage = _usage(messages)
  232. for key in ("input_tokens", "output_tokens", "total_tokens"):
  233. self.usage[key] = int(self.usage[key]) + int(usage[key])
  234. self.usage["cost"] = round(float(self.usage["cost"]) + float(usage["cost"]), 8)
  235. process = {
  236. "messages": [_message_dict(message) for message in messages],
  237. "events": _events(messages),
  238. "usage": usage,
  239. "iterations": sum(isinstance(message, AIMessage) for message in messages),
  240. "tool_calls_made": sum(
  241. len(message.tool_calls) for message in messages if isinstance(message, AIMessage)
  242. ),
  243. }
  244. last_ai = next((message for message in reversed(messages) if isinstance(message, AIMessage)), None)
  245. content = str(last_ai.content if last_ai is not None else "")
  246. observation.record_react(output=process, ok=True)
  247. observation.set_output({"agent输出": content, **process}, ok=True)
  248. return NodeRun(
  249. node=node,
  250. round_index=round_index,
  251. content=content,
  252. iterations=int(process["iterations"]),
  253. tool_calls_made=int(process["tool_calls_made"]),
  254. )
  255. async def run_evaluation(
  256. self,
  257. *,
  258. round_index: int,
  259. system_prompt: str,
  260. user_content: str,
  261. slots: tuple[InputSlot, ...] = (),
  262. branch_key: str = "",
  263. tools: Iterable[ToolFn] = (),
  264. ) -> tuple[NodeRun, list[dict[str, Any]]]:
  265. """Get one structured decision batch; the graph owns validation and persistence."""
  266. model_name = self.models_by_role.get("evaluator", self.default_model)
  267. langchain_tools = [_as_langchain_tool(fn) for fn in tools]
  268. agent = create_agent(
  269. model=self._model("evaluator"),
  270. tools=langchain_tools,
  271. system_prompt=system_prompt,
  272. middleware=self._middleware("evaluator"),
  273. response_format=EvaluationBatch,
  274. name="find_agent_v2_evaluator",
  275. )
  276. with self.observer.node(node="evaluator", branch_key=branch_key) as observation:
  277. actual_user_content = observation.declare(
  278. fallback=user_content,
  279. system_prompt=system_prompt,
  280. slots=slots,
  281. tools=tuple(langchain_tools),
  282. model=model_name,
  283. )
  284. result = await agent.ainvoke(
  285. {"messages": [{"role": "user", "content": actual_user_content}]},
  286. config={"recursion_limit": 24},
  287. )
  288. messages: list[BaseMessage] = list(result.get("messages") or [])
  289. usage = _usage(messages)
  290. for key in ("input_tokens", "output_tokens", "total_tokens"):
  291. self.usage[key] = int(self.usage[key]) + int(usage[key])
  292. self.usage["cost"] = round(float(self.usage["cost"]) + float(usage["cost"]), 8)
  293. structured = result.get("structured_response")
  294. batch = structured if isinstance(structured, EvaluationBatch) else EvaluationBatch.model_validate(structured)
  295. items = [item.model_dump(exclude_none=True) for item in batch.items]
  296. process = {
  297. "messages": [_message_dict(message) for message in messages],
  298. "events": _events(messages),
  299. "usage": usage,
  300. "iterations": sum(isinstance(message, AIMessage) for message in messages),
  301. "tool_calls_made": sum(
  302. len(message.tool_calls)
  303. for message in messages
  304. if isinstance(message, AIMessage)
  305. ),
  306. "structured_items": items,
  307. }
  308. observation.record_react(output=process, ok=True)
  309. observation.set_output({"结构化评估": items, **process}, ok=True)
  310. last_ai = next((message for message in reversed(messages) if isinstance(message, AIMessage)), None)
  311. run = NodeRun(
  312. node="evaluator",
  313. round_index=round_index,
  314. content=str(last_ai.content if last_ai is not None else ""),
  315. iterations=int(process["iterations"]),
  316. tool_calls_made=int(process["tool_calls_made"]),
  317. )
  318. return run, items
  319. async def _run_supervisor_structured(
  320. self,
  321. *,
  322. round_index: int,
  323. system_prompt: str,
  324. user_content: str,
  325. slots: tuple[InputSlot, ...],
  326. response_format: type[BaseModel],
  327. output_enricher: Callable[[BaseModel], dict[str, Any]] | None = None,
  328. ) -> tuple[NodeRun, BaseModel]:
  329. """Run planner/supervisor with a validated host-owned output contract."""
  330. node = "supervisor"
  331. model_name = self.models_by_role.get(node, self.default_model)
  332. agent = create_agent(
  333. model=self._model(node),
  334. tools=[],
  335. system_prompt=system_prompt,
  336. middleware=self._middleware(node),
  337. response_format=response_format,
  338. name="find_agent_v2_supervisor",
  339. )
  340. with self.observer.node(node=node) as observation:
  341. actual_user_content = observation.declare(
  342. fallback=user_content,
  343. system_prompt=system_prompt,
  344. slots=slots,
  345. tools=(),
  346. model=model_name,
  347. )
  348. result = await agent.ainvoke(
  349. {"messages": [{"role": "user", "content": actual_user_content}]},
  350. config={"recursion_limit": 24},
  351. )
  352. messages: list[BaseMessage] = list(result.get("messages") or [])
  353. usage = _usage(messages)
  354. for key in ("input_tokens", "output_tokens", "total_tokens"):
  355. self.usage[key] = int(self.usage[key]) + int(usage[key])
  356. self.usage["cost"] = round(float(self.usage["cost"]) + float(usage["cost"]), 8)
  357. raw_structured = result.get("structured_response")
  358. structured = (
  359. raw_structured
  360. if isinstance(raw_structured, response_format)
  361. else response_format.model_validate(raw_structured)
  362. )
  363. structured_payload = structured.model_dump(mode="json", exclude_none=True)
  364. enriched_output = output_enricher(structured) if output_enricher else {}
  365. process = {
  366. "messages": [_message_dict(message) for message in messages],
  367. "events": _events(messages),
  368. "usage": usage,
  369. "iterations": sum(isinstance(message, AIMessage) for message in messages),
  370. "tool_calls_made": 0,
  371. "structured_output": structured_payload,
  372. **enriched_output,
  373. }
  374. content = json.dumps(structured_payload, ensure_ascii=False)
  375. observation.record_react(output=process, ok=True)
  376. observation.set_output({"结构化输出": structured_payload, **process}, ok=True)
  377. return NodeRun(node, round_index, content, int(process["iterations"]), 0), structured
  378. async def run_planning(
  379. self,
  380. *,
  381. round_index: int,
  382. system_prompt: str,
  383. user_content: str,
  384. slots: tuple[InputSlot, ...] = (),
  385. ) -> tuple[NodeRun, PlanningDecision]:
  386. run, structured = await self._run_supervisor_structured(
  387. round_index=round_index,
  388. system_prompt=system_prompt,
  389. user_content=user_content,
  390. slots=slots,
  391. response_format=PlanningDecision,
  392. )
  393. return run, PlanningDecision.model_validate(structured)
  394. async def run_supervision(
  395. self,
  396. *,
  397. round_index: int,
  398. system_prompt: str,
  399. user_content: str,
  400. slots: tuple[InputSlot, ...] = (),
  401. output_enricher: Callable[[SupervisorDecision], dict[str, Any]] | None = None,
  402. ) -> tuple[NodeRun, SupervisorDecision]:
  403. run, structured = await self._run_supervisor_structured(
  404. round_index=round_index,
  405. system_prompt=system_prompt,
  406. user_content=user_content,
  407. slots=slots,
  408. response_format=SupervisorDecision,
  409. output_enricher=output_enricher,
  410. )
  411. return run, SupervisorDecision.model_validate(structured)
  412. def normalize_models(
  413. *,
  414. model: str | None = None,
  415. planning: str | None = None,
  416. search: str | None = None,
  417. evidence: str | None = None,
  418. evaluation: str | None = None,
  419. report: str | None = None,
  420. ) -> dict[str, str]:
  421. base = model or "google/gemini-3-flash-preview"
  422. return {
  423. "supervisor": planning or base,
  424. "search": search or base,
  425. "evidence": evidence or base,
  426. "evaluator": evaluation or base,
  427. "report": report or base,
  428. }