♻️ 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>
This commit is contained in:
Rumio
2026-07-16 09:09:59 +08:00
committed by GitHub
co-authored by webjoin111
parent 52f7dbdedf
commit cd5fa065d3
13 changed files with 416 additions and 300 deletions
+51 -38
View File
@@ -2,18 +2,16 @@ from __future__ import annotations
import hashlib
import json
from typing import Any, cast
from typing import Any, ClassVar, cast
from zhenxun.models.user_console import UserConsole
from zhenxun.services.ai.config import get_llm_config
from zhenxun.services.ai.core.exceptions import (
AbortException,
ControlFlowExit,
GuardrailViolationError,
LLMException,
ModelRetry,
ResponseParseException,
SchemaParseError,
ToolFatalError,
ToolFinishException,
ToolRetryError,
@@ -279,18 +277,52 @@ class ReflexionCapability(AbstractCapability):
"""自愈反思与验证引擎 (Reflexion Engine)。
统一处理结构化解析失败和语义护栏拦截。"""
_PROMPT_TEMPLATES: ClassVar[dict[str, str]] = {
"schema_validation_error": (
"### ⚠️ [数据内容校验失败]\n"
"你输出的 JSON 格式完全正确,但部分字段的内容未能通过业务规则约束。\n\n"
"**失败详情:**\n"
"> {error_msg}\n\n"
"**修正要求:** 请仔细阅读上述失败详情,你必须打破先前的部分指令限制以满足上述规则," # noqa: E501
"调整报错字段的值并重新输出."
),
"schema_parse_error": (
"### ❌ [格式解析失败]\n"
"你输出的结构化数据(JSON)格式损坏或字段类型不匹配,未能通过校验。\n\n"
"**解析错误报告:**\n"
"> {error_msg}\n\n"
"**修正要求:** 请仔细检查缺失的必填字段、错误的数据类型或未闭合的括号,"
"严格参考你可用的 Schema 定义,重新输出正确格式的数据。"
),
"guardrail_violation": (
"### 🛡️ [业务护栏违规]\n"
"你输出的数据格式完全正确,但在业务逻辑层触发了合规/风控护栏。\n\n"
"**拦截原因报告:**\n"
"> {error_msg}\n\n"
"**修正要求:** 请结合上述反馈报告,反思你的决策逻辑或内容生成,"
"在保持数据格式正确的前提下,重新生成符合护栏规范的内容。"
),
"default_retry": (
"### ❌ [输出内容或格式验证失败]\n"
"你的上一次输出未能通过系统的校验与规则检查。请立即启动修正流程:\n\n"
"**错误反馈报告:**\n"
"> {error_msg}\n\n"
"**修正要求:** 请结合反馈报告,仔细反思你的输出内容或格式,"
"并重新生成正确的数据以满足所有的规则与规范。"
),
}
async def wrap_tool_execute(self, context, tool_name, arguments, handler):
try:
return await handler(arguments)
except Exception as error:
from zhenxun.services.ai.core.engine.structured_parser import (
DEFAULT_IVR_TEMPLATE,
)
from zhenxun.services.ai.tools.models import ToolResult
if isinstance(error, ToolRetryError | ModelRetry):
error_msg = getattr(error, "message", str(error))
feedback_prompt = DEFAULT_IVR_TEMPLATE.format(error_msg=error_msg)
feedback_prompt = self._PROMPT_TEMPLATES["default_retry"].format(
error_msg=error_msg
)
context.run.add_system_prompt(feedback_prompt)
return ToolResult(
output=f"执行失败:{error_msg}",
@@ -337,41 +369,22 @@ class ReflexionCapability(AbstractCapability):
self, e: Exception, error_msg: str, error_template: str | None
) -> str:
"""
根据不同的异常类型生成针对性的自愈反思提示词 (Feedback Prompt)
基于异常多态与模板字典生成自愈反思提示词
"""
if isinstance(e, SchemaParseError):
if "数据内容未通过规则校验" in error_msg:
return f"""### ⚠️ [数据内容校验失败]
你输出的 JSON 格式完全正确,但部分字段的内容未能通过业务规则约束。
if error_template:
return error_template.format(error_msg=error_msg)
**失败详情:**
> {error_msg}
template_name = (
e.get_template_name() if isinstance(e, ModelRetry) else "default_retry"
)
**修正要求:** 请仔细阅读上述失败详情,你必须打破先前的部分指令限制以满足上述规则,调整报错字段的值并重新输出。""" # noqa: E501
else:
return f"""### ❌ [格式解析失败]
你输出的结构化数据(JSON)格式损坏或字段类型不匹配,未能通过校验。
template = self._PROMPT_TEMPLATES.get(
template_name, self._PROMPT_TEMPLATES["default_retry"]
)
context_data = e.get_feedback_context() if isinstance(e, ModelRetry) else {}
context_data["error_msg"] = error_msg
**解析错误报告:**
> {error_msg}
**修正要求:** 请仔细检查缺失的必填字段、错误的数据类型或未闭合的括号,严格参考你可用的 Schema 定义,重新输出正确格式的数据。""" # noqa: E501
elif isinstance(e, GuardrailViolationError):
return f"""### 🛡️ [业务护栏违规]
你输出的数据格式完全正确,但在业务逻辑层触发了合规/风控护栏。
**拦截原因报告:**
> {error_msg}
**修正要求:** 请结合上述反馈报告,反思你的决策逻辑或内容生成,在保持数据格式正确的前提下,重新生成符合护栏规范的内容。""" # noqa: E501
else:
if error_template:
return error_template.format(error_msg=error_msg)
from zhenxun.services.ai.core.engine.structured_parser import (
DEFAULT_IVR_TEMPLATE,
)
return DEFAULT_IVR_TEMPLATE.format(error_msg=error_msg)
return template.format(**context_data)
async def wrap_model_request(
self,
@@ -4,6 +4,7 @@ from typing import Any, Generic, TypeVar
from pydantic import BaseModel, Field
from zhenxun.services.ai.config import get_llm_config
from zhenxun.services.ai.context.memory.models import MemoryConfig
from zhenxun.services.ai.core.engine.token_counter import token_counter
from zhenxun.services.ai.core.messages import (
@@ -15,7 +16,10 @@ from zhenxun.services.ai.core.messages import (
TextPart,
VideoPart,
)
from zhenxun.services.ai.core.models import ModelCapabilities
from zhenxun.services.ai.llm.api import chat, generate_structured
from zhenxun.services.ai.llm.manager import get_default_model
from zhenxun.services.ai.llm.system.capabilities import get_model_capabilities
from zhenxun.services.ai.utils.logger import log_memory as logger
from zhenxun.utils.pydantic_compat import model_copy
@@ -396,8 +400,6 @@ class LLMSummarizerReducer(AbstractSummarizerReducer):
prompt_text += f"[{speaker}]: {c_str}\n"
prompt_text += "</需要合并的旧对话记录>\n"
from zhenxun.services.ai.llm.api import chat
try:
model_to_use = self.summarization_model or get_default_model("chat")
response = await chat(
@@ -474,8 +476,6 @@ class StructuredSummaryReducer(AbstractSummarizerReducer, Generic[_T_Summary]):
prev_summary=prev_summary, dialogue=dialogue_text
)
from zhenxun.services.ai.llm.api import generate_structured
try:
model_to_use = self.summarization_model or get_default_model("chat")
summary_obj = await generate_structured(
@@ -513,15 +513,19 @@ class CondenserPipeline:
@classmethod
def create_from_configs(
cls, memory_config: MemoryConfig | None, model_name: str
cls,
memory_config: MemoryConfig | None,
capabilities: ModelCapabilities | None,
model_name: str,
) -> "CondenserPipeline":
"""基于全局和局部配置组装压缩管线工厂方法"""
from zhenxun.services.ai.config import get_llm_config
from zhenxun.services.ai.llm.system.capabilities import get_model_capabilities
config = get_llm_config().context_settings
pipeline_reducers = []
caps = get_model_capabilities(model_name)
caps = (
capabilities
if capabilities is not None
else get_model_capabilities(model_name)
)
vw = config.vision_window_size
if memory_config and memory_config.compression.vision_window is not None:
+7 -2
View File
@@ -3,6 +3,7 @@ from typing import cast
from zhenxun.services.ai.core.engine.context_renderer import ContextConverter
from zhenxun.services.ai.core.messages import AgentMessage
from zhenxun.services.ai.core.models import ModelCapabilities
from zhenxun.services.ai.run.context import RunContext
from zhenxun.services.ai.utils.logger import log_memory as logger
from zhenxun.utils.pydantic_compat import model_copy
@@ -42,7 +43,9 @@ class SessionMemoryContext:
async def read(
self,
model_name: str,
capabilities: ModelCapabilities | None = None,
override_history: Sequence[AgentMessage] | None = None,
base_overhead: int = 0,
) -> list[AgentMessage]:
"""
拉取短期对话历史,并执行 Token 压缩与管线修剪。
@@ -69,14 +72,16 @@ class SessionMemoryContext:
)
pipeline = CondenserPipeline.create_from_configs(
self.memory_config, model_name
self.memory_config, capabilities, model_name
)
if pipeline.reducers:
flattened_to_reduce = ContextConverter.flatten_to_llm_messages(
current_history
)
new_history, changed = await pipeline.run(
flattened_to_reduce, model_name=model_name, base_overhead=0
flattened_to_reduce,
model_name=model_name,
base_overhead=base_overhead,
)
if changed:
await chat_context.set_messages(self.session_meta, new_history)
@@ -8,21 +8,12 @@ from pydantic import BaseModel, Field, ValidationError, create_model
from zhenxun.services.ai.core.exceptions import (
SchemaParseError,
SchemaValidationError,
)
from zhenxun.services.ai.core.messages.types import OutputDataT
from zhenxun.services.ai.utils.logger import log_core as logger
from zhenxun.utils.pydantic_compat import model_json_schema, model_validate
DEFAULT_IVR_TEMPLATE = (
"### ❌ [输出内容或格式验证失败]\n"
"你的上一次输出未能通过系统的校验与规则检查。请立即启动修正流程:\n\n"
"**错误反馈报告:**\n"
"> {error_msg}\n\n"
"**修正要求:** 请结合反馈报告,"
"仔细反思你的输出内容或格式,\n"
"并重新生成正确的数据以满足所有的规则与规范。"
)
class BaseOutputProcessor(Generic[OutputDataT]):
"""
@@ -47,7 +38,7 @@ class BaseOutputProcessor(Generic[OutputDataT]):
如果不为 None 则跳过根据 response_model 生成,默认 None。
""" # noqa: E501
self.original_model = response_model
self.error_template = error_template or DEFAULT_IVR_TEMPLATE
self.error_template = error_template
self.raw_schema = raw_schema
self.target_model = None
self.is_union_wrapped = False
@@ -120,7 +111,7 @@ class BaseOutputProcessor(Generic[OutputDataT]):
msg = err.get("msg", "")
error_msgs.append(f"字段 `{loc}`: {msg}")
clean_error_str = "\n".join(error_msgs)
raise SchemaParseError(
raise SchemaValidationError(
f"数据内容未通过规则校验:\n{clean_error_str}"
)
@@ -3,6 +3,7 @@ LLM Token 动态预估与上下文管理模块
"""
from collections.abc import Sequence
import json
import math
import re
from typing import Any
@@ -56,17 +57,11 @@ class TokenCounter:
@classmethod
def count_tools_schema(cls, obj: dict | list | str | Any) -> int:
"""递归计算 JSON Schema 结构在被大模型作为工具时的 Token 开销。"""
if isinstance(obj, dict):
cost = len(obj.keys()) * 12
for k, v in obj.items():
if k == "description" and isinstance(v, str):
cost += int(len(v) * 0.3)
else:
cost += cls.count_tools_schema(v)
return cost
elif isinstance(obj, list):
return sum(cls.count_tools_schema(item) for item in obj)
return 0
try:
json_str = json.dumps(obj, ensure_ascii=False)
return int(cls._count_text(json_str) * 1.2) + 15
except Exception:
return 50
@classmethod
def count_message(cls, msg: LLMMessage, model_name: str) -> int:
+35 -6
View File
@@ -8,41 +8,70 @@ from typing import Any
class ModelRetry(Exception):
"""用于通知大模型修正并重试的异常"""
def __init__(self, message: str):
def __init__(self, message: str, payload: dict[str, Any] | None = None):
"""
初始化用于通知大模型重试的异常。
参数:
message: 用于提示大模型的具体重试和自我纠错信息。
payload: 携带产生错误时的上下文状态字典,用于渲染反馈提示词。
"""
self.message = message
super().__init__(message)
self.payload = payload or {}
super().__init__(self.message)
def get_template_name(self) -> str:
"""多态:子类告诉渲染引擎,自己应该使用哪个 Prompt 模板标识"""
return "default_retry"
def get_feedback_context(self) -> dict[str, Any]:
"""多态:子类提供渲染模板所需的领域数据上下文"""
return {"error_msg": self.message, **self.payload}
class SchemaParseError(ModelRetry):
"""格式解析异常。当大模型返回的 JSON 损坏或不符合 Schema 时抛出。"""
def __init__(self, message: str):
def __init__(self, message: str, payload: dict[str, Any] | None = None):
"""
初始化 Schema 格式解析错误异常。
参数:
message: 详细的 JSON 解析失败或 Schema 校验报错信息。
"""
super().__init__(message)
super().__init__(message, payload)
def get_template_name(self) -> str:
return "schema_parse_error"
class SchemaValidationError(ModelRetry):
"""数据校验异常。当大模型返回的 JSON 格式正确,但业务字段约束不满足时抛出。"""
def __init__(self, message: str, payload: dict[str, Any] | None = None):
"""
初始化 Schema 业务字段约束验证错误异常。
"""
super().__init__(message, payload)
def get_template_name(self) -> str:
return "schema_validation_error"
class GuardrailViolationError(ModelRetry):
"""护栏违规异常。当大模型返回的数据格式正确,但违反业务规则时抛出。"""
def __init__(self, message: str):
def __init__(self, message: str, payload: dict[str, Any] | None = None):
"""
初始化安全护栏校验未通过的异常。
参数:
message: 触发业务护栏违规拦截的详细原因说明。
"""
super().__init__(message)
super().__init__(message, payload)
def get_template_name(self) -> str:
return "guardrail_violation"
class ControlFlowExit(BaseException):
+29 -8
View File
@@ -10,6 +10,7 @@ from zhenxun.services.ai.config import get_llm_config
from zhenxun.services.ai.context.knowledge.base import BaseKnowledge
from zhenxun.services.ai.context.memory.builder import MemoryBuilder
from zhenxun.services.ai.context.memory.models import MemoryConfig
from zhenxun.services.ai.core.engine.token_counter import token_counter
from zhenxun.services.ai.core.exceptions import (
ControlFlowExit,
)
@@ -29,6 +30,7 @@ from zhenxun.services.ai.flow.core.base import BaseRunnable
from zhenxun.services.ai.flow.core.models import InterventionPolicy
from zhenxun.services.ai.guardrails import GuardrailSource, parse_guardrails
from zhenxun.services.ai.llm.builder import IntentBuilder
from zhenxun.services.ai.llm.manager import resolve_model_capabilities
from zhenxun.services.ai.message_builder import MessageBuilder
from zhenxun.services.ai.run import (
AgentRunResult,
@@ -700,6 +702,10 @@ class Agent(
self.model_name() if callable(self.model_name) else self.model_name
)
resources.model_capabilities = await resolve_model_capabilities(
context.run.current_model
)
return state, resources
async def on_context_build(
@@ -752,10 +758,29 @@ class Agent(
if tool_payload.injected_prompts:
static_prompts_list.extend(tool_payload.injected_prompts)
final_tools = await ToolBuilder.prepare_effective_tools(
effective_tools, context, self.tool_filters, run_scoped_cap
)
base_overhead = 0
for sp in static_prompts_list:
if sp:
base_overhead += token_counter._count_text(str(sp))
for m in dynamic_messages:
base_overhead += token_counter.count_message(
m, context.run.current_model or ""
)
for t in final_tools:
t_def = getattr(t, "_dynamic_def", None)
if t_def and getattr(t_def, "parameters", None):
base_overhead += token_counter.count_tools_schema(t_def.parameters)
messages_for_run = (
await memory_context.read(
model_name=context.run.current_model or "",
capabilities=resources.model_capabilities,
override_history=resources.config.message_history,
base_overhead=base_overhead,
)
if memory_context
else []
@@ -772,17 +797,14 @@ class Agent(
if resources.memory_context:
await resources.memory_context.write([msgs[-1]])
final_tools = await ToolBuilder.prepare_effective_tools(
effective_tools, context, self.tool_filters, run_scoped_cap
)
context.session.append_only_manager.build(static_prompts_list, final_tools)
context.session.append_only_manager.sync_messages(messages_for_run)
state.messages = messages_for_run
context.run.messages = messages_for_run
state.tools = final_tools
state.static_system_prompt = static_prompts_list
state.dynamic_system_messages = dynamic_messages
state.origin_msg_len = len(messages_for_run)
state._origin_msg_len = len(messages_for_run)
async def on_execute(
self, state: AgentState, resources: AgentRunResources
@@ -793,7 +815,7 @@ class Agent(
for tk in resources.toolkits:
if hasattr(tk, "before_llm_request"):
await DependencyInjector.invoke(
tk.before_llm_request, {"messages": state.messages}, context
tk.before_llm_request, {"messages": context.run.messages}, context
)
config_exec = resources.config.executor if resources.config else None
@@ -802,12 +824,11 @@ class Agent(
or self.executor
or StandardAgentExecutor(directive_handlers=self.directive_handlers)
)
resources.model_name = context.run.current_model
raw_result: AgentRunResult[Any] = await executor.run(
state=state, resources=resources
)
new_msgs = raw_result.messages[state.origin_msg_len :]
new_msgs = raw_result.messages[state._origin_msg_len :]
if resources.memory_context:
await resources.memory_context.write(new_msgs)
@@ -65,10 +65,10 @@ async def handle_submit_structured(
)
logger.debug("✅ 拦截到结构化结果提交,结束循环。")
state.is_finished = True
state.final_result = model_construct(
state.pending_result = model_construct(
AgentRunResult,
output=None,
messages=state.messages,
messages=list(resources.run_context.run.messages),
structured_data=parsed_obj,
usage=state.usage,
)
@@ -85,10 +85,10 @@ async def handle_end_run(
)
logger.debug("✅ 捕获到工具发出的终止信号,提前结束推理循环。")
state.is_finished = True
state.final_result = model_construct(
state.pending_result = model_construct(
AgentRunResult,
output=output,
messages=state.messages,
messages=list(resources.run_context.run.messages),
usage=state.usage,
)
@@ -106,10 +106,10 @@ async def handle_handoff(
output_text = f"已触发控制权移交 -> {handoff.target}。原因: {handoff.reason}"
logger.info(f"✅ 拦截到移交(Handoff)信号: 移交给 -> {handoff.target}。结束循环。")
state.is_finished = True
state.final_result = model_construct(
state.pending_result = model_construct(
AgentRunResult,
output=output_text,
messages=state.messages,
messages=list(resources.run_context.run.messages),
usage=state.usage,
handoff=handoff,
)
+131 -168
View File
@@ -1,8 +1,7 @@
from abc import ABC, abstractmethod
import asyncio
from collections.abc import Awaitable, Callable
import json
from typing import Any, Literal, cast
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
@@ -11,22 +10,15 @@ from zhenxun.services.ai.core.engine.token_counter import (
token_counter,
)
from zhenxun.services.ai.core.exceptions import (
ControlFlowExit,
UpstreamServerException,
)
from zhenxun.services.ai.core.messages import (
AgentMessage,
AssistantMessage,
AudioPart,
ChatRequest,
ChatResponse,
FilePart,
ImagePart,
LLMMessage,
TextPart,
ToolCallPart,
ToolReturnPart,
VideoPart,
)
from zhenxun.services.ai.core.models import CancellationToken, LLMContext, ToolChoice
from zhenxun.services.ai.core.options import GenerationConfig
@@ -40,11 +32,11 @@ 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 SessionInfo, session_manager
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 dump_json_safely, model_construct
from zhenxun.utils.pydantic_compat import model_construct
from .directive import (
DirectiveHandlerFunc,
@@ -68,29 +60,42 @@ class BaseAgentExecutor(ABC):
await self.on_start(state, resources)
try:
for cycle_index in range(resources.config.max_cycles):
state.current_cycle = cycle_index
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.final_result is not None
return state.final_result
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.final_result is not None
return state.final_result
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.final_result is not None
return state.final_result
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:
@@ -179,23 +184,45 @@ class StandardAgentExecutor(BaseAgentExecutor):
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
def _check_follow_up(
self, state: AgentState, resources: AgentRunResources, session_info: SessionInfo
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:
state.messages.append(LLMMessage.user(f"💬 [用户追加指示]:{fm}"))
resources.run_context.run.messages.append(
LLMMessage.user(f"💬 [用户追加指示]:{fm}")
)
resources.run_context.session.append_only_manager.sync_messages(
state.messages
resources.run_context.run.messages
)
state.is_finished = False
state.final_result = None
state.pending_result = None
state.should_reset_cycle = True
return True
return False
@@ -224,13 +251,13 @@ class StandardAgentExecutor(BaseAgentExecutor):
await run_context.run.emit(
LLMStartEvent(
model_name=resources.model_name or "unknown",
model_name=resources.run_context.run.current_model or "unknown",
messages=flattened_messages,
)
)
response = await self._execute_model_request(
model_name=resources.model_name,
model_name=resources.run_context.run.current_model,
messages=flattened_messages,
config=resources.generation_config or GenerationConfig(),
run_context=run_context,
@@ -266,8 +293,16 @@ class StandardAgentExecutor(BaseAgentExecutor):
if usage_obj.completion_tokens > 0:
assistant_message.token_cost = usage_obj.completion_tokens
state.messages.append(assistant_message)
run_context.session.append_only_manager.sync_messages(state.messages)
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
@@ -308,97 +343,53 @@ class StandardAgentExecutor(BaseAgentExecutor):
)
async def on_start(self, state: AgentState, resources: AgentRunResources) -> None:
resources.run_context.run.messages = state.messages
async def _execute_phase(
self,
phase_func: Callable[[AgentState, AgentRunResources], Awaitable[None]],
state: AgentState,
resources: AgentRunResources,
session_info: SessionInfo,
) -> Literal["NEXT", "RESET_CYCLE", "FINISH"]:
"""统一处理阶段执行与状态机完成/追加检查的高阶函数"""
await phase_func(state, resources)
if state.is_finished:
if self._check_follow_up(state, resources, session_info):
return "RESET_CYCLE"
return "FINISH"
return "NEXT"
async def run(
self, state: AgentState, resources: AgentRunResources
) -> AgentRunResult[OutputDataT]:
"""覆盖基类的模板方法,实现灵活的 while 控制流和 FOLLOW_UP 合并"""
await self.on_start(state, resources)
session_info = await session_manager.get_or_create(
resources.run_context.session_id or "default_session"
)
try:
cycle_count = 0
while cycle_count < resources.config.max_cycles:
state.current_cycle = cycle_count
await self.on_cycle_start(state, resources)
await self.build_llm_request(state, resources)
await self.execute_llm(state, resources)
status = await self._execute_phase(
self.handle_llm_response, state, resources, session_info
)
if status == "RESET_CYCLE":
cycle_count = 0
continue
if status == "FINISH":
assert state.final_result is not None
return state.final_result
status = await self._execute_phase(
self.filter_tool_calls, state, resources, session_info
)
if status == "RESET_CYCLE":
cycle_count = 0
continue
if status == "FINISH":
assert state.final_result is not None
return state.final_result
await self.execute_tools(state, resources)
status = await self._execute_phase(
self.handle_tool_results, state, resources, session_info
)
if status == "RESET_CYCLE":
cycle_count = 0
continue
if status == "FINISH":
assert state.final_result is not None
return state.final_result
cycle_count += 1
if self._check_follow_up(state, resources, session_info):
return await self.run(state, resources)
return await self.on_fallback(state, resources)
except Exception as e:
raise e
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:
est_tokens = token_counter.count_context(
state.messages, resources.model_name or "", base_overhead=0
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.model_name or 'Unknown'})"
f"(Model: {resources.run_context.run.current_model or 'Unknown'})"
)
except Exception:
pass
@@ -413,8 +404,12 @@ class StandardAgentExecutor(BaseAgentExecutor):
steer_msgs = session_info.steer_queue.drain()
if steer_msgs:
for sm in steer_msgs:
state.messages.append(LLMMessage.user(f"💬 [用户实时修正指示]:{sm}"))
run_context.session.append_only_manager.sync_messages(state.messages)
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:
@@ -439,7 +434,7 @@ class StandardAgentExecutor(BaseAgentExecutor):
if prompt_text and prompt_text.strip():
messages_to_send.append(LLMMessage.system(prompt_text))
messages_to_send.extend(state.messages)
messages_to_send.extend(run_context.run.messages)
state.current_request_messages = messages_to_send
@@ -466,13 +461,16 @@ class StandardAgentExecutor(BaseAgentExecutor):
if not response.tool_calls:
logger.debug("✅ 模型未请求工具调用,推理循环结束。")
state.is_finished = True
state.final_result = model_construct(
state.pending_result = model_construct(
AgentRunResult,
output=response.text,
messages=state.messages,
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:
@@ -522,14 +520,16 @@ class StandardAgentExecutor(BaseAgentExecutor):
logger.info("✅ 无本地客户端工具需执行,推理循环平滑结束。")
state.is_finished = True
state.final_result = model_construct(
state.pending_result = model_construct(
AgentRunResult,
output=response.text,
messages=state.messages,
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
@@ -565,51 +565,6 @@ class StandardAgentExecutor(BaseAgentExecutor):
tool_results = await asyncio.gather(*exec_tasks, return_exceptions=True)
state.current_tool_results = tool_results
def _assemble_tool_message(
self,
original_call: ToolCallPart,
res_or_exc: BaseException | tuple[ToolCallPart, ToolResult],
tool_res: ToolResult | None,
state: AgentState,
) -> LLMMessage:
"""负责处理异常、解析多模态、序列化,并装配为最终的工具消息载体"""
media_parts = []
final_content = "Success"
if isinstance(res_or_exc, BaseException):
if isinstance(res_or_exc, ControlFlowExit):
raise res_or_exc
final_content = json.dumps(
{"error": str(res_or_exc), "status": "failed"},
ensure_ascii=False,
)
elif tool_res is not None:
if isinstance(tool_res.output, list):
texts = []
for item in tool_res.output:
if isinstance(item, ImagePart | AudioPart | VideoPart | FilePart):
media_parts.append(item)
elif isinstance(item, TextPart):
texts.append(item.text)
else:
texts.append(str(item))
final_content = " ".join(texts) if texts else "Success"
elif isinstance(tool_res.output, str):
final_content = tool_res.output
else:
final_content = dump_json_safely(tool_res.output, ensure_ascii=False)
tool_usage = getattr(tool_res, "usage", None)
if tool_usage is not None:
state.usage += tool_usage
msg = LLMMessage.tool_response(
original_call.id, original_call.tool_name, final_content
)
if media_parts:
msg.content.extend(media_parts)
return msg
async def handle_tool_results(
self, state: AgentState, resources: AgentRunResources
) -> None:
@@ -627,10 +582,12 @@ class StandardAgentExecutor(BaseAgentExecutor):
_, raw_tool_res = res_or_exc
tool_res = raw_tool_res
msg = self._assemble_tool_message(
original_call, res_or_exc, tool_res, state
msg, usage = ToolExecutor.assemble_tool_message(
original_call, res_or_exc, tool_res
)
state.messages.append(msg)
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")
@@ -642,8 +599,10 @@ class StandardAgentExecutor(BaseAgentExecutor):
await handler(state, resources, tool_res)
if state.is_finished:
resources.run_context.session.append_only_manager.sync_messages(
state.messages
resources.run_context.run.messages
)
if state.is_finished:
await self._check_follow_up(state, resources)
return
else:
logger.warning(
@@ -651,7 +610,11 @@ class StandardAgentExecutor(BaseAgentExecutor):
f"的指令处理器 (Namespace: {ns})"
)
resources.run_context.session.append_only_manager.sync_messages(state.messages)
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
@@ -680,12 +643,12 @@ class StandardAgentExecutor(BaseAgentExecutor):
"请**诚实地**向用户总结:你目前进行到了哪一步?遇到了什么困难导致循环耗尽?还有哪些预期步骤未能完成?\n"
"**绝对禁止**对用户撒谎声声称你已经完成了任务。严禁再次尝试调用任何工具!请直接输出纯文本结果。"
)
state.messages.append(fallback_msg)
run_context.run.messages.append(fallback_msg)
fallback_response = await self._invoke_and_record_llm(
state=state,
resources=resources,
messages=state.messages,
messages=run_context.run.messages,
tools=[],
tool_choice="none",
)
@@ -695,7 +658,7 @@ class StandardAgentExecutor(BaseAgentExecutor):
model_construct(
AgentRunResult,
output=fallback_response.text,
messages=state.messages,
messages=list(run_context.run.messages),
structured_data=None,
usage=state.usage,
),
+23 -23
View File
@@ -4,7 +4,7 @@ from collections.abc import Sequence
from pathlib import Path
from typing import Any
from pydantic import BaseModel, ConfigDict, Field
from pydantic import BaseModel, ConfigDict, Field, PrivateAttr
from zhenxun.services.ai.capabilities import CapabilitySource, CombinedCapability
from zhenxun.services.ai.context.memory.builder import MemoryBuilder
@@ -18,10 +18,11 @@ from zhenxun.services.ai.core.messages import (
ToolCallPart,
UsageInfo,
)
from zhenxun.services.ai.core.models import ModelCapabilities
from zhenxun.services.ai.core.options import GenerationConfig
from zhenxun.services.ai.flow.core.models import BaseRuntimeConfig
from zhenxun.services.ai.run import RunContext
from zhenxun.services.ai.run.models import AgentRunResult, AgentTask, HandoffPayload
from zhenxun.services.ai.run.models import AgentRunResult, AgentTask
from zhenxun.services.ai.tools.core.toolkit import BaseToolkit
from zhenxun.services.ai.tools.engine.registry import ToolCollection
from zhenxun.utils.pydantic_compat import model_copy
@@ -121,27 +122,15 @@ class AgentState(BaseModel):
tools: ToolCollection | None = None
"""当前轮次生效的、已完成鉴权和过滤的工具集合"""
messages: list[AgentMessage] = Field(default_factory=list)
"""大模型将看到的完整历史消息列表 (执行历史)"""
usage: UsageInfo = Field(default_factory=UsageInfo)
"""累计的 Token 消耗"""
structured_result: Any | None = None
"""拦截到的结构化输出结果"""
early_result_output: Any | None = None
"""拦截到的早期终止输出结果"""
should_terminate: bool = False
"""标记是否应提前终止循环"""
handoff_triggered: HandoffPayload | None = None
"""标记是否触发了移交"""
is_finished: bool = False
"""标记大模型循环是否彻底结束"""
final_result: AgentRunResult[Any] | None = None
"""最终的运行结果 (AgentRunResult)"""
origin_msg_len: int = 0
"""初始进入循环时的消息历史长度 (用于增量保存记忆)"""
metadata: dict[str, Any] = Field(default_factory=dict)
"""供第三方开发者或生命周期钩子使用的自定义插槽,
在单个 Agent FSM 循环中存取临时变量"""
_origin_msg_len: int = PrivateAttr(default=0)
"""内部变量:初始进入循环时的消息历史长度 (用于增量保存记忆,防止被外部业务篡改)"""
current_cycle: int = 0
"""当前思考循环的轮次索引"""
current_request_messages: list[AgentMessage] = Field(default_factory=list)
"""当前即将发往大模型的实际请求消息"""
current_request_extra: dict[str, Any] = Field(default_factory=dict)
@@ -153,6 +142,17 @@ class AgentState(BaseModel):
current_tool_results: list[Any] = Field(default_factory=list)
"""当前轮次工具执行的结果或异常收集"""
should_reset_cycle: bool = False
"""标记是否需要重置当前思考循环(例如由于外部干预插入了新消息)"""
is_finished: bool = False
"""标记大模型循环是否彻底结束"""
pending_result: AgentRunResult[Any] | None = None
"""单一数据源:等待返回的最终运行结果 (SSOT)"""
usage: UsageInfo = Field(default_factory=UsageInfo)
"""累计的 Token 消耗"""
token_drift: int = 0
"""动态 Token 校准偏移量 (真实 - 预估)"""
class AgentRunResources(BaseModel):
"""大模型执行过程中的全局静态资源与配置载体"""
@@ -175,5 +175,5 @@ class AgentRunResources(BaseModel):
"""Agent 全局与运行时的统一策略配置"""
generation_config: GenerationConfig | None = None
"""大模型生成配置"""
model_name: str | None = None
"""当前实际调用的模型名称"""
model_capabilities: ModelCapabilities | None = None
"""当前底层大模型的能力配置(合并了用户自定义覆盖)"""
+36 -2
View File
@@ -11,11 +11,11 @@ from zhenxun.services.ai.config import (
get_llm_config,
)
from zhenxun.services.ai.core.exceptions import ConfigurationException
from zhenxun.services.ai.core.models import ModelDetail
from zhenxun.services.ai.core.models import ModelCapabilities, ModelDetail
from zhenxun.services.ai.core.options import GenerationConfig
from zhenxun.services.ai.utils.logger import log_llm as logger
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
from zhenxun.utils.pydantic_compat import model_dump
from zhenxun.utils.pydantic_compat import model_copy, model_dump
from .system.cache import clear_model_cache, get_or_create_model
from .system.capabilities import get_model_capabilities
@@ -191,6 +191,40 @@ def get_default_model(task: str = "chat") -> str | None:
return getattr(config.default_models, task, None)
async def resolve_model_capabilities(
provider_model_name: str | None = None, task: str = "chat"
) -> ModelCapabilities:
"""解析并合并带有用户自定义覆盖(如 max_input_tokens) 的模型能力。"""
resolved_name = provider_model_name
if resolved_name is None:
resolved_name = get_default_model(task)
if resolved_name is None:
avail = list_available_models()
if not avail:
return get_model_capabilities("unknown")
resolved_name = avail[0]["full_name"]
group_name = _get_group_name(resolved_name)
if group_name is not None:
model_names = _resolve_model_group(group_name)
if model_names:
resolved_name = model_names[0]
prov_name, mod_name = parse_provider_model_string(resolved_name)
caps = get_model_capabilities(mod_name or resolved_name)
if prov_name and mod_name:
config_tuple = find_model_config(prov_name, mod_name)
if config_tuple:
_, model_detail = config_tuple
if model_detail.max_input_tokens is not None:
caps = model_copy(
caps, update={"max_input_tokens": model_detail.max_input_tokens}
)
return caps
async def get_key_usage_stats() -> dict[str, Any]:
"""获取所有 Provider 的 Key 使用统计。"""
providers = get_configured_providers()
+12 -6
View File
@@ -7,6 +7,7 @@ from zhenxun.services.ai.core.models import (
)
from zhenxun.utils.pydantic_compat import model_copy
CTX_1_05M = 1_050_000
CTX_1M = 1_000_000
CTX_400K = 400_000
CTX_256K = 256_000
@@ -283,7 +284,17 @@ _ROUTING_TABLE: list[tuple[list[str], ModelCapabilities, int]] = [
(["*MiniMax-M3*"], CAP_MINIMAX_MULTIMODAL, CTX_1M),
(["mimo-v2.5-pro*", "mimo-v2-pro*", "mimo-v2-flash*"], CAP_MIMO_TEXT, CTX_1M),
(["mimo-v2.5", "mimo-v2-omni*"], CAP_MIMO_MULTIMODAL, CTX_1M),
(["gpt-5.5*", "gpt-5.4*"], CAP_OPENAI_MULTIMODAL, CTX_1M),
(
["gpt-5.3-chat*", "gpt-5.3-instant*", "*codex-spark*"],
CAP_OPENAI_MULTIMODAL,
CTX_128K,
),
(
["gpt-[5-9]*mini*", "gpt-[5-9]*nano*", "gpt-[5-9]*-instant*", "*codex*"],
CAP_OPENAI_MULTIMODAL,
CTX_400K,
),
(["gpt-[5-9]*"], CAP_OPENAI_MULTIMODAL, CTX_1_05M),
(["gemini-3*pro*"], CAP_GEMINI_3_PRO, CTX_1M),
(["gemini-3*"], CAP_GEMINI_3_FLASH, CTX_1M),
(
@@ -291,11 +302,6 @@ _ROUTING_TABLE: list[tuple[list[str], ModelCapabilities, int]] = [
CAP_GEMINI_2_5,
CTX_1M,
),
(
["gpt-5*", "gpt-5-mini*", "gpt-5-nano*", "*codex*"],
CAP_OPENAI_MULTIMODAL,
CTX_400K,
),
(
["kimi-k2.7*", "kimi-k2.6*", "kimi-k2.5*"],
DEFAULT_PERMISSIVE_CAPABILITIES,
+64 -9
View File
@@ -17,7 +17,18 @@ from zhenxun.services.ai.core.exceptions import (
ToolFatalError,
ToolRetryError,
)
from zhenxun.services.ai.core.messages import AnyLLMMessage, LLMMessage, ToolCallPart
from zhenxun.services.ai.core.messages import (
AnyLLMMessage,
AudioPart,
FilePart,
ImagePart,
LLMMessage,
TextPart,
ToolCallPart,
ToolMessage,
UsageInfo,
VideoPart,
)
from zhenxun.services.ai.core.stream_events import (
EventBus,
ToolCallEndEvent,
@@ -37,6 +48,7 @@ from zhenxun.services.ai.tools.models import (
ValidatedToolCall,
)
from zhenxun.services.ai.utils.logger import log_tool as logger
from zhenxun.utils.pydantic_compat import dump_json_safely
from .registry import ToolCollection
@@ -51,6 +63,49 @@ class ToolExecutor:
"""初始化工具执行器。"""
pass
@staticmethod
def assemble_tool_message(
original_call: ToolCallPart,
res_or_exc: BaseException | tuple[ToolCallPart, ToolResult],
tool_res: ToolResult | None,
) -> tuple[ToolMessage, UsageInfo | None]:
"""负责处理异常、解析多模态、序列化,并装配为最终的工具消息载体"""
media_parts = []
final_content = "Success"
usage = None
if isinstance(res_or_exc, BaseException):
if isinstance(res_or_exc, ControlFlowExit):
raise res_or_exc
final_content = json.dumps(
{"error": str(res_or_exc), "status": "failed"},
ensure_ascii=False,
)
elif tool_res is not None:
if isinstance(tool_res.output, list):
texts = []
for item in tool_res.output:
if isinstance(item, ImagePart | AudioPart | VideoPart | FilePart):
media_parts.append(item)
elif isinstance(item, TextPart):
texts.append(item.text)
else:
texts.append(str(item))
final_content = " ".join(texts) if texts else "Success"
elif isinstance(tool_res.output, str):
final_content = tool_res.output
else:
final_content = dump_json_safely(tool_res.output, ensure_ascii=False)
usage = getattr(tool_res, "usage", None)
msg = LLMMessage.tool_response(
original_call.id, original_call.tool_name, final_content
)
if media_parts:
msg.content.extend(media_parts)
return msg, usage
def _get_combined_capability(
self, executable: Any, context: RunContext
) -> CombinedCapability:
@@ -396,15 +451,15 @@ class ToolExecutor:
ToolResult(output=f"Crash: {result_pair}").as_error(),
)
tool_call_result = cast(tuple[ToolCallPart, ToolResult], result_pair)
_, tool_result = tool_call_result
tool_messages.append(
LLMMessage.tool_response(
tool_call_id=original_call.id,
function_name=func_name,
result=tool_result.output,
)
tool_res = None
if not isinstance(result_pair, BaseException):
_, raw_tool_res = cast(tuple[ToolCallPart, ToolResult], result_pair)
tool_res = raw_tool_res
msg, _ = ToolExecutor.assemble_tool_message(
original_call, result_pair, tool_res
)
tool_messages.append(msg)
return tool_messages