| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465 |
- """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 Callable, 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, PlanningDecision, SupervisorDecision
- from find_agent_v2.tools import ToolFn
- from find_agent_v2.tools import EvaluationBatch
- 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,
- ) -> 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=(InputSlot(
- "委派任务", request.task, "delegate_assignment",
- "delegate_agents_v2.requests", False,
- ),),
- 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,
- ))
- 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"]),
- )
- async def run_evaluation(
- self,
- *,
- round_index: int,
- system_prompt: str,
- user_content: str,
- slots: tuple[InputSlot, ...] = (),
- branch_key: str = "",
- tools: Iterable[ToolFn] = (),
- ) -> tuple[NodeRun, list[dict[str, Any]]]:
- """Get one structured decision batch; the graph owns validation and persistence."""
- model_name = self.models_by_role.get("evaluator", self.default_model)
- langchain_tools = [_as_langchain_tool(fn) for fn in tools]
- agent = create_agent(
- model=self._model("evaluator"),
- tools=langchain_tools,
- system_prompt=system_prompt,
- middleware=self._middleware("evaluator"),
- response_format=EvaluationBatch,
- name="find_agent_v2_evaluator",
- )
- with self.observer.node(node="evaluator", 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,
- )
- result = await agent.ainvoke(
- {"messages": [{"role": "user", "content": actual_user_content}]},
- config={"recursion_limit": 24},
- )
- 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)
- structured = result.get("structured_response")
- batch = structured if isinstance(structured, EvaluationBatch) else EvaluationBatch.model_validate(structured)
- items = [item.model_dump(exclude_none=True) for item in batch.items]
- 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)
- ),
- "structured_items": items,
- }
- observation.record_react(output=process, ok=True)
- observation.set_output({"结构化评估": items, **process}, ok=True)
- last_ai = next((message for message in reversed(messages) if isinstance(message, AIMessage)), None)
- run = NodeRun(
- node="evaluator",
- round_index=round_index,
- content=str(last_ai.content if last_ai is not None else ""),
- iterations=int(process["iterations"]),
- tool_calls_made=int(process["tool_calls_made"]),
- )
- return run, items
- async def _run_supervisor_structured(
- self,
- *,
- round_index: int,
- system_prompt: str,
- user_content: str,
- slots: tuple[InputSlot, ...],
- response_format: type[BaseModel],
- output_enricher: Callable[[BaseModel], dict[str, Any]] | None = None,
- ) -> tuple[NodeRun, BaseModel]:
- """Run planner/supervisor with a validated host-owned output contract."""
- node = "supervisor"
- model_name = self.models_by_role.get(node, self.default_model)
- agent = create_agent(
- model=self._model(node),
- tools=[],
- system_prompt=system_prompt,
- middleware=self._middleware(node),
- response_format=response_format,
- name="find_agent_v2_supervisor",
- )
- with self.observer.node(node=node) as observation:
- actual_user_content = observation.declare(
- fallback=user_content,
- system_prompt=system_prompt,
- slots=slots,
- tools=(),
- model=model_name,
- )
- result = await agent.ainvoke(
- {"messages": [{"role": "user", "content": actual_user_content}]},
- config={"recursion_limit": 24},
- )
- 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)
- raw_structured = result.get("structured_response")
- structured = (
- raw_structured
- if isinstance(raw_structured, response_format)
- else response_format.model_validate(raw_structured)
- )
- structured_payload = structured.model_dump(mode="json", exclude_none=True)
- enriched_output = output_enricher(structured) if output_enricher else {}
- 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": 0,
- "structured_output": structured_payload,
- **enriched_output,
- }
- content = json.dumps(structured_payload, ensure_ascii=False)
- observation.record_react(output=process, ok=True)
- observation.set_output({"结构化输出": structured_payload, **process}, ok=True)
- return NodeRun(node, round_index, content, int(process["iterations"]), 0), structured
- async def run_planning(
- self,
- *,
- round_index: int,
- system_prompt: str,
- user_content: str,
- slots: tuple[InputSlot, ...] = (),
- ) -> tuple[NodeRun, PlanningDecision]:
- run, structured = await self._run_supervisor_structured(
- round_index=round_index,
- system_prompt=system_prompt,
- user_content=user_content,
- slots=slots,
- response_format=PlanningDecision,
- )
- return run, PlanningDecision.model_validate(structured)
- async def run_supervision(
- self,
- *,
- round_index: int,
- system_prompt: str,
- user_content: str,
- slots: tuple[InputSlot, ...] = (),
- output_enricher: Callable[[SupervisorDecision], dict[str, Any]] | None = None,
- ) -> tuple[NodeRun, SupervisorDecision]:
- run, structured = await self._run_supervisor_structured(
- round_index=round_index,
- system_prompt=system_prompt,
- user_content=user_content,
- slots=slots,
- response_format=SupervisorDecision,
- output_enricher=output_enricher,
- )
- return run, SupervisorDecision.model_validate(structured)
- 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,
- }
|