| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302 |
- """Independent LangChain ReAct host for find_agent_v2.
- The host owns model construction, retries, usage accounting, tool adaptation and
- controlled nested-agent delegation. It never imports the legacy find-agent package.
- """
- from __future__ import annotations
- import asyncio
- import json
- import os
- from collections.abc import Iterable
- from typing import Any
- from langchain.agents import create_agent
- from langchain.agents.middleware import ModelFallbackMiddleware, ModelRetryMiddleware, ToolRetryMiddleware
- from langchain_core.messages import AIMessage, BaseMessage, ToolMessage
- from langchain_core.tools import StructuredTool
- from langchain_openai import ChatOpenAI
- from pydantic import BaseModel, Field
- from find_agent_v2.observability import InputSlot, ObagentObserver
- from find_agent_v2.state import NodeRun
- from find_agent_v2.tools import ToolFn
- from supply_agent.config import Settings, get_settings
- class DelegateRequest(BaseModel):
- task: str = Field(description="交给子 Agent 的完整、独立任务")
- class DelegateArgs(BaseModel):
- requests: list[DelegateRequest] = Field(
- min_length=1, max_length=8,
- description="可并发执行的子 Agent 任务;有依赖的任务不要放在同一批",
- )
- def _tool_name(fn: ToolFn) -> str:
- return str(getattr(fn, "_tool_name", fn.__name__))
- def _as_langchain_tool(fn: ToolFn) -> StructuredTool:
- kwargs = {
- "name": _tool_name(fn),
- "description": str(getattr(fn, "_tool_description", fn.__doc__ or "")),
- }
- if asyncio.iscoroutinefunction(fn):
- return StructuredTool.from_function(coroutine=fn, **kwargs)
- return StructuredTool.from_function(func=fn, **kwargs)
- def _message_dict(message: BaseMessage) -> dict[str, Any]:
- data = message.model_dump(mode="json")
- data["type"] = message.type
- return data
- def _usage(messages: list[BaseMessage]) -> dict[str, int | float]:
- totals: dict[str, int | float] = {
- "input_tokens": 0, "output_tokens": 0, "total_tokens": 0, "cost": 0.0,
- }
- for message in messages:
- usage = getattr(message, "usage_metadata", None) or {}
- totals["input_tokens"] += int(usage.get("input_tokens") or 0)
- totals["output_tokens"] += int(usage.get("output_tokens") or 0)
- totals["total_tokens"] += int(usage.get("total_tokens") or 0)
- response = getattr(message, "response_metadata", None) or {}
- cost = response.get("cost") or (response.get("usage") or {}).get("cost")
- if cost is not None:
- totals["cost"] += float(cost)
- totals["cost"] = round(float(totals["cost"]), 8)
- return totals
- def _events(messages: list[BaseMessage]) -> list[dict[str, Any]]:
- events: list[dict[str, Any]] = []
- for message in messages:
- if isinstance(message, AIMessage):
- events.append({
- "type": "llm_output",
- "content": message.content,
- "tool_calls": message.tool_calls,
- "usage": message.usage_metadata,
- })
- elif isinstance(message, ToolMessage):
- events.append({
- "type": "tool_call",
- "name": message.name,
- "tool_call_id": message.tool_call_id,
- "result": message.content,
- "status": message.status,
- })
- return events
- class FindAgentNodeHost:
- """Create a fresh LangChain Agent for every workflow or delegated node."""
- def __init__(
- self,
- *,
- settings: Settings | None = None,
- models_by_role: dict[str, str] | None = None,
- default_model: str = "google/gemini-3-flash-preview",
- observer: ObagentObserver | None = None,
- ) -> None:
- self.settings = settings or get_settings()
- self.models_by_role = dict(models_by_role or {})
- self.default_model = default_model
- self.observer = observer or ObagentObserver()
- self.usage = {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0, "cost": 0.0}
- def reset_usage(self) -> None:
- self.usage = {
- "input_tokens": 0,
- "output_tokens": 0,
- "total_tokens": 0,
- "cost": 0.0,
- }
- def _model(self, role: str) -> ChatOpenAI:
- model_name = self.models_by_role.get(role, self.default_model)
- return ChatOpenAI(
- model=model_name,
- api_key=self.settings.openrouter_api_key,
- base_url=self.settings.openrouter_base_url,
- timeout=self.settings.openrouter_timeout_seconds,
- temperature=0.2,
- max_retries=0, # retries are observable LangChain middleware below
- default_headers={
- "HTTP-Referer": self.settings.openrouter_site_url,
- "X-Title": self.settings.openrouter_site_name,
- },
- )
- def _middleware(self, role: str):
- items: list[Any] = [
- ModelRetryMiddleware(max_retries=2, on_failure="error"),
- ToolRetryMiddleware(max_retries=2, on_failure="continue"),
- ]
- fallback = os.getenv("FIND_AGENT_V2_FALLBACK_MODEL", "").strip()
- if fallback and fallback != self.models_by_role.get(role, self.default_model):
- items.insert(0, ModelFallbackMiddleware(self._model_for_name(fallback)))
- return items
- def _model_for_name(self, model_name: str) -> ChatOpenAI:
- return ChatOpenAI(
- model=model_name,
- api_key=self.settings.openrouter_api_key,
- base_url=self.settings.openrouter_base_url,
- timeout=self.settings.openrouter_timeout_seconds,
- temperature=0.2,
- max_retries=0,
- )
- def _delegate_tool(
- self,
- *,
- parent_node: str,
- round_index: int,
- system_prompt: str,
- tools: tuple[ToolFn, ...],
- max_iterations: int,
- slots: tuple[InputSlot, ...],
- ) -> StructuredTool:
- async def delegate_agents_v2(requests: list[DelegateRequest]) -> str:
- """并发委派多个相互独立的任务给同阶段子 Agent。"""
- normalized = [
- item if isinstance(item, DelegateRequest) else DelegateRequest.model_validate(item)
- for item in requests
- ]
- async def run_one(index: int, request: DelegateRequest) -> dict[str, Any]:
- result = await self.run_node(
- node=parent_node,
- round_index=round_index,
- system_prompt=system_prompt,
- user_content=request.task,
- tools=tools,
- max_iterations=max_iterations,
- slots=slots,
- branch_key=f"delegate-{index}",
- allow_delegation=False,
- )
- return {
- "branch": index,
- "content": result.content,
- "iterations": result.iterations,
- "tool_calls_made": result.tool_calls_made,
- }
- results = await asyncio.gather(*(
- run_one(index, request) for index, request in enumerate(normalized, start=1)
- ), return_exceptions=True)
- output = []
- for index, result in enumerate(results, start=1):
- if isinstance(result, BaseException):
- output.append({"branch": index, "error": f"{type(result).__name__}: {result}"})
- else:
- output.append(result)
- return json.dumps({"delegated": output}, ensure_ascii=False)
- return StructuredTool.from_function(
- coroutine=delegate_agents_v2,
- name="delegate_agents_v2",
- description="将多个无依赖任务并发委派给拥有相同阶段工具权限的子 Agent。",
- args_schema=DelegateArgs,
- )
- async def run_node(
- self,
- *,
- node: str,
- round_index: int,
- system_prompt: str,
- user_content: str,
- tools: Iterable[ToolFn] = (),
- max_iterations: int = 12,
- slots: tuple[InputSlot, ...] = (),
- branch_key: str = "",
- allow_delegation: bool = True,
- ) -> NodeRun:
- tool_functions = tuple(tools)
- model_name = self.models_by_role.get(node, self.default_model)
- langchain_tools = [_as_langchain_tool(fn) for fn in tool_functions]
- if allow_delegation and node in {"search", "evidence", "evaluator"}:
- langchain_tools.append(self._delegate_tool(
- parent_node=node,
- round_index=round_index,
- system_prompt=system_prompt,
- tools=tool_functions,
- max_iterations=max_iterations,
- slots=slots,
- ))
- agent = create_agent(
- model=self._model(node),
- tools=langchain_tools,
- system_prompt=system_prompt,
- middleware=self._middleware(node),
- name=f"find_agent_v2_{node}",
- )
- with self.observer.node(node=node, branch_key=branch_key) as observation:
- actual_user_content = observation.declare(
- fallback=user_content,
- system_prompt=system_prompt,
- slots=slots,
- tools=tuple(langchain_tools),
- model=model_name,
- refs=({"delegate": "same-stage-worker"} if allow_delegation else None),
- )
- result = await agent.ainvoke(
- {"messages": [{"role": "user", "content": actual_user_content}]},
- # create_agent counts model and tool nodes separately, and middleware
- # retries also consume supersteps. Keep this budget distinct from
- # the business-level no-progress guard in graph.py.
- config={"recursion_limit": max(64, max_iterations * 6)},
- )
- messages: list[BaseMessage] = list(result.get("messages") or [])
- usage = _usage(messages)
- for key in ("input_tokens", "output_tokens", "total_tokens"):
- self.usage[key] = int(self.usage[key]) + int(usage[key])
- self.usage["cost"] = round(float(self.usage["cost"]) + float(usage["cost"]), 8)
- process = {
- "messages": [_message_dict(message) for message in messages],
- "events": _events(messages),
- "usage": usage,
- "iterations": sum(isinstance(message, AIMessage) for message in messages),
- "tool_calls_made": sum(
- len(message.tool_calls) for message in messages if isinstance(message, AIMessage)
- ),
- }
- last_ai = next((message for message in reversed(messages) if isinstance(message, AIMessage)), None)
- content = str(last_ai.content if last_ai is not None else "")
- observation.record_react(output=process, ok=True)
- observation.set_output({"agent输出": content, **process}, ok=True)
- return NodeRun(
- node=node,
- round_index=round_index,
- content=content,
- iterations=int(process["iterations"]),
- tool_calls_made=int(process["tool_calls_made"]),
- )
- def normalize_models(
- *,
- model: str | None = None,
- planning: str | None = None,
- search: str | None = None,
- evidence: str | None = None,
- evaluation: str | None = None,
- report: str | None = None,
- ) -> dict[str, str]:
- base = model or "google/gemini-3-flash-preview"
- return {
- "supervisor": planning or base,
- "search": search or base,
- "evidence": evidence or base,
- "evaluator": evaluation or base,
- "report": report or base,
- }
|