mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-11 15:00:00 +08:00
♻️ 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:
co-authored by
webjoin111
pre-commit-ci[bot]
parent
922d092650
commit
52f7dbdedf
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user