runtime.py 23 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588
  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 (
  19. EvidenceAssignment,
  20. NodeRun,
  21. PlanningDecision,
  22. SearchAssignment,
  23. SupervisorDecision,
  24. )
  25. from find_agent_v2.tools import (
  26. EvaluationBatch,
  27. ToolFn,
  28. fetch_candidate_details_v2,
  29. fetch_candidate_portraits_v2,
  30. search_videos_v2,
  31. )
  32. from supply_agent.config import Settings, get_settings
  33. class DelegateRequest(BaseModel):
  34. task: str = Field(description="交给子 Agent 的完整、独立任务")
  35. class DelegateArgs(BaseModel):
  36. requests: list[DelegateRequest] = Field(
  37. min_length=1, max_length=8,
  38. description="可并发执行的子 Agent 任务;有依赖的任务不要放在同一批",
  39. )
  40. def _tool_name(fn: ToolFn) -> str:
  41. return str(getattr(fn, "_tool_name", fn.__name__))
  42. def _as_langchain_tool(fn: ToolFn) -> StructuredTool:
  43. kwargs = {
  44. "name": _tool_name(fn),
  45. "description": str(getattr(fn, "_tool_description", fn.__doc__ or "")),
  46. }
  47. if asyncio.iscoroutinefunction(fn):
  48. return StructuredTool.from_function(coroutine=fn, **kwargs)
  49. return StructuredTool.from_function(func=fn, **kwargs)
  50. def _message_dict(message: BaseMessage) -> dict[str, Any]:
  51. data = message.model_dump(mode="json")
  52. data["type"] = message.type
  53. if isinstance(message, ToolMessage) and _tool_result_error_code(message.content):
  54. data["status"] = "error"
  55. return data
  56. def _tool_result_error_code(content: Any) -> str:
  57. if not isinstance(content, str):
  58. return ""
  59. try:
  60. payload = json.loads(content)
  61. except (TypeError, ValueError):
  62. return ""
  63. if not isinstance(payload, dict):
  64. return ""
  65. return str(payload.get("error_code") or "")
  66. def _usage(messages: list[BaseMessage]) -> dict[str, int | float]:
  67. totals: dict[str, int | float] = {
  68. "input_tokens": 0, "output_tokens": 0, "total_tokens": 0, "cost": 0.0,
  69. }
  70. for message in messages:
  71. usage = getattr(message, "usage_metadata", None) or {}
  72. totals["input_tokens"] += int(usage.get("input_tokens") or 0)
  73. totals["output_tokens"] += int(usage.get("output_tokens") or 0)
  74. totals["total_tokens"] += int(usage.get("total_tokens") or 0)
  75. response = getattr(message, "response_metadata", None) or {}
  76. cost = (
  77. response.get("cost")
  78. or (response.get("token_usage") or {}).get("cost")
  79. or (response.get("usage") or {}).get("cost")
  80. )
  81. if cost is not None:
  82. totals["cost"] += float(cost)
  83. totals["cost"] = round(float(totals["cost"]), 8)
  84. return totals
  85. def _events(messages: list[BaseMessage]) -> list[dict[str, Any]]:
  86. events: list[dict[str, Any]] = []
  87. for message in messages:
  88. if isinstance(message, AIMessage):
  89. events.append({
  90. "type": "llm_output",
  91. "content": message.content,
  92. "tool_calls": message.tool_calls,
  93. "usage": message.usage_metadata,
  94. })
  95. elif isinstance(message, ToolMessage):
  96. error_code = _tool_result_error_code(message.content)
  97. events.append({
  98. "type": "tool_call",
  99. "name": message.name,
  100. "tool_call_id": message.tool_call_id,
  101. "result": message.content,
  102. "status": "error" if error_code else message.status,
  103. "error_code": error_code or None,
  104. })
  105. return events
  106. class FindAgentNodeHost:
  107. """Create a fresh LangChain Agent for every workflow or delegated node."""
  108. def __init__(
  109. self,
  110. *,
  111. settings: Settings | None = None,
  112. models_by_role: dict[str, str] | None = None,
  113. default_model: str = "google/gemini-3-flash-preview",
  114. observer: ObagentObserver | None = None,
  115. ) -> None:
  116. self.settings = settings or get_settings()
  117. self.models_by_role = dict(models_by_role or {})
  118. self.default_model = default_model
  119. self.observer = observer or ObagentObserver()
  120. self.usage = {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0, "cost": 0.0}
  121. def reset_usage(self) -> None:
  122. self.usage = {
  123. "input_tokens": 0,
  124. "output_tokens": 0,
  125. "total_tokens": 0,
  126. "cost": 0.0,
  127. }
  128. def _model(self, role: str) -> ChatOpenAI:
  129. model_name = self.models_by_role.get(role, self.default_model)
  130. return ChatOpenAI(
  131. model=model_name,
  132. api_key=self.settings.openrouter_api_key,
  133. base_url=self.settings.openrouter_base_url,
  134. timeout=self.settings.openrouter_timeout_seconds,
  135. temperature=0.2,
  136. max_retries=0, # retries are observable LangChain middleware below
  137. default_headers={
  138. "HTTP-Referer": self.settings.openrouter_site_url,
  139. "X-Title": self.settings.openrouter_site_name,
  140. },
  141. )
  142. def _middleware(self, role: str):
  143. items: list[Any] = [
  144. ModelRetryMiddleware(max_retries=2, on_failure="error"),
  145. ToolRetryMiddleware(max_retries=2, on_failure="continue"),
  146. ]
  147. fallback = os.getenv("FIND_AGENT_V2_FALLBACK_MODEL", "").strip()
  148. if fallback and fallback != self.models_by_role.get(role, self.default_model):
  149. items.insert(0, ModelFallbackMiddleware(self._model_for_name(fallback)))
  150. return items
  151. def _model_for_name(self, model_name: str) -> ChatOpenAI:
  152. return ChatOpenAI(
  153. model=model_name,
  154. api_key=self.settings.openrouter_api_key,
  155. base_url=self.settings.openrouter_base_url,
  156. timeout=self.settings.openrouter_timeout_seconds,
  157. temperature=0.2,
  158. max_retries=0,
  159. )
  160. @staticmethod
  161. def _json_result(raw: str) -> dict[str, Any]:
  162. try:
  163. value = json.loads(raw)
  164. except (TypeError, ValueError):
  165. return {"raw": raw}
  166. return value if isinstance(value, dict) else {"result": value}
  167. async def run_search_assignment(
  168. self,
  169. *,
  170. round_index: int,
  171. assignment: SearchAssignment,
  172. system_prompt: str,
  173. user_content: str,
  174. slots: tuple[InputSlot, ...] = (),
  175. ) -> NodeRun:
  176. """Execute an already-planned search assignment without another LLM hop."""
  177. with self.observer.node(node="search") as observation:
  178. observation.declare(
  179. fallback=user_content,
  180. system_prompt=system_prompt,
  181. slots=slots,
  182. tools=(search_videos_v2,),
  183. model="host/deterministic",
  184. )
  185. raw = await search_videos_v2(
  186. run_id=assignment.run_id,
  187. round_index=assignment.round_index,
  188. searches=[item.model_dump(mode="json") for item in assignment.tasks],
  189. )
  190. payload = self._json_result(raw)
  191. observation.set_output({"宿主直执行": payload}, ok=not bool(payload.get("error")))
  192. return NodeRun("search", round_index, raw, 0, 1)
  193. async def run_evidence_assignment(
  194. self,
  195. *,
  196. round_index: int,
  197. assignment: EvidenceAssignment,
  198. system_prompt: str,
  199. user_content: str,
  200. slots: tuple[InputSlot, ...] = (),
  201. branch_key: str = "",
  202. ) -> NodeRun:
  203. """Execute a host-owned evidence shard without asking a model to select its tool."""
  204. tool_fn = (
  205. fetch_candidate_details_v2
  206. if assignment.evidence_type == "detail"
  207. else fetch_candidate_portraits_v2
  208. )
  209. with self.observer.node(node="evidence", branch_key=branch_key) as observation:
  210. observation.declare(
  211. fallback=user_content,
  212. system_prompt=system_prompt,
  213. slots=slots,
  214. tools=(tool_fn,),
  215. model="host/deterministic",
  216. )
  217. raw = await tool_fn(
  218. run_id=assignment.run_id,
  219. candidate_ids=assignment.candidate_ids,
  220. )
  221. payload = self._json_result(raw)
  222. observation.set_output({"宿主直执行": payload}, ok=not bool(payload.get("error")))
  223. return NodeRun("evidence", round_index, raw, 0, 1)
  224. async def run_host_supervision(
  225. self,
  226. *,
  227. round_index: int,
  228. decision: SupervisorDecision,
  229. system_prompt: str,
  230. user_content: str,
  231. slots: tuple[InputSlot, ...] = (),
  232. output_enricher: Callable[[SupervisorDecision], dict[str, Any]] | None = None,
  233. ) -> tuple[NodeRun, SupervisorDecision]:
  234. """Record a deterministic routing decision without invoking an LLM."""
  235. payload = decision.model_dump(mode="json", exclude_none=True)
  236. enriched = output_enricher(decision) if output_enricher else {}
  237. with self.observer.node(node="supervisor") as observation:
  238. observation.declare(
  239. fallback=user_content,
  240. system_prompt=system_prompt,
  241. slots=slots,
  242. tools=(),
  243. model="host/deterministic",
  244. )
  245. observation.set_output({"宿主确定性决策": payload, **enriched}, ok=True)
  246. content = json.dumps(payload, ensure_ascii=False)
  247. return NodeRun("supervisor", round_index, content, 0, 0), decision
  248. def _delegate_tool(
  249. self,
  250. *,
  251. parent_node: str,
  252. round_index: int,
  253. system_prompt: str,
  254. tools: tuple[ToolFn, ...],
  255. max_iterations: int,
  256. ) -> StructuredTool:
  257. async def delegate_agents_v2(requests: list[DelegateRequest]) -> str:
  258. """并发委派多个相互独立的任务给同阶段子 Agent。"""
  259. normalized = [
  260. item if isinstance(item, DelegateRequest) else DelegateRequest.model_validate(item)
  261. for item in requests
  262. ]
  263. async def run_one(index: int, request: DelegateRequest) -> dict[str, Any]:
  264. result = await self.run_node(
  265. node=parent_node,
  266. round_index=round_index,
  267. system_prompt=system_prompt,
  268. user_content=request.task,
  269. tools=tools,
  270. max_iterations=max_iterations,
  271. slots=(InputSlot(
  272. "委派任务", request.task, "delegate_assignment",
  273. "delegate_agents_v2.requests", False,
  274. ),),
  275. branch_key=f"delegate-{index}",
  276. allow_delegation=False,
  277. )
  278. return {
  279. "branch": index,
  280. "content": result.content,
  281. "iterations": result.iterations,
  282. "tool_calls_made": result.tool_calls_made,
  283. }
  284. results = await asyncio.gather(*(
  285. run_one(index, request) for index, request in enumerate(normalized, start=1)
  286. ), return_exceptions=True)
  287. output = []
  288. for index, result in enumerate(results, start=1):
  289. if isinstance(result, BaseException):
  290. output.append({"branch": index, "error": f"{type(result).__name__}: {result}"})
  291. else:
  292. output.append(result)
  293. return json.dumps({"delegated": output}, ensure_ascii=False)
  294. return StructuredTool.from_function(
  295. coroutine=delegate_agents_v2,
  296. name="delegate_agents_v2",
  297. description="将多个无依赖任务并发委派给拥有相同阶段工具权限的子 Agent。",
  298. args_schema=DelegateArgs,
  299. )
  300. async def run_node(
  301. self,
  302. *,
  303. node: str,
  304. round_index: int,
  305. system_prompt: str,
  306. user_content: str,
  307. tools: Iterable[ToolFn] = (),
  308. max_iterations: int = 12,
  309. slots: tuple[InputSlot, ...] = (),
  310. branch_key: str = "",
  311. allow_delegation: bool = True,
  312. ) -> NodeRun:
  313. tool_functions = tuple(tools)
  314. model_name = self.models_by_role.get(node, self.default_model)
  315. langchain_tools = [_as_langchain_tool(fn) for fn in tool_functions]
  316. if allow_delegation and node in {"search", "evidence", "evaluator"}:
  317. langchain_tools.append(self._delegate_tool(
  318. parent_node=node,
  319. round_index=round_index,
  320. system_prompt=system_prompt,
  321. tools=tool_functions,
  322. max_iterations=max_iterations,
  323. ))
  324. agent = create_agent(
  325. model=self._model(node),
  326. tools=langchain_tools,
  327. system_prompt=system_prompt,
  328. middleware=self._middleware(node),
  329. name=f"find_agent_v2_{node}",
  330. )
  331. with self.observer.node(node=node, branch_key=branch_key) as observation:
  332. actual_user_content = observation.declare(
  333. fallback=user_content,
  334. system_prompt=system_prompt,
  335. slots=slots,
  336. tools=tuple(langchain_tools),
  337. model=model_name,
  338. refs=({"delegate": "same-stage-worker"} if allow_delegation else None),
  339. )
  340. result = await agent.ainvoke(
  341. {"messages": [{"role": "user", "content": actual_user_content}]},
  342. # create_agent counts model and tool nodes separately, and middleware
  343. # retries also consume supersteps. Keep this budget distinct from
  344. # the business-level no-progress guard in graph.py.
  345. config={"recursion_limit": max(64, max_iterations * 6)},
  346. )
  347. messages: list[BaseMessage] = list(result.get("messages") or [])
  348. usage = _usage(messages)
  349. for key in ("input_tokens", "output_tokens", "total_tokens"):
  350. self.usage[key] = int(self.usage[key]) + int(usage[key])
  351. self.usage["cost"] = round(float(self.usage["cost"]) + float(usage["cost"]), 8)
  352. process = {
  353. "messages": [_message_dict(message) for message in messages],
  354. "events": _events(messages),
  355. "usage": usage,
  356. "iterations": sum(isinstance(message, AIMessage) for message in messages),
  357. "tool_calls_made": sum(
  358. len(message.tool_calls) for message in messages if isinstance(message, AIMessage)
  359. ),
  360. }
  361. last_ai = next((message for message in reversed(messages) if isinstance(message, AIMessage)), None)
  362. content = str(last_ai.content if last_ai is not None else "")
  363. observation.record_react(output=process, ok=True)
  364. observation.set_output({"agent输出": content, **process}, ok=True)
  365. return NodeRun(
  366. node=node,
  367. round_index=round_index,
  368. content=content,
  369. iterations=int(process["iterations"]),
  370. tool_calls_made=int(process["tool_calls_made"]),
  371. )
  372. async def run_evaluation(
  373. self,
  374. *,
  375. round_index: int,
  376. system_prompt: str,
  377. user_content: str,
  378. slots: tuple[InputSlot, ...] = (),
  379. branch_key: str = "",
  380. tools: Iterable[ToolFn] = (),
  381. ) -> tuple[NodeRun, list[dict[str, Any]]]:
  382. """Get one structured decision batch; the graph owns validation and persistence."""
  383. model_name = self.models_by_role.get("evaluator", self.default_model)
  384. langchain_tools = [_as_langchain_tool(fn) for fn in tools]
  385. agent = create_agent(
  386. model=self._model("evaluator"),
  387. tools=langchain_tools,
  388. system_prompt=system_prompt,
  389. middleware=self._middleware("evaluator"),
  390. response_format=EvaluationBatch,
  391. name="find_agent_v2_evaluator",
  392. )
  393. with self.observer.node(node="evaluator", branch_key=branch_key) as observation:
  394. actual_user_content = observation.declare(
  395. fallback=user_content,
  396. system_prompt=system_prompt,
  397. slots=slots,
  398. tools=tuple(langchain_tools),
  399. model=model_name,
  400. )
  401. result = await agent.ainvoke(
  402. {"messages": [{"role": "user", "content": actual_user_content}]},
  403. config={"recursion_limit": 24},
  404. )
  405. messages: list[BaseMessage] = list(result.get("messages") or [])
  406. usage = _usage(messages)
  407. for key in ("input_tokens", "output_tokens", "total_tokens"):
  408. self.usage[key] = int(self.usage[key]) + int(usage[key])
  409. self.usage["cost"] = round(float(self.usage["cost"]) + float(usage["cost"]), 8)
  410. structured = result.get("structured_response")
  411. batch = structured if isinstance(structured, EvaluationBatch) else EvaluationBatch.model_validate(structured)
  412. items = [item.model_dump(exclude_none=True) for item in batch.items]
  413. process = {
  414. "messages": [_message_dict(message) for message in messages],
  415. "events": _events(messages),
  416. "usage": usage,
  417. "iterations": sum(isinstance(message, AIMessage) for message in messages),
  418. "tool_calls_made": sum(
  419. len(message.tool_calls)
  420. for message in messages
  421. if isinstance(message, AIMessage)
  422. ),
  423. "structured_items": items,
  424. }
  425. observation.record_react(output=process, ok=True)
  426. observation.set_output({"结构化评估": items, **process}, ok=True)
  427. last_ai = next((message for message in reversed(messages) if isinstance(message, AIMessage)), None)
  428. run = NodeRun(
  429. node="evaluator",
  430. round_index=round_index,
  431. content=str(last_ai.content if last_ai is not None else ""),
  432. iterations=int(process["iterations"]),
  433. tool_calls_made=int(process["tool_calls_made"]),
  434. )
  435. return run, items
  436. async def _run_supervisor_structured(
  437. self,
  438. *,
  439. round_index: int,
  440. system_prompt: str,
  441. user_content: str,
  442. slots: tuple[InputSlot, ...],
  443. response_format: type[BaseModel],
  444. output_enricher: Callable[[BaseModel], dict[str, Any]] | None = None,
  445. ) -> tuple[NodeRun, BaseModel]:
  446. """Run planner/supervisor with a validated host-owned output contract."""
  447. node = "supervisor"
  448. model_name = self.models_by_role.get(node, self.default_model)
  449. agent = create_agent(
  450. model=self._model(node),
  451. tools=[],
  452. system_prompt=system_prompt,
  453. middleware=self._middleware(node),
  454. response_format=response_format,
  455. name="find_agent_v2_supervisor",
  456. )
  457. with self.observer.node(node=node) as observation:
  458. actual_user_content = observation.declare(
  459. fallback=user_content,
  460. system_prompt=system_prompt,
  461. slots=slots,
  462. tools=(),
  463. model=model_name,
  464. )
  465. result = await agent.ainvoke(
  466. {"messages": [{"role": "user", "content": actual_user_content}]},
  467. config={"recursion_limit": 24},
  468. )
  469. messages: list[BaseMessage] = list(result.get("messages") or [])
  470. usage = _usage(messages)
  471. for key in ("input_tokens", "output_tokens", "total_tokens"):
  472. self.usage[key] = int(self.usage[key]) + int(usage[key])
  473. self.usage["cost"] = round(float(self.usage["cost"]) + float(usage["cost"]), 8)
  474. raw_structured = result.get("structured_response")
  475. structured = (
  476. raw_structured
  477. if isinstance(raw_structured, response_format)
  478. else response_format.model_validate(raw_structured)
  479. )
  480. structured_payload = structured.model_dump(mode="json", exclude_none=True)
  481. enriched_output = output_enricher(structured) if output_enricher else {}
  482. process = {
  483. "messages": [_message_dict(message) for message in messages],
  484. "events": _events(messages),
  485. "usage": usage,
  486. "iterations": sum(isinstance(message, AIMessage) for message in messages),
  487. "tool_calls_made": 0,
  488. "structured_output": structured_payload,
  489. **enriched_output,
  490. }
  491. content = json.dumps(structured_payload, ensure_ascii=False)
  492. observation.record_react(output=process, ok=True)
  493. observation.set_output({"结构化输出": structured_payload, **process}, ok=True)
  494. return NodeRun(node, round_index, content, int(process["iterations"]), 0), structured
  495. async def run_planning(
  496. self,
  497. *,
  498. round_index: int,
  499. system_prompt: str,
  500. user_content: str,
  501. slots: tuple[InputSlot, ...] = (),
  502. ) -> tuple[NodeRun, PlanningDecision]:
  503. run, structured = await self._run_supervisor_structured(
  504. round_index=round_index,
  505. system_prompt=system_prompt,
  506. user_content=user_content,
  507. slots=slots,
  508. response_format=PlanningDecision,
  509. )
  510. return run, PlanningDecision.model_validate(structured)
  511. async def run_supervision(
  512. self,
  513. *,
  514. round_index: int,
  515. system_prompt: str,
  516. user_content: str,
  517. slots: tuple[InputSlot, ...] = (),
  518. output_enricher: Callable[[SupervisorDecision], dict[str, Any]] | None = None,
  519. ) -> tuple[NodeRun, SupervisorDecision]:
  520. run, structured = await self._run_supervisor_structured(
  521. round_index=round_index,
  522. system_prompt=system_prompt,
  523. user_content=user_content,
  524. slots=slots,
  525. response_format=SupervisorDecision,
  526. output_enricher=output_enricher,
  527. )
  528. return run, SupervisorDecision.model_validate(structured)
  529. def normalize_models(
  530. *,
  531. model: str | None = None,
  532. planning: str | None = None,
  533. search: str | None = None,
  534. evidence: str | None = None,
  535. evaluation: str | None = None,
  536. report: str | None = None,
  537. ) -> dict[str, str]:
  538. base = model or "google/gemini-3-flash-preview"
  539. return {
  540. "supervisor": planning or base,
  541. "search": search or base,
  542. "evidence": evidence or base,
  543. "evaluator": evaluation or base,
  544. "report": report or base,
  545. }