♻️ refactor(core): 重构 AI 编排框架与记忆及 RAG 子系统 (#2149)

* ♻️ refactor(core): 重构 AI 编排框架与记忆及 RAG 子系统

- 【重构】重构 `BaseRunnable` 并引入统一的 `RunIntent` 意图载体,规范 Agent、Team 和 Workflow 的执行流
- 【解耦】将中期记忆槽和长期向量记忆从 `MemoryConfig` 中解耦,转为独立的能力组件与工具箱进行管理
- 【记忆】移除 `MemoryReader` 和 `MemoryWriter`,统一封装为 `SessionMemoryContext` 会话记忆门面
- 【RAG】重构检索器与存储后端接口,统一采用 `QueryRequest` 进行多维度联合检索,并引入 `InMemoryScorer` 提升打分性能
- 【事件】优化 `EventBus` 异步事件分发机制,引入队列机制确保事件按序处理,避免并发竞态问题
- 【依赖注入】移除 `memory` 注入项,优化 `DependencyInjector` 的签名解析缓存以提升性能

* 🚨 auto fix by pre-commit hooks

---------

Co-authored-by: webjoin111 <455457521@qq.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
Rumio
2026-07-14 16:48:33 +08:00
committed by GitHub
co-authored by webjoin111 pre-commit-ci[bot]
parent 922d092650
commit 52f7dbdedf
66 changed files with 2131 additions and 2353 deletions
+65 -218
View File
@@ -1,6 +1,4 @@
import asyncio
from collections.abc import AsyncIterator, Callable, Sequence
import contextlib
from pathlib import Path
from typing import Any, Generic, cast
@@ -11,17 +9,11 @@ from zhenxun.services.ai.capabilities import (
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.capabilities import (
AgenticMemoryCapability,
SlotMemoryCapability,
)
from zhenxun.services.ai.context.memory.models import MemoryConfig
from zhenxun.services.ai.core.exceptions import (
ConcurrencyInterruptException,
ControlFlowExit,
)
from zhenxun.services.ai.core.messages import (
LLMMessage,
PromptInput,
UsageInfo,
)
@@ -33,8 +25,8 @@ from zhenxun.services.ai.core.options import (
from zhenxun.services.ai.core.protocols.tool import ToolExecutable, ToolResolvable
from zhenxun.services.ai.core.stream_events import AgentStreamEvent, EventBus
from zhenxun.services.ai.core.templates import PromptTemplate
from zhenxun.services.ai.flow.base import BaseRunnable, ConcurrencyPolicy
from zhenxun.services.ai.flow.concurrency import apply_concurrency_policy
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.message_builder import MessageBuilder
@@ -47,14 +39,9 @@ from zhenxun.services.ai.run.context import AgentDepsT
from zhenxun.services.ai.run.di import DependencyInjector
from zhenxun.services.ai.run.models import (
AgentRunEnd,
AgentRunError,
AgentRunStart,
OutputDataT,
StreamedRunResult,
)
from zhenxun.services.ai.run.subscribers import (
DefaultUISubscriber,
TelemetrySubscriber,
RunIntent,
)
from zhenxun.services.ai.tools.bridges.delegate import DelegateTool
from zhenxun.services.ai.tools.core.tool import BaseTool, FunctionTool
@@ -65,8 +52,6 @@ from zhenxun.services.ai.tools.providers.skills.capabilities import (
SkillCapability,
)
from zhenxun.services.ai.tools.providers.skills.models import Skill, SkillSource
from zhenxun.services.ai.utils import ContextUtils
from zhenxun.services.ai.utils.logger import log_agent as logger
from zhenxun.utils.pydantic_compat import (
model_construct,
model_copy,
@@ -104,8 +89,8 @@ class AgentBuilder(Generic[AgentDepsT, OutputDataT]):
def __init__(self, name: str):
self._kwargs: dict[str, Any] = {"name": name}
self._config: AgentConfig | dict | None = None
self._executor: Any | None = None
self._directive_handlers: dict[str, Any] = {}
self._executor: "BaseAgentExecutor | None" = None
self._directive_handlers: dict[str, Callable] = {}
def with_instruction(
self, instruction: str | PromptTemplate
@@ -203,7 +188,7 @@ class AgentBuilder(Generic[AgentDepsT, OutputDataT]):
配置对话记忆与上下文管理策略。
参数:
memory: 是否开启长期记忆与上下文压缩,支持布尔值或显式配置对象。
memory: 是否开启短期记忆与上下文压缩,支持布尔值或显式配置对象。
"""
self._kwargs["memory"] = memory
return self
@@ -220,7 +205,9 @@ class AgentBuilder(Generic[AgentDepsT, OutputDataT]):
self._kwargs["generation_config"] = config
return self
def with_intervention(self, policy: Any) -> "AgentBuilder[AgentDepsT, OutputDataT]":
def with_intervention(
self, policy: "InterventionPolicy | None"
) -> "AgentBuilder[AgentDepsT, OutputDataT]":
"""配置运行时消息干预策略。"""
if self._config is None:
self._config = AgentConfig()
@@ -349,7 +336,7 @@ class Agent(
name: str,
instruction: str | PromptTemplate = "",
description: str | None = None,
persona: Persona | dict | None = None,
persona: Persona | None = None,
model: str | Callable[[], str] | None = None,
tools: Sequence[ToolSource] | None = None,
skills: Sequence[str | Path | Skill | SkillSource] | None = None,
@@ -361,7 +348,7 @@ class Agent(
guardrails: list[GuardrailSource] | None = None,
capabilities: list[CapabilitySource] | None = None,
executor: BaseAgentExecutor | None = None,
directive_handlers: dict[str, Any] | None = None,
directive_handlers: dict[str, Callable] | None = None,
):
"""
初始化 Agent。
@@ -370,13 +357,13 @@ class Agent(
name: Agent 名称,用于日志、事件和链路标识。
instruction: 静态系统指令,可为普通字符串或模板字符串。
description: 智能体描述,用于外部路由节点决定是否调用。
persona: 角色设定配置,传入 dict 会自动构造成 Persona。
persona: 角色设定配置 (Persona 实例)。
model: 默认模型名称 (如 Provider/Model) 或返回模型名的回调。
tools: 初始工具定义列表,支持混合使用工具对象与字符串工具名。
skills: 注入的领域知识技能,支持 ID、目录 Path、Skill 对象或动态源。
generation_config: 默认生成配置,支持 GenerationConfig、IntentBuilder 或 dict。
response_model: 结构化输出模型,若为空则按纯文本输出。
memory: 是否开启长期记忆与上下文压缩,支持布尔值或 MemoryBuilder/Config。
memory: 是否开启短期记忆与上下文压缩,支持布尔值或 MemoryBuilder/Config。
knowledge: 挂载的知识库,支持单个或列表,底层自动将其注册入工具链。
config: 统一配置,合并了全局与单次运行策略,支持字典。
guardrails: 护栏定义列表,支持可调用对象、规则字符串或护栏实例。
@@ -388,18 +375,12 @@ class Agent(
if description:
self.description = description
elif persona:
p_obj = persona if isinstance(persona, Persona) else Persona(**persona)
self.description = f"角色:{p_obj.role},目标:{p_obj.goal}"
else:
self.description = str(instruction)[:150] if instruction else "AI Agent"
self.instruction = instruction
if isinstance(persona, dict):
self.persona = Persona(**persona)
else:
self.persona = persona
self.persona = persona
self.model_name = model
self.namespace = infer_plugin_namespace() or "unknown"
@@ -465,16 +446,6 @@ class Agent(
self.capabilities: list[CapabilitySource] = []
if self.memory_config.long_term.enable and self.memory_config.long_term.agentic:
self.capabilities.append(
AgenticMemoryCapability(self.memory_config, self.namespace)
)
if self.memory_config.slots.enable:
self.capabilities.append(
SlotMemoryCapability(self.memory_config, self.namespace)
)
if capabilities:
self.capabilities.extend(capabilities)
@@ -492,7 +463,7 @@ class Agent(
*,
name: str | None = None,
description: str | None = None,
settings: Any | None = None,
settings: ToolOptions | None = None,
):
"""
实例级工具注册装饰器。
@@ -556,7 +527,7 @@ class Agent(
return decorator if func is None else decorator(func)
def guardrail(self, func: Callable | str | Any | None = None):
def guardrail(self, func: GuardrailSource | None = None):
"""护栏装饰器/注册器 (支持传入函数或自然语言风控规则字符串)"""
if func is None:
@@ -614,26 +585,20 @@ class Agent(
**kwargs,
)
@contextlib.asynccontextmanager
async def run_stream(
async def _execute_stream(
self,
prompt: PromptInput | AgentTask | None = None,
*,
config: AgentConfig | dict | None = None,
deps: AgentDepsT | None = None,
context: RunContext[AgentDepsT] | None = None,
event_bus: EventBus | None = None,
intent: RunIntent,
context: RunContext[AgentDepsT],
cancel_token: CancellationToken,
event_bus: EventBus,
**kwargs: Any,
) -> AsyncIterator[StreamedRunResult[OutputDataT]]:
"""
智能体流式运行入口。
返回上下文管理器,可安全、解耦地获取底层事件或纯净文本结果。
"""
override_conf = (
AgentConfig(**config)
if isinstance(config, dict)
else (config or AgentConfig())
)
) -> AsyncIterator[AgentStreamEvent]:
raw_config = kwargs.pop("config", None)
if isinstance(raw_config, dict):
override_conf = AgentConfig(**raw_config)
else:
override_conf = raw_config or AgentConfig()
effective_config = self.config.merge_with(override_conf)
if effective_config.skills:
@@ -644,9 +609,6 @@ class Agent(
skills=effective_config.skills, namespace=infer_plugin_namespace()
)
)
bus = event_bus or EventBus()
TelemetrySubscriber().attach(bus)
if self._event_listeners:
for ev_type, callbacks in self._event_listeners.items():
@@ -655,130 +617,30 @@ class Agent(
def _make_handler(callback_func: Callable) -> Callable:
async def _di_handler(event: AgentStreamEvent):
await DependencyInjector.invoke(
callback_func, {"stream_event": event}, safe_context
callback_func, {"stream_event": event}, context
)
return _di_handler
bus.subscribe(ev_type, _make_handler(cb))
event_bus.subscribe(ev_type, _make_handler(cb))
if context is None:
explicit_session_id = kwargs.get("session_id")
safe_context = RunContext[AgentDepsT](session_id=explicit_session_id)
if deps is not None:
safe_context.deps = cast(AgentDepsT, deps)
else:
safe_context = context
if deps is not None and safe_context.deps is None:
safe_context.deps = cast(AgentDepsT, deps)
if safe_context.get_bot() and safe_context.get_event():
verbose_ui = effective_config.verbose_ui
DefaultUISubscriber(safe_context, verbose=verbose_ui).attach(bus)
policy = getattr(self.config, "concurrency_policy", None)
if policy is None:
policy = (
ConcurrencyPolicy.ALLOW
if getattr(self.config, "stateless", True)
else ConcurrencyPolicy.QUEUE
)
intervention_policy = getattr(self.config, "intervention_policy", None)
lock_id = ContextUtils.extract_concurrency_lock_id(
safe_context,
getattr(self.config, "concurrency_scope", None),
safe_context.session_id or "default_session",
)
async def _execution_task():
cancel_token = safe_context.run.cancellation_token or CancellationToken()
safe_context.run.cancellation_token = cancel_token
try:
async with apply_concurrency_policy(
session_id=safe_context.session_id or "default_session",
lock_id=lock_id,
policy=policy,
cancel_token=cancel_token,
intervention_policy=intervention_policy,
message=prompt,
):
await bus.emit(AgentRunStart(agent_name=self.name))
result = await self._run_step(
prompt=prompt,
context=safe_context,
config=effective_config,
cancellation_token=cancel_token,
event_bus=bus,
**kwargs,
)
await bus.emit(AgentRunEnd(result=result))
except ControlFlowExit as e:
await bus.emit(AgentRunError(error=e))
except asyncio.CancelledError:
logger.debug(f"Agent {self.name} 执行被并发策略中断取消。")
await bus.emit(
AgentRunError(
error=ConcurrencyInterruptException("任务已被新请求打断并接管")
)
)
except Exception as e:
await bus.emit(AgentRunError(error=e))
finally:
await bus.end()
task = asyncio.create_task(_execution_task())
result_obj = StreamedRunResult[OutputDataT](bus)
try:
yield result_obj
finally:
if not task.done():
task.cancel()
def _parse_task_prompt(
self, prompt: PromptInput | AgentTask | None
) -> tuple[AgentTask | None, Any | None, list[Any], Any, list[Any]]:
"""解析输入意图,提取数据契约 (AgentTask)"""
task_obj = None
final_prompt_payload = None
extra_tools = []
run_output_type = self.response_model
task_guardrails = []
if isinstance(prompt, AgentTask):
task_obj = prompt
if task_obj.response_model:
run_output_type = task_obj.response_model
if task_obj.tools:
extra_tools.extend(task_obj.tools)
if hasattr(task_obj, "_parsed_guardrails"):
task_guardrails.extend(task_obj._parsed_guardrails)
prompt_parts = [
f"### 📋 [任务指令]\n{task_obj.description}",
f"### 🎯 [预期产出要求]\n{task_obj.expected_output}",
]
final_prompt_payload = "\n\n".join(prompt_parts)
elif prompt is not None:
final_prompt_payload = prompt
return (
task_obj,
final_prompt_payload,
extra_tools,
run_output_type,
task_guardrails,
yield AgentRunStart(agent_name=self.name)
result = await self._run_step(
intent=intent,
context=context,
config=effective_config,
cancellation_token=cancel_token,
event_bus=event_bus,
**kwargs,
)
yield AgentRunEnd(result=result)
async def on_state_init(
self,
prompt: PromptInput | AgentTask | None = None,
intent: RunIntent,
context: RunContext[AgentDepsT] | None = None,
config: AgentConfig | None = None,
cancellation_token: Any = None,
cancellation_token: CancellationToken | None = None,
event_bus: EventBus | None = None,
**kwargs: Any,
) -> tuple[AgentState, AgentRunResources]:
@@ -789,19 +651,17 @@ class Agent(
if config is None:
config = AgentConfig()
(
task_obj,
final_prompt_payload,
extra_tools,
run_output_type,
task_guardrails,
) = self._parse_task_prompt(prompt)
task_obj = intent.task_obj
final_prompt_payload = intent.payload_to_render
extra_tools = intent.extra_tools
run_output_type = intent.response_model or self.response_model
task_guardrails = intent.guardrails
effective_memory = AgentProfileResolver.resolve_memory(
self.memory_config, config.memory
)
session_metadata, reader, writer = SessionBuilder.build_session_and_memory(
session_metadata, memory_context = SessionBuilder.build_session_and_memory(
context, self.namespace, self.name, effective_memory
)
@@ -821,8 +681,7 @@ class Agent(
resources = AgentRunResources(
run_context=context,
session_meta=session_metadata,
memory_reader=reader,
memory_writer=writer,
memory_context=memory_context,
run_scoped_cap=run_scoped_cap,
task_obj=task_obj,
config=config,
@@ -831,13 +690,7 @@ class Agent(
state.current_request_extra["final_prompt_payload"] = final_prompt_payload
state.current_request_extra["extra_tools"] = extra_tools
if final_prompt_payload is not None:
if isinstance(final_prompt_payload, str):
context.run.user_input = final_prompt_payload
elif hasattr(final_prompt_payload, "extract_plain_text"):
context.run.user_input = final_prompt_payload.extract_plain_text()
else:
context.run.user_input = str(final_prompt_payload)
context.run.user_input = intent.text
context.run.agent_name = self.name
context.run.cancellation_token = cancellation_token
@@ -855,7 +708,7 @@ class Agent(
"""装配记忆与提示词上下文、解析可用工具集"""
context = resources.run_context
reader = resources.memory_reader
memory_context = resources.memory_context
run_scoped_cap = (
resources.run_scoped_cap
if isinstance(resources.run_scoped_cap, CombinedCapability)
@@ -874,16 +727,6 @@ class Agent(
persona=cast(Persona | None, self.persona),
)
if reader:
long_term_fact = await reader.get_long_term_context(
context.run.user_input or ""
)
if long_term_fact:
dynamic_messages.append(LLMMessage.system(long_term_fact))
slots_fact = await reader.get_slots_context()
if slots_fact:
dynamic_messages.append(LLMMessage.system(slots_fact))
tool_payload = await ToolBuilder.resolve_tools(
tool_definitions=self.tool_definitions,
toolset_funcs=getattr(self, "toolset_funcs", []),
@@ -910,11 +753,11 @@ class Agent(
static_prompts_list.extend(tool_payload.injected_prompts)
messages_for_run = (
await reader.get_short_term_context(
await memory_context.read(
model_name=context.run.current_model or "",
override_history=resources.config.message_history,
)
if reader
if memory_context
else []
)
@@ -926,8 +769,8 @@ class Agent(
final_prompt_payload, bot=context.get_bot(), event=context.get_event()
):
messages_for_run.append(msgs[-1])
if resources.memory_writer:
await resources.memory_writer.save_new_messages([msgs[-1]])
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
@@ -960,14 +803,18 @@ class Agent(
or StandardAgentExecutor(directive_handlers=self.directive_handlers)
)
resources.model_name = context.run.current_model
raw_result: Any = await executor.run(state=state, resources=resources)
raw_result: AgentRunResult[Any] = await executor.run(
state=state, resources=resources
)
new_msgs = raw_result.messages[state.origin_msg_len :]
if resources.memory_writer:
await resources.memory_writer.save_new_messages(new_msgs)
if resources.memory_context:
await resources.memory_context.write(new_msgs)
final_output = getattr(raw_result, "output", None) or (
raw_result.messages[-1].extract_text if raw_result.messages else ""
getattr(raw_result.messages[-1], "extract_text", "")
if raw_result.messages
else ""
)
return cast(
@@ -984,17 +831,17 @@ class Agent(
async def _run_step(
self,
prompt: PromptInput | AgentTask | None = None,
intent: RunIntent,
*,
context: RunContext[AgentDepsT],
config: AgentConfig,
cancellation_token: Any = None,
cancellation_token: CancellationToken | None = None,
event_bus: EventBus | None = None,
**kwargs: Any,
) -> AgentRunResult[OutputDataT]:
"""原子步总管:将具体的生命周期方法编织为洋葱模型管道"""
state, resources = await self.on_state_init(
prompt,
intent,
context,
config,
cancellation_token,
-151
View File
@@ -1,151 +0,0 @@
import asyncio
from typing import Any, Generic, cast
from typing_extensions import TypeVar
import uuid
from nonebot.adapters import Bot, Event
from nonebot_plugin_alconna.uniseg import UniMessage
from zhenxun.services.ai.core.exceptions import (
ConcurrencyInterruptException,
ConcurrencyRejectException,
ControlFlowExit,
InterventionHandledException,
)
from zhenxun.services.ai.core.messages import UsageInfo
from zhenxun.services.ai.flow.base import BaseRunnable
from zhenxun.services.ai.run import AgentRunResult, RunContext
from zhenxun.services.ai.run.models import AgentRunEnd, AgentRunError
from zhenxun.services.ai.run.ui import UIController
from zhenxun.services.ai.utils.logger import log_agent as logger
from zhenxun.utils.message import MessageUtils
from zhenxun.utils.platform import PlatformUtils
T_Deps = TypeVar("T_Deps", default=Any)
T_Out = TypeVar("T_Out", default=str)
class AgentRunner(Generic[T_Out]):
"""
智能体运行器。
负责将大模型的纯净数据流包装为平台交互动作(发消息、UI渲染)。
自带 ContextVars 隐式上下文提取魔法。
"""
def __init__(
self,
runnable: BaseRunnable,
context: RunContext | None = None,
**kwargs: Any,
):
self.runnable = runnable
self.context = context or RunContext(**kwargs)
is_stateless = (
getattr(self.runnable.runtime_config, "stateless", True)
if hasattr(self.runnable, "runtime_config")
else True
)
if is_stateless and self.context.session_id:
if not self.context.session_id.startswith("stateless_"):
self.context.session_id = (
f"stateless_{self.context.session_id}_{uuid.uuid4().hex[:8]}"
)
if self.context.session:
self.context.session.session_id = self.context.session_id
@property
def _bot(self) -> Bot | None:
return self.context.get_bot()
@property
def _event(self) -> Event | None:
return self.context.get_event()
async def reply(
self, prompt: Any = None, reply_to: bool = False, **kwargs: Any
) -> AgentRunResult[T_Out]:
"""交互式执行:将 Agent 运行过程中的工具调用状态和最终结果自动发送给用户。"""
final_result = None
profile = kwargs.pop("profile", None)
try:
async with self.runnable.run_stream(
prompt=prompt,
context=self.context,
profile=profile,
**kwargs,
) as stream_result:
async for stream_event in stream_result.stream_events():
if isinstance(stream_event, AgentRunEnd):
final_result = stream_event.result
elif isinstance(stream_event, AgentRunError):
raise stream_event.error
except ControlFlowExit as e:
if isinstance(e, InterventionHandledException):
logger.info(f"✨ {self.runnable.name} 触发运行时干预: {e.message}")
if e.display_content and self._bot and self._event:
await MessageUtils.build_message(str(e.display_content)).send(
reply_to=reply_to
)
return cast(
AgentRunResult[T_Out], AgentRunResult(output="", usage=UsageInfo())
)
if isinstance(e, ConcurrencyRejectException):
logger.warning(
f"⏳ {self.runnable.name} 触发并发拒绝 (REJECT): {e.message}"
)
return cast(
AgentRunResult[T_Out], AgentRunResult(output="", usage=UsageInfo())
)
if isinstance(e, ConcurrencyInterruptException):
logger.warning(
f"🛑 {self.runnable.name} 触发并发中断 (INTERRUPT): {e.message}"
)
return cast(
AgentRunResult[T_Out], AgentRunResult(output="", usage=UsageInfo())
)
logger.debug(
f"{self.runnable.name} 控制流正常中断: {type(e).__name__} - {e}"
)
await UIController.handle_control_flow_exit_display(
e, self.context, reply_to
)
raise asyncio.CancelledError()
except Exception as e:
logger.error(f"{self.runnable.name} 运行失败: {e}", e=e)
if self._bot and self._event:
await MessageUtils.build_message(f"❌ 运行发生错误: {e}").send()
raise e
if final_result and final_result.output and self._bot:
msg_to_send = (
final_result.output
if isinstance(final_result.output, UniMessage)
else MessageUtils.build_message(str(final_result.output))
)
if self._event:
await msg_to_send.send(self._event, bot=self._bot, reply_to=reply_to)
else:
target = PlatformUtils.get_target(
user_id=self.context.get_user_id(),
group_id=self.context.get_group_id(),
)
if target:
await msg_to_send.send(target=target, bot=self._bot)
if isinstance(final_result.output, UniMessage):
final_result.output = final_result.output.extract_plain_text()
if final_result is None:
raise RuntimeError("智能体运行流异常结束:未返回最终结果。")
return cast(AgentRunResult[T_Out], final_result)
+98 -10
View File
@@ -1,24 +1,108 @@
import asyncio
from typing import Any, cast
from zhenxun.services.ai.capabilities import AbstractCapability
from zhenxun.services.ai.capabilities import AbstractCapability, WrapRunHandler
from zhenxun.services.ai.capabilities.base import CapabilityOrdering
from zhenxun.services.ai.core.engine.structured_parser import (
BaseOutputProcessor,
SubmitFinalResultExecutable,
)
from zhenxun.services.ai.core.exceptions import UpstreamServerException
from zhenxun.services.ai.core.exceptions import (
ControlFlowExit,
ModelRetry,
SchemaParseError,
UpstreamServerException,
)
from zhenxun.services.ai.core.messages import TaskLifecycleEvent
from zhenxun.services.ai.core.models import ToolDefinition
from zhenxun.services.ai.core.options import BaseOutputDefinition, ToolOutput
from zhenxun.services.ai.guardrails import parse_guardrails
from zhenxun.services.ai.guardrails import (
BaseGuardrail,
GuardrailSource,
parse_guardrails,
)
from zhenxun.services.ai.run import AgentRunResult, AgentTask, RunContext
from zhenxun.services.ai.tools.core.tool import BaseTool
from zhenxun.services.ai.tools.models import StructuredSubmissionResult, ToolResult
from zhenxun.services.ai.utils.logger import log_agent as logger
class SubmitFinalResultExecutable(BaseTool):
"""
动态生成的提交最终结果工具。
用于将大模型的结构化输出拦截并终止 AgentExecutor 的循环。
"""
def __init__(
self,
output_processor: BaseOutputProcessor,
guardrails: list[BaseGuardrail] | None = None,
):
"""
初始化提交最终结果的动态执行工具。
参数:
output_processor: 绑定的结构化输出处理器,用于验证提交的最终结果。
guardrails: 用于在结果输出前进行安全合规拦截的护栏中间件列表,默认 None。
"""
super().__init__(
name="submit_final_result",
description=(
"当你完成所有必要的调查 and 思考后,"
"必须且只能调用此工具来提交最终的结构化结果。"
"提交后任务将立刻结束。"
),
)
self.output_processor = output_processor
self.guardrails = guardrails or []
async def get_definition(
self, context: RunContext | None = None
) -> ToolDefinition | None:
if getattr(self, "_dynamic_def", None) is not None:
return self._dynamic_def
schema = self.output_processor.get_json_schema()
return ToolDefinition(
name=self.name,
description=self.description,
parameters=schema,
)
async def execute(self, context: RunContext | None = None, **kwargs) -> ToolResult:
parse_target = kwargs
if isinstance(kwargs, dict):
if "kwargs" in kwargs and len(kwargs) == 1:
parse_target = kwargs["kwargs"]
elif "result" in kwargs and len(kwargs) == 1:
parse_target = kwargs["result"]
try:
json_str = __import__("json").dumps(parse_target, ensure_ascii=False)
final_obj = await self.output_processor.validate_and_parse(
json_str, context=context
)
from zhenxun.services.ai.guardrails import GuardrailPipeline
pipeline = GuardrailPipeline(self.guardrails)
json_str, final_obj = await pipeline.run_output_pipeline(
json_str, final_obj, context
)
return StructuredSubmissionResult(
output="结构化数据已成功提交", parsed_obj=final_obj
)
except ControlFlowExit as e:
raise e
except ModelRetry as e:
raise e
except Exception as e:
error_msg = f"系统捕获到解析异常:\n{e}"
raise SchemaParseError(error_msg)
class OutputValidationCapability(AbstractCapability):
"""输出拦截与校验能力组件 (支持纯文本及结构化护栏)"""
def get_ordering(self) -> Any:
def get_ordering(self) -> CapabilityOrdering | None:
from zhenxun.services.ai.capabilities.builtin import (
ReflexionCapability,
)
@@ -27,8 +111,8 @@ class OutputValidationCapability(AbstractCapability):
def __init__(
self,
output_type: Any | None = None,
guardrails: list[Any] | None = None,
output_type: type[Any] | BaseOutputDefinition | None = None,
guardrails: list[GuardrailSource] | None = None,
raw_schema: dict[str, Any] | None = None,
):
self.output_type = output_type
@@ -83,7 +167,7 @@ class OutputValidationCapability(AbstractCapability):
]
return []
async def get_tools(self, context: RunContext) -> list[Any]:
async def get_tools(self, context: RunContext) -> list[BaseTool]:
"""动态挂载提交最终结果的工具"""
if self.submit_tool:
return [self.submit_tool]
@@ -95,7 +179,9 @@ class OutputValidationCapability(AbstractCapability):
llm_context.request.extra["guardrails"] = self.guardrails
return await handler(llm_context)
async def wrap_run(self, context: RunContext, handler: Any) -> AgentRunResult[Any]:
async def wrap_run(
self, context: RunContext, handler: WrapRunHandler
) -> AgentRunResult[Any]:
"""运行结束后,校验是否成功提取了结构化数据"""
result = await handler()
if self.output_type is not None or self.raw_schema is not None:
@@ -117,7 +203,9 @@ class TaskTrackingCapability(AbstractCapability):
self.task = task
self.agent_name = agent_name
async def wrap_run(self, context: RunContext, handler: Any) -> AgentRunResult[Any]:
async def wrap_run(
self, context: RunContext, handler: WrapRunHandler
) -> AgentRunResult[Any]:
"""任务生命周期追踪"""
task_name = self.task.name or self.task.id[:8]
@@ -9,25 +9,26 @@ from zhenxun.services.ai.capabilities import (
CombinedCapability,
)
from zhenxun.services.ai.context.memory.builder import MemoryBuilder
from zhenxun.services.ai.context.memory.engine import MemoryReader, MemoryWriter
from zhenxun.services.ai.context.memory.engine import SessionMemoryContext
from zhenxun.services.ai.context.memory.models import MemoryConfig
from zhenxun.services.ai.context.memory.types import SessionMetadata
from zhenxun.services.ai.core.messages import LLMMessage, TextPart
from zhenxun.services.ai.core.options import GenerationConfig
from zhenxun.services.ai.core.options import BaseOutputDefinition, GenerationConfig
from zhenxun.services.ai.core.protocols.tool import ToolExecutable
from zhenxun.services.ai.core.templates import PromptTemplate
from zhenxun.services.ai.flow.agent.capabilities import (
OutputValidationCapability,
TaskTrackingCapability,
)
from zhenxun.services.ai.flow.agent.models import Persona
from zhenxun.services.ai.run import RunContext
from zhenxun.services.ai.run import AgentTask, RunContext
from zhenxun.services.ai.run.di import DependencyInjector
from zhenxun.services.ai.tools.engine.registry import (
ToolCollection,
tool_provider_manager,
)
from zhenxun.services.ai.tools.models import ResolvedToolPayload
from zhenxun.services.ai.utils.scope import ScopeSelector
from zhenxun.services.ai.utils.runtime import ContextUtils
from zhenxun.utils.pydantic_compat import model_copy
@@ -36,7 +37,8 @@ class AgentProfileResolver:
@staticmethod
def resolve_memory(
agent_memory_config: MemoryConfig, override_memory: Any | None
agent_memory_config: MemoryConfig,
override_memory: bool | MemoryConfig | MemoryBuilder | None,
) -> MemoryConfig:
"""
解析并合并 Memory 记忆域的配置。
@@ -86,11 +88,11 @@ class CapabilityBuilder:
async def build_for_run(
agent_name: str,
namespace: str,
output_type: Any | None,
output_type: type[Any] | BaseOutputDefinition | None,
raw_schema: dict | None,
agent_guardrails: list,
task_guardrails: list,
task_obj: Any | None,
task_obj: AgentTask | None,
agent_capabilities: list,
profile_capabilities: list | None,
context: RunContext,
@@ -160,11 +162,11 @@ class ContextBuilder:
@staticmethod
async def build_prompts(
instruction: str | PromptTemplate,
system_prompts: list[Any],
system_prompts: list[Callable],
run_context: RunContext,
run_scoped_cap: CombinedCapability,
persona: Persona | None = None,
) -> tuple[str, list[Any]]:
) -> tuple[str, list[LLMMessage]]:
"""
解析、合并并渲染 Agent 的系统提示词和上下文记忆。
包含对依赖参数的动态注入和 Jinja 模板渲染。
@@ -291,9 +293,9 @@ class ToolBuilder:
@staticmethod
async def resolve_tools(
tool_definitions: list[Any],
toolset_funcs: list[Any],
system_tools: list[Any],
tool_definitions: list[ToolExecutable | Callable | dict[str, Any] | str],
toolset_funcs: list[Callable],
system_tools: list[ToolExecutable | Callable | dict[str, Any] | str],
namespace: str,
run_context: RunContext,
run_scoped_cap: CombinedCapability,
@@ -359,7 +361,7 @@ class ToolBuilder:
@staticmethod
async def prepare_effective_tools(
effective_tools: list[Any],
effective_tools: list[ToolExecutable],
context: RunContext,
tool_filters: list[Callable],
run_scoped_cap: CombinedCapability,
@@ -407,11 +409,15 @@ class ToolBuilder:
final_defs_map = {d.name.lower(): d for d in current_tool_defs if d}
final_effective_tools_list = []
seen_names = set()
for t_exec in effective_tools:
t_name = getattr(t_exec, "name", "unknown")
if t_name.lower() in final_defs_map:
t_name_lower = t_name.lower()
if t_name_lower in final_defs_map and t_name_lower not in seen_names:
seen_names.add(t_name_lower)
cloned_tool = copy.copy(t_exec)
cloned_tool._dynamic_def = final_defs_map[t_name.lower()]
setattr(cloned_tool, "_dynamic_def", final_defs_map[t_name_lower])
final_effective_tools_list.append(cloned_tool)
return ToolCollection(final_effective_tools_list)
@@ -425,7 +431,7 @@ class SessionBuilder:
namespace: str,
agent_name: str,
effective_memory: MemoryConfig,
) -> tuple[Any, Any, Any]:
) -> tuple[SessionMetadata, SessionMemoryContext]:
"""
根据当前用户、群组和平台标识,动态隔离前缀,计算并构建会话与记忆存储的读写门面。
@@ -436,78 +442,21 @@ class SessionBuilder:
effective_memory: 运行时最终生效的 MemoryConfig 配置对象。
返回:
tuple[Any, Any, Any]: 包含 (SessionMetadata 会话元数据, MemoryReader 记忆读取器, MemoryWriter 记忆写入器) 的元组。
tuple[SessionMetadata, SessionMemoryContext]: 包含 (会话元数据, 会话记忆门面上下文) 的元组。
""" # noqa: E501
bot_id = None
bot_inst = context.get_bot()
if bot_inst and hasattr(bot_inst, "self_id"):
bot_id = str(bot_inst.self_id)
selector = ScopeSelector(
user_id=context.get_user_id(),
group_id=context.get_group_id(),
platform=context.get_platform(),
bot_id=bot_id,
namespace=namespace,
agent_name=agent_name,
)
all_scopes = {"/"}
scope_name_mapping = {}
target_builder = None
if effective_memory.short_term and effective_memory.short_term.isolation:
sel = effective_memory.short_term.isolation.resolve(
deps=context.deps,
prefix="",
default_namespace=namespace,
default_agent=agent_name,
)
all_scopes.add(sel.scope_prefix)
target_builder = effective_memory.short_term.isolation
for config_part in [effective_memory.slots, effective_memory.long_term]:
if config_part and hasattr(config_part, "scopes") and config_part.scopes:
for name, builder in config_part.scopes.items():
sel = builder.resolve(
deps=context.deps,
prefix="",
default_namespace=namespace,
default_agent=agent_name,
)
all_scopes.add(sel.scope_prefix)
scope_name_mapping[sel.scope_prefix] = name
parts = selector.get_scope_parts()
for i in range(len(parts)):
all_scopes.add("/" + "/".join(parts[: i + 1]))
accessible_scopes = sorted(all_scopes, key=lambda x: len(x.split("/")))
short_term_builder = (
effective_memory.short_term.isolation
if effective_memory.short_term
else effective_memory.base_isolation
)
short_term_selector = short_term_builder.resolve(
deps=context.deps,
prefix="",
default_namespace=namespace,
default_agent=agent_name,
session_metadata = ContextUtils.build_session_meta(
context=context, target_builder=target_builder, custom_namespace=namespace
)
context.session.session_meta = session_metadata
session_metadata = SessionMetadata(
session_id=short_term_selector.scope_prefix,
selector=selector,
scope_prefix=selector.scope_prefix,
accessible_scopes=accessible_scopes,
scope_name_mapping=scope_name_mapping,
)
reader = MemoryReader(
session_meta=session_metadata, memory_config=effective_memory
)
writer = MemoryWriter(
memory_context = SessionMemoryContext(
session_meta=session_metadata,
memory_config=effective_memory,
context=context,
)
return session_metadata, reader, writer
return session_metadata, memory_context
@@ -1,7 +1,8 @@
from abc import ABC, abstractmethod
import asyncio
from collections.abc import Awaitable, Callable
import json
from typing import Any
from typing import Any, Literal, cast
from zhenxun.services.ai.capabilities import CombinedCapability
from zhenxun.services.ai.core.engine.context_renderer import ContextConverter
@@ -27,8 +28,9 @@ from zhenxun.services.ai.core.messages import (
ToolReturnPart,
VideoPart,
)
from zhenxun.services.ai.core.models import LLMContext
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,
@@ -37,7 +39,8 @@ from zhenxun.services.ai.core.stream_events import (
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.session import session_manager
from zhenxun.services.ai.run.models import OutputDataT
from zhenxun.services.ai.run.session import SessionInfo, 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
@@ -57,7 +60,7 @@ class BaseAgentExecutor(ABC):
async def run(
self, state: AgentState, resources: AgentRunResources
) -> AgentRunResult[Any]:
) -> AgentRunResult[OutputDataT]:
"""
核心模板方法 (Template Method)。
组织整个大模型推导与工具调用的生命周期循环。如无必要,请勿重写此方法。
@@ -156,7 +159,7 @@ class BaseAgentExecutor(ABC):
@abstractmethod
async def on_fallback(
self, state: AgentState, resources: AgentRunResources
) -> AgentRunResult[Any]:
) -> AgentRunResult[OutputDataT]:
"""生命周期: 当大模型思考循环达到 max_cycles 时触发,执行兜底策略。"""
pass
@@ -181,7 +184,7 @@ class StandardAgentExecutor(BaseAgentExecutor):
return result.is_retryable
def _check_follow_up(
self, state: AgentState, resources: AgentRunResources, session_info: Any
self, state: AgentState, resources: AgentRunResources, session_info: SessionInfo
) -> bool:
"""检查追加队列,排空并合并数据到上下文,返回是否发现新消息"""
follow_ups = session_info.follow_up_queue.drain()
@@ -201,8 +204,8 @@ class StandardAgentExecutor(BaseAgentExecutor):
state: AgentState,
resources: AgentRunResources,
messages: list[AgentMessage],
tools: list[Any] | None,
tool_choice: Any = None,
tools: list[ToolExecutable | dict[str, Any]] | None,
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
) -> ChatResponse:
"""执行 LLM 请求,处理基础指标遥测统计,并将新对话上下文追加到状态流"""
run_context = resources.run_context
@@ -274,10 +277,10 @@ class StandardAgentExecutor(BaseAgentExecutor):
messages: list[LLMMessage],
config: GenerationConfig,
run_context: RunContext,
tools: list[Any] | None = None,
tool_choice: Any = None,
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: Any = None,
cancellation_token: CancellationToken | None = None,
) -> ChatResponse:
request = ChatRequest(
messages=messages,
@@ -307,9 +310,24 @@ 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[Any]:
) -> AgentRunResult[OutputDataT]:
"""覆盖基类的模板方法,实现灵活的 while 控制流和 FOLLOW_UP 合并"""
await self.on_start(state, resources)
session_info = await session_manager.get_or_create(
@@ -325,29 +343,35 @@ class StandardAgentExecutor(BaseAgentExecutor):
await self.build_llm_request(state, resources)
await self.execute_llm(state, resources)
await self.handle_llm_response(state, resources)
if state.is_finished:
if self._check_follow_up(state, resources, session_info):
cycle_count = 0
continue
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
await self.filter_tool_calls(state, resources)
if state.is_finished:
if self._check_follow_up(state, resources, session_info):
cycle_count = 0
continue
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)
await self.handle_tool_results(state, resources)
if state.is_finished:
if self._check_follow_up(state, resources, session_info):
cycle_count = 0
continue
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
@@ -544,7 +568,7 @@ class StandardAgentExecutor(BaseAgentExecutor):
def _assemble_tool_message(
self,
original_call: ToolCallPart,
res_or_exc: Any,
res_or_exc: BaseException | tuple[ToolCallPart, ToolResult],
tool_res: ToolResult | None,
state: AgentState,
) -> LLMMessage:
@@ -631,7 +655,7 @@ class StandardAgentExecutor(BaseAgentExecutor):
async def on_fallback(
self, state: AgentState, resources: AgentRunResources
) -> AgentRunResult[Any]:
) -> AgentRunResult[OutputDataT]:
run_context = resources.run_context
if not resources.config.enable_fallback_summary:
@@ -666,10 +690,13 @@ class StandardAgentExecutor(BaseAgentExecutor):
tool_choice="none",
)
return model_construct(
AgentRunResult,
output=fallback_response.text,
messages=state.messages,
structured_data=None,
usage=state.usage,
return cast(
AgentRunResult[OutputDataT],
model_construct(
AgentRunResult,
output=fallback_response.text,
messages=state.messages,
structured_data=None,
usage=state.usage,
),
)
+6 -8
View File
@@ -8,7 +8,7 @@ from pydantic import BaseModel, ConfigDict, Field
from zhenxun.services.ai.capabilities import CapabilitySource, CombinedCapability
from zhenxun.services.ai.context.memory.builder import MemoryBuilder
from zhenxun.services.ai.context.memory.engine import MemoryReader, MemoryWriter
from zhenxun.services.ai.context.memory.engine import SessionMemoryContext
from zhenxun.services.ai.context.memory.models import MemoryConfig
from zhenxun.services.ai.context.memory.types import SessionMetadata
from zhenxun.services.ai.core.messages import (
@@ -19,7 +19,7 @@ from zhenxun.services.ai.core.messages import (
UsageInfo,
)
from zhenxun.services.ai.core.options import GenerationConfig
from zhenxun.services.ai.flow.base import BaseRuntimeConfig
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.tools.core.toolkit import BaseToolkit
@@ -43,7 +43,7 @@ class Persona(BaseModel):
class AgentConfig(BaseRuntimeConfig):
"""统一的智能体全局与单次运行配置 (Unification of Settings & Profile)"""
"""统一的智能体全局与单次运行配置"""
model_config = ConfigDict(arbitrary_types_allowed=True)
@@ -66,7 +66,7 @@ class AgentConfig(BaseRuntimeConfig):
"""初始化的底层对话历史记录。"""
memory: MemoryConfig | MemoryBuilder | bool | None = Field(default=None)
"""单次运行级别的记忆门面覆盖 (支持 bool, MemoryConfig, MemoryBuilder)。"""
"""单次运行级别的记忆门面覆盖"""
generation_config: GenerationConfig | None = Field(default=None)
"""单次运行覆盖的大模型生成配置。"""
capabilities: list[CapabilitySource] | None = Field(default=None)
@@ -163,10 +163,8 @@ class AgentRunResources(BaseModel):
"""保留依赖注入(DI)与黑板引用的全局运行时上下文"""
session_meta: SessionMetadata | None = None
"""隔离会话的元信息(Session ID, 命名空间, 权限等)"""
memory_reader: MemoryReader | None = None
"""用于读取短/中/长期上下文记忆的读取器"""
memory_writer: MemoryWriter | None = None
"""用于将对话历史安全落盘的写入器"""
memory_context: SessionMemoryContext | None = None
"""统一处理对话历史读写、压缩与清洗的会话记忆门面"""
run_scoped_cap: CombinedCapability | None = None
"""聚合了 Agent/AgentTask/全局 的复合能力拦截器 (CombinedCapability)"""
task_obj: AgentTask | None = None