Files
zhenxun_bot/zhenxun/services/ai/flow/agent/engine/executor.py
T
Rumioandwebjoin111 cd5fa065d3 ♻️ refactor(agent): 重构 Agent 状态管理与执行器流程,优化 Token 预估与自愈反思机制 (#2150)
- 统一使用 `run_context.run.messages` 作为消息历史的单一数据源,清理 `AgentState` 冗余字段
- 将工具消息装配逻辑 `assemble_tool_message` 提取并重构至 `ToolExecutor`
- 引入 `token_drift` 动态校准偏移量,并精确计算工具与系统提示词的 Token 开销
- 重构 `ReflexionCapability` 自愈反思引擎,基于异常多态与模板字典动态生成反馈提示词
- 支持通过 `resolve_model_capabilities` 解析并合并用户自定义的模型能力覆盖
- 在执行器循环中支持 `should_reset_cycle`,以优雅处理外部干预(如用户追加指示)
- 扩展 `capabilities` 中对 `gpt-[5-9]*` 等新型号模型的能力定义与上下文限制

Co-authored-by: webjoin111 <455457521@qq.com>
2026-07-16 09:09:59 +08:00

666 lines
24 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from abc import ABC, abstractmethod
import asyncio
from collections.abc import Iterable
from typing import Any, cast
from zhenxun.services.ai.capabilities import CombinedCapability
from zhenxun.services.ai.core.engine.context_renderer import ContextConverter
from zhenxun.services.ai.core.engine.token_counter import (
parse_usage_info,
token_counter,
)
from zhenxun.services.ai.core.exceptions import (
UpstreamServerException,
)
from zhenxun.services.ai.core.messages import (
AgentMessage,
AssistantMessage,
ChatRequest,
ChatResponse,
LLMMessage,
ToolReturnPart,
)
from zhenxun.services.ai.core.models import CancellationToken, LLMContext, ToolChoice
from zhenxun.services.ai.core.options import GenerationConfig
from zhenxun.services.ai.core.protocols.tool import ToolExecutable
from zhenxun.services.ai.core.stream_events import (
LLMEndEvent,
LLMStartEvent,
ToolStreamChunkEvent,
)
from zhenxun.services.ai.flow.agent.models import AgentRunResources, AgentState
from zhenxun.services.ai.llm.engine.router import LLMOrchestrator
from zhenxun.services.ai.run import AgentRunResult, RunContext
from zhenxun.services.ai.run.models import OutputDataT
from zhenxun.services.ai.run.session import session_manager
from zhenxun.services.ai.tools.engine.executor import ToolExecutor
from zhenxun.services.ai.tools.models import ToolResult
from zhenxun.services.ai.utils.logger import log_agent as logger
from zhenxun.utils.pydantic_compat import model_construct
from .directive import (
DirectiveHandlerFunc,
directive_manager,
)
class BaseAgentExecutor(ABC):
"""
Agent 核心执行器基类 (Template Method Pattern)。
定义了基于生命周期的大模型控制流。第三方开发者可通过重写特定钩子,
"""
async def run(
self, state: AgentState, resources: AgentRunResources
) -> AgentRunResult[OutputDataT]:
"""
核心模板方法 (Template Method)。
组织整个大模型推导与工具调用的生命周期循环。如无必要,请勿重写此方法。
"""
await self.on_start(state, resources)
try:
cycle_count = 0
while cycle_count < resources.config.max_cycles:
state.current_cycle = cycle_count
state.should_reset_cycle = False
await self.on_cycle_start(state, resources)
await self.build_llm_request(state, resources)
await self.execute_llm(state, resources)
await self.handle_llm_response(state, resources)
if state.should_reset_cycle:
cycle_count = 0
continue
if state.is_finished:
assert state.pending_result is not None
return state.pending_result
await self.filter_tool_calls(state, resources)
if state.should_reset_cycle:
cycle_count = 0
continue
if state.is_finished:
assert state.pending_result is not None
return state.pending_result
await self.execute_tools(state, resources)
await self.handle_tool_results(state, resources)
if state.should_reset_cycle:
cycle_count = 0
continue
if state.is_finished:
assert state.pending_result is not None
return state.pending_result
cycle_count += 1
return await self.on_fallback(state, resources)
except Exception as e:
raise e
@abstractmethod
async def on_start(self, state: AgentState, resources: AgentRunResources) -> None:
"""生命周期: Agent 启动时调用,用于初始化状态或资源。"""
pass
@abstractmethod
async def on_cycle_start(
self, state: AgentState, resources: AgentRunResources
) -> None:
"""生命周期: 每次推理循环开始时调用。可用于 Token 预估或防死循环检测。"""
pass
@abstractmethod
async def build_llm_request(
self, state: AgentState, resources: AgentRunResources
) -> None:
"""生命周期: 构造请求大模型的 Messages 上下文和 Extra 参数。"""
pass
@abstractmethod
async def execute_llm(
self, state: AgentState, resources: AgentRunResources
) -> None:
"""生命周期: 触发大模型 API 请求并返回响应。"""
pass
@abstractmethod
async def handle_llm_response(
self, state: AgentState, resources: AgentRunResources
) -> None:
"""
生命周期: 处理大模型返回的结果,解析 Token 用量,
并将模型回复追加至对话历史。
"""
pass
@abstractmethod
async def filter_tool_calls(
self, state: AgentState, resources: AgentRunResources
) -> None:
"""生命周期: 从大模型的响应中提取并过滤出需要在本地客户端执行的工具调用请求。"""
pass
@abstractmethod
async def execute_tools(
self, state: AgentState, resources: AgentRunResources
) -> None:
"""生命周期: 并发执行提取出的工具,并收集结果或异常。"""
pass
@abstractmethod
async def handle_tool_results(
self, state: AgentState, resources: AgentRunResources
) -> None:
"""
生命周期: 处理工具返回的结果。
包括异常拦截、UI 渲染、Handoff 移交指令以及将结果追加至对话历史。
"""
pass
@abstractmethod
async def on_fallback(
self, state: AgentState, resources: AgentRunResources
) -> AgentRunResult[OutputDataT]:
"""生命周期: 当大模型思考循环达到 max_cycles 时触发,执行兜底策略。"""
pass
class StandardAgentExecutor(BaseAgentExecutor):
"""
LLM 任务执行器(核心推理引擎)。
负责:生命周期回调触发、工具循环调用、
错误反思(Reflexion)、Token消耗追踪。
"""
def __init__(
self, directive_handlers: dict[str, DirectiveHandlerFunc] | None = None
):
self.tool_executor = ToolExecutor()
self._directive_handlers: dict[str, DirectiveHandlerFunc] = (
directive_handlers or {}
)
@staticmethod
def _calculate_tool_overhead(tools: Iterable[Any] | None) -> int:
"""辅助方法:计算工具集合的 Schema Token 开销"""
if not tools:
return 0
overhead = 0
for t in tools:
if isinstance(t, dict) and "function" in t:
overhead += token_counter.count_tools_schema(
t["function"].get("parameters", {})
)
elif hasattr(t, "get_definition"):
t_def = getattr(t, "_dynamic_def", None)
if t_def and getattr(t_def, "parameters", None):
overhead += token_counter.count_tools_schema(t_def.parameters)
return overhead
def _can_retry_via_llm(self, result: ToolResult) -> bool:
"""通过新版的专属字段直接判断是否允许重试"""
return result.is_retryable
async def _check_follow_up(
self, state: AgentState, resources: AgentRunResources
) -> bool:
"""检查追加队列,排空并合并数据到上下文,更新状态机标志位"""
session_id = resources.run_context.session_id or "default_session"
session_info = await session_manager.get_or_create(session_id)
follow_ups = session_info.follow_up_queue.drain()
if follow_ups:
for fm in follow_ups:
resources.run_context.run.messages.append(
LLMMessage.user(f"💬 [用户追加指示]:{fm}")
)
resources.run_context.session.append_only_manager.sync_messages(
resources.run_context.run.messages
)
state.is_finished = False
state.pending_result = None
state.should_reset_cycle = True
return True
return False
async def _invoke_and_record_llm(
self,
state: AgentState,
resources: AgentRunResources,
messages: list[AgentMessage],
tools: list[ToolExecutable | dict[str, Any]] | None,
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
) -> ChatResponse:
"""执行 LLM 请求,处理基础指标遥测统计,并将新对话上下文追加到状态流"""
run_context = resources.run_context
cancellation_token = run_context.run.cancellation_token
current_extra = run_context.state.copy()
current_extra["__global_max_cycles__"] = getattr(
resources.config, "global_max_cycles", None
)
current_extra["__sys_capabilities"] = getattr(run_context, "capabilities", [])
current_extra["run_context"] = run_context
flattened_messages = ContextConverter.flatten_to_llm_messages(
messages, run_context
)
await run_context.run.emit(
LLMStartEvent(
model_name=resources.run_context.run.current_model or "unknown",
messages=flattened_messages,
)
)
response = await self._execute_model_request(
model_name=resources.run_context.run.current_model,
messages=flattened_messages,
config=resources.generation_config or GenerationConfig(),
run_context=run_context,
tools=tools,
tool_choice=tool_choice,
extra=current_extra,
cancellation_token=cancellation_token,
)
await run_context.run.emit(LLMEndEvent(response=response))
assistant_content = (
response.content_parts if response.content_parts else response.text
)
if response.thought_signature and isinstance(assistant_content, list):
for part in assistant_content:
if part.type == "thought":
if part.metadata is None:
part.metadata = {}
part.metadata["thought_signature"] = response.thought_signature
break
assistant_message = AssistantMessage(content=response.content_parts)
if hasattr(response, "parsed_obj") and response.parsed_obj is not None:
if not isinstance(response.parsed_obj, str):
if assistant_message.metadata is None:
assistant_message.metadata = {}
assistant_message.metadata["parsed_obj"] = response.parsed_obj
usage_obj = parse_usage_info(response.usage_info)
state.usage += usage_obj
if usage_obj.completion_tokens > 0:
assistant_message.token_cost = usage_obj.completion_tokens
if usage_obj.prompt_tokens > 0:
est_prompt_tokens = token_counter.count_context(
messages, resources.run_context.run.current_model or "", base_overhead=0
)
tool_overhead = self._calculate_tool_overhead(tools)
est_total = est_prompt_tokens + tool_overhead
state.token_drift = usage_obj.prompt_tokens - est_total
run_context.run.messages.append(assistant_message)
run_context.session.append_only_manager.sync_messages(run_context.run.messages)
return response
async def _execute_model_request(
self,
model_name: str | None,
messages: list[LLMMessage],
config: GenerationConfig,
run_context: RunContext,
tools: list[ToolExecutable | dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
extra: dict[str, Any] | None = None,
cancellation_token: CancellationToken | None = None,
) -> ChatResponse:
request = ChatRequest(
messages=messages,
config=config,
tools=tools,
tool_choice=tool_choice,
extra=extra or {},
)
sys_caps = request.extra.pop("__sys_capabilities", [])
llm_context = LLMContext(request=request, cancellation_token=cancellation_token)
combined_cap = CombinedCapability(sys_caps)
async def inner_handler(ctx: LLMContext[Any, Any]) -> ChatResponse:
return await LLMOrchestrator.invoke(
request=ctx.request,
model_name=model_name,
task="chat",
override_config=config,
cancellation_token=ctx.cancellation_token,
)
return await combined_cap.wrap_model_request(
run_context, llm_context, inner_handler
)
async def on_start(self, state: AgentState, resources: AgentRunResources) -> None:
pass
async def on_cycle_start(
self, state: AgentState, resources: AgentRunResources
) -> None:
state.current_request_messages = []
state.current_response = None
state.current_tool_calls = []
state.current_tool_results = []
state.should_reset_cycle = False
cancellation_token = resources.run_context.run.cancellation_token
if cancellation_token:
cancellation_token.raise_if_cancelled()
try:
base_overhead = 0
if state.static_system_prompt:
sp_list = (
state.static_system_prompt
if isinstance(state.static_system_prompt, list)
else [state.static_system_prompt]
)
for sp in sp_list:
if sp:
base_overhead += token_counter._count_text(str(sp))
for m in state.dynamic_system_messages:
base_overhead += token_counter.count_message(
m, resources.run_context.run.current_model or ""
)
base_overhead += self._calculate_tool_overhead(state.tools)
est_tokens = (
token_counter.count_context(
resources.run_context.run.messages,
resources.run_context.run.current_model or "",
base_overhead=base_overhead,
)
+ state.token_drift
)
est_tokens = max(est_tokens, 0)
logger.debug(
f"(Iter {state.current_cycle + 1}) "
f"预估将消耗 {est_tokens} Token "
f"(Model: {resources.run_context.run.current_model or 'Unknown'})"
)
except Exception:
pass
async def build_llm_request(
self, state: AgentState, resources: AgentRunResources
) -> None:
run_context = resources.run_context
session_info = await session_manager.get_or_create(
run_context.session_id or "default_session"
)
steer_msgs = session_info.steer_queue.drain()
if steer_msgs:
for sm in steer_msgs:
run_context.run.messages.append(
LLMMessage.user(f"💬 [用户实时修正指示]:{sm}")
)
run_context.session.append_only_manager.sync_messages(
run_context.run.messages
)
messages_to_send = []
if state.static_system_prompt:
if isinstance(state.static_system_prompt, list):
for sp in state.static_system_prompt:
if sp and sp.strip():
messages_to_send.append(LLMMessage.system(sp))
else:
if state.static_system_prompt and state.static_system_prompt.strip():
messages_to_send.append(
LLMMessage.system(state.static_system_prompt)
)
if state.dynamic_system_messages:
messages_to_send.extend(state.dynamic_system_messages)
if (
hasattr(run_context.run, "dynamic_prompts")
and run_context.run.dynamic_prompts
):
for prompt_text in run_context.run.dynamic_prompts.values():
if prompt_text and prompt_text.strip():
messages_to_send.append(LLMMessage.system(prompt_text))
messages_to_send.extend(run_context.run.messages)
state.current_request_messages = messages_to_send
async def execute_llm(
self, state: AgentState, resources: AgentRunResources
) -> None:
tools = state.tools
state.current_response = await self._invoke_and_record_llm(
state=state,
resources=resources,
messages=state.current_request_messages,
tools=list(tools) if tools else None,
tool_choice=None,
)
async def handle_llm_response(
self, state: AgentState, resources: AgentRunResources
) -> None:
response = state.current_response
if not response:
return
if not response.tool_calls:
logger.debug("✅ 模型未请求工具调用,推理循环结束。")
state.is_finished = True
state.pending_result = model_construct(
AgentRunResult,
output=response.text,
messages=list(resources.run_context.run.messages),
usage=state.usage,
)
if state.is_finished:
await self._check_follow_up(state, resources)
async def filter_tool_calls(
self, state: AgentState, resources: AgentRunResources
) -> None:
response = state.current_response
if not response:
return
tools = state.tools
event_bus = resources.run_context.run.event_bus
completed_call_ids = {
p.tool_call_id
for p in response.content_parts
if isinstance(p, ToolReturnPart)
}
client_tool_calls = []
for call in response.tool_calls:
tool_inst = tools.get(call.tool_name) if tools else None
is_server_side = call.id in completed_call_ids or (
tool_inst and getattr(tool_inst, "execution_side", "client") == "server"
)
if is_server_side:
logger.debug(
f"☁️ 检测到云端工具调用: {call.tool_name},已跳过本地执行。"
)
async with self.tool_executor._tool_stream_scope(
event_bus,
call.tool_name,
call.args if isinstance(call.args, dict) else {},
getattr(call, "intent", None),
) as box:
return_part = next(
(
p
for p in response.content_parts
if isinstance(p, ToolReturnPart)
and p.tool_call_id == call.id
),
None,
)
if return_part:
box["result"] = ToolResult(output=return_part.output)
else:
client_tool_calls.append(call)
if not client_tool_calls:
logger.info("✅ 无本地客户端工具需执行,推理循环平滑结束。")
state.is_finished = True
state.pending_result = model_construct(
AgentRunResult,
output=response.text,
messages=list(resources.run_context.run.messages),
usage=state.usage,
)
state.current_tool_calls = client_tool_calls
if state.is_finished:
await self._check_follow_up(state, resources)
async def execute_tools(
self, state: AgentState, resources: AgentRunResources
) -> None:
run_context = resources.run_context
tools = state.tools
event_bus = run_context.run.event_bus
tool_calls = state.current_tool_calls
if not tool_calls:
return
val_tasks = [
self.tool_executor.validate_tool_call(
call,
tools,
run_context,
event_bus=event_bus,
)
for call in tool_calls
]
validated_calls = await asyncio.gather(*val_tasks)
exec_tasks = [
self.tool_executor.execute_tool_call(
val_call,
tools,
run_context,
event_bus=event_bus,
)
for val_call in validated_calls
]
tool_results = await asyncio.gather(*exec_tasks, return_exceptions=True)
state.current_tool_results = tool_results
async def handle_tool_results(
self, state: AgentState, resources: AgentRunResources
) -> None:
"""处理所有工具执行结果,调度副作用指令并装配对话回传报文。"""
tool_calls = state.current_tool_calls
tool_results = state.current_tool_results
if not tool_calls or not tool_results:
return
for i, res_or_exc in enumerate(tool_results):
original_call = tool_calls[i]
tool_res = None
if not isinstance(res_or_exc, BaseException):
_, raw_tool_res = res_or_exc
tool_res = raw_tool_res
msg, usage = ToolExecutor.assemble_tool_message(
original_call, res_or_exc, tool_res
)
if usage:
state.usage += usage
resources.run_context.run.messages.append(msg)
if tool_res and getattr(tool_res, "directive", None):
ns = getattr(resources.run_context.session, "namespace", "global")
handler = directive_manager.get_handler(
tool_res.directive.name, namespace=ns
)
if handler:
await handler(state, resources, tool_res)
if state.is_finished:
resources.run_context.session.append_only_manager.sync_messages(
resources.run_context.run.messages
)
if state.is_finished:
await self._check_follow_up(state, resources)
return
else:
logger.warning(
f"⚠️ 未能找到名为 '{tool_res.directive.name}' "
f"的指令处理器 (Namespace: {ns})"
)
resources.run_context.session.append_only_manager.sync_messages(
resources.run_context.run.messages
)
if state.is_finished:
await self._check_follow_up(state, resources)
async def on_fallback(
self, state: AgentState, resources: AgentRunResources
) -> AgentRunResult[OutputDataT]:
run_context = resources.run_context
if not resources.config.enable_fallback_summary:
raise UpstreamServerException(
f"超过最大工具调用循环次数 ({resources.config.max_cycles})。",
)
logger.warning(
f"达到最大循环次数 ({resources.config.max_cycles}),触发兜底总结机制。"
)
await run_context.run.emit(
ToolStreamChunkEvent(
tool_name="System",
content="⏳ 思考过程过于复杂,正在强制生成最终总结...",
)
)
fallback_msg = LLMMessage.user(
"### 🚨 [系统强制指令]\n"
"你的任务执行已达到最大工具调用循环次数上限,当前思考流已被框架强制中断。\n"
"请**诚实地**向用户总结:你目前进行到了哪一步?遇到了什么困难导致循环耗尽?还有哪些预期步骤未能完成?\n"
"**绝对禁止**对用户撒谎声声称你已经完成了任务。严禁再次尝试调用任何工具!请直接输出纯文本结果。"
)
run_context.run.messages.append(fallback_msg)
fallback_response = await self._invoke_and_record_llm(
state=state,
resources=resources,
messages=run_context.run.messages,
tools=[],
tool_choice="none",
)
return cast(
AgentRunResult[OutputDataT],
model_construct(
AgentRunResult,
output=fallback_response.text,
messages=list(run_context.run.messages),
structured_data=None,
usage=state.usage,
),
)