diff --git a/zhenxun/services/ai/capabilities/__init__.py b/zhenxun/services/ai/capabilities/__init__.py
index 1fe4464b..7aa07458 100644
--- a/zhenxun/services/ai/capabilities/__init__.py
+++ b/zhenxun/services/ai/capabilities/__init__.py
@@ -21,3 +21,5 @@ __all__ = [
"WrapperCapability",
"capability",
]
+
+from . import builtin # noqa: F401
diff --git a/zhenxun/services/ai/capabilities/manager.py b/zhenxun/services/ai/capabilities/manager.py
index ce84382b..470ac82d 100644
--- a/zhenxun/services/ai/capabilities/manager.py
+++ b/zhenxun/services/ai/capabilities/manager.py
@@ -193,11 +193,7 @@ def capability(
"""装饰器内部函数,实现类注册"""
final_name = name or cls.__name__
final_tags = tags or []
- ns = (
- namespace
- if namespace is not None
- else infer_plugin_namespace(default="global")
- )
+ ns = namespace if namespace is not None else infer_plugin_namespace()
capability_manager.register(
cls, name=final_name, namespace=ns, tags=final_tags, auto_apply=auto_apply
)
diff --git a/zhenxun/services/ai/context/__init__.py b/zhenxun/services/ai/context/__init__.py
index 6c951fa7..fa4a5caf 100644
--- a/zhenxun/services/ai/context/__init__.py
+++ b/zhenxun/services/ai/context/__init__.py
@@ -1,17 +1,15 @@
"""
-Zhenxun AI - 上下文、记忆与知识管理子系统门面 (Context, Memory & Knowledge Facade)
+Zhenxun AI - 上下文、记忆与知识管理子系统门面
"""
from .knowledge import FileSystemKnowledge, VectorKnowledge
from .memory import (
- AgentSessionFacade,
MemoryBuilder,
memory_manager,
)
from .rag import RAGBuilder
__all__ = [
- "AgentSessionFacade",
"FileSystemKnowledge",
"MemoryBuilder",
"RAGBuilder",
diff --git a/zhenxun/services/ai/context/knowledge/vector.py b/zhenxun/services/ai/context/knowledge/vector.py
index cbeb5015..beed7ace 100644
--- a/zhenxun/services/ai/context/knowledge/vector.py
+++ b/zhenxun/services/ai/context/knowledge/vector.py
@@ -5,6 +5,7 @@ import anyio
from nonebot.adapters import Bot, Event
from pydantic import BaseModel, Field
+from zhenxun.services.ai.context.rag.backends import StorageBackend
from zhenxun.services.ai.context.rag.engine import ScopedRAGClient
from zhenxun.services.ai.context.rag.models import BaseRecord
from zhenxun.services.ai.core.messages import LLMMessage
@@ -55,7 +56,7 @@ class VectorKnowledge(BaseKnowledge):
"{knowledge_text}"
)
- _global_storage: Any = None
+ _global_storage: StorageBackend | None = None
def __init__(
self,
@@ -194,7 +195,7 @@ class VectorKnowledge(BaseKnowledge):
return super().get_instructions()
async def before_llm_request(
- self, context: RunContext, messages: list[Any]
+ self, context: RunContext, messages: list[LLMMessage]
) -> None:
"""
生命周期钩子:在向底层 LLM 发起请求前触发。
@@ -285,13 +286,27 @@ class VectorKnowledge(BaseKnowledge):
return await self.rag_client.ingest([doc])
async def add_directory(self, dir_path: str | Path) -> int:
- """扫描目录并注入所有支持的文件"""
- total_chunks = 0
+ """扫描目录并批量注入所有支持的文件"""
aio_path = anyio.Path(dir_path)
+ docs_to_ingest = []
+
async for p in aio_path.rglob("*"):
if await p.is_file():
- total_chunks += await self.add_file(Path(p))
- return total_chunks
+ std_path = Path(p)
+ ext = std_path.suffix.lower()
+ reader = self.readers.get(ext)
+ if not reader:
+ logger.warning(f"当前知识库未配置支持解析文件后缀: {ext}")
+ continue
+
+ doc = await reader.read(std_path)
+ if doc:
+ docs_to_ingest.append(doc)
+
+ if not docs_to_ingest:
+ return 0
+
+ return await self.rag_client.ingest(docs_to_ingest)
@tool(
name="search_knowledge",
diff --git a/zhenxun/services/ai/context/memory/__init__.py b/zhenxun/services/ai/context/memory/__init__.py
index 28419e54..7fdc2dcd 100644
--- a/zhenxun/services/ai/context/memory/__init__.py
+++ b/zhenxun/services/ai/context/memory/__init__.py
@@ -1,9 +1,7 @@
from .builder import MemoryBuilder
from .compression import MemoryPolicy
-from .facades import AgentSessionFacade
from .manager import memory_manager
from .models import (
- BaseMemoryIngestionMiddleware,
MemoryConfig,
)
from .types import (
@@ -12,8 +10,6 @@ from .types import (
)
__all__ = [
- "AgentSessionFacade",
- "BaseMemoryIngestionMiddleware",
"Isolation",
"MemoryBuilder",
"MemoryConfig",
diff --git a/zhenxun/services/ai/context/memory/builder.py b/zhenxun/services/ai/context/memory/builder.py
index cc7092d8..9ce0502d 100644
--- a/zhenxun/services/ai/context/memory/builder.py
+++ b/zhenxun/services/ai/context/memory/builder.py
@@ -6,27 +6,19 @@ from typing_extensions import Self
from pydantic import BaseModel
-from zhenxun.services.ai.context.memory.compression import MemoryPolicy
-from zhenxun.services.ai.context.memory.models import (
+from zhenxun.services.ai.utils.scope import ScopeBuilder
+
+from .compression import MemoryPolicy
+from .models import (
ContextCompressionConfig,
IngestionConfig,
- LongTermConfig,
MemoryConfig,
- MemorySlot,
ShortTermConfig,
- SlotMemoryConfig,
)
-from zhenxun.services.ai.context.memory.storage.interfaces import (
+from .storage.interfaces import (
BaseChatContext,
BaseMemoryIngestionMiddleware,
- BaseSlotContext,
)
-from zhenxun.services.ai.context.memory.types import (
- AutoRecallPolicy,
-)
-from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend
-from zhenxun.services.ai.context.rag.engine import ScopedRAGClient
-from zhenxun.services.ai.utils.scope import ScopeBuilder
class MemoryBuilder:
@@ -42,8 +34,6 @@ class MemoryBuilder:
"""
self._config = MemoryConfig(
short_term=ShortTermConfig(enable=False),
- slots=SlotMemoryConfig(enable=False),
- long_term=LongTermConfig(enable=False),
compression=ContextCompressionConfig(),
ingestion=IngestionConfig(),
)
@@ -68,17 +58,10 @@ class MemoryBuilder:
if isinstance(memory, bool):
return MemoryConfig(
short_term=ShortTermConfig(enable=memory),
- long_term=LongTermConfig(enable=memory),
)
return MemoryConfig(short_term=ShortTermConfig(enable=False))
- def with_base_isolation(self, isolation: ScopeBuilder) -> Self:
- """设置顶层基准隔离级别,短期/中期/长期记忆将默认继承此级别"""
- self._config.base_isolation = isolation
- self._config.short_term.isolation = isolation
- return self
-
def with_short_term(
self,
enable: bool = True,
@@ -95,77 +78,11 @@ class MemoryBuilder:
"""
self._config.short_term.enable = enable
if isolation is not None:
- self._config.base_isolation = isolation
self._config.short_term.isolation = isolation
if backend is not None:
self._config.short_term.backend = backend
return self
- def with_slots(
- self,
- enable: bool = True,
- scopes: dict[str, ScopeBuilder] | None = None,
- default_slots: list[MemorySlot] | None = None,
- backend: str | BaseSlotContext | None = None,
- instructions: str | None = None,
- ) -> Self:
- """
- 配置核心槽位记忆 (Memory Slots)。
-
- 参数:
- enable: 是否启用槽位记忆。
- scopes: 语义化作用域映射字典。如果只有一个键值对,则大模型不可见该参数。
- default_slots: 首次初始化时自动写入的默认槽位列表。
- backend: 槽位记忆存储后端,如果为 None 则使用全局默认后端。
- instructions: 覆写内置槽位管理工具箱的默认系统提示词规则。
- """
- self._config.slots.enable = enable
- if scopes is not None:
- self._config.slots.scopes = scopes
- if default_slots is not None:
- self._config.slots.default_slots = default_slots
- if backend is not None:
- self._config.slots.backend = backend
- if instructions is not None:
- self._config.slots.instructions = instructions
- return self
-
- def with_long_term(
- self,
- enable: bool = True,
- scopes: dict[str, ScopeBuilder] | None = None,
- engine: ScopedRAGClient | None = None,
- backend: str | StorageBackend | None = None,
- embedder: Embedder | str | None = None,
- agentic: bool = True,
- auto_recall: AutoRecallPolicy = False,
- instructions: str | None = None,
- ) -> Self:
- """
- 配置长期向量记忆与 RAG 设定。
-
- 参数:
- enable: 是否启用长期记忆。
- scopes: 语义化作用域映射字典。如果只有一个键值对,则大模型不可见该参数。
- engine: 高级 RAG 检索引擎实例 (推荐)。若提供,将接管记忆的底层检索、混合与重排。
- backend: 长期记忆存储后端。
- embedder: 用于向量化的文本嵌入模型实例。
- agentic: 是否开启主动智能体记忆管理 (增删改查工具自动注入)。
- auto_recall: 长期记忆的自动召回策略,支持 bool 或 Callable 函数。
- instructions: 覆写内置长期记忆工具箱的默认系统提示词规则。
- """ # noqa: E501
- self._config.long_term.enable = enable
- self._config.long_term.engine = engine
- if scopes is not None:
- self._config.long_term.scopes = scopes
- self._config.long_term.backend = backend
- self._config.long_term.embedder = embedder
- self._config.long_term.agentic = agentic
- self._config.long_term.auto_recall = auto_recall
- if instructions is not None:
- self._config.long_term.instructions = instructions
- return self
-
def with_multimodal_window(self, window_size: int = 5) -> Self:
"""
配置多模态历史视窗大小。
@@ -261,8 +178,4 @@ class MemoryBuilder:
"""
生成最终构建好的 MemoryConfig 配置对象。
"""
- if not self._config.slots.scopes:
- self._config.slots.scopes = {"私有": self._config.base_isolation}
- if not self._config.long_term.scopes:
- self._config.long_term.scopes = {"私有": self._config.base_isolation}
return self._config
diff --git a/zhenxun/services/ai/context/memory/capabilities.py b/zhenxun/services/ai/context/memory/capabilities.py
index 6808dbd6..b3f19201 100644
--- a/zhenxun/services/ai/context/memory/capabilities.py
+++ b/zhenxun/services/ai/context/memory/capabilities.py
@@ -1,52 +1,270 @@
+import inspect
from typing import Any
from zhenxun.services.ai.capabilities.base import AbstractCapability
-from zhenxun.services.ai.context.memory.models import MemoryConfig
+from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend
+from zhenxun.services.ai.context.rag.engine import ScopedRAGClient
from zhenxun.services.ai.run.context import RunContext
+from zhenxun.services.ai.tools.core.toolkit import BaseToolkit
from zhenxun.services.ai.tools.providers.builtin.memory import MemoryManagementToolkit
+from zhenxun.services.ai.utils.logger import log_memory as logger
+from zhenxun.services.ai.utils.runtime import ContextUtils
+from zhenxun.services.ai.utils.scope import ScopeBuilder
+
+from .manager import memory_manager
+from .models import MemoryScoringConfig
+from .storage.backends import MemoryScope
+from .storage.interfaces import BaseSlotContext
+from .types import (
+ AutoRecallPolicy,
+ Isolation,
+ MemorySlot,
+ SessionMetadata,
+)
-class AgenticMemoryCapability(AbstractCapability):
+class LongTermMemoryCapability(AbstractCapability):
"""
- 智能体主动记忆管理能力 (Agentic Memory Management)。
- 当 `MemoryConfig.long_term.enable == True` 且 `agentic == True` 时隐式挂载,
- 在运行时动态组装并向大模型提供 `MemoryManagementToolkit` 工具箱。
+ 长期向量记忆 (RAG) 核心能力组件。
+ 负责静默执行自动召回 (Auto Recall),并在必要时提供读写 RAG 数据库的工具链。
"""
- def __init__(self, memory_config: MemoryConfig, namespace: str):
- self.memory_config = memory_config
+ def __init__(
+ self,
+ engine: ScopedRAGClient | None = None,
+ storage_backend: StorageBackend | None = None,
+ embedder: Embedder | str | None = None,
+ scopes: dict[str, ScopeBuilder] | ScopeBuilder | None = None,
+ toolkit: bool | BaseToolkit = True,
+ scoring_config: MemoryScoringConfig | None = None,
+ auto_recall: AutoRecallPolicy = False,
+ recall_limit: int = 5,
+ recall_threshold: float = 0.5,
+ namespace: str | None = None,
+ ):
+ """
+ 初始化长期记忆能力组件。
+
+ 参数:
+ engine: RAG 客户端引擎实例。若为 None 则在运行时按需构建。
+ storage_backend: RAG 向量存储后端。若为 None 则从管理器按命名空间提取。
+ embedder: 嵌入模型实例或模型名称。
+ scopes: 控制长期记忆的隔离级别。支持单 ScopeBuilder 或 映射字典。
+ toolkit: 是否启用记忆管理工具箱,或传入自定义的工具箱实例。
+ scoring_config: 记忆检索打分配置(包含时间衰减等参数)。
+ auto_recall: 自动召回策略,可以是布尔值或自定义的回调函数。
+ recall_limit: 自动召回的记忆条数限制。
+ recall_threshold: 自动召回的相似度分数阈值。
+ namespace: 指定的命名空间,用于自动路由存储后端及构建器。
+ """
+ self.engine = engine
+ self.storage_backend = storage_backend
+ self.embedder = embedder
+ if isinstance(scopes, ScopeBuilder):
+ self.scopes = {"默认": scopes}
+ else:
+ self.scopes = scopes or {"私有": Isolation.AGENT_USER()}
+ self.toolkit = toolkit
+ self.scoring_config = scoring_config or MemoryScoringConfig()
+ self.auto_recall = auto_recall
+ self.recall_limit = recall_limit
+ self.recall_threshold = recall_threshold
self.namespace = namespace
- async def get_tools(self, context: RunContext) -> list[Any]:
- kwargs = self.memory_config.long_term.toolkit_kwargs.copy()
- kwargs["memory_config"] = self.memory_config
- kwargs["namespace"] = self.namespace
- if self.memory_config.long_term.instructions is not None:
- kwargs["instructions"] = self.memory_config.long_term.instructions
+ def _build_session_meta(self, context: RunContext) -> SessionMetadata:
+ scope_builder = next(iter(self.scopes.values())) if self.scopes else None
+ return ContextUtils.build_session_meta(
+ context=context,
+ target_builder=scope_builder,
+ extra_scopes=self.scopes,
+ custom_namespace=self.namespace,
+ )
- toolkit = MemoryManagementToolkit(**kwargs)
- return [toolkit]
+ def _ensure_engine(self, context: RunContext) -> ScopedRAGClient | None:
+ if self.engine is not None:
+ return self.engine
+
+ ns = self.namespace or getattr(context.session, "namespace", "global")
+
+ storage_instance = self.storage_backend
+ if not storage_instance:
+ factory = memory_manager._storage_factories.get(
+ ns
+ ) or memory_manager._storage_factories.get("global")
+ if factory:
+ storage_instance = factory()
+
+ if not storage_instance:
+ return None
+
+ embedder_instance = self.embedder
+ if isinstance(embedder_instance, str):
+ from zhenxun.services.ai.context.rag.backends.embedders import (
+ DefaultEmbedder,
+ )
+
+ embedder_instance = DefaultEmbedder(model_name=embedder_instance)
+
+ from zhenxun.services.ai.context.rag.builder import RAGBuilder
+
+ builder = RAGBuilder(storage_instance).with_scope("/")
+ if embedder_instance:
+ builder.with_embedder(embedder_instance)
+
+ builder.enable_lifecycle_scoring(
+ half_life_days=self.scoring_config.recency_half_life_days,
+ decay_weight=self.scoring_config.recency_weight,
+ semantic_weight=self.scoring_config.semantic_weight,
+ importance_weight=self.scoring_config.importance_weight,
+ reinforcement_weight=self.scoring_config.reinforcement_weight,
+ )
+
+ self.engine = builder.build()
+ return self.engine
+
+ async def get_system_prompts(self, context: RunContext) -> list[str]:
+ engine = self._ensure_engine(context)
+ user_input = context.run.user_input
+ if not user_input or not engine:
+ return []
+
+ should_recall = False
+ session_meta = self._build_session_meta(context)
+
+ if isinstance(self.auto_recall, bool):
+ should_recall = self.auto_recall
+ elif callable(self.auto_recall):
+ try:
+ res = self.auto_recall(user_input, session_meta)
+ if inspect.isawaitable(res):
+ should_recall = await res
+ else:
+ should_recall = bool(res)
+ except Exception as e:
+ logger.error(
+ f"[LongTermMemoryCapability] auto_recall 函数执行失败: {e}"
+ )
+ should_recall = False
+
+ if not should_recall:
+ return []
+
+ scope = MemoryScope(rag_client=engine)
+ matches = await scope.recall(
+ session=session_meta,
+ query=user_input,
+ limit=self.recall_limit,
+ )
+ if matches:
+ valid_matches = [m for m in matches if m.score >= self.recall_threshold]
+ if valid_matches:
+ fact_str = "\n".join(f"- {m.record.content}" for m in valid_matches)
+ return [f"[系统补充:有关用户的长期记忆设定]\n{fact_str}"]
+ return []
+
+ async def get_tools(self, context: RunContext) -> list[Any]:
+ if self.toolkit is False:
+ return []
+
+ engine = self._ensure_engine(context)
+ if not engine:
+ return []
+
+ if isinstance(self.toolkit, BaseToolkit):
+ tk = self.toolkit.clone_with(
+ rag_client=engine, scopes=self.scopes, _namespace=self.namespace
+ )
+ return [tk]
+
+ return [
+ MemoryManagementToolkit(
+ rag_client=engine, scopes=self.scopes, namespace=self.namespace
+ )
+ ]
class SlotMemoryCapability(AbstractCapability):
"""
- 槽位记忆能力组件。
- 当 `MemoryConfig.slots.enable == True` 时隐式挂载,
- 在运行时动态组装并向大模型提供 `MemorySlotToolkit` 工具箱。
+ 独立的槽位记忆 (Memory Slots) 能力组件。
+ 直接作为插件挂载至 Agent 的 capabilities 列表中。
"""
- def __init__(self, memory_config: MemoryConfig, namespace: str):
- self.memory_config = memory_config
+ def __init__(
+ self,
+ scopes: dict[str, ScopeBuilder] | ScopeBuilder | None = None,
+ default_slots: list[MemorySlot] | None = None,
+ backend: BaseSlotContext | None = None,
+ toolkit: bool | BaseToolkit = True,
+ namespace: str | None = None,
+ ):
+ """
+ 初始化槽位记忆能力组件。
+
+ 参数:
+ scopes: 控制槽位记忆的隔离级别。支持单 ScopeBuilder 或 映射字典。
+ default_slots: 默认记忆槽列表,初始时自动创建未存在的槽位。
+ backend: 中期记忆槽持久化后端。若为 None 则从管理器按命名空间提取。
+ toolkit: 是否启用记忆槽管理工具箱,或传入自定义的工具箱实例。
+ namespace: 指定的命名空间,用于自动路由后端及工具箱。
+ """
+ if isinstance(scopes, ScopeBuilder):
+ self.scopes = {"默认": scopes}
+ else:
+ self.scopes = scopes or {"私有": Isolation.AGENT_USER()}
+ self.default_slots = default_slots or []
+ self.backend = backend
+ self.toolkit = toolkit
self.namespace = namespace
+ async def _get_slot_ctx_and_meta(self, context: RunContext):
+ ns = self.namespace or getattr(context.session, "namespace", "global")
+ slot_ctx = self.backend or memory_manager.get_backend("slots", namespace=ns)
+
+ target_builder = next(iter(self.scopes.values())) if self.scopes else None
+ session_meta = ContextUtils.build_session_meta(
+ context=context,
+ target_builder=target_builder,
+ extra_scopes=self.scopes,
+ custom_namespace=self.namespace,
+ )
+ return slot_ctx, session_meta
+
+ async def get_system_prompts(self, context: RunContext) -> list[str]:
+ slot_ctx, session_meta = await self._get_slot_ctx_and_meta(context)
+ if not slot_ctx:
+ return []
+
+ for default_slot in self.default_slots:
+ if not await slot_ctx.get_slot(session_meta, default_slot.label):
+ await slot_ctx.set_slot(session_meta, default_slot)
+
+ slots = await slot_ctx.list_pinned_slots(session_meta)
+ if not slots:
+ return []
+
+ xml_parts = [""]
+ for slot in slots:
+ semantic_name = session_meta.scope_name_mapping.get(slot.scope, "未知")
+ xml_parts.append(
+ f' \n'
+ f" {slot.content}\n"
+ " "
+ )
+ xml_parts.append("")
+ return ["\n".join(xml_parts)]
+
async def get_tools(self, context: RunContext) -> list[Any]:
from zhenxun.services.ai.tools.providers.builtin.slots import MemorySlotToolkit
- kwargs = self.memory_config.slots.toolkit_kwargs.copy()
- kwargs["memory_config"] = self.memory_config
- kwargs["namespace"] = self.namespace
- if self.memory_config.slots.instructions is not None:
- kwargs["instructions"] = self.memory_config.slots.instructions
+ if self.toolkit is False:
+ return []
- toolkit = MemorySlotToolkit(**kwargs)
+ if isinstance(self.toolkit, BaseToolkit):
+ toolkit = self.toolkit.clone_with(
+ scopes=self.scopes, backend=self.backend, _namespace=self.namespace
+ )
+ else:
+ toolkit = MemorySlotToolkit(
+ scopes=self.scopes, backend=self.backend, namespace=self.namespace
+ )
return [toolkit]
diff --git a/zhenxun/services/ai/context/memory/compression.py b/zhenxun/services/ai/context/memory/compression.py
index 05dd7f0d..343ecadb 100644
--- a/zhenxun/services/ai/context/memory/compression.py
+++ b/zhenxun/services/ai/context/memory/compression.py
@@ -4,6 +4,7 @@ from typing import Any, Generic, TypeVar
from pydantic import BaseModel, Field
+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 (
AudioPart,
@@ -386,7 +387,7 @@ class LLMSummarizerReducer(AbstractSummarizerReducer):
) -> LLMMessage | None:
prompt_text = f"### 📋 [对话摘要任务]\n{self.summarization_prompt}\n\n"
if prev_summary:
- prompt_text += "#### önceki_summary (参考先前的快照):\n"
+ prompt_text += "#### prev_summary (参考先前的快照):\n"
prompt_text += f"> {prev_summary}\n\n"
prompt_text += "#### 待处理的历史消息流:\n"
for m in to_summarize:
@@ -512,7 +513,7 @@ class CondenserPipeline:
@classmethod
def create_from_configs(
- cls, memory_config: Any, model_name: str
+ cls, memory_config: MemoryConfig | None, model_name: str
) -> "CondenserPipeline":
"""基于全局和局部配置组装压缩管线工厂方法"""
from zhenxun.services.ai.config import get_llm_config
@@ -602,7 +603,7 @@ class CondenserPipeline:
class MemoryPolicy:
"""
- 记忆策略工厂 (Strategy Factory Facade)。
+ 记忆策略工厂。
为开发者提供开箱即用的上下文压缩管线组装方案。
"""
diff --git a/zhenxun/services/ai/context/memory/engine.py b/zhenxun/services/ai/context/memory/engine.py
index 150e73bd..65b9a7bd 100644
--- a/zhenxun/services/ai/context/memory/engine.py
+++ b/zhenxun/services/ai/context/memory/engine.py
@@ -1,8 +1,9 @@
from collections.abc import Sequence
-from typing import Any, cast
+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.run.context import RunContext
from zhenxun.services.ai.utils.logger import log_memory as logger
from zhenxun.utils.pydantic_compat import model_copy
@@ -14,144 +15,37 @@ from .models import MemoryConfig
from .types import SessionMetadata
-class MemoryReader:
+class SessionMemoryContext:
"""
- 记忆读取器 (Memory Reader)。
- 负责从数据库中提取短期上下文历史,召回长期的背景知识,并执行自动压缩。
+ 会话记忆门面。
+ 封装了当前会话记忆的读写操作、上下文压缩管线以及入库清洗中间件。
"""
def __init__(
- self, session_meta: SessionMetadata, memory_config: MemoryConfig | None
+ self,
+ session_meta: SessionMetadata,
+ memory_config: MemoryConfig | None,
+ context: RunContext,
):
"""
- 初始化记忆读取器。
+ 初始化会话记忆门面。
参数:
session_meta: 会话元数据,包含 Namespace 与作用域映射等上下文信息。
- memory_config: 记忆系统的配置对象,控制长期、短期及槽位记忆的启用与逻辑。
+ memory_config: 记忆系统的配置对象,控制长期、短期记忆的启用与逻辑。
+ context: 必填的运行时上下文环境 (RunContext),供中间件进行依赖注入。
"""
self.session_meta = session_meta
self.memory_config = memory_config
+ self.context = context
- async def get_long_term_context(self, user_input: str) -> str:
- """
- 基于用户输入召回长期记忆(RAG),返回格式化后的背景提示词。
- """
- if (
- not self.memory_config
- or not self.memory_config.long_term.enable
- or not user_input
- ):
- return ""
-
- policy = self.memory_config.long_term.auto_recall
- should_recall = False
-
- if isinstance(policy, bool):
- should_recall = policy
- elif callable(policy):
- import inspect
-
- try:
- res = policy(user_input, self.session_meta)
- if inspect.isawaitable(res):
- should_recall = await res
- else:
- should_recall = bool(res)
- except Exception as e:
- logger.error(f"[MemoryReader] 自定义 auto_recall 函数执行失败: {e}")
- should_recall = False
-
- if not should_recall:
- return ""
-
- ltm_scope = memory_manager.get_long_term_memory(
- self.memory_config,
- namespace=self.session_meta.selector.namespace or "global",
- )
- if not ltm_scope:
- return ""
-
- matches = await ltm_scope.recall(session=self.session_meta, query=user_input)
- if matches:
- logger.debug(f"🧠 [MemoryReader] 长期记忆召回详情 (Query: '{user_input}'):")
- for i, m in enumerate(matches):
- logger.debug(
- f" [{i + 1}] 得分: {m.score:.4f} | 内容: {m.record.content}"
- )
-
- threshold = self.memory_config.long_term.recall_threshold
- valid_matches = [m for m in matches if m.score >= threshold]
- if not valid_matches:
- logger.debug("🧠 [MemoryReader] 召回的记忆均未达到相关性阈值,已丢弃。")
- return ""
-
- fact_str = "\n".join(f"- {m.record.content}" for m in valid_matches)
- logger.debug(
- f"🧠 [MemoryReader]"
- f"成功截取并注入 {len(valid_matches)} 条高价值长期记忆。"
- )
- return f"[系统补充:有关用户的长期记忆设定]\n{fact_str}"
- return ""
-
- async def get_slots_context(self) -> str:
- """
- 读取并组装核心槽位记忆 (Memory Slots),返回 XML 格式字符串供大模型使用。
- """
- if not self.memory_config or not self.memory_config.slots.enable:
- return ""
- slot_ctx = memory_manager.get_slot_context(
- self.memory_config,
- namespace=self.session_meta.selector.namespace or "global",
- )
- if not slot_ctx:
- return ""
-
- if self.memory_config.slots.default_slots:
- for default_slot in self.memory_config.slots.default_slots:
- existing = await slot_ctx.get_slot(
- self.session_meta, default_slot.label
- )
- if not existing:
- await slot_ctx.set_slot(self.session_meta, default_slot)
-
- slots = await slot_ctx.list_pinned_slots(self.session_meta)
- if not slots:
- return ""
-
- show_scope = False
- if (
- self.memory_config
- and self.memory_config.slots.scopes
- and len(self.memory_config.slots.scopes) > 1
- ):
- show_scope = True
-
- xml_parts = [""]
- for slot in slots:
- if show_scope:
- semantic_name = self.session_meta.scope_name_mapping.get(
- slot.scope, "未知"
- )
- xml_parts.append(
- f' \n'
- f" {slot.content}\n"
- " "
- )
- else:
- xml_parts.append(
- f' \n {slot.content}\n '
- )
- xml_parts.append("")
- return "\n".join(xml_parts)
-
- async def get_short_term_context(
+ async def read(
self,
model_name: str,
override_history: Sequence[AgentMessage] | None = None,
) -> list[AgentMessage]:
"""
- 拉取短期对话历史,并执行 Token 压缩。
+ 拉取短期对话历史,并执行 Token 压缩与管线修剪。
"""
current_history: list[AgentMessage] = []
if override_history is not None:
@@ -187,43 +81,18 @@ class MemoryReader:
if changed:
await chat_context.set_messages(self.session_meta, new_history)
logger.info(
- "💾 [MemoryReader] 压缩截断完毕,已同步覆写数据库。"
+ "💾 [SessionMemory] 压缩截断完毕,已同步覆写数据库。"
f"压缩后条数: {len(new_history)}"
)
current_history = cast(list[AgentMessage], new_history)
return current_history
-
-class MemoryWriter:
- """
- 记忆写入器 (Memory Writer)。
- 负责将对话增量安全地写入数据库。
- """
-
- def __init__(
- self,
- session_meta: SessionMetadata,
- memory_config: MemoryConfig | None,
- context: Any = None,
- ):
- """
- 初始化记忆写入器。
-
- 参数:
- session_meta: 会话元数据,包含 Namespace 与作用域映射等上下文信息。
- memory_config: 记忆系统的配置对象,控制记忆存入的逻辑。
- context: 运行时上下文环境,作为可选参数传入,供中间件使用,默认 None。
- """
- self.session_meta = session_meta
- self.memory_config = memory_config
- self.context = context
-
- async def save_new_messages(
+ async def write(
self,
new_messages: Sequence[AgentMessage],
- ):
- """将新产生的对话增量保存到数据库"""
+ ) -> None:
+ """将新产生的对话增量,经过入库中间件清洗后保存到数据库"""
if not new_messages:
return
diff --git a/zhenxun/services/ai/context/memory/facades.py b/zhenxun/services/ai/context/memory/facades.py
deleted file mode 100644
index 1fb07cfe..00000000
--- a/zhenxun/services/ai/context/memory/facades.py
+++ /dev/null
@@ -1,129 +0,0 @@
-from __future__ import annotations
-
-from collections.abc import Sequence
-from typing import Literal
-
-from zhenxun.services.ai.core.messages import AgentMessage, LLMMessage
-
-from .manager import GlobalMemoryManager
-from .storage.interfaces import (
- BaseChatContext,
- BaseSlotContext,
-)
-from .types import (
- MemorySlot,
- SessionMetadata,
-)
-
-
-class ChatHistoryFacade:
- """短期对话历史门面"""
-
- def __init__(self, manager: GlobalMemoryManager, session_meta: SessionMetadata):
- self.manager = manager
- self.session_meta = session_meta
-
- @property
- def _backend(self) -> BaseChatContext | None:
- return self.manager.get_chat_context(
- None, self.session_meta.namespace or "global"
- )
-
- async def get(self, limit: int | None = None) -> list[LLMMessage]:
- """获取当前会话的历史消息"""
- if not self._backend:
- return []
- msgs = await self._backend.get_messages(self.session_meta)
- return msgs[-limit:] if limit else msgs
-
- async def add(self, messages: Sequence[AgentMessage] | AgentMessage) -> None:
- """向当前会话追加一条或多条历史消息"""
- if not self._backend:
- return
- from zhenxun.services.ai.core.engine.context_renderer import ContextConverter
-
- msgs = messages if isinstance(messages, Sequence) else [messages]
- flattened = ContextConverter.flatten_to_llm_messages(msgs)
- if flattened:
- await self._backend.add_messages(self.session_meta, flattened)
-
- async def clear(self) -> None:
- """清空当前会话的短期对话历史"""
- if not self._backend:
- return
- await self._backend.clear(self.session_meta)
-
-
-class SlotFacade:
- """中期记忆槽门面"""
-
- def __init__(self, manager: GlobalMemoryManager, session_meta: SessionMetadata):
- self.manager = manager
- self.session_meta = session_meta
-
- @property
- def _backend(self) -> BaseSlotContext | None:
- """获取底层槽位存储后端"""
- return self.manager.get_slot_context(
- None, self.session_meta.namespace or "global"
- )
-
- async def get(self, label: str) -> str | None:
- """获取指定标识的槽位记忆内容"""
- if not self._backend:
- return None
- slot = await self._backend.get_slot(self.session_meta, label)
- return slot.content if slot else None
-
- async def set(
- self,
- label: str,
- content: str,
- scope: Literal["session", "global"] = "session",
- size_limit: int = 2000,
- pinned: bool = True,
- ) -> None:
- """设置或更新指定的槽位记忆"""
- if not self._backend:
- return
- slot = MemorySlot(
- label=label,
- content=content,
- scope=scope,
- size_limit=size_limit,
- pinned=pinned,
- )
- await self._backend.set_slot(self.session_meta, slot)
-
- async def delete(
- self, label: str, scope: Literal["session", "global"] = "session"
- ) -> None:
- """删除指定的槽位记忆"""
- if not self._backend:
- return
- await self._backend.delete_slot(self.session_meta, label, scope)
-
- async def list_all(self) -> dict[str, str]:
- """获取当前会话所有被置顶的槽位记忆"""
- if not self._backend:
- return {}
- slots = await self._backend.list_pinned_slots(self.session_meta)
- return {s.label: s.content for s in slots}
-
-
-class AgentSessionFacade:
- """
- 提供给第三方开发者的会话记忆访问聚合门面 (Facade)。
- """
-
- def __init__(self, manager: GlobalMemoryManager, session_meta: SessionMetadata):
- self.manager = manager
- self.session_meta = session_meta
- self.history = ChatHistoryFacade(manager, session_meta)
- self.slots = SlotFacade(manager, session_meta)
-
- async def clear_all(self) -> None:
- """一键清空当前会话下的短期对话历史与记忆槽"""
- cleaner = self.manager.cleaner().session(self.session_meta.session_id)
- await cleaner.clear_short_term()
- await cleaner.clear_slots()
diff --git a/zhenxun/services/ai/context/memory/manager.py b/zhenxun/services/ai/context/memory/manager.py
index 0b9ba851..e32df3bb 100644
--- a/zhenxun/services/ai/context/memory/manager.py
+++ b/zhenxun/services/ai/context/memory/manager.py
@@ -1,7 +1,8 @@
+from collections import defaultdict
from collections.abc import Callable
from typing import Any, cast
-from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend
+from zhenxun.services.ai.context.rag.backends import StorageBackend
from zhenxun.services.ai.utils.logger import log_memory as logger
from zhenxun.services.ai.utils.scope import BaseScopeBuilder
from zhenxun.utils.utils import infer_plugin_namespace
@@ -9,11 +10,10 @@ from zhenxun.utils.utils import infer_plugin_namespace
from .models import MemoryConfig
from .storage.backends import (
InMemoryChatContext,
- MemoryScope,
)
from .storage.interfaces import (
BaseChatContext,
- BaseSlotContext,
+ IClearableBackend,
)
@@ -33,41 +33,46 @@ class MemoryCleaner(BaseScopeBuilder["MemoryCleaner"]):
self._config = cfg.build() if hasattr(cfg, "build") else cfg
return self
+ async def clear_target(self, target_name: str):
+ """底层派发器:定向清理指定注册名称的泛型扩展后端数据"""
+ ns_dict = self.manager._backends.get(target_name, {})
+ for backend in ns_dict.values():
+ if isinstance(backend, IClearableBackend):
+ await backend.clear_by_query(self._selector)
+ else:
+ logger.warning(
+ f"后端 {backend.__class__.__name__}"
+ "未实现 IClearableBackend 协议,已跳过清理。"
+ )
+
async def clear_short_term(self):
"""一键清理目标范围下的短期对话历史记忆"""
- if self._config and self._config.short_term.backend:
+ if isinstance(self._config, MemoryConfig) and isinstance(
+ self._config.short_term.backend, IClearableBackend
+ ):
await self._config.short_term.backend.clear_by_query(self._selector)
else:
- for backend in self.manager._chat_backends.values():
- await backend.clear_by_query(self._selector)
+ await self.clear_target("chat")
async def clear_slots(self):
- """一键清理目标范围下的中期记忆槽 (Memory Slots)"""
- if self._config and self._config.slots.backend:
- await self._config.slots.backend.clear_by_query(self._selector)
- else:
- for backend in self.manager._slot_backends.values():
- await backend.clear_by_query(self._selector)
+ """一键清理目标范围下的记忆槽数据"""
+ await self.clear_target("slots")
async def clear_long_term(self):
"""一键清理目标范围下的长期向量记忆 (RAG Vector Database)"""
- if self._config and self._config.long_term.backend:
- from zhenxun.services.ai.context.rag.backends import StorageBackend
-
- storage = cast(StorageBackend, self._config.long_term.backend)
- await storage.clear_by_query(self._selector)
- else:
- for factory in self.manager._storage_factories.values():
- storage = factory()
- if hasattr(storage, "clear_by_query"):
- await storage.clear_by_query(self._selector)
- else:
- await storage.delete(scope_prefix=self._selector.scope_prefix)
+ for factory in self.manager._storage_factories.values():
+ storage = factory()
+ if isinstance(storage, IClearableBackend):
+ await storage.clear_by_query(self._selector)
async def clear_all(self):
- """一键清理指定范围下的所有生命周期记忆(对话、槽位、RAG)"""
- await self.clear_short_term()
- await self.clear_slots()
+ """一键清理指定范围下的所有生命周期记忆(对话、记忆槽、RAG、及其他泛型扩展后端)"""
+ if isinstance(self._config, MemoryConfig) and isinstance(
+ self._config.short_term.backend, IClearableBackend
+ ):
+ await self._config.short_term.backend.clear_by_query(self._selector)
+ for target_name in self.manager._backends:
+ await self.clear_target(target_name)
await self.clear_long_term()
logger.info(
f"🧹 成功清理作用域 '{self._selector.scope_prefix}'下的所有记忆痕迹!"
@@ -81,10 +86,8 @@ class GlobalMemoryManager:
"""
def __init__(self):
- self._chat_backends: dict[str, BaseChatContext] = {
- "global": InMemoryChatContext()
- }
- self._slot_backends: dict[str, BaseSlotContext] = {}
+ self._backends: dict[str, dict[str, Any]] = defaultdict(dict)
+ self.register_backend("chat", InMemoryChatContext(), "global")
from zhenxun.services.ai.context.rag.backends import DictStorageBackend
@@ -92,19 +95,23 @@ class GlobalMemoryManager:
"global": lambda: DictStorageBackend()
}
+ def register_backend(
+ self, backend_type: str, backend: Any, scope: str | None = None
+ ) -> None:
+ """泛型注册:注册任意类型的存储后端"""
+ ns = scope if scope is not None else infer_plugin_namespace()
+ self._backends[backend_type][ns] = backend
+
+ def get_backend(self, backend_type: str, namespace: str = "global") -> Any | None:
+ """泛型获取:获取任意类型的存储后端"""
+ backends = self._backends.get(backend_type, {})
+ return backends.get(namespace) or backends.get("global")
+
def register_chat_backend(
self, backend: BaseChatContext, scope: str | None = None
) -> None:
"""注册特定命名空间的短期记忆存储后端。"""
- ns = scope if scope is not None else infer_plugin_namespace()
- self._chat_backends[ns] = backend
-
- def register_slot_backend(
- self, backend: BaseSlotContext, scope: str | None = None
- ) -> None:
- """注册特定命名空间的中期记忆槽存储后端。"""
- ns = scope if scope is not None else infer_plugin_namespace()
- self._slot_backends[ns] = backend
+ self.register_backend("chat", backend, scope)
def register_storage_factory(
self, factory: Callable[[], StorageBackend], scope: str | None = None
@@ -117,20 +124,6 @@ class GlobalMemoryManager:
"""获取声明式记忆清理构建器,供第三方开发者极速清理指定记忆"""
return MemoryCleaner(self)
- def get_embedder(self, embedder_val: "Embedder | str | None") -> Embedder | None:
- """获取向量化引擎实例。如果传入的是字符串,则视为 API 模型名称。"""
- if not embedder_val:
- return None
-
- if isinstance(embedder_val, str):
- from zhenxun.services.ai.context.rag.backends.embedders import (
- DefaultEmbedder,
- )
-
- return DefaultEmbedder(model_name=embedder_val)
-
- return embedder_val
-
def get_chat_context(
self, config: MemoryConfig | None, namespace: str = "global"
) -> BaseChatContext | None:
@@ -142,69 +135,7 @@ class GlobalMemoryManager:
if backend_cfg is not None:
return cast(BaseChatContext, backend_cfg)
- return self._chat_backends.get(namespace) or self._chat_backends["global"]
-
- def get_slot_context(
- self, config: MemoryConfig | None, namespace: str = "global"
- ) -> BaseSlotContext | None:
- """根据配置分配对应的槽位记忆实例"""
- if not config or not config.slots.enable:
- return None
-
- backend_cfg = config.slots.backend
- if backend_cfg is not None:
- return cast(BaseSlotContext, backend_cfg)
-
- return self._slot_backends.get(namespace) or self._slot_backends["global"]
-
- def get_long_term_memory(
- self, config: MemoryConfig | None, namespace: str = "global"
- ) -> MemoryScope | None:
- """根据声明式配置动态组装长期向量记忆实例"""
- if not config or not config.long_term.enable:
- return None
-
- if config.long_term.engine is not None:
- return MemoryScope(
- rag_client=config.long_term.engine,
- )
-
- storage_instance = None
- backend_cfg = config.long_term.backend
- if backend_cfg is not None:
- storage_instance = cast(StorageBackend, backend_cfg)
- else:
- factory = (
- self._storage_factories.get(namespace)
- or self._storage_factories["global"]
- )
- storage_instance = factory()
-
- embedder = self.get_embedder(config.long_term.embedder)
-
- from zhenxun.services.ai.context.rag.builder import RAGBuilder
-
- builder = RAGBuilder(storage_instance).with_scope("/")
- if embedder:
- builder.with_embedder(embedder)
-
- from .models import MemoryScoringConfig
-
- scoring_cfg = MemoryScoringConfig()
-
- builder.enable_lifecycle_scoring(
- half_life_days=scoring_cfg.recency_half_life_days,
- decay_weight=scoring_cfg.recency_weight,
- semantic_weight=scoring_cfg.semantic_weight,
- importance_weight=scoring_cfg.importance_weight,
- reinforcement_weight=scoring_cfg.reinforcement_weight,
- )
-
- client = builder.build()
-
- return MemoryScope(
- rag_client=client,
- )
+ return self.get_backend("chat", namespace)
memory_manager = GlobalMemoryManager()
diff --git a/zhenxun/services/ai/context/memory/models.py b/zhenxun/services/ai/context/memory/models.py
index d5f97b05..dc87eef6 100644
--- a/zhenxun/services/ai/context/memory/models.py
+++ b/zhenxun/services/ai/context/memory/models.py
@@ -2,50 +2,21 @@
记忆域类型定义
"""
-from typing import Any
-
from pydantic import BaseModel, ConfigDict, Field
-from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend
-from zhenxun.services.ai.context.rag.engine import ScopedRAGClient
from zhenxun.services.ai.utils.scope import ScopeBuilder
from .storage.interfaces import (
BaseChatContext,
BaseMemoryIngestionMiddleware,
BaseMemoryReducer,
- BaseSlotContext,
)
from .types import (
- AutoRecallPolicy,
Isolation,
- MemorySlot,
SessionMetadata,
)
-class SlotMemoryConfig(BaseModel):
- """槽位记忆 (Memory Slots) 配置"""
-
- model_config = ConfigDict(arbitrary_types_allowed=True)
-
- enable: bool = Field(default=False)
- """是否启用中期记忆槽"""
- scopes: dict[str, ScopeBuilder] | None = Field(default=None)
- """语义化作用域映射字典,供大模型作为 Literal 选择。如果只有一个,则自动隐藏参数"""
- default_slots: list[MemorySlot] = Field(default_factory=list)
- """首次初始化时自动写入的默认槽位列表"""
- backend: str | BaseSlotContext | None = Field(default=None)
- """
- 指定底层槽位记忆数据库注册名称,或直接传入 BaseSlotContext 实例。
- 为空则使用全局默认
- """
- instructions: str | None = Field(default=None)
- """覆写内置槽位管理工具箱的系统提示词"""
- toolkit_kwargs: dict[str, Any] = Field(default_factory=dict)
- """透传给底层 MemorySlotToolkit 的高级参数 (如 prefix, exclude, shared_options)"""
-
-
class MemoryScoringConfig(BaseModel):
"""长期记忆的复合打分与检索配置"""
@@ -57,7 +28,6 @@ class MemoryScoringConfig(BaseModel):
"""重要性权重"""
recency_half_life_days: int = Field(default=30)
"""时间衰减的半衰期(天)"""
-
reinforcement_weight: float = Field(default=0.2)
"""访问强化的加权权重 (被检索越多得分越高)"""
@@ -78,44 +48,6 @@ class ShortTermConfig(BaseModel):
"""单一的记忆隔离级别 (ScopeBuilder),决定短期记忆存储边界"""
-class LongTermConfig(BaseModel):
- """长期向量记忆配置"""
-
- model_config = ConfigDict(arbitrary_types_allowed=True)
-
- enable: bool = Field(default=False)
- """是否启用长期记忆(开启后自动赋予 Agent 存取记忆的工具,并附加 RAG 召回能力)"""
- engine: ScopedRAGClient | None = Field(default=None)
- """
- [推荐] 指定底层的高级 RAG 检索引擎实例。若传入此项,将覆盖默认的 backend
- 和 embedder 配置。
- """
- backend: str | StorageBackend | None = Field(default=None)
- """
- 指定底层长期向量数据库 (Storage) 注册名称,或直接传入 StorageBackend 实例。
- 为空则使用全局默认
- """
- scopes: dict[str, ScopeBuilder] | None = Field(default=None)
- """语义化作用域映射字典,决定长期记忆存储边界。如果只有一个,则自动隐藏参数"""
- embedder: str | Embedder | None = Field(default=None)
- """
- 指定底层向量化引擎 (Embedder) 实例,若为字符串则视为 API 模型名称。
- 为空则使用全局默认
- """
-
- agentic: bool = Field(default=True)
- """是否赋予大模型主动管理记忆的能力 (Agentic Memory)"""
- auto_recall: AutoRecallPolicy = Field(default=False)
- """长期记忆的自动召回策略,默认 False (从不自动召回),由大模型自主
- 决定调用搜索工具"""
- recall_threshold: float = Field(default=0.5)
- """长期记忆召回的最低余弦相似度要求"""
- instructions: str | None = Field(default=None)
- """覆写内置长期记忆管理工具箱的系统提示词"""
- toolkit_kwargs: dict[str, Any] = Field(default_factory=dict)
- """透传给底层 MemoryManagementToolkit 的高级参数"""
-
-
class ContextCompressionConfig(BaseModel):
"""上下文压缩与管理配置"""
@@ -144,14 +76,8 @@ class MemoryConfig(BaseModel):
"""统一的记忆配置项声明"""
model_config = ConfigDict(arbitrary_types_allowed=True)
- base_isolation: ScopeBuilder = Field(default_factory=Isolation.AGENT_USER)
- """顶层基准隔离级别,短期/中期/长期记忆将默认继承此级别"""
short_term: ShortTermConfig = Field(default_factory=ShortTermConfig)
"""短期对话记忆配置"""
- slots: SlotMemoryConfig = Field(default_factory=SlotMemoryConfig)
- """槽位记忆配置"""
- long_term: LongTermConfig = Field(default_factory=LongTermConfig)
- """长期向量记忆配置"""
compression: ContextCompressionConfig = Field(
default_factory=ContextCompressionConfig
)
@@ -161,12 +87,10 @@ class MemoryConfig(BaseModel):
__all__ = [
- "AutoRecallPolicy",
"BaseMemoryIngestionMiddleware",
"ContextCompressionConfig",
"IngestionConfig",
"Isolation",
- "LongTermConfig",
"MemoryConfig",
"MemoryScoringConfig",
"SessionMetadata",
diff --git a/zhenxun/services/ai/context/memory/storage/backends.py b/zhenxun/services/ai/context/memory/storage/backends.py
index 4e7da1d0..fb0252e9 100644
--- a/zhenxun/services/ai/context/memory/storage/backends.py
+++ b/zhenxun/services/ai/context/memory/storage/backends.py
@@ -95,7 +95,9 @@ class DBMessageSerializer:
return content_parts
@staticmethod
- def serialize_content(content_payload: Any) -> list[dict[str, Any]]:
+ def serialize_content(
+ content_payload: list[LLMContentPart] | str,
+ ) -> list[dict[str, Any]]:
"""将 LLMMessage 消息内容序列化为可存储于数据库的 JSON 格式"""
if isinstance(content_payload, str):
return [{"type": "text", "text": content_payload}]
@@ -134,7 +136,7 @@ class MemoryScope:
):
"""初始化长期记忆作用域与 RAG 客户端"""
self.rag_client = rag_client
- self._background_tasks: set[Any] = set()
+ self._background_tasks: set[asyncio.Task[Any]] = set()
async def remember(
self,
@@ -492,3 +494,13 @@ def get_orm_slot_context(model_class: type[AbstractSlotRecord]) -> TortoiseSlotC
[工厂方法] 供第三方开发者调用,将 Tortoise ORM 表直接包装为记忆槽存储系统。
"""
return TortoiseSlotContext(model_class=model_class)
+
+
+__all__ = [
+ "AbstractMemoryRecord",
+ "AbstractSlotRecord",
+ "InMemoryChatContext",
+ "MemoryScope",
+ "TortoiseChatContext",
+ "TortoiseSlotContext",
+]
diff --git a/zhenxun/services/ai/context/memory/storage/interfaces.py b/zhenxun/services/ai/context/memory/storage/interfaces.py
index 978269f3..de96a64c 100644
--- a/zhenxun/services/ai/context/memory/storage/interfaces.py
+++ b/zhenxun/services/ai/context/memory/storage/interfaces.py
@@ -1,5 +1,6 @@
from abc import ABC, abstractmethod
from collections.abc import Sequence
+from typing import Protocol, runtime_checkable
from zhenxun.services.ai.context.memory.types import (
MemorySlot,
@@ -10,7 +11,14 @@ from zhenxun.services.ai.run.context import RunContext
from zhenxun.services.ai.utils.scope import ScopeSelector
-class BaseChatContext(ABC):
+@runtime_checkable
+class IClearableBackend(Protocol):
+ """支持声明式清理的作用域后端协议"""
+
+ async def clear_by_query(self, query: ScopeSelector) -> int | None: ...
+
+
+class BaseChatContext(IClearableBackend, ABC):
"""短期对话历史记忆接口"""
@abstractmethod
@@ -44,13 +52,8 @@ class BaseChatContext(ABC):
"""清空当前会话的历史消息。"""
...
- @abstractmethod
- async def clear_by_query(self, query: ScopeSelector) -> None:
- """根据条件领域查询对象清理对话历史。"""
- ...
-
-class BaseSlotContext(ABC):
+class BaseSlotContext(IClearableBackend, ABC):
"""中期记忆槽持久化接口"""
@abstractmethod
@@ -80,11 +83,6 @@ class BaseSlotContext(ABC):
"""列出当前会话的所有记忆槽(包括未置顶的)。"""
...
- @abstractmethod
- async def clear_by_query(self, query: ScopeSelector) -> None:
- """根据条件领域查询对象清理记忆槽。"""
- ...
-
class BaseMemoryReducer(ABC):
"""记忆压缩器基类"""
@@ -97,7 +95,19 @@ class BaseMemoryReducer(ABC):
model_name: str,
base_overhead: int = 0,
) -> tuple[list[LLMMessage], bool, int]:
- """对消息列表进行压缩处理。"""
+ """
+ 执行记忆压缩处理,精简或提炼对话上下文以降低 Token 消耗。
+
+ 参数:
+ messages: 需要进行压缩的原始 LLM 消息历史列表。
+ current_tokens: 压缩前消息列表的当前 Token 总数。
+ model_name: 用于判定压缩阈值或计算 Token 的底层大模型名称。
+ base_overhead: 基础系统提示词等静态开销的 Token 计数。
+
+ 返回:
+ tuple[list[LLMMessage], bool, int]: 包含压缩后的新消息历史列表、
+ 本次是否实际触发了压缩的布尔标记、以及压缩后的新 Token 总数。
+ """
...
diff --git a/zhenxun/services/ai/context/rag/backends/embedders.py b/zhenxun/services/ai/context/rag/backends/embedders.py
index 6391fef6..5a1049f2 100644
--- a/zhenxun/services/ai/context/rag/backends/embedders.py
+++ b/zhenxun/services/ai/context/rag/backends/embedders.py
@@ -4,6 +4,7 @@ import threading
from typing import Any, Literal, Protocol, runtime_checkable
from zhenxun.services.ai.core.messages import EmbedBatch
+from zhenxun.services.ai.core.options import LLMEmbeddingConfig
from zhenxun.services.ai.llm.api import embed as api_embed
from zhenxun.services.ai.message_builder import MessageBuilder
from zhenxun.services.ai.utils.logger import log_rag as logger
@@ -16,7 +17,7 @@ EmbedTaskType = Literal[
@runtime_checkable
class Embedder(Protocol):
"""
- 向量化引擎协议 (Callable Protocol)。
+ 向量化引擎协议。
任何实现了异步 __call__ 的对象或闭包函数均可作为 Embedder。
"""
@@ -32,7 +33,9 @@ class Embedder(Protocol):
class DefaultEmbedder(Embedder):
"""系统默认的向量化引擎,调用大模型底座 API"""
- def __init__(self, model_name: str | None = None, config: Any = None):
+ def __init__(
+ self, model_name: str | None = None, config: LLMEmbeddingConfig | None = None
+ ):
self.model_name = model_name
self.config = config
@@ -57,10 +60,28 @@ class BaseLocalEmbedder(Embedder, ABC):
def __init__(self, model_name: str):
self.model_name = model_name
self._model_lock = threading.Lock()
+ self._model: Any | None = None
+
+ def _ensure_model_loaded(self) -> None:
+ """线程安全的懒加载机制"""
+ if self._model is None:
+ with self._model_lock:
+ if self._model is None:
+ logger.info(
+ f"正在后台加载本地向量模型: {self.model_name} ... "
+ "(首次加载可能需要较长时间下载)"
+ )
+ self._model = self._load_model_impl()
+ logger.info(f"本地向量模型 {self.model_name} 加载完毕!")
+
+ @abstractmethod
+ def _load_model_impl(self) -> Any:
+ """子类实现:执行具体的依赖导入与模型实例化,并返回模型对象。"""
+ pass
@abstractmethod
def _encode_texts(self, texts: list[str]) -> list[list[float]]:
- """子类只需实现此同步的批量文本向量化方法即可。"""
+ """子类实现:执行同步的批量文本向量化方法。"""
pass
async def __call__(
@@ -93,7 +114,6 @@ class FastEmbedder(BaseLocalEmbedder):
def __init__(self, model_name: str | None = None):
super().__init__(model_name or "BAAI/bge-small-zh-v1.5")
- self.model = None
import importlib.util
@@ -102,29 +122,19 @@ class FastEmbedder(BaseLocalEmbedder):
"⚠️ 使用 FastEmbed 需要额外依赖,请在终端执行: pip install fastembed"
)
- def _ensure_model_loaded(self):
- """线程安全的懒加载机制"""
- if self.model is None:
- with self._model_lock:
- if self.model is None:
- try:
- from fastembed import TextEmbedding
- except ImportError:
- raise ImportError(
- "⚠️ 使用 FastEmbed 需要额外依赖,"
- "请在终端执行: pip install fastembed"
- )
- logger.info(
- f"正在后台加载 FastEmbed 本地模型: {self.model_name} ... "
- "(首次加载可能需要极长时间下载)"
- )
- self.model = TextEmbedding(model_name=self.model_name)
- logger.info(f"FastEmbed 模型 {self.model_name} 加载完毕!")
+ def _load_model_impl(self) -> Any:
+ try:
+ from fastembed import TextEmbedding
+ except ImportError:
+ raise ImportError(
+ "⚠️ 使用 FastEmbed 需要额外依赖,请在终端执行: pip install fastembed"
+ )
+ return TextEmbedding(model_name=self.model_name)
def _encode_texts(self, texts: list[str]) -> list[list[float]]:
self._ensure_model_loaded()
- assert self.model is not None
- return [vec.tolist() for vec in self.model.embed(texts)]
+ assert self._model is not None
+ return [vec.tolist() for vec in self._model.embed(texts)]
class SentenceTransformerEmbedder(BaseLocalEmbedder):
@@ -135,7 +145,6 @@ class SentenceTransformerEmbedder(BaseLocalEmbedder):
def __init__(self, model_name: str | None = None):
super().__init__(model_name or "BAAI/bge-small-zh-v1.5")
- self.model = None
import importlib.util
@@ -145,32 +154,18 @@ class SentenceTransformerEmbedder(BaseLocalEmbedder):
"请在终端执行: pip install sentence-transformers"
)
- def _ensure_model_loaded(self):
- """线程安全的懒加载机制"""
- if self.model is None:
- with self._model_lock:
- if self.model is None:
- try:
- from sentence_transformers import (
- SentenceTransformer,
- )
- except ImportError:
- raise ImportError(
- "⚠️ 使用 SentenceTransformers 需要额外依赖,"
- "请在终端执行: pip install sentence-transformers"
- )
- logger.info(
- "正在后台加载 SentenceTransformer "
- f"本地模型: {self.model_name} ... "
- "(首次加载可能需要极长时间下载)"
- )
- self.model = SentenceTransformer(self.model_name)
- logger.info(
- f"SentenceTransformer 模型 {self.model_name} 加载完毕!"
- )
+ def _load_model_impl(self) -> Any:
+ try:
+ from sentence_transformers import SentenceTransformer
+ except ImportError:
+ raise ImportError(
+ "⚠️ 使用 SentenceTransformers 需要额外依赖,"
+ "请在终端执行: pip install sentence-transformers"
+ )
+ return SentenceTransformer(self.model_name)
def _encode_texts(self, texts: list[str]) -> list[list[float]]:
self._ensure_model_loaded()
- assert self.model is not None
- embeddings = self.model.encode(texts)
+ assert self._model is not None
+ embeddings = self._model.encode(texts)
return embeddings.tolist()
diff --git a/zhenxun/services/ai/context/rag/backends/storages.py b/zhenxun/services/ai/context/rag/backends/storages.py
index 18c21d17..bb4ac116 100644
--- a/zhenxun/services/ai/context/rag/backends/storages.py
+++ b/zhenxun/services/ai/context/rag/backends/storages.py
@@ -11,6 +11,10 @@ from zhenxun.services.ai.context.rag.models import (
SearchResult,
)
from zhenxun.services.ai.context.rag.retrieval import FilterEvaluator
+from zhenxun.services.ai.context.rag.utils import (
+ InMemoryScorer,
+ normalize_vector,
+)
from zhenxun.services.ai.utils.logger import log_rag as logger
from zhenxun.services.ai.utils.scope import ScopeSelector
from zhenxun.services.db_context import Model
@@ -24,9 +28,7 @@ class StorageBackend(Protocol):
"""保存或更新数据块"""
...
- async def search(
- self, query: QueryRequest, scopes: list[str] | None = None
- ) -> list[SearchResult]:
+ async def search(self, request: QueryRequest) -> list[SearchResult]:
"""按向量和前缀检索数据块"""
...
@@ -49,22 +51,6 @@ class StorageBackend(Protocol):
...
-def normalize_vector(vec: list[float] | np.ndarray) -> np.ndarray:
- """将一维向量转化为 float32 数组并进行 L2 归一化"""
- v = np.array(vec, dtype=np.float32)
- norm = np.linalg.norm(v)
- if norm == 0:
- return v
- return v / norm
-
-
-def normalize_matrix(mat: np.ndarray) -> np.ndarray:
- """将二维矩阵的每一行进行 L2 归一化"""
- norms = np.linalg.norm(mat, axis=1, keepdims=True)
- norms[norms == 0] = 1.0
- return mat / norms
-
-
class DictStorageBackend(StorageBackend):
"""基于内存字典的轻量级纯净 RAG 存储实现"""
@@ -83,62 +69,36 @@ class DictStorageBackend(StorageBackend):
else:
self._vectors.pop(r.id, None)
- async def search(
- self, query: QueryRequest, scopes: list[str] | None = None
- ) -> list[SearchResult]:
+ async def search(self, request: QueryRequest) -> list[SearchResult]:
candidate_ids = []
for record in self._records.values():
- if scopes is not None:
- if record.metadata.get("scope", "/") not in scopes:
+ if request.scopes is not None:
+ if record.metadata.get("scope", "/") not in request.scopes:
continue
- if not FilterEvaluator.evaluate(record.metadata, query.metadata_filters):
+ if not FilterEvaluator.evaluate(record.metadata, request.metadata_filters):
continue
- if not query.embedding and query.text and query.text not in record.content:
+ if (
+ not request.embedding
+ and request.text
+ and request.text not in record.content
+ ):
continue
candidate_ids.append(record.id)
if not candidate_ids:
return []
- results = []
- if query.search_type == "sparse":
- import jieba
+ records = [self._records[r_id] for r_id in candidate_ids]
- tokens = set(jieba.lcut_for_search(query.text.lower()))
- for r_id in candidate_ids:
- record = self._records[r_id]
- content = record.content.lower()
- matched_count = sum(1 for t in tokens if t in content)
- if matched_count > 0:
- score = matched_count / len(tokens)
- results.append(SearchResult(record=record, score=score))
- elif query.search_type == "dense" and query.embedding:
- q_vec = normalize_vector(query.embedding)
- valid_ids = [r_id for r_id in candidate_ids if r_id in self._vectors]
-
- if valid_ids:
- try:
- mat = np.array([self._vectors[r_id] for r_id in valid_ids])
- scores = mat @ q_vec
- for r_id, score in zip(valid_ids, scores):
- results.append(
- SearchResult(record=self._records[r_id], score=float(score))
- )
- except ValueError as e:
- logger.warning(
- "⚠️ DictStorage 中缓存的向量维度与当前查询维度不匹配,"
- f"跳过向量检索。原因: {e}"
- )
-
- missing_ids = [r_id for r_id in candidate_ids if r_id not in self._vectors]
- for r_id in missing_ids:
- results.append(SearchResult(record=self._records[r_id], score=0.1))
+ if request.search_type == "sparse":
+ results = InMemoryScorer.calculate_sparse_scores(request.text, records)
+ elif request.search_type == "dense" and request.embedding:
+ results = InMemoryScorer.calculate_dense_scores(request.embedding, records)
else:
- for r_id in candidate_ids:
- results.append(SearchResult(record=self._records[r_id], score=0.1))
+ results = [SearchResult(record=r, score=0.1) for r in records]
results.sort(key=lambda x: x.score, reverse=True)
- return results[: query.limit]
+ return results[: request.limit]
async def update(self, record: BaseRecord) -> None:
if record.id in self._records:
@@ -212,94 +172,48 @@ class TortoiseStorageBackend(StorageBackend):
},
)
- async def search(
- self, query: QueryRequest, scopes: list[str] | None = None
- ) -> list[SearchResult]:
+ async def search(self, request: QueryRequest) -> list[SearchResult]:
query_orm = self.model_class.all()
- if scopes is not None:
- query_orm = query_orm.filter(scope__in=scopes)
+ if request.scopes is not None:
+ query_orm = query_orm.filter(scope__in=request.scopes)
- if query.search_type == "sparse" and query.text:
+ if request.search_type == "sparse" and request.text:
import jieba
from tortoise.expressions import Q
tokens = [
- t for t in jieba.lcut_for_search(query.text) if len(t.strip()) > 1
- ] or [query.text]
+ t for t in jieba.lcut_for_search(request.text) if len(t.strip()) > 1
+ ] or [request.text]
q_expr = Q()
for token in tokens:
q_expr |= Q(content__icontains=token)
query_orm = query_orm.filter(q_expr)
- elif query.search_type == "dense" and not query.embedding and query.text:
- query_orm = query_orm.filter(content__icontains=query.text)
+ elif request.search_type == "dense" and not request.embedding and request.text:
+ query_orm = query_orm.filter(content__icontains=request.text)
rows = await query_orm
valid_rows = []
for row in rows:
row_meta = row.meta_data if isinstance(row.meta_data, dict) else {}
- if not FilterEvaluator.evaluate(row_meta, query.metadata_filters):
+ if not FilterEvaluator.evaluate(row_meta, request.metadata_filters):
continue
valid_rows.append(row)
if not valid_rows:
return []
- results = []
- if query.search_type == "sparse":
- import jieba
+ records = [self._to_base_record(row) for row in valid_rows]
- tokens = set(jieba.lcut_for_search(query.text.lower()))
- for row in valid_rows:
- content = row.content.lower()
- matched_count = sum(1 for t in tokens if t in content)
- score = matched_count / len(tokens) if tokens else 0.1
- results.append(
- SearchResult(record=self._to_base_record(row), score=score)
- )
- elif query.search_type == "dense" and query.embedding:
- q_vec = normalize_vector(query.embedding)
- vec_rows = []
- missing_rows = []
-
- for row in valid_rows:
- if isinstance(row.embedding, list):
- vec_rows.append(row)
- else:
- missing_rows.append(row)
-
- if vec_rows:
- try:
- raw_mat = np.array(
- [r.embedding for r in vec_rows], dtype=np.float32
- )
- norm_mat = normalize_matrix(raw_mat)
- scores = norm_mat @ q_vec
-
- for row, score in zip(vec_rows, scores):
- results.append(
- SearchResult(
- record=self._to_base_record(row), score=float(score)
- )
- )
- except ValueError as e:
- logger.warning(
- "⚠️ 数据库中缓存的向量维度与当前模型查询维度不匹配,"
- f"已安全跳过向量检索(降级为稀疏匹配)。原因: {e}"
- )
-
- for row in missing_rows:
- results.append(
- SearchResult(record=self._to_base_record(row), score=0.1)
- )
+ if request.search_type == "sparse":
+ results = InMemoryScorer.calculate_sparse_scores(request.text, records)
+ elif request.search_type == "dense" and request.embedding:
+ results = InMemoryScorer.calculate_dense_scores(request.embedding, records)
else:
- for row in valid_rows:
- results.append(
- SearchResult(record=self._to_base_record(row), score=0.1)
- )
+ results = [SearchResult(record=r, score=0.1) for r in records]
results.sort(key=lambda x: x.score, reverse=True)
- return results[: query.limit]
+ return results[: request.limit]
async def update(self, record: BaseRecord) -> None:
await self.model_class.filter(id=record.id).update(
@@ -386,50 +300,50 @@ class QdrantStorageBackend(StorageBackend):
)
await self.client.upsert(collection_name=self.collection_name, points=points)
- async def search(
- self, query: QueryRequest, scopes: list[str] | None = None
- ) -> list[SearchResult]:
- if query.search_type == "dense" and not query.embedding:
+ async def search(self, request: QueryRequest) -> list[SearchResult]:
+ if request.search_type == "dense" and not request.embedding:
return []
- if query.embedding:
- await self._ensure_collection(len(query.embedding))
+ if request.embedding:
+ await self._ensure_collection(len(request.embedding))
from qdrant_client.models import FieldCondition, Filter, MatchText, MatchValue
must_conditions = []
- if scopes is not None:
+ if request.scopes is not None:
try:
from qdrant_client.models import MatchAny
must_conditions.append(
- FieldCondition(key="metadata.scope", match=MatchAny(any=scopes))
+ FieldCondition(
+ key="metadata.scope", match=MatchAny(any=request.scopes)
+ )
)
except ImportError:
scope_conditions = [
FieldCondition(key="metadata.scope", match=MatchValue(value=s))
- for s in scopes
+ for s in request.scopes
]
must_conditions.append(Filter(should=scope_conditions))
- if query.metadata_filters:
- for k, v in query.metadata_filters.items():
+ if request.metadata_filters:
+ for k, v in request.metadata_filters.items():
must_conditions.append(
FieldCondition(key=f"metadata.{k}", match=MatchValue(value=v))
)
- if query.search_type == "sparse":
+ if request.search_type == "sparse":
must_conditions.append(
- FieldCondition(key="content", match=MatchText(text=query.text))
+ FieldCondition(key="content", match=MatchText(text=request.text))
)
query_filter = Filter(must=must_conditions) if must_conditions else None
- if query.search_type == "sparse":
+ if request.search_type == "sparse":
results = await self.client.scroll(
collection_name=self.collection_name,
scroll_filter=query_filter,
- limit=query.limit,
+ limit=request.limit,
with_payload=True,
)
return [
@@ -446,8 +360,8 @@ class QdrantStorageBackend(StorageBackend):
results = await self.client.search( # type: ignore
collection_name=self.collection_name,
- query_vector=query.embedding,
- limit=query.limit,
+ query_vector=request.embedding,
+ limit=request.limit,
query_filter=query_filter,
)
@@ -559,27 +473,25 @@ class LanceDBStorageBackend(StorageBackend):
else:
self.db.open_table(self.table_name).add(data)
- async def search(
- self, query: QueryRequest, scopes: list[str] | None = None
- ) -> list[SearchResult]:
+ async def search(self, request: QueryRequest) -> list[SearchResult]:
if self.table_name not in self.db.table_names():
return []
- if query.search_type == "dense" and not query.embedding:
+ if request.search_type == "dense" and not request.embedding:
return []
tbl = self.db.open_table(self.table_name)
- if query.search_type == "sparse":
+ if request.search_type == "sparse":
try:
results = (
- tbl.search(query.text, query_type="fts")
- .limit(query.limit)
+ tbl.search(request.text, query_type="fts")
+ .limit(request.limit)
.to_list()
)
except Exception as e:
logger.warning(f"LanceDB FTS 检索失败(可能是由于尚未创建FTS索引): {e}")
return []
else:
- results = tbl.search(query.embedding).limit(query.limit).to_list()
+ results = tbl.search(request.embedding).limit(request.limit).to_list()
import ast
diff --git a/zhenxun/services/ai/context/rag/builder.py b/zhenxun/services/ai/context/rag/builder.py
index 753d5c96..fcd71847 100644
--- a/zhenxun/services/ai/context/rag/builder.py
+++ b/zhenxun/services/ai/context/rag/builder.py
@@ -2,7 +2,7 @@ from typing import Any
from zhenxun.services.ai.utils.logger import log_rag as logger
-from .backends import StorageBackend
+from .backends import Embedder, StorageBackend
from .configs import RAGConfig
from .engine import ScopedRAGClient
from .ingestion import (
@@ -29,18 +29,25 @@ from .retrieval import (
class RAGBuilder:
"""
- RAG 管线组装构建器 (Fluent Builder Pattern)。
- 使用内部状态驱动模式 (The Memory Pattern),内部维护私有的 RAGConfig 实例。
+ RAG 管线组装构建器。
+ 使用内部状态驱动模式,内部维护私有的 RAGConfig 实例。
"""
def __init__(
self, storage: StorageBackend | None = None, config: RAGConfig | None = None
):
+ """
+ 初始化 RAG 管线组装构建器。
+
+ 参数:
+ storage: (可选) 底层存储后端实例。
+ config: (可选) RAG 配置对象。若不提供,将新建默认的 RAGConfig 实例。
+ """
self._config = config or RAGConfig()
if storage is not None:
self._config.storage = storage
- def with_embedder(self, embedder: Any) -> "RAGBuilder":
+ def with_embedder(self, embedder: Embedder) -> "RAGBuilder":
"""
设置向量化引擎。
diff --git a/zhenxun/services/ai/context/rag/engine.py b/zhenxun/services/ai/context/rag/engine.py
index bad85aed..89ff6e87 100644
--- a/zhenxun/services/ai/context/rag/engine.py
+++ b/zhenxun/services/ai/context/rag/engine.py
@@ -11,11 +11,13 @@ from .ingestion import (
)
from .models import (
BaseRecord,
+ QueryRequest,
SearchResult,
)
from .retrieval import (
BaseRetriever,
)
+from .utils import normalize_query_text
class ScopedRAGClient:
@@ -55,7 +57,15 @@ class ScopedRAGClient:
self.scope_prefix = self.scopes[0] if self.scopes else "/"
async def ingest(self, records: list[BaseRecord]) -> int:
- """通过 RAG Ingestion Pipeline 处理并入库数据"""
+ """
+ 通过 RAG Ingestion Pipeline 处理并入库数据。
+
+ 参数:
+ records: 待导入的数据记录列表。
+
+ 返回:
+ int: 成功导入并存入的记录数量。
+ """
for r in records:
if "scope" not in r.metadata:
r.metadata["scope"] = self.scope_prefix
@@ -73,7 +83,16 @@ class ScopedRAGClient:
"""
多作用域联合切片视图检索 (Union Search)。
并发向多个独立的作用域发起检索,并对结果进行合并、去重和重排。
- """
+
+ 参数:
+ query: 检索的查询对象,可以是文本字符串、向量或高级查询对象。
+ limit: 最大返回结果条数。
+ scopes: (可选) 自定义检索的作用域。若不提供,则使用客户端初始化时设定的 scopes。
+ **kwargs: 传递给底层检索器的其他关键字参数。
+
+ 返回:
+ list[SearchResult]: 检索到的相似度排序后的结果列表。
+ """ # noqa: E501
target_scopes = self.scopes
if scopes is not None:
target_scopes = (
@@ -85,13 +104,40 @@ class ScopedRAGClient:
if not target_scopes:
return []
- kwargs["scopes"] = target_scopes
- return await self.retriever.retrieve(query, limit=limit, **kwargs)
+ text_query = normalize_query_text(query)
+ metadata_filters = kwargs.pop("metadata_filters", None)
+
+ request = QueryRequest(
+ text=text_query,
+ limit=limit,
+ scopes=target_scopes,
+ metadata_filters=metadata_filters,
+ extra=kwargs,
+ )
+
+ return await self.retriever.retrieve(request)
async def update(self, record: BaseRecord) -> None:
+ """
+ 更新指定的已有数据记录。
+ 会自动强行附加当前客户端的单作用域前缀 (scope_prefix)。
+
+ 参数:
+ record: 待更新的完整数据记录对象。
+ """
record.metadata["scope"] = self.scope_prefix
await self.storage.update(record)
async def delete(self, record_ids: list[str] | None = None, **kwargs: Any) -> int:
+ """
+ 删除指定的已有数据记录,操作限定在当前客户端的单作用域前缀下。
+
+ 参数:
+ record_ids: (可选) 待删除的记录 ID 列表。
+ **kwargs: 传递给底层存储后端的其他过滤或删除参数。
+
+ 返回:
+ int: 成功删除的记录条数。
+ """
kwargs["scope_prefix"] = self.scope_prefix
return await self.storage.delete(record_ids=record_ids, **kwargs)
diff --git a/zhenxun/services/ai/context/rag/models.py b/zhenxun/services/ai/context/rag/models.py
index 674a3ed8..521f2e36 100644
--- a/zhenxun/services/ai/context/rag/models.py
+++ b/zhenxun/services/ai/context/rag/models.py
@@ -42,6 +42,10 @@ class QueryRequest(BaseModel):
"""元数据精确匹配字典"""
limit: int = Field(default=10)
"""返回的最大条数"""
+ scopes: list[str] | None = Field(default=None)
+ """检索的数据隔离作用域列表"""
+ extra: dict[str, Any] = Field(default_factory=dict)
+ """透传参数逃生舱"""
StorageConfigType = dict[str, Any]
diff --git a/zhenxun/services/ai/context/rag/retrieval.py b/zhenxun/services/ai/context/rag/retrieval.py
index fff2e05b..dff38e36 100644
--- a/zhenxun/services/ai/context/rag/retrieval.py
+++ b/zhenxun/services/ai/context/rag/retrieval.py
@@ -4,22 +4,15 @@ import time
from typing import TYPE_CHECKING, Any, Protocol, cast, runtime_checkable
from zhenxun.services.ai.utils.logger import log_rag as logger
+from zhenxun.utils.pydantic_compat import model_copy
+from .backends.embedders import Embedder
from .models import QueryRequest, SearchResult
if TYPE_CHECKING:
from .backends.storages import StorageBackend
-def normalize_query_text(query: Any) -> str:
- """辅助函数:提取各种输入形式(如字符串、平台Message对象)的纯文本用于检索"""
- if isinstance(query, str):
- return query
- if hasattr(query, "extract_plain_text"):
- return query.extract_plain_text()
- return str(query) if query is not None else ""
-
-
@runtime_checkable
class BaseRetriever(Protocol):
"""
@@ -29,9 +22,7 @@ class BaseRetriever(Protocol):
"""
@abstractmethod
- async def retrieve(
- self, query: Any, limit: int = 10, **kwargs: Any
- ) -> list[SearchResult]: ...
+ async def retrieve(self, request: QueryRequest) -> list[SearchResult]: ...
@runtime_checkable
@@ -72,7 +63,7 @@ class VectorDBRetriever(BaseRetriever):
def __init__(
self,
storage: "StorageBackend",
- embedder: Any,
+ embedder: Embedder,
scope_prefix: str | None = None,
score_threshold: float = 0.4,
):
@@ -90,29 +81,25 @@ class VectorDBRetriever(BaseRetriever):
self.scope_prefix = scope_prefix
self.score_threshold = score_threshold
- async def retrieve(
- self, query: Any, limit: int = 10, **kwargs: Any
- ) -> list[SearchResult]:
- text_query = normalize_query_text(query)
+ async def retrieve(self, request: QueryRequest) -> list[SearchResult]:
+ if not request.embedding and request.text:
+ vecs = await self.embedder(request.text, task="query")
+ request.embedding = vecs[0] if vecs else None
- vecs = await self.embedder(query, task="query")
- query_vec = vecs[0] if vecs else None
-
- if not text_query.strip() and not query_vec:
+ if not request.text.strip() and not request.embedding:
return []
- req = QueryRequest(
- text=text_query,
- embedding=query_vec,
- limit=limit * 2,
- search_type="dense",
- metadata_filters=kwargs.get("metadata_filters"),
- )
- effective_scopes = kwargs.get(
- "scopes", [self.scope_prefix] if self.scope_prefix else None
- )
- results = await self.storage.search(req, scopes=effective_scopes)
- return [r for r in results if r.score >= self.score_threshold][:limit]
+ request.search_type = "dense"
+ if not request.scopes and self.scope_prefix:
+ request.scopes = [self.scope_prefix]
+
+ original_limit = request.limit
+ request.limit = original_limit * 2
+
+ results = await self.storage.search(request)
+
+ request.limit = original_limit
+ return [r for r in results if r.score >= self.score_threshold][: request.limit]
class DatabaseSparseRetriever(BaseRetriever):
@@ -136,24 +123,20 @@ class DatabaseSparseRetriever(BaseRetriever):
self.scope_prefix = scope_prefix
self.score_threshold = score_threshold
- async def retrieve(
- self, query: Any, limit: int = 10, **kwargs: Any
- ) -> list[SearchResult]:
- text_query = normalize_query_text(query)
- if not text_query.strip():
+ async def retrieve(self, request: QueryRequest) -> list[SearchResult]:
+ if not request.text.strip():
return []
- req = QueryRequest(
- text=text_query,
- limit=limit * 2,
- search_type="sparse",
- metadata_filters=kwargs.get("metadata_filters"),
- )
- effective_scopes = kwargs.get(
- "scopes", [self.scope_prefix] if self.scope_prefix else None
- )
- results = await self.storage.search(req, scopes=effective_scopes)
- return [r for r in results if r.score > self.score_threshold][:limit]
+ request.search_type = "sparse"
+ if not request.scopes and self.scope_prefix:
+ request.scopes = [self.scope_prefix]
+
+ original_limit = request.limit
+ request.limit = original_limit * 2
+
+ results = await self.storage.search(request)
+ request.limit = original_limit
+ return [r for r in results if r.score > self.score_threshold][: request.limit]
class RerankRetriever(BaseRetriever):
@@ -183,19 +166,17 @@ class RerankRetriever(BaseRetriever):
self.oversample_factor = oversample_factor
self.min_oversample = min_oversample
- async def retrieve(
- self, query: Any, limit: int = 10, **kwargs: Any
- ) -> list[SearchResult]:
- oversample_limit = max(limit * self.oversample_factor, self.min_oversample)
- initial_results = await self.base_retriever.retrieve(
- query, limit=oversample_limit, **kwargs
+ async def retrieve(self, request: QueryRequest) -> list[SearchResult]:
+ req_clone = model_copy(request, deep=True)
+ req_clone.limit = max(
+ request.limit * self.oversample_factor, self.min_oversample
)
+ initial_results = await self.base_retriever.retrieve(req_clone)
+
if not initial_results:
return []
- text_query = normalize_query_text(query)
-
docs: list[str | dict[str, str]] = [
res.record.content for res in initial_results
]
@@ -204,14 +185,14 @@ class RerankRetriever(BaseRetriever):
try:
reranked = await rerank(
- query=text_query,
+ query=request.text,
documents=docs,
- top_n=min(limit, self.top_n),
+ top_n=min(request.limit, self.top_n),
model=self.model_name,
)
except Exception as e:
logger.warning(f"Rerank 重排请求失败,将降级返回初筛结果: {e}")
- return initial_results[:limit]
+ return initial_results[: request.limit]
final_results = []
for rr in reranked:
@@ -243,15 +224,11 @@ class PipelineRetriever(BaseRetriever):
self.post_processors = post_processors or []
self.pre_processors = pre_processors or []
- async def retrieve(
- self, query: Any, limit: int = 10, **kwargs: Any
- ) -> list[SearchResult]:
- text_query = normalize_query_text(query)
+ async def retrieve(self, request: QueryRequest) -> list[SearchResult]:
+ requests_to_search = [request]
- queries_to_search = [query]
-
- if text_query.strip():
- processed_texts = [text_query]
+ if request.text.strip():
+ processed_texts = [request.text]
for pp in self.pre_processors:
new_texts = []
for t in processed_texts:
@@ -259,15 +236,23 @@ class PipelineRetriever(BaseRetriever):
processed_texts = new_texts
if len(processed_texts) > 1 or (
- len(processed_texts) == 1 and processed_texts[0] != text_query
+ len(processed_texts) == 1 and processed_texts[0] != request.text
):
- queries_to_search.extend(processed_texts)
+ requests_to_search = []
+ for pt in processed_texts:
+ new_req = model_copy(request, deep=True)
+ new_req.text = pt
+ requests_to_search.append(new_req)
all_results = []
seen_ids = set()
- for q in queries_to_search:
- res = await self.base_retriever.retrieve(q, limit=limit * 2, **kwargs)
+ for req in requests_to_search:
+ original_limit = req.limit
+ req.limit = original_limit * 2
+ res = await self.base_retriever.retrieve(req)
+ req.limit = original_limit
+
for r in res:
if r.record.id not in seen_ids:
seen_ids.add(r.record.id)
@@ -276,9 +261,9 @@ class PipelineRetriever(BaseRetriever):
results = sorted(all_results, key=lambda x: x.score, reverse=True)
for pp in self.post_processors:
- results = await pp.process(results, query)
+ results = await pp.process(results, request.text)
- return results[:limit]
+ return results[: request.limit]
class LifecyclePostProcessor(PostProcessor):
@@ -367,14 +352,22 @@ class HybridRetriever(BaseRetriever):
self.oversample_factor = oversample_factor
self.min_oversample = min_oversample
- async def retrieve(
- self, query: Any, limit: int = 10, **kwargs: Any
- ) -> list[SearchResult]:
- oversample_limit = max(limit * self.oversample_factor, self.min_oversample)
+ async def retrieve(self, request: QueryRequest) -> list[SearchResult]:
+ oversample_limit = max(
+ request.limit * self.oversample_factor, self.min_oversample
+ )
+
+ dense_req = model_copy(request, deep=True)
+ dense_req.limit = oversample_limit
+ dense_req.search_type = "dense"
+
+ sparse_req = model_copy(request, deep=True)
+ sparse_req.limit = oversample_limit
+ sparse_req.search_type = "sparse"
results = await asyncio.gather(
- self.dense_retriever.retrieve(query, limit=oversample_limit, **kwargs),
- self.sparse_retriever.retrieve(query, limit=oversample_limit, **kwargs),
+ self.dense_retriever.retrieve(dense_req),
+ self.sparse_retriever.retrieve(sparse_req),
return_exceptions=True,
)
@@ -435,6 +428,6 @@ class HybridRetriever(BaseRetriever):
logger.debug(
f"⚖️ [HybridSearch] 融合完成: "
f"Dense({len(dense_res)}) + Sparse({len(sparse_res)}) "
- f"-> Merged({len(final_results)}), 截取 Top {limit}"
+ f"-> Merged({len(final_results)}), 截取 Top {request.limit}"
)
- return final_results[:limit]
+ return final_results[: request.limit]
diff --git a/zhenxun/services/ai/context/rag/utils.py b/zhenxun/services/ai/context/rag/utils.py
index 02c5ea33..2c0329df 100644
--- a/zhenxun/services/ai/context/rag/utils.py
+++ b/zhenxun/services/ai/context/rag/utils.py
@@ -1,5 +1,10 @@
+from typing import Any
+
import numpy as np
+from zhenxun.services.ai.context.rag.models import BaseRecord, SearchResult
+from zhenxun.services.ai.utils.logger import log_rag as logger
+
def cosine_similarity(vec1: list[float], vec2: list[float]) -> float:
"""使用 numpy 计算两组向量的余弦相似度"""
@@ -10,3 +15,79 @@ def cosine_similarity(vec1: list[float], vec2: list[float]) -> float:
if norm1 == 0 or norm2 == 0:
return 0.0
return float(np.dot(v1, v2) / (norm1 * norm2))
+
+
+def normalize_vector(vec: list[float] | np.ndarray) -> np.ndarray:
+ """将一维向量转化为 float32 数组并进行 L2 归一化"""
+ v = np.array(vec, dtype=np.float32)
+ norm = np.linalg.norm(v)
+ if norm == 0:
+ return v
+ return v / norm
+
+
+def normalize_matrix(mat: np.ndarray) -> np.ndarray:
+ """将二维矩阵的每一行进行 L2 归一化"""
+ norms = np.linalg.norm(mat, axis=1, keepdims=True)
+ norms[norms == 0] = 1.0
+ return mat / norms
+
+
+def normalize_query_text(query: Any) -> str:
+ """辅助函数:提取各种输入形式(如字符串、平台Message对象)的纯文本用于检索"""
+ if isinstance(query, str):
+ return query
+ if hasattr(query, "extract_plain_text"):
+ return query.extract_plain_text()
+ return str(query) if query is not None else ""
+
+
+class InMemoryScorer:
+ """提供内存级别的纯 Python 向量打分与 BM25 稀疏打分工具类"""
+
+ @staticmethod
+ def calculate_sparse_scores(
+ query: str, candidate_records: list[BaseRecord]
+ ) -> list[SearchResult]:
+ import jieba
+
+ tokens = set(jieba.lcut_for_search(query.lower()))
+ if not tokens:
+ return [SearchResult(record=r, score=0.1) for r in candidate_records]
+
+ results = []
+ for r in candidate_records:
+ content_lower = r.content.lower()
+ matched_count = sum(1 for t in tokens if t in content_lower)
+ score = matched_count / len(tokens) if tokens else 0.1
+ results.append(SearchResult(record=r, score=score))
+ return results
+
+ @staticmethod
+ def calculate_dense_scores(
+ query_embedding: list[float], candidate_records: list[BaseRecord]
+ ) -> list[SearchResult]:
+ if not candidate_records or not query_embedding:
+ return []
+
+ q_vec = normalize_vector(query_embedding)
+ vec_records = [r for r in candidate_records if r.embedding]
+ missing_records = [r for r in candidate_records if not r.embedding]
+
+ results = []
+ if vec_records:
+ try:
+ raw_mat = np.array([r.embedding for r in vec_records], dtype=np.float32)
+ norm_mat = normalize_matrix(raw_mat)
+ scores = norm_mat @ q_vec
+ for r, score in zip(vec_records, scores):
+ results.append(SearchResult(record=r, score=float(score)))
+ except ValueError as e:
+ logger.warning(
+ "⚠️ 维度不匹配,已安全跳过向量检索(降级)。原因: " + str(e)
+ )
+
+ for r in missing_records:
+ results.append(SearchResult(record=r, score=0.1))
+
+ return results
diff --git a/zhenxun/services/ai/core/engine/structured_parser.py b/zhenxun/services/ai/core/engine/structured_parser.py
index 7c40ff16..da18c35b 100644
--- a/zhenxun/services/ai/core/engine/structured_parser.py
+++ b/zhenxun/services/ai/core/engine/structured_parser.py
@@ -1,3 +1,4 @@
+import json
import types
from typing import Any, Generic, Union, cast, get_origin
@@ -6,14 +7,9 @@ from nonebot.compat import type_validate_json
from pydantic import BaseModel, Field, ValidationError, create_model
from zhenxun.services.ai.core.exceptions import (
- ControlFlowExit,
- ModelRetry,
SchemaParseError,
)
-from zhenxun.services.ai.core.models import ToolDefinition
-from zhenxun.services.ai.run.models import OutputDataT
-from zhenxun.services.ai.tools.core.tool import BaseTool
-from zhenxun.services.ai.tools.models import StructuredSubmissionResult, ToolResult
+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
@@ -96,8 +92,6 @@ class BaseOutputProcessor(Generic[OutputDataT]):
def _parse_and_validate(self, text: str) -> Any:
"""[私有方法] 执行带有容错修复的 JSON 解析与模型验证"""
if self.raw_schema is not None:
- import json
-
try:
return json.loads(text)
except Exception:
@@ -152,74 +146,3 @@ class BaseOutputProcessor(Generic[OutputDataT]):
return final_obj
except Exception as e:
raise e
-
-
-class SubmitFinalResultExecutable(BaseTool):
- """
- 动态生成的提交最终结果工具。
- 用于将大模型的结构化输出拦截并终止 AgentExecutor 的循环。
- """
-
- def __init__(
- self,
- output_processor: BaseOutputProcessor,
- guardrails: list[Any] | 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: Any | 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: Any | 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)
diff --git a/zhenxun/services/ai/core/messages/types.py b/zhenxun/services/ai/core/messages/types.py
index c1676f8f..94ba9351 100644
--- a/zhenxun/services/ai/core/messages/types.py
+++ b/zhenxun/services/ai/core/messages/types.py
@@ -25,6 +25,8 @@ RoleT = TypeVar("RoleT", default=str, covariant=True)
ContentT = TypeVar("ContentT", default=LLMContentPart, covariant=True)
"""泛型:多模态片段数组的元素内容类型变量"""
+OutputDataT = TypeVar("OutputDataT", default=str)
+
from .context_events import AgentEvent
from .models import (
@@ -53,6 +55,7 @@ __all__ = [
"AnyLLMMessage",
"ContentT",
"LLMContentPart",
+ "OutputDataT",
"PromptInput",
"RoleT",
"UserContentUnion",
diff --git a/zhenxun/services/ai/core/stream_events.py b/zhenxun/services/ai/core/stream_events.py
index 1bf8da23..729e8446 100644
--- a/zhenxun/services/ai/core/stream_events.py
+++ b/zhenxun/services/ai/core/stream_events.py
@@ -106,6 +106,7 @@ class EventBus:
defaultdict(list)
)
self._background_tasks = set()
+ self._async_handlers_queues: dict[Callable, asyncio.Queue] = {}
def subscribe(
self, event_type: type[T_Event], handler: Callable[[T_Event], Any]
@@ -119,6 +120,33 @@ class EventBus:
"""
self._subscribers[event_type].append(handler)
+ if (
+ is_coroutine_callable(handler)
+ and handler not in self._async_handlers_queues
+ ):
+ q = asyncio.Queue(maxsize=1000)
+ self._async_handlers_queues[handler] = q
+
+ async def _worker(h=handler, queue=q):
+ while True:
+ evt = await queue.get()
+ if evt is None:
+ queue.task_done()
+ break
+ try:
+ await h(evt)
+ except asyncio.CancelledError:
+ queue.task_done()
+ break
+ except Exception as err:
+ logger.error(f"EventBus 订阅者执行异常: {err}")
+ finally:
+ queue.task_done()
+
+ task = asyncio.create_task(_worker())
+ self._background_tasks.add(task)
+ task.add_done_callback(self._background_tasks.discard)
+
async def emit(self, event: AgentStreamEvent) -> None:
"""
发布事件,触发所有匹配的订阅者,并将事件放入迭代队列中。
@@ -133,21 +161,14 @@ class EventBus:
for handler in handlers:
if is_coroutine_callable(handler):
-
- async def _run_handler(h=handler, e=event):
- try:
- await h(e)
- except Exception as err:
- logger.error(f"EventBus 订阅者执行异常: {err}")
-
- task = asyncio.create_task(_run_handler())
- self._background_tasks.add(task)
- task.add_done_callback(self._background_tasks.discard)
+ q = self._async_handlers_queues.get(handler)
+ if q:
+ await q.put(event)
else:
try:
handler(event)
except Exception as err:
- logger.error(f"EventBus 订阅者执行异常: {err}")
+ logger.error(f"EventBus 同步订阅者执行异常: {err}")
if not self._finished:
await self._queue.put(event)
@@ -156,11 +177,17 @@ class EventBus:
"""
结束事件总线。等待所有后台任务执行完毕,并向队列中投放结束标记以终止异步迭代。
"""
+ self._finished = True
+ await self._queue.put(None)
+
+ for q in self._async_handlers_queues.values():
+ await q.put(None)
+
if self._background_tasks:
tasks = list(self._background_tasks)
await asyncio.gather(*tasks, return_exceptions=True)
- self._finished = True
- await self._queue.put(None)
+ self._background_tasks.clear()
+ self._async_handlers_queues.clear()
async def __aiter__(self):
"""
diff --git a/zhenxun/services/ai/flow/agent/agent.py b/zhenxun/services/ai/flow/agent/agent.py
index 598f87f6..11a0b94b 100644
--- a/zhenxun/services/ai/flow/agent/agent.py
+++ b/zhenxun/services/ai/flow/agent/agent.py
@@ -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,
diff --git a/zhenxun/services/ai/flow/agent/capabilities.py b/zhenxun/services/ai/flow/agent/capabilities.py
index 756c1836..8c40aea1 100644
--- a/zhenxun/services/ai/flow/agent/capabilities.py
+++ b/zhenxun/services/ai/flow/agent/capabilities.py
@@ -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]
diff --git a/zhenxun/services/ai/flow/agent/engine/builders.py b/zhenxun/services/ai/flow/agent/engine/builders.py
index 86a8b393..20267b63 100644
--- a/zhenxun/services/ai/flow/agent/engine/builders.py
+++ b/zhenxun/services/ai/flow/agent/engine/builders.py
@@ -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
diff --git a/zhenxun/services/ai/flow/agent/engine/executor.py b/zhenxun/services/ai/flow/agent/engine/executor.py
index 37417bdd..af7b8327 100644
--- a/zhenxun/services/ai/flow/agent/engine/executor.py
+++ b/zhenxun/services/ai/flow/agent/engine/executor.py
@@ -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,
+ ),
)
diff --git a/zhenxun/services/ai/flow/agent/models.py b/zhenxun/services/ai/flow/agent/models.py
index 43b0a020..4516cb0c 100644
--- a/zhenxun/services/ai/flow/agent/models.py
+++ b/zhenxun/services/ai/flow/agent/models.py
@@ -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
diff --git a/zhenxun/services/ai/flow/base.py b/zhenxun/services/ai/flow/base.py
deleted file mode 100644
index 8f53e2ed..00000000
--- a/zhenxun/services/ai/flow/base.py
+++ /dev/null
@@ -1,151 +0,0 @@
-from __future__ import annotations
-
-from abc import ABC, abstractmethod
-import asyncio
-from collections.abc import AsyncIterator
-import contextlib
-from enum import Enum
-from typing import TYPE_CHECKING, Any, Generic, TypeVar, cast
-
-from pydantic import BaseModel, Field
-
-from zhenxun.services.ai.core.exceptions import ControlFlowExit
-from zhenxun.services.ai.run.context import RunContext
-from zhenxun.services.ai.run.models import StreamedRunResult
-from zhenxun.services.ai.run.ui import UIController
-from zhenxun.services.ai.utils.logger import log_flow as logger
-
-if TYPE_CHECKING:
- from .agent.models import Persona
-
-from zhenxun.services.ai.core.messages import PromptInput
-
-T_RunResult = TypeVar("T_RunResult")
-
-
-class ConcurrencyPolicy(str, Enum):
- """并发执行策略枚举"""
-
- ALLOW = "allow"
- """允许并发:不做任何限制(适用于无状态或绝对独立任务)"""
- REJECT = "reject"
- """拒绝新请求:当前有任务在执行时,直接丢弃新任务并提醒"""
- QUEUE = "queue"
- """排队等待:当前有任务在执行时,新任务排队等待(先进先出)"""
- INTERRUPT = "interrupt"
- """中断旧任务:新任务到达时,立即强制取消并覆盖正在执行的旧任务"""
-
-
-class ConcurrencyScope(str, Enum):
- """并发作用域枚举(决定锁的粒度,解耦于会话隔离)"""
-
- GLOBAL = "global"
- """全局互斥:整个系统同一时间只能执行一个该任务"""
- GROUP = "group"
- """群组互斥:同一群组内串行排队(私聊退化为用户级),防止抢话刷屏"""
- USER = "user"
- """用户互斥:同一用户发起的任务串行排队(允许同群不同人并行)"""
- SESSION = "session"
- """会话互斥:跟随记忆 SessionID 进行物理锁隔离"""
-
-
-class InterventionPolicy(str, Enum):
- """运行时消息干预策略枚举"""
-
- IGNORE = "ignore"
- """忽略干预:丢弃在任务执行期间收到的额外消息(默认)"""
- STEER = "steer"
- """动态转向:将额外消息立即注入到下一轮大模型推理历史中,影响其思考方向"""
- FOLLOW_UP = "follow_up"
- """追加执行:将额外消息放入队列,在当前大模型意图(所有工具等)执行完毕后追加推理"""
-
-
-class BaseRuntimeConfig(BaseModel):
- """所有可执行实体(Agent/Team/Workflow)的通用基础运行时配置"""
-
- stateless: bool = Field(default=True)
- """是否使用临时会话,不持久化历史记录"""
- concurrency_policy: ConcurrencyPolicy | None = Field(default=None)
- """并发执行策略。如果未显式指定,无状态(stateless=True)默认为ALLOW,有状态(stateless=False)默认为QUEUE。"""
- concurrency_scope: ConcurrencyScope | None = Field(default=None)
- """并发作用域,决定锁的粒度。如果未显式指定,默认为 GROUP 级排队。"""
- intervention_policy: InterventionPolicy | None = Field(default=None)
- """运行时干预策略,决定在大模型执行期间接收到新消息时该如何处理数据流合并。"""
-
-
-class BaseRunnable(ABC, Generic[T_RunResult]):
- """
- 所有可执行 AI 编排实体的统一基类 (Composite Pattern)。
- 统一了 Agent, Team, Workflow 的核心契约,支持物理上的任意嵌套。
- """
-
- name: str
- """可执行实体的名称标识"""
-
- description: str
- """可执行实体的详细描述。用于外部路由(Router)或上层智能体(DelegateTool)决定是否调用它"""
-
- persona: "Persona | dict | None" = None
- """(可选) 实体的角色设定 (Persona)。包含 role 和 goal,
- 在多智能体路由移交时优先级最高"""
-
- runtime_config: BaseRuntimeConfig
- """运行时配置,如是否无状态、UI输出模式等"""
-
- def bind(self, **kwargs: Any) -> Any:
- """DI 注入语法糖:返回 Depends,自动绑定当前上下文"""
- from nonebot.params import Depends
-
- from .agent.bridge import AgentRunner
-
- async def _dependency() -> AgentRunner[Any]:
- return AgentRunner[Any](self, **kwargs)
-
- return Depends(_dependency)
-
- async def reply(
- self,
- prompt: PromptInput | None = None,
- reply_to: bool = False,
- *,
- context: RunContext | None = None,
- **kwargs: Any,
- ) -> T_RunResult:
- """交互执行语法糖,自动渲染流式进度并最终将结果回复给终端用户"""
- from .agent.bridge import AgentRunner
-
- runner = AgentRunner(self, context=context, **kwargs)
- return cast(T_RunResult, await runner.reply(prompt=prompt, reply_to=reply_to))
-
- async def run(
- self,
- prompt: PromptInput | None = None,
- *,
- context: RunContext | None = None,
- **kwargs: Any,
- ) -> T_RunResult:
- """阻塞式核心运行入口,安全捕获内部抛出的静默退出信号"""
-
- try:
- async with self.run_stream(
- prompt=prompt, context=context, **kwargs
- ) as stream_result:
- return cast(T_RunResult, await stream_result.get_run_result())
- except ControlFlowExit as e:
- logger.info(f"[{self.name}] 触发底层控制流,已安全退出: {e}")
-
- await UIController.handle_control_flow_exit_display(e, context)
-
- raise asyncio.CancelledError()
-
- @abstractmethod
- @contextlib.asynccontextmanager
- async def run_stream(
- self,
- prompt: PromptInput | None = None,
- *,
- context: RunContext | None = None,
- **kwargs: Any,
- ) -> AsyncIterator[StreamedRunResult[Any]]:
- """流式运行入口,返回上下文管理器,用于消费底层执行流事件 (StreamedRunResult)"""
- yield cast(Any, None)
diff --git a/zhenxun/services/ai/flow/core/__init__.py b/zhenxun/services/ai/flow/core/__init__.py
new file mode 100644
index 00000000..e0725f89
--- /dev/null
+++ b/zhenxun/services/ai/flow/core/__init__.py
@@ -0,0 +1,17 @@
+from .base import BaseRunnable
+from .models import (
+ BaseRuntimeConfig,
+ ConcurrencyPolicy,
+ ConcurrencyScope,
+ InterventionPolicy,
+)
+from .runner import FlowRunner
+
+__all__ = [
+ "BaseRunnable",
+ "BaseRuntimeConfig",
+ "ConcurrencyPolicy",
+ "ConcurrencyScope",
+ "FlowRunner",
+ "InterventionPolicy",
+]
diff --git a/zhenxun/services/ai/flow/core/base.py b/zhenxun/services/ai/flow/core/base.py
new file mode 100644
index 00000000..68f06ba3
--- /dev/null
+++ b/zhenxun/services/ai/flow/core/base.py
@@ -0,0 +1,214 @@
+from __future__ import annotations
+
+from abc import ABC, abstractmethod
+import asyncio
+from collections.abc import AsyncIterator
+import contextlib
+from typing import TYPE_CHECKING, Any, Generic, TypeVar, cast
+
+from nonebot.params import Depends
+
+from zhenxun.services.ai.core.exceptions import (
+ ConcurrencyInterruptException,
+ ControlFlowExit,
+)
+from zhenxun.services.ai.core.models import CancellationToken
+from zhenxun.services.ai.core.stream_events import AgentStreamEvent, EventBus
+from zhenxun.services.ai.run.context import RunContext
+from zhenxun.services.ai.run.models import AgentRunError, RunIntent, StreamedRunResult
+from zhenxun.services.ai.run.subscribers import DefaultUISubscriber, TelemetrySubscriber
+from zhenxun.services.ai.run.ui import UIController
+from zhenxun.services.ai.utils import ContextUtils
+from zhenxun.services.ai.utils.logger import log_flow as logger
+
+if TYPE_CHECKING:
+ from zhenxun.services.ai.flow.agent.models import Persona
+
+from zhenxun.services.ai.core.messages import PromptInput
+
+from .models import (
+ BaseRuntimeConfig,
+ ConcurrencyPolicy,
+)
+
+T_RunResult = TypeVar("T_RunResult")
+
+
+class BaseRunnable(ABC, Generic[T_RunResult]):
+ """
+ 所有可执行 AI 编排实体的统一基类
+ 统一了 Agent, Team, Workflow 的核心契约,支持物理上的任意嵌套。
+ """
+
+ name: str
+ """可执行实体的名称标识"""
+
+ description: str
+ """可执行实体的详细描述。用于外部路由(Router)或上层智能体(DelegateTool)决定是否调用它"""
+
+ persona: "Persona | None" = None
+ """(可选) 实体的角色设定 (Persona)。包含 role 和 goal,
+ 在多智能体路由移交时优先级最高"""
+
+ runtime_config: BaseRuntimeConfig
+ """运行时配置,如是否无状态、UI输出模式等"""
+
+ @property
+ def profile_summary(self) -> str:
+ """获取该实体的标准化简要画像/描述,供上层路由和规划决策使用"""
+ if self.persona:
+ return f"角色:{self.persona.role},目标:{self.persona.goal}"
+ return self.description or "处理节点"
+
+ def bind(self, **kwargs: Any) -> Any:
+ """DI 注入语法糖:返回 Depends,自动绑定当前上下文"""
+
+ from .runner import FlowRunner
+
+ async def _dependency() -> FlowRunner[Any]:
+ return FlowRunner[Any](self, **kwargs)
+
+ return Depends(_dependency)
+
+ async def reply(
+ self,
+ prompt: PromptInput | None = None,
+ reply_to: bool = False,
+ *,
+ context: RunContext | None = None,
+ **kwargs: Any,
+ ) -> T_RunResult:
+ """交互执行语法糖,自动渲染流式进度并最终将结果回复给终端用户"""
+ from .runner import FlowRunner
+
+ runner = FlowRunner(self, context=context, **kwargs)
+ return cast(T_RunResult, await runner.reply(prompt=prompt, reply_to=reply_to))
+
+ async def run(
+ self,
+ prompt: PromptInput | None = None,
+ *,
+ context: RunContext | None = None,
+ **kwargs: Any,
+ ) -> T_RunResult:
+ """阻塞式核心运行入口,安全捕获内部抛出的静默退出信号"""
+
+ try:
+ async with self.run_stream(
+ prompt=prompt, context=context, **kwargs
+ ) as stream_result:
+ return cast(T_RunResult, await stream_result.get_run_result())
+ except ControlFlowExit as e:
+ logger.info(f"[{self.name}] 触发底层控制流,已安全退出: {e}")
+
+ await UIController.handle_control_flow_exit_display(e, context)
+
+ raise asyncio.CancelledError()
+
+ @contextlib.asynccontextmanager
+ async def run_stream(
+ self,
+ prompt: PromptInput | None = None,
+ *,
+ context: RunContext | None = None,
+ deps: Any = None,
+ event_bus: EventBus | None = None,
+ **kwargs: Any,
+ ) -> AsyncIterator[StreamedRunResult[Any]]:
+ """统一的流式运行入口,负责生命周期调度、并发锁管理和事件总线挂载。"""
+ from .concurrency import apply_concurrency_policy
+
+ intent = RunIntent.from_input(prompt)
+ bus = event_bus or EventBus()
+
+ if context is None:
+ safe_context = RunContext(session_id=kwargs.get("session_id"))
+ if deps is not None:
+ safe_context.deps = deps
+ else:
+ safe_context = context
+ if deps is not None and safe_context.deps is None:
+ safe_context.deps = deps
+
+ is_root = not safe_context.state.get("__is_root_run_executed__", False)
+ if is_root:
+ safe_context.state["__is_root_run_executed__"] = True
+ TelemetrySubscriber().attach(bus)
+ if (
+ safe_context.get_bot()
+ and safe_context.get_event()
+ and safe_context.run.delegate_depth == 0
+ ):
+ config_obj = getattr(self, "config", self.runtime_config)
+ verbose_ui = getattr(config_obj, "verbose_ui", False)
+ DefaultUISubscriber(safe_context, verbose=verbose_ui).attach(bus)
+
+ policy = getattr(self.runtime_config, "concurrency_policy", None)
+ if policy is None:
+ policy = (
+ ConcurrencyPolicy.ALLOW
+ if getattr(self.runtime_config, "stateless", True)
+ else ConcurrencyPolicy.QUEUE
+ )
+
+ intervention_policy = getattr(self.runtime_config, "intervention_policy", None)
+ lock_id = ContextUtils.extract_concurrency_lock_id(
+ safe_context,
+ getattr(self.runtime_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,
+ intent=intent,
+ ):
+ async for event in self._execute_stream(
+ intent=intent,
+ context=safe_context,
+ cancel_token=cancel_token,
+ event_bus=bus,
+ **kwargs,
+ ):
+ await bus.emit(event)
+ except ControlFlowExit as e:
+ await bus.emit(AgentRunError(error=e))
+ except asyncio.CancelledError:
+ logger.debug(f"[{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[Any](bus)
+ try:
+ yield result_obj
+ finally:
+ if not task.done():
+ task.cancel()
+
+ @abstractmethod
+ async def _execute_stream(
+ self,
+ intent: RunIntent,
+ context: RunContext,
+ cancel_token: CancellationToken,
+ event_bus: EventBus,
+ **kwargs: Any,
+ ) -> AsyncIterator[AgentStreamEvent]:
+ """核心执行流(由子类实现),通过 yield 返回执行事件。"""
+ if False:
+ yield cast(Any, None)
diff --git a/zhenxun/services/ai/flow/concurrency.py b/zhenxun/services/ai/flow/core/concurrency.py
similarity index 88%
rename from zhenxun/services/ai/flow/concurrency.py
rename to zhenxun/services/ai/flow/core/concurrency.py
index 0af17685..9b254cfa 100644
--- a/zhenxun/services/ai/flow/concurrency.py
+++ b/zhenxun/services/ai/flow/core/concurrency.py
@@ -1,17 +1,16 @@
import asyncio
from contextlib import asynccontextmanager
-from typing import Any
from zhenxun.services.ai.core.exceptions import (
ConcurrencyRejectException,
InterventionHandledException,
)
from zhenxun.services.ai.core.models import CancellationToken
-from zhenxun.services.ai.run.models import AgentTask
+from zhenxun.services.ai.run.models import RunIntent
from zhenxun.services.ai.run.session import LockContext, session_manager
from zhenxun.services.ai.utils.logger import log_flow as logger
-from .base import ConcurrencyPolicy, InterventionPolicy
+from .models import ConcurrencyPolicy, InterventionPolicy
@asynccontextmanager
@@ -20,8 +19,8 @@ async def apply_concurrency_policy(
lock_id: str,
policy: ConcurrencyPolicy,
cancel_token: CancellationToken,
- intervention_policy: Any = None,
- message: Any = None,
+ intervention_policy: InterventionPolicy | None = None,
+ intent: RunIntent | None = None,
):
"""
异步上下文管理器:对大模型执行流应用特定的并发及消息干预调度策略。
@@ -56,21 +55,14 @@ async def apply_concurrency_policy(
):
session = await session_manager.get_or_create(session_id)
- actual_msg = message
-
- if isinstance(message, AgentTask):
- actual_msg = message.description
- elif hasattr(message, "extract_plain_text"):
- actual_msg = message.extract_plain_text()
-
if intervention_policy == InterventionPolicy.STEER:
- session.steer_queue.enqueue(str(actual_msg))
+ session.steer_queue.enqueue(intent.text if intent else "")
raise InterventionHandledException(
"Steer successful",
display_content="💬 已将您的补充信息传递给正在思考的 AI...",
)
elif intervention_policy == InterventionPolicy.FOLLOW_UP:
- session.follow_up_queue.enqueue(str(actual_msg))
+ session.follow_up_queue.enqueue(intent.text if intent else "")
raise InterventionHandledException(
"Follow-up successful",
display_content="📝 已记录,AI 处理完当前任务后即刻执行...",
diff --git a/zhenxun/services/ai/flow/core/models.py b/zhenxun/services/ai/flow/core/models.py
new file mode 100644
index 00000000..b9d1ce57
--- /dev/null
+++ b/zhenxun/services/ai/flow/core/models.py
@@ -0,0 +1,53 @@
+from enum import Enum
+
+from pydantic import BaseModel, Field
+
+
+class ConcurrencyPolicy(str, Enum):
+ """并发执行策略枚举"""
+
+ ALLOW = "allow"
+ """允许并发:不做任何限制(适用于无状态或独立任务)"""
+ REJECT = "reject"
+ """拒绝新请求:当前有任务在执行时,直接丢弃新任务并提醒"""
+ QUEUE = "queue"
+ """排队等待:当前有任务在执行时,新任务排队等待(先进先出)"""
+ INTERRUPT = "interrupt"
+ """中断旧任务:新任务到达时,立即强制取消并覆盖正在执行的旧任务"""
+
+
+class ConcurrencyScope(str, Enum):
+ """并发作用域枚举(决定锁的粒度,解耦于会话隔离)"""
+
+ GLOBAL = "global"
+ """全局互斥:整个系统同一时间只能执行一个该任务"""
+ GROUP = "group"
+ """群组互斥:同一群组内串行排队(私聊退化为用户级),防止抢话刷屏"""
+ USER = "user"
+ """用户互斥:同一用户发起的任务串行排队(允许同群不同人并行)"""
+ SESSION = "session"
+ """会话互斥:跟随记忆 SessionID 进行物理锁隔离"""
+
+
+class InterventionPolicy(str, Enum):
+ """运行时消息干预策略枚举"""
+
+ IGNORE = "ignore"
+ """忽略干预:丢弃在任务执行期间收到的额外消息(默认)"""
+ STEER = "steer"
+ """动态转向:将额外消息立即注入到下一轮大模型推理历史中,影响其思考方向"""
+ FOLLOW_UP = "follow_up"
+ """追加执行:将额外消息放入队列,在当前大模型意图(所有工具等)执行完毕后追加推理"""
+
+
+class BaseRuntimeConfig(BaseModel):
+ """所有可执行实体(Agent/Team/Workflow)的通用基础运行时配置"""
+
+ stateless: bool = Field(default=True)
+ """是否使用临时会话,不持久化历史记录"""
+ concurrency_policy: ConcurrencyPolicy | None = Field(default=None)
+ """并发执行策略。如果未显式指定,无状态(stateless=True)默认为ALLOW,有状态(stateless=False)默认为QUEUE。"""
+ concurrency_scope: ConcurrencyScope | None = Field(default=None)
+ """并发作用域,决定锁的粒度。如果未显式指定,默认为 GROUP 级排队。"""
+ intervention_policy: InterventionPolicy | None = Field(default=None)
+ """运行时干预策略,决定在大模型执行期间接收到新消息时该如何处理数据流合并。"""
diff --git a/zhenxun/services/ai/flow/agent/bridge.py b/zhenxun/services/ai/flow/core/runner.py
similarity index 93%
rename from zhenxun/services/ai/flow/agent/bridge.py
rename to zhenxun/services/ai/flow/core/runner.py
index 6fd2ad2a..c20a5227 100644
--- a/zhenxun/services/ai/flow/agent/bridge.py
+++ b/zhenxun/services/ai/flow/core/runner.py
@@ -12,12 +12,12 @@ from zhenxun.services.ai.core.exceptions import (
ControlFlowExit,
InterventionHandledException,
)
-from zhenxun.services.ai.core.messages import UsageInfo
-from zhenxun.services.ai.flow.base import BaseRunnable
+from zhenxun.services.ai.core.messages import PromptInput, UsageInfo
+from zhenxun.services.ai.flow.core.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.models import AgentRunEnd, AgentRunError, AgentTask
from zhenxun.services.ai.run.ui import UIController
-from zhenxun.services.ai.utils.logger import log_agent as logger
+from zhenxun.services.ai.utils.logger import log_flow as logger
from zhenxun.utils.message import MessageUtils
from zhenxun.utils.platform import PlatformUtils
@@ -25,9 +25,9 @@ T_Deps = TypeVar("T_Deps", default=Any)
T_Out = TypeVar("T_Out", default=str)
-class AgentRunner(Generic[T_Out]):
+class FlowRunner(Generic[T_Out]):
"""
- 智能体运行器。
+ 执行流交互运行器。
负责将大模型的纯净数据流包装为平台交互动作(发消息、UI渲染)。
自带 ContextVars 隐式上下文提取魔法。
"""
@@ -63,7 +63,10 @@ class AgentRunner(Generic[T_Out]):
return self.context.get_event()
async def reply(
- self, prompt: Any = None, reply_to: bool = False, **kwargs: Any
+ self,
+ prompt: PromptInput | AgentTask | None = None,
+ reply_to: bool = False,
+ **kwargs: Any,
) -> AgentRunResult[T_Out]:
"""交互式执行:将 Agent 运行过程中的工具调用状态和最终结果自动发送给用户。"""
final_result = None
diff --git a/zhenxun/services/ai/flow/team/capabilities.py b/zhenxun/services/ai/flow/team/capabilities.py
index 66dbfec5..926c1324 100644
--- a/zhenxun/services/ai/flow/team/capabilities.py
+++ b/zhenxun/services/ai/flow/team/capabilities.py
@@ -1,13 +1,11 @@
from collections.abc import Callable, Mapping, Sequence
-import inspect
-from typing import Any, cast
-
-from nonebot.utils import is_coroutine_callable
from zhenxun.services.ai.capabilities import AbstractCapability
+from zhenxun.services.ai.flow.core.base import BaseRunnable
from zhenxun.services.ai.run import RunContext
from zhenxun.services.ai.run.di import DependencyInjector
from zhenxun.services.ai.tools.bridges.handoff import HandoffTool
+from zhenxun.services.ai.tools.core.tool import BaseTool
from .models import Transition
@@ -18,8 +16,8 @@ class TeamRoutingCapability(AbstractCapability):
def __init__(
self,
team_name: str,
- members: list[Any],
- state_flow: Mapping[str, Sequence[Any]] | Callable | None = None,
+ members: list[BaseRunnable],
+ state_flow: Mapping[str, Sequence[Transition | str]] | Callable | None = None,
max_handoffs: int = 3,
):
self.team_name = team_name
@@ -27,7 +25,9 @@ class TeamRoutingCapability(AbstractCapability):
self.state_flow = state_flow
self.max_handoffs = max_handoffs
- async def _get_allowed_transitions(self, context: RunContext) -> list[Any] | None:
+ async def _get_allowed_transitions(
+ self, context: RunContext
+ ) -> list[Transition] | None:
"""核心FSM解析:解析静态字典或动态执行函数获取允许的 Transition 列表"""
if self.state_flow is None:
return None
@@ -45,16 +45,10 @@ class TeamRoutingCapability(AbstractCapability):
]
if callable(self.state_flow):
- sig = inspect.signature(self.state_flow)
- kwargs = await DependencyInjector.resolve_all(
- sig, call_kwargs={}, context=context
+ result = await DependencyInjector.invoke(
+ self.state_flow, call_kwargs={}, context=context
)
- if is_coroutine_callable(self.state_flow):
- result = await cast(Callable, self.state_flow)(**kwargs)
- else:
- result = cast(Callable, self.state_flow)(**kwargs)
-
if result is None:
return None
@@ -62,7 +56,7 @@ class TeamRoutingCapability(AbstractCapability):
return None
- async def get_tools(self, context: RunContext) -> list[Any]:
+ async def get_tools(self, context: RunContext) -> list[BaseTool]:
tools = []
allowed_transitions = await self._get_allowed_transitions(context)
@@ -81,10 +75,7 @@ class TeamRoutingCapability(AbstractCapability):
if transition is None:
continue
- if getattr(m, "persona", None):
- desc = f"角色:{m.persona.role},目标:{m.persona.goal}"
- else:
- desc = getattr(m, "description", "") or "处理节点"
+ desc = m.profile_summary
if transition and getattr(transition, "description", ""):
desc += f" 【移交条件】:{transition.description}"
diff --git a/zhenxun/services/ai/flow/team/models.py b/zhenxun/services/ai/flow/team/models.py
index 2afa8dae..41ecc2db 100644
--- a/zhenxun/services/ai/flow/team/models.py
+++ b/zhenxun/services/ai/flow/team/models.py
@@ -7,9 +7,10 @@ import uuid
from pydantic import BaseModel, ConfigDict, Field
-from zhenxun.services.ai.core.messages import AgentMessage
+from zhenxun.services.ai.core.messages import AgentMessage, PromptInput
from zhenxun.services.ai.core.options import BaseOutputDefinition
-from zhenxun.services.ai.flow.base import BaseRunnable, BaseRuntimeConfig
+from zhenxun.services.ai.flow.core.base import BaseRunnable
+from zhenxun.services.ai.flow.core.models import BaseRuntimeConfig
from zhenxun.services.ai.run import AgentTask
@@ -67,7 +68,7 @@ class CallAction(TeamAction):
agent: str | BaseRunnable[Any]
"""目标 Agent 的名称(字符串)或动态生成的 Agent 实例"""
- task: str | AgentTask
+ task: PromptInput | AgentTask
"""派发给该 Agent 的具体任务或提示词"""
history: Sequence[AgentMessage] | None = None
"""需要传递给该 Agent 的上下文历史记录(可选)"""
diff --git a/zhenxun/services/ai/flow/team/router.py b/zhenxun/services/ai/flow/team/router.py
index f1d0055b..ce1fe73a 100644
--- a/zhenxun/services/ai/flow/team/router.py
+++ b/zhenxun/services/ai/flow/team/router.py
@@ -1,18 +1,18 @@
from abc import ABC, abstractmethod
-from collections.abc import Awaitable, Callable, Mapping, Sequence
-import inspect
+from collections.abc import Callable, Mapping, Sequence
import re
-from typing import Any, cast
-
-from nonebot.utils import is_coroutine_callable
from zhenxun.services.ai.core.messages import AgentMessage
from zhenxun.services.ai.core.templates import PromptTemplate
-from zhenxun.services.ai.run import AgentTask, RunContext
+from zhenxun.services.ai.flow.agent.agent import Agent, ToolSource
+from zhenxun.services.ai.flow.agent.models import AgentConfig
+from zhenxun.services.ai.flow.core.base import BaseRunnable
+from zhenxun.services.ai.run import RunContext, RunIntent
from zhenxun.services.ai.run.di import DependencyInjector
from zhenxun.services.ai.utils.logger import log_team as logger
-from .models import RouteDecision, Transition
+from .capabilities import TeamRoutingCapability
+from .models import RouteDecision, TeamRuntimeConfig, Transition
class BaseRouter(ABC):
@@ -23,7 +23,7 @@ class BaseRouter(ABC):
self,
context: RunContext,
history: Sequence[AgentMessage],
- prompt: str | AgentTask | None = None,
+ intent: RunIntent,
) -> RouteDecision | None:
"""核心路由方法"""
pass
@@ -32,7 +32,9 @@ class BaseRouter(ABC):
class FunctionRouter(BaseRouter):
"""基于纯函数的极速路由器"""
- def __init__(self, selector_func: Callable[..., Any], target: str | None = None):
+ def __init__(
+ self, selector_func: Callable[..., str | bool | None], target: str | None = None
+ ):
"""
初始化基于函数的极速路由器。
@@ -47,27 +49,21 @@ class FunctionRouter(BaseRouter):
self,
context: RunContext,
history: Sequence[AgentMessage],
- prompt: str | AgentTask | None = None,
+ intent: RunIntent,
) -> RouteDecision | None:
- sig = inspect.signature(self.selector_func)
- call_kwargs = {"prompt": prompt, "context": context, "history": history}
- if isinstance(prompt, AgentTask):
- call_kwargs["agent_task"] = prompt
- call_kwargs["task"] = prompt
-
- kwargs_resolved = await DependencyInjector.resolve_all(
- sig, call_kwargs, context
- )
- filtered_kwargs = {
- k: v for k, v in kwargs_resolved.items() if k in sig.parameters
+ call_kwargs = {
+ "intent": intent,
+ "prompt": intent.original_input,
+ "context": context,
+ "history": history,
}
+ if intent.task_obj:
+ call_kwargs["agent_task"] = intent.task_obj
+ call_kwargs["task"] = intent.task_obj
- if is_coroutine_callable(self.selector_func):
- _async_func = cast(Callable[..., Awaitable[Any]], self.selector_func)
- selected_target = await _async_func(**filtered_kwargs)
- else:
- _sync_func = cast(Callable[..., Any], self.selector_func)
- selected_target = _sync_func(**filtered_kwargs)
+ selected_target = await DependencyInjector.invoke(
+ self.selector_func, call_kwargs, context
+ )
if isinstance(selected_target, bool):
if selected_target and self.target:
@@ -97,13 +93,9 @@ class RegexRouter(BaseRouter):
self,
context: RunContext,
history: Sequence[AgentMessage],
- prompt: str | AgentTask | None = None,
+ intent: RunIntent,
) -> RouteDecision | None:
- text_to_match = (
- prompt.description
- if isinstance(prompt, AgentTask)
- else (prompt or context.run.user_input or "")
- )
+ text_to_match = intent.text or context.run.user_input or ""
if self.pattern.search(text_to_match):
logger.debug(f"命中正则极速路由 -> {self.target}")
@@ -127,10 +119,10 @@ class ChainRouter(BaseRouter):
self,
context: RunContext,
history: Sequence[AgentMessage],
- prompt: str | AgentTask | None = None,
+ intent: RunIntent,
) -> RouteDecision | None:
for router in self.routers:
- decision = await router.route(context, history, prompt)
+ decision = await router.route(context, history, intent)
if decision is not None:
return decision
return None
@@ -142,11 +134,11 @@ class LLMRouter(BaseRouter):
def __init__(
self,
team_name: str,
- members: list[Any],
+ members: list[BaseRunnable],
leader_model: str | None = None,
- leader_tools: list[Any] | None = None,
+ leader_tools: list[ToolSource] | None = None,
state_flow: Mapping[str, Sequence[Transition | str]] | Callable | None = None,
- runtime_config: Any = None,
+ runtime_config: TeamRuntimeConfig | None = None,
custom_prompt: str | None = None,
allowed_transitions: list[Transition] | None = None,
max_handoffs: int = 3,
@@ -179,13 +171,8 @@ class LLMRouter(BaseRouter):
self,
context: RunContext,
history: Sequence[AgentMessage],
- prompt: str | AgentTask | None = None,
+ intent: RunIntent,
) -> RouteDecision | None:
- from zhenxun.services.ai.flow.agent.agent import Agent
- from zhenxun.services.ai.flow.agent.models import AgentConfig
-
- from .capabilities import TeamRoutingCapability
-
default_system_prompt = """## 角色与目标
你是一个高级任务路由器 (所在团队: {{ team_name }})。
请根据用户的输入意图,立刻调用相应的移交工具 (transfer_to_...)
@@ -217,13 +204,6 @@ class LLMRouter(BaseRouter):
)
target_model = self.leader_model
- if not target_model:
- for m in self.members:
- if m_model := getattr(m, "model_name", None) or getattr(
- m, "model", None
- ):
- target_model = m_model
- break
router_agent = Agent(
name=f"{self.team_name}_Router",
@@ -240,7 +220,7 @@ class LLMRouter(BaseRouter):
logger.debug("🤖 [LLMRouter] 启动 LLM 思考路由决策...")
res = await router_agent.run(
- prompt=prompt,
+ prompt=intent.original_input,
context=sub_context,
config=AgentConfig(message_history=history),
)
diff --git a/zhenxun/services/ai/flow/team/runner.py b/zhenxun/services/ai/flow/team/runner.py
index e23ef5d4..55971f88 100644
--- a/zhenxun/services/ai/flow/team/runner.py
+++ b/zhenxun/services/ai/flow/team/runner.py
@@ -1,6 +1,9 @@
import asyncio
from collections.abc import AsyncGenerator
-from typing import Any
+from typing import TYPE_CHECKING, Any
+
+if TYPE_CHECKING:
+ from .team import Team
from zhenxun.services.ai.core.exceptions import (
AbortException,
@@ -8,19 +11,19 @@ from zhenxun.services.ai.core.exceptions import (
LLMException,
)
from zhenxun.services.ai.core.messages import UsageInfo
+from zhenxun.services.ai.core.stream_events import AgentStreamEvent
from zhenxun.services.ai.flow.agent.models import AgentConfig
-from zhenxun.services.ai.run import AgentRunResult, RunContext
+from zhenxun.services.ai.run import AgentRunResult, RunContext, RunIntent
from zhenxun.services.ai.run.models import AgentRunEnd
from zhenxun.services.ai.utils.logger import log_team as logger
from zhenxun.utils.pydantic_compat import model_construct
-from .capabilities import TeamRoutingCapability
from .models import (
CallAction,
ConcurrentCallAction,
FinishAction,
)
-from .strategy import BaseTeamStrategy, RouteStrategy
+from .strategy import BaseTeamStrategy
class TeamRunner:
@@ -28,10 +31,33 @@ class TeamRunner:
多智能体团队核心执行引擎。
"""
- def __init__(self, team: Any, strategy: BaseTeamStrategy):
+ def __init__(self, team: "Team", strategy: BaseTeamStrategy):
self.team = team
self.strategy = strategy
+ async def _consume_event_queue(
+ self,
+ queue: asyncio.Queue,
+ tasks: list[asyncio.Task],
+ expected_results: int,
+ results_box: dict[int, tuple[str, AgentRunResult]],
+ ) -> AsyncGenerator[AgentStreamEvent, None]:
+ """辅助方法:统一消费队列中的流事件、异常和结果"""
+ try:
+ while len(results_box) < expected_results:
+ msg_type, *payload = await queue.get()
+ if msg_type == "yield_event":
+ yield payload[0]
+ elif msg_type == "control_flow_error":
+ raise payload[0]
+ elif msg_type == "result":
+ idx, agent_name, agent_res = payload
+ results_box[idx] = (agent_name, agent_res)
+ finally:
+ for task in tasks:
+ if not task.done():
+ task.cancel()
+
async def _execute_call_action_to_queue(
self,
index: int,
@@ -69,14 +95,9 @@ class TeamRunner:
sub_context = context.clone_for_member(target_agent.name)
sub_context.capabilities = list(sub_context.capabilities)
- if isinstance(self.strategy, RouteStrategy):
- routing_cap = TeamRoutingCapability(
- team_name=self.team.name,
- members=self.team.members,
- state_flow=getattr(self.strategy, "state_flow", None),
- max_handoffs=getattr(self.strategy, "max_handoffs", 3),
- )
- sub_context.capabilities.append(routing_cap)
+ sub_context.capabilities.extend(
+ self.strategy.get_member_capabilities(self.team, target_agent)
+ )
logger.debug(f"🚀 **专员 👨💼`{target_agent.name}`** 开始执行子任务...")
@@ -125,7 +146,7 @@ class TeamRunner:
target_name = agent_res.handoff.target
reason = agent_res.handoff.reason
- logger.info(
+ logger.debug(
f"🛣️ **路由决策**: 委派给专员 👨💼`{target_name}` (理由: {reason})"
)
@@ -138,15 +159,37 @@ class TeamRunner:
await queue.put(("result", index, target_agent.name, agent_res))
+ async def _dispatch_actions(
+ self,
+ actions: list[CallAction],
+ context: RunContext,
+ session_id: str,
+ results_container: list,
+ ) -> AsyncGenerator[AgentStreamEvent, None]:
+ """统一的任务派发、流事件转译与结果回传调度器"""
+ queue = asyncio.Queue()
+ tasks = [
+ asyncio.create_task(
+ self._execute_call_action_to_queue(i, act, context, session_id, queue)
+ )
+ for i, act in enumerate(actions)
+ ]
+ results_box: dict[int, tuple[str, AgentRunResult]] = {}
+ async for event in self._consume_event_queue(
+ queue, tasks, len(actions), results_box
+ ):
+ yield event
+ results_container.extend([results_box[i] for i in range(len(actions))])
+
async def run_stream(
- self, prompt: Any, context: RunContext, **kwargs: Any
- ) -> AsyncGenerator[Any, None]:
+ self, intent: RunIntent, context: RunContext, **kwargs: Any
+ ) -> AsyncGenerator[AgentStreamEvent, None]:
session_id = context.session_id or "default_team_session"
- task_desc = getattr(prompt, "description", str(prompt))
+ task_desc = intent.text
- logger.info(f"🤝 **团队 [{self.team.name}] 开始协作**: `{task_desc}`")
+ logger.debug(f"🤝 **团队 [{self.team.name}] 开始协作**: `{task_desc}`")
- plan_gen = self.strategy.generate_plan(self.team, prompt, context, **kwargs)
+ plan_gen = self.strategy.generate_plan(self.team, intent, context, **kwargs)
send_value = None
final_result = None
@@ -160,61 +203,23 @@ class TeamRunner:
break
if isinstance(action, CallAction):
- queue = asyncio.Queue()
- task = asyncio.create_task(
- self._execute_call_action_to_queue(
- 0, action, context, session_id, queue
- )
- )
- try:
- while True:
- msg_type, *payload = await queue.get()
- if msg_type == "yield_event":
- yield payload[0]
- elif msg_type == "control_flow_error":
- raise payload[0]
- elif msg_type == "result":
- idx, agent_name, agent_res = payload
- send_value = agent_res
- cumulative_usage += agent_res.usage
- break
- finally:
- if not task.done():
- task.cancel()
+ res_container = []
+ async for event in self._dispatch_actions(
+ [action], context, session_id, res_container
+ ):
+ yield event
+ _, send_value = res_container[0]
+ cumulative_usage += send_value.usage
elif isinstance(action, ConcurrentCallAction):
- queue = asyncio.Queue()
- tasks = []
- for i, act in enumerate(action.actions):
- tasks.append(
- asyncio.create_task(
- self._execute_call_action_to_queue(
- i, act, context, session_id, queue
- )
- )
- )
-
- results_dict = {}
- try:
- while len(results_dict) < len(action.actions):
- msg_type, *payload = await queue.get()
- if msg_type == "yield_event":
- yield payload[0]
- elif msg_type == "control_flow_error":
- for t in tasks:
- t.cancel()
- raise payload[0]
- elif msg_type == "result":
- idx, agent_name, agent_res = payload
- results_dict[idx] = (agent_name, agent_res)
- cumulative_usage += agent_res.usage
- send_value = [
- results_dict[i] for i in range(len(action.actions))
- ]
- finally:
- for task in tasks:
- if not task.done():
- task.cancel()
+ res_container = []
+ async for event in self._dispatch_actions(
+ action.actions, context, session_id, res_container
+ ):
+ yield event
+ send_value = res_container
+ for _, res in res_container:
+ cumulative_usage += res.usage
elif isinstance(action, FinishAction):
final_result = action.result
@@ -225,7 +230,7 @@ class TeamRunner:
except Exception as e:
raise e
- logger.info(f"🏁 **团队 [{self.team.name}]** 协作圆满结束!")
+ logger.debug(f"🏁 **团队 [{self.team.name}]** 协作圆满结束!")
if not isinstance(final_result, AgentRunResult):
final_result = model_construct(
diff --git a/zhenxun/services/ai/flow/team/strategy.py b/zhenxun/services/ai/flow/team/strategy.py
index 3c5fd665..93c22603 100644
--- a/zhenxun/services/ai/flow/team/strategy.py
+++ b/zhenxun/services/ai/flow/team/strategy.py
@@ -1,16 +1,20 @@
from abc import ABC
from collections.abc import AsyncGenerator, Callable, Mapping, Sequence
+import json
+import re
from typing import TYPE_CHECKING, Any, cast
from pydantic import BaseModel
+from zhenxun.services.ai.capabilities import AbstractCapability
from zhenxun.services.ai.core.exceptions import AbortException
from zhenxun.services.ai.core.messages import AgentMessage, LLMMessage
from zhenxun.services.ai.core.stream_events import ToolStreamChunkEvent
from zhenxun.services.ai.core.templates import PromptTemplate
from zhenxun.services.ai.flow.agent.agent import Agent, ToolSource
from zhenxun.services.ai.flow.agent.models import AgentConfig
-from zhenxun.services.ai.run import AgentTask, RunContext
+from zhenxun.services.ai.flow.core.base import BaseRunnable
+from zhenxun.services.ai.run import RunContext, RunIntent
from zhenxun.services.ai.run.blackboard import BlackboardManager
from zhenxun.services.ai.tools.bridges.delegate import DelegateTool
from zhenxun.services.ai.tools.providers.builtin.blackboard import BlackboardToolkit
@@ -25,7 +29,7 @@ from .models import (
TeamAction,
Transition,
)
-from .router import BaseRouter
+from .router import BaseRouter, ChainRouter, FunctionRouter, LLMRouter
from .task_tools import TaskPlanningToolkit
if TYPE_CHECKING:
@@ -50,10 +54,16 @@ class BaseTeamStrategy(ABC):
template = self.custom_prompt or self.default_system_prompt
return PromptTemplate(template).render(**kwargs)
+ def get_member_capabilities(
+ self, team: "Team", member: BaseRunnable
+ ) -> list[AbstractCapability]:
+ """获取派发给子成员时需要动态注入的能力组件"""
+ return []
+
async def generate_plan(
self,
team: "Team",
- prompt: str | AgentTask | None,
+ intent: RunIntent,
context: RunContext,
**kwargs,
) -> AsyncGenerator[TeamAction, Any]:
@@ -66,36 +76,39 @@ class BaseTeamStrategy(ABC):
2. 使用 `yield FinishAction(...)` 结束团队协作。
"""
- yield FinishAction(
- result="The Strategy has not implemented generate_plan() yet."
- )
+ yield FinishAction(result="该策略尚未实现 generate_plan() 方法。")
def _build_leader_agent(
- self, team: "Team", name: str, instruction: str, tools: list[ToolSource]
+ self,
+ team: "Team",
+ role_name: str,
+ extra_instruction: str = "",
+ extra_tools: list[ToolSource] | None = None,
) -> Agent:
"""
- 统一的团队 Leader / Planner 装配工厂。
- 自动处理无状态配置以及 HITL 状态继承。
+ 统一的团队 Leader / Planner / Broadcaster 装配工厂。
+ 自动合并默认指令、追加指令、基类工具和策略专属工具,
+ 并处理无状态配置以及 HITL 状态继承。
"""
+ instruction = self.get_prompt()
+ if extra_instruction:
+ instruction += f"\n\n{extra_instruction}"
+
+ tools = getattr(self, "leader_tools", []).copy()
+ if extra_tools:
+ tools.extend(extra_tools)
+
leader_config = AgentConfig(
stateless=team.runtime_config.stateless if team.runtime_config else True,
enable_hitl=getattr(team.runtime_config, "leader_enable_hitl", False),
)
- target_model = getattr(self, "leader_model", None) or getattr(
- team, "model", None
- )
- if not target_model:
- for m in team.members:
- if m_model := getattr(m, "model_name", None) or getattr(
- m, "model", None
- ):
- target_model = m_model
- break
+ target_model = getattr(self, "leader_model", None) or team.default_model
return Agent(
- name=name,
+ name=f"{team.name}_{role_name}",
instruction=instruction,
+ persona=team.persona,
model=target_model,
tools=tools,
config=leader_config,
@@ -107,7 +120,7 @@ class RouteStrategy(BaseTeamStrategy):
def __init__(
self,
- state_flow: "Mapping[str, Sequence[str | Any]] | Callable | None" = None,
+ state_flow: Mapping[str, Sequence[str | Any]] | Callable | None = None,
selector_func: Callable[..., str | None] | None = None,
router: BaseRouter | None = None,
leader_model: str | None = None,
@@ -148,17 +161,79 @@ class RouteStrategy(BaseTeamStrategy):
else:
self.state_flow = state_flow
+ def get_member_capabilities(
+ self, team: "Team", member: BaseRunnable
+ ) -> list[AbstractCapability]:
+ from .capabilities import TeamRoutingCapability
+
+ return [
+ TeamRoutingCapability(
+ team_name=team.name,
+ members=team.members,
+ state_flow=self.state_flow,
+ max_handoffs=self.max_handoffs,
+ )
+ ]
+
+ def _build_handoff_history(
+ self, handoff_reason: str, context_data: Any
+ ) -> list[AgentMessage]:
+ """将移交数据格式化为系统的引导历史消息"""
+ handoff_history_messages = []
+ upstream_info = []
+ if handoff_reason:
+ upstream_info.append(f"【移交说明】\n{handoff_reason}")
+ if context_data:
+ if isinstance(context_data, dict):
+ formatted_data = json.dumps(context_data, ensure_ascii=False, indent=2)
+ upstream_info.append(
+ f"【结构化上下文载荷】\n```json\n{formatted_data}\n```"
+ )
+ else:
+ upstream_info.append(f"【核心上下文数据】\n{context_data}")
+
+ if combined_info := "\n\n".join(upstream_info):
+ handoff_msg = LLMMessage.system(
+ f"### 🔄 [来自上游节点的移交数据]\n{combined_info}"
+ )
+ handoff_history_messages.append(handoff_msg)
+ return handoff_history_messages
+
+ def _check_fast_route(
+ self, current_target: str, output_str: str
+ ) -> tuple[bool, str]:
+ """检查并执行快速硬路由,返回(是否触发路由, 新目标名称)"""
+ if (
+ not isinstance(self.state_flow, dict)
+ or current_target not in self.state_flow
+ ):
+ return False, current_target
+
+ for t in self.state_flow[current_target]:
+ if trigger_regex := getattr(t, "trigger_regex", None):
+ if re.search(trigger_regex, output_str):
+ return True, getattr(t, "target", current_target)
+
+ if trigger_func := getattr(t, "trigger_func", None):
+ try:
+ res = trigger_func(output_str)
+ if res:
+ if isinstance(res, str):
+ return True, res
+ return True, getattr(t, "target", current_target)
+ except Exception:
+ pass
+ return False, current_target
+
async def generate_plan(
self,
team: "Team",
- prompt: str | AgentTask | None,
+ intent: RunIntent,
context: RunContext,
**kwargs,
) -> AsyncGenerator[TeamAction, Any]:
router = self.router
if not router:
- from .router import ChainRouter, FunctionRouter, LLMRouter
-
routers = []
if self.selector_func:
routers.append(FunctionRouter(self.selector_func))
@@ -166,7 +241,7 @@ class RouteStrategy(BaseTeamStrategy):
LLMRouter(
team_name=team.name,
members=team.members,
- leader_model=self.leader_model,
+ leader_model=self.leader_model or team.default_model,
leader_tools=self.leader_tools,
state_flow=self.state_flow,
runtime_config=getattr(team, "runtime_config", None),
@@ -182,7 +257,7 @@ class RouteStrategy(BaseTeamStrategy):
logger.debug(f"🛣️ '{team.name}' 正在获取初始路由决策...")
- decision = await router.route(context, [], prompt)
+ decision = await router.route(context, [], intent)
if not decision:
logger.warning(f"🚨 Team '{team.name}' 的所有路由策略未能命中目标。")
raise AbortException(
@@ -210,33 +285,13 @@ class RouteStrategy(BaseTeamStrategy):
display="🚨 团队协作陷入死循环,已被系统强制中断。",
)
- handoff_history_messages: list[AgentMessage] = []
- upstream_info = []
- if handoff_reason:
- upstream_info.append(f"【移交说明】\n{handoff_reason}")
- if context_data:
- if isinstance(context_data, dict):
- import json
-
- formatted_data = json.dumps(
- context_data, ensure_ascii=False, indent=2
- )
- upstream_info.append(
- f"【结构化上下文载荷】\n```json\n{formatted_data}\n```"
- )
- else:
- upstream_info.append(f"【核心上下文数据】\n{context_data}")
- combined_info = "\n\n".join(upstream_info)
-
- if combined_info:
- handoff_msg = LLMMessage.system(
- f"### 🔄 [来自上游节点的移交数据]\n{combined_info}"
- )
- handoff_history_messages.append(handoff_msg)
+ handoff_history_messages = self._build_handoff_history(
+ handoff_reason, context_data
+ )
run_result = yield CallAction(
agent=current_target,
- task=prompt or "",
+ task=intent.original_input or "",
history=handoff_history_messages,
kwargs=kwargs,
)
@@ -249,32 +304,11 @@ class RouteStrategy(BaseTeamStrategy):
output_str = str(run_result.output)
- fast_routed = False
- if isinstance(self.state_flow, dict) and current_target in self.state_flow:
- for t in self.state_flow[current_target]:
- trigger_regex = getattr(t, "trigger_regex", None)
- if trigger_regex:
- import re
-
- if re.search(trigger_regex, output_str):
- current_target = getattr(t, "target", current_target)
- handoff_reason = ""
- context_data = output_str
- fast_routed = True
- break
- trigger_func = getattr(t, "trigger_func", None)
- if trigger_func:
- try:
- if trigger_func(output_str):
- current_target = getattr(t, "target", current_target)
- handoff_reason = ""
- context_data = output_str
- fast_routed = True
- break
- except Exception:
- pass
-
+ fast_routed, new_target = self._check_fast_route(current_target, output_str)
if fast_routed:
+ current_target = new_target
+ handoff_reason = ""
+ context_data = output_str
logger.debug(
f"🛣️ **路由决策**: 委派给专员 👨💼`{current_target}`"
"(系统拦截:正则/函数状态流发生转移)"
@@ -317,16 +351,13 @@ class CoordinateStrategy(BaseTeamStrategy):
async def generate_plan(
self,
team: "Team",
- prompt: str | AgentTask | None,
+ intent: RunIntent,
context: RunContext,
**kwargs,
) -> AsyncGenerator[TeamAction, Any]:
delegation_tools = []
for m in team.members:
- persona = getattr(m, "persona", None)
- desc = getattr(m, "description", "") or "处理节点"
- if persona and not isinstance(persona, dict):
- desc = f"角色:{persona.role},目标:{persona.goal}"
+ desc = m.profile_summary
delegation_tools.append(
DelegateTool(
@@ -337,14 +368,10 @@ class CoordinateStrategy(BaseTeamStrategy):
)
)
- leader_tools = self.leader_tools.copy()
- leader_tools.extend(delegation_tools)
-
leader_agent = self._build_leader_agent(
team=team,
- name=f"{team.name}_Leader",
- instruction=self.get_prompt(),
- tools=leader_tools,
+ role_name="Leader",
+ extra_tools=delegation_tools,
)
logger.debug(f"✨ **团队 [{team.name}] Leader** 正在汇总各方报告...")
@@ -358,7 +385,9 @@ class CoordinateStrategy(BaseTeamStrategy):
logger.debug(f"👨💼 [CoordinateStrategy] '{team.name}' 正在启动协调推理循环...")
- leader_res = yield CallAction(agent=leader_agent, task=prompt or "")
+ leader_res = yield CallAction(
+ agent=leader_agent, task=intent.original_input or ""
+ )
yield FinishAction(result=leader_res)
@@ -391,13 +420,11 @@ class BroadcastStrategy(BaseTeamStrategy):
async def generate_plan(
self,
team: "Team",
- prompt: str | AgentTask | None,
+ intent: RunIntent,
context: RunContext,
**kwargs,
) -> AsyncGenerator[TeamAction, Any]:
- task_desc_str = (
- prompt.description if isinstance(prompt, AgentTask) else (prompt or "")
- )
+ task_desc_str = intent.text
await context.run.emit(
ToolStreamChunkEvent(
@@ -406,7 +433,10 @@ class BroadcastStrategy(BaseTeamStrategy):
)
)
- actions = [CallAction(agent=m.name, task=task_desc_str) for m in team.members]
+ actions = [
+ CallAction(agent=m.name, task=intent.original_input or "")
+ for m in team.members
+ ]
results = yield ConcurrentCallAction(actions=actions)
await context.run.emit(
@@ -430,9 +460,7 @@ class BroadcastStrategy(BaseTeamStrategy):
leader_agent = self._build_leader_agent(
team=team,
- name=f"{team.name}_Leader",
- instruction=self.get_prompt(),
- tools=self.leader_tools,
+ role_name="Leader",
)
leader_res = yield CallAction(agent=leader_agent, task=synthesize_prompt)
@@ -497,33 +525,36 @@ class TaskStrategy(BaseTeamStrategy):
schema=schema, initial_state=initial_state
)
self.bb_toolkit = BlackboardToolkit(self.blackboard)
- self.leader_tools.append(self.bb_toolkit)
+
+ def get_member_capabilities(
+ self, team: "Team", member: BaseRunnable
+ ) -> list[AbstractCapability]:
+ caps = super().get_member_capabilities(team, member)
+ if self.bb_toolkit:
+
+ class BlackboardInjectCapability(AbstractCapability):
+ def __init__(self, tk):
+ self.toolkit = tk
+
+ async def get_tools(self, context: RunContext) -> list[Any]:
+ return [self.toolkit]
+
+ caps.append(BlackboardInjectCapability(self.bb_toolkit))
+ return caps
async def generate_plan(
self,
team: "Team",
- prompt: str | AgentTask | None,
+ intent: RunIntent,
context: RunContext,
**kwargs,
) -> AsyncGenerator[TeamAction, Any]:
if self.blackboard is not None:
context.session.blackboard = self.blackboard
- if self.bb_toolkit:
- for m in team.members:
- if not hasattr(m, "tool_definitions"):
- setattr(m, "tool_definitions", [])
-
- m_tools = getattr(m, "tool_definitions")
- if self.bb_toolkit not in m_tools:
- m_tools.append(self.bb_toolkit)
-
member_infos = []
for m in team.members:
- desc = getattr(m, "description", "") or "处理节点"
- persona = getattr(m, "persona", None)
- if persona and not isinstance(persona, dict):
- desc = f"角色:{persona.role},目标:{persona.goal}"
+ desc = m.profile_summary
member_infos.append(
f'\n'
f" Description: {desc}\n"
@@ -532,18 +563,11 @@ class TaskStrategy(BaseTeamStrategy):
members_xml = "\n" + "\n".join(member_infos) + "\n"
- final_instruction = self.get_prompt() + "\n\n" + members_xml
-
- task_toolkit = TaskPlanningToolkit(members=team.members)
-
- leader_tools = self.leader_tools.copy()
- leader_tools.append(task_toolkit)
-
leader_agent = self._build_leader_agent(
team=team,
- name=f"{team.name}_Planner",
- instruction=final_instruction,
- tools=leader_tools,
+ role_name="Planner",
+ extra_instruction=members_xml,
+ extra_tools=[TaskPlanningToolkit(members=team.members)],
)
logger.debug(
@@ -555,7 +579,7 @@ class TaskStrategy(BaseTeamStrategy):
board = cast(TaskBoardState, context.session.shared_state["__task_board__"])
max_iterations = self.max_iterations
- planner_prompt = prompt
+ planner_prompt = intent.original_input
for iteration in range(max_iterations):
if board.is_goal_complete:
@@ -567,12 +591,7 @@ class TaskStrategy(BaseTeamStrategy):
if not available_tasks:
if iteration > 0:
board_str = board.render_board_to_string()
- goal_str = getattr(prompt, "description", None) or (
- str(prompt) if prompt else ""
- )
- expected_out = getattr(prompt, "expected_output", None)
- if expected_out:
- goal_str += f"\n\n### 🎯 [预期产出要求]\n{expected_out}"
+ goal_str = intent.text
planner_prompt = f"""### 🎯 用户的终极目标 (Original Goal)
{goal_str}
@@ -637,8 +656,6 @@ class TaskStrategy(BaseTeamStrategy):
)
if task.metadata:
- import json
-
meta_str = json.dumps(task.metadata, ensure_ascii=False)
task_prompt += f"\n\n### ⚙️ 附加系统元数据约束:\n{meta_str}"
diff --git a/zhenxun/services/ai/flow/team/task_tools.py b/zhenxun/services/ai/flow/team/task_tools.py
index 8d9a3a83..f004068a 100644
--- a/zhenxun/services/ai/flow/team/task_tools.py
+++ b/zhenxun/services/ai/flow/team/task_tools.py
@@ -2,7 +2,7 @@ from typing import Annotated, Any
from pydantic import Field
-from zhenxun.services.ai.flow.base import BaseRunnable
+from zhenxun.services.ai.flow.core.base import BaseRunnable
from zhenxun.services.ai.run.context import RunContext
from zhenxun.services.ai.tools.core.decorators import tool
from zhenxun.services.ai.tools.core.toolkit import BaseToolkit
diff --git a/zhenxun/services/ai/flow/team/team.py b/zhenxun/services/ai/flow/team/team.py
index efa952bc..cff4afcf 100644
--- a/zhenxun/services/ai/flow/team/team.py
+++ b/zhenxun/services/ai/flow/team/team.py
@@ -1,6 +1,4 @@
-import asyncio
-from collections.abc import Callable, Mapping, Sequence
-import contextlib
+from collections.abc import AsyncIterator, Callable, Mapping, Sequence
from pathlib import Path
from typing import Any
from typing_extensions import Self
@@ -12,24 +10,20 @@ from zhenxun.services.ai.capabilities import (
CapabilitySource,
DynamicCapability,
)
-from zhenxun.services.ai.core.exceptions import ConcurrencyInterruptException
from zhenxun.services.ai.core.messages import PromptInput
from zhenxun.services.ai.core.models import CancellationToken
-from zhenxun.services.ai.core.stream_events import EventBus
+from zhenxun.services.ai.core.stream_events import AgentStreamEvent, EventBus
from zhenxun.services.ai.flow.agent.agent import ToolSource
from zhenxun.services.ai.flow.agent.models import Persona
-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.run import (
AgentRunResult,
AgentTask,
RunContext,
- StreamedRunResult,
+ RunIntent,
)
-from zhenxun.services.ai.run.models import AgentRunError
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.utils.utils import infer_plugin_namespace
from .models import TeamRuntimeConfig, Transition
@@ -52,7 +46,7 @@ class Team(BaseRunnable[AgentRunResult[Any]]):
model: str | Callable[[], str] | None = None,
strategy: BaseTeamStrategy | None = None,
description: str | None = None,
- persona: Persona | dict | None = None,
+ persona: Persona | None = None,
runtime_config: TeamRuntimeConfig | dict | None = None,
capabilities: list[CapabilitySource] | None = None,
skills: Sequence[str | Path | Skill | SkillSource] | None = None,
@@ -74,15 +68,17 @@ class Team(BaseRunnable[AgentRunResult[Any]]):
self.members = members
self.model = model
self.strategy = strategy
- self.description = (
- description
- or f"一个名为 {self.name} 的协作团队,包含 {len(self.members)} 个处理节点。"
- )
self.persona = persona
self.namespace = infer_plugin_namespace() or "unknown"
- self.capabilities: list[Any] = []
+ if description:
+ self.description = description
+ else:
+ self.description = f"一个名为 {self.name} 的协作团队,"
+ f"包含 {len(self.members)} 个处理节点。"
+
+ self.capabilities: list[AbstractCapability] = []
if capabilities:
for cap in capabilities:
if isinstance(cap, AbstractCapability):
@@ -107,6 +103,18 @@ class Team(BaseRunnable[AgentRunResult[Any]]):
getattr(strategy, "selector_func", None) if strategy else None
)
+ @property
+ def default_model(self) -> str | None:
+ """获取当前团队默认调用的可用大模型。优先取自身配置,其次遍历成员寻找可用模型。"""
+ if getattr(self, "model", None):
+ return str(self.model() if callable(self.model) else self.model)
+
+ for m in self.members:
+ m_model = getattr(m, "model_name", None) or getattr(m, "model", None)
+ if m_model:
+ return str(m_model() if callable(m_model) else m_model)
+ return None
+
def with_strategy(self, strategy: BaseTeamStrategy) -> Self:
"""
挂载自定义的团队协作策略。
@@ -122,9 +130,7 @@ class Team(BaseRunnable[AgentRunResult[Any]]):
def with_routing(
self,
- state_flow: (
- Mapping[str, Sequence[Transition | str | Any]] | Callable | None
- ) = None,
+ state_flow: (Mapping[str, Sequence[Transition | str]] | Callable | None) = None,
selector_func: Callable[..., str | None] | None = None,
router: BaseRouter | None = None,
leader_model: str | None = None,
@@ -287,20 +293,18 @@ class Team(BaseRunnable[AgentRunResult[Any]]):
prompt=prompt, context=context, capabilities=capabilities, **kwargs
)
- @contextlib.asynccontextmanager
- async def run_stream(
+ async def _execute_stream(
self,
- prompt: PromptInput | AgentTask | None = None,
- *,
- context: "RunContext | None" = None,
- capabilities: list[CapabilitySource] | None = None,
- skills: Sequence[str | Path | Skill | SkillSource] | None = None,
+ intent: RunIntent,
+ context: RunContext,
+ cancel_token: CancellationToken,
+ event_bus: EventBus,
**kwargs: Any,
- ):
+ ) -> AsyncIterator[AgentStreamEvent]:
self._ensure_strategy()
- if context is None:
- context = RunContext()
+ capabilities = kwargs.pop("capabilities", None)
+ skills = kwargs.pop("skills", None)
if not hasattr(context, "capabilities"):
context.capabilities = []
@@ -310,6 +314,10 @@ class Team(BaseRunnable[AgentRunResult[Any]]):
if skills:
capabilities = list(capabilities) if capabilities else []
+ from zhenxun.services.ai.tools.providers.skills.capabilities import (
+ SkillCapability,
+ )
+
capabilities.append(
SkillCapability(skills=skills, namespace=self.namespace)
)
@@ -323,54 +331,8 @@ class Team(BaseRunnable[AgentRunResult[Any]]):
from .runner import TeamRunner
- event_bus = EventBus()
- context.run.event_bus = event_bus
assert self.strategy is not None
runner = TeamRunner(self, self.strategy)
- policy = getattr(self.runtime_config, "concurrency_policy", None)
- if policy is None:
- policy = (
- ConcurrencyPolicy.ALLOW
- if getattr(self.runtime_config, "stateless", True)
- else ConcurrencyPolicy.QUEUE
- )
-
- intervention_policy = getattr(self.runtime_config, "intervention_policy", None)
-
- lock_id = ContextUtils.extract_concurrency_lock_id(
- context,
- getattr(self.runtime_config, "concurrency_scope", None),
- context.session_id or "default_session",
- )
-
- async def _execution_task():
- cancel_token = context.run.cancellation_token or CancellationToken()
- context.run.cancellation_token = cancel_token
-
- try:
- async with apply_concurrency_policy(
- session_id=context.session_id or "default_session",
- lock_id=lock_id,
- policy=policy,
- cancel_token=cancel_token,
- intervention_policy=intervention_policy,
- message=prompt,
- ):
- async for event in runner.run_stream(prompt, context, **kwargs):
- await event_bus.emit(event)
- except BaseException as e:
- if isinstance(e, asyncio.CancelledError):
- e = ConcurrencyInterruptException("团队执行已被新请求打断并接管")
- await event_bus.emit(AgentRunError(error=e))
- finally:
- await event_bus.end()
-
- task = asyncio.create_task(_execution_task())
- result_obj = StreamedRunResult[Any](event_bus)
-
- try:
- yield result_obj
- finally:
- if not task.done():
- task.cancel()
+ async for event in runner.run_stream(intent, context, **kwargs):
+ yield event
diff --git a/zhenxun/services/ai/flow/workflow/base.py b/zhenxun/services/ai/flow/workflow/base.py
index f1f3c197..625bb03d 100644
--- a/zhenxun/services/ai/flow/workflow/base.py
+++ b/zhenxun/services/ai/flow/workflow/base.py
@@ -1,13 +1,14 @@
from abc import ABC, abstractmethod
import asyncio
from collections.abc import AsyncIterator
-from typing import Any
+from typing import cast
from zhenxun.services.ai.core.exceptions import (
AbortException,
ControlFlowExit,
ToolFatalError,
)
+from zhenxun.services.ai.core.stream_events import AgentStreamEvent
from zhenxun.services.ai.run import RunContext
from zhenxun.services.ai.utils.logger import log_flow as logger
@@ -23,6 +24,21 @@ from .types import (
)
+class StreamCapturer:
+ """内部辅助类:透传工作流节点的内部流事件,并捕获最终的 StepOutput 产出值"""
+
+ def __init__(self, stream: AsyncIterator[AgentStreamEvent | StepOutput]):
+ self._stream = stream
+ self.output: StepOutput | None = None
+
+ async def __aiter__(self) -> AsyncIterator[AgentStreamEvent]:
+ async for event in self._stream:
+ if isinstance(event, StepOutput):
+ self.output = event
+ else:
+ yield event
+
+
class BaseNode(ABC):
"""工作流节点统一抽象基类"""
@@ -50,7 +66,7 @@ class BaseNode(ABC):
async def _handle_execution_failure(
self, e: BaseException, step_input: StepInput, context: RunContext, attempt: int
- ) -> tuple[str, StepOutput | None, StepInput | None, Any]:
+ ) -> tuple[str, StepOutput | None, StepInput | None, "BaseNode | None"]:
"""
解析执行异常并应用容错策略
"""
@@ -128,19 +144,10 @@ class BaseNode(ABC):
@abstractmethod
async def run_stream(
self, step_input: StepInput, context: RunContext
- ) -> AsyncIterator[Any]:
+ ) -> AsyncIterator[AgentStreamEvent | StepOutput]:
"""子类必须实现的核心流式执行逻辑"""
- yield None
-
- async def _forward_stream(
- self, stream: AsyncIterator[Any], output_box: list[StepOutput]
- ) -> AsyncIterator[Any]:
- """辅助方法:转发内部流事件,并将最终的 StepOutput 拦截放入 output_box 列表中"""
- async for event in stream:
- if isinstance(event, StepOutput):
- output_box.append(event)
- else:
- yield event
+ if False:
+ yield cast(StepOutput, None)
async def aexecute(self, step_input: StepInput, context: RunContext) -> StepOutput:
"""非流式执行(聚合流并返回最终结果),子类无需重写"""
@@ -160,7 +167,7 @@ class BaseNode(ABC):
async def aexecute_stream(
self, step_input: StepInput, context: RunContext
- ) -> AsyncIterator[Any]:
+ ) -> AsyncIterator[AgentStreamEvent | StepOutput]:
"""标准化模板方法:处理缓存快进、授权挂起、异常熔断与生命周期事件分发"""
logger.debug(f" ⚙️ [节点] `{self.name}` 开始执行...")
@@ -222,4 +229,5 @@ class BaseNode(ABC):
yield evt
break
+ assert output is not None
yield output
diff --git a/zhenxun/services/ai/flow/workflow/engine.py b/zhenxun/services/ai/flow/workflow/engine.py
index 3c29ccbb..23f9df23 100644
--- a/zhenxun/services/ai/flow/workflow/engine.py
+++ b/zhenxun/services/ai/flow/workflow/engine.py
@@ -1,7 +1,5 @@
-import asyncio
from collections.abc import AsyncIterator
-import contextlib
-from typing import TYPE_CHECKING, Any
+from typing import TYPE_CHECKING, Any, cast
import uuid
from nonebot.params import Depends
@@ -12,15 +10,16 @@ if TYPE_CHECKING:
from zhenxun.services.ai.core.exceptions import ControlFlowExit, ToolRetryError
from zhenxun.services.ai.core.messages import PromptInput, UsageInfo
-from zhenxun.services.ai.core.stream_events import EventBus
-from zhenxun.services.ai.flow.base import BaseRunnable, BaseRuntimeConfig
+from zhenxun.services.ai.core.models import CancellationToken
+from zhenxun.services.ai.core.stream_events import AgentStreamEvent, EventBus
+from zhenxun.services.ai.flow.core.base import BaseRunnable
+from zhenxun.services.ai.flow.core.models import BaseRuntimeConfig
from zhenxun.services.ai.run.blackboard import BlackboardManager
from zhenxun.services.ai.run.context import RunContext
from zhenxun.services.ai.run.models import (
AgentRunEnd,
- AgentRunError,
AgentRunResult,
- StreamedRunResult,
+ RunIntent,
)
from zhenxun.services.ai.tools.core.tool import FunctionTool
from zhenxun.services.ai.utils.logger import log_flow as logger
@@ -176,96 +175,37 @@ class Workflow(BaseRunnable[WorkflowRunResult]):
返回:
WorkflowRunResult: 包含执行状态、断点快照、各节点产出的全量工作流结果对象。
"""
- session_id = (
- context.session_id if context and context.session_id else f"wf_{self.id}"
- )
- safe_context = context or RunContext(session_id=session_id)
+ async with self.run_stream(
+ prompt=prompt, context=context, **kwargs
+ ) as stream_result:
+ res = await stream_result.get_run_result()
+ return cast(WorkflowRunResult, res.structured_data)
- if self.blackboard_schema and not safe_context.session.blackboard:
- safe_context.session.blackboard = BlackboardManager(
+ async def _execute_stream(
+ self,
+ intent: RunIntent,
+ context: RunContext,
+ cancel_token: CancellationToken,
+ event_bus: EventBus,
+ **kwargs: Any,
+ ) -> AsyncIterator[AgentStreamEvent]:
+ """统一核心流,不再自己维护 Task 和 EventBus"""
+
+ if self.blackboard_schema and not context.session.blackboard:
+ context.session.blackboard = BlackboardManager(
schema=self.blackboard_schema,
initial_state=self.initial_blackboard_state,
)
logger.debug(f"🏭 **工作流 [{self.name}] 启动**")
- initial_input = StepInput(input=prompt)
- if kwargs:
- initial_input.additional_data.update(kwargs)
-
- try:
- final_output = await self.root_steps.aexecute(initial_input, safe_context)
-
- logger.debug(f"🏭 **工作流 [{self.name}] 运行结束**")
-
- return self._build_result(initial_input, safe_context, final_output)
-
- except BaseException as e:
- if isinstance(e, ControlFlowExit):
- logger.debug(f"⏭️ 工作流执行被业务控制流安全中止: {e}")
- dummy_output = StepOutput(content=str(e), success=False)
- return self._build_result(initial_input, safe_context, dummy_output)
-
- raise e
-
- @contextlib.asynccontextmanager
- async def run_stream(
- self,
- prompt: PromptInput | None = None,
- *,
- context: RunContext | None = None,
- **kwargs: Any,
- ) -> AsyncIterator["StreamedRunResult[Any]"]:
- """对齐 BaseRunnable 接口的流式上下文管理器"""
- event_bus = EventBus()
- if context:
- context.run.event_bus = event_bus
-
- async def _execution_task():
- try:
- async for event in self._internal_stream(prompt, context, **kwargs):
- await event_bus.emit(event)
- except BaseException as e:
- await event_bus.emit(AgentRunError(error=e))
- finally:
- await event_bus.end()
-
- task = asyncio.create_task(_execution_task())
- try:
- yield StreamedRunResult[Any](event_bus)
- finally:
- if not task.done():
- task.cancel()
-
- async def _internal_stream(
- self,
- prompt: PromptInput | None = None,
- context: RunContext | None = None,
- **kwargs: Any,
- ) -> AsyncIterator[Any]:
- """流式执行工作流节点树的内部实现"""
- session_id = (
- context.session_id if context and context.session_id else f"wf_{self.id}"
- )
- safe_context = context or RunContext(session_id=session_id)
-
- if self.blackboard_schema and not safe_context.session.blackboard:
- safe_context.session.blackboard = BlackboardManager(
- schema=self.blackboard_schema,
- initial_state=self.initial_blackboard_state,
- )
-
- logger.debug(f"🏭 **工作流 [{self.name}] 启动**")
-
- initial_input = StepInput(input=prompt)
+ initial_input = StepInput(input=intent.original_input, intent=intent)
if kwargs:
initial_input.additional_data.update(kwargs)
try:
final_output = None
- async for event in self.root_steps.aexecute_stream(
- initial_input, safe_context
- ):
+ async for event in self.root_steps.aexecute_stream(initial_input, context):
if isinstance(event, StepOutput):
final_output = event
else:
@@ -274,17 +214,26 @@ class Workflow(BaseRunnable[WorkflowRunResult]):
if final_output:
logger.debug(f"🏭 **工作流 [{self.name}] 运行结束**")
- wf_result = self._build_result(
- initial_input, safe_context, final_output
- )
+ wf_result = self._build_result(initial_input, context, final_output)
agent_res = AgentRunResult(
output=wf_result.last_step_content,
structured_data=wf_result,
usage=UsageInfo(),
)
yield AgentRunEnd(result=agent_res)
- except Exception:
- pass
+ except BaseException as e:
+ if isinstance(e, ControlFlowExit):
+ logger.debug(f"⏭️ 工作流执行被业务控制流安全中止: {e}")
+ dummy_output = StepOutput(content=str(e), success=False)
+ wf_result = self._build_result(initial_input, context, dummy_output)
+ agent_res = AgentRunResult(
+ output=wf_result.last_step_content,
+ structured_data=wf_result,
+ usage=UsageInfo(),
+ )
+ yield AgentRunEnd(result=agent_res)
+ else:
+ raise e
def as_tool(self, tool_name: str | None = None) -> FunctionTool:
"""将工作流封装并导出为可供 Agent 直接调用的 FunctionTool 实例"""
diff --git a/zhenxun/services/ai/flow/workflow/nodes.py b/zhenxun/services/ai/flow/workflow/nodes.py
index 755cfe24..78a69edd 100644
--- a/zhenxun/services/ai/flow/workflow/nodes.py
+++ b/zhenxun/services/ai/flow/workflow/nodes.py
@@ -1,16 +1,18 @@
+from abc import ABC, abstractmethod
import asyncio
-from collections.abc import AsyncIterator, Callable, Sequence
+from collections.abc import AsyncIterator, Awaitable, Callable, Sequence
import copy
-from typing import Any, cast
+from typing import cast
from zhenxun.services.ai.core.messages import PromptInput
-from zhenxun.services.ai.flow.base import BaseRunnable
-from zhenxun.services.ai.run import AgentTask, RunContext
+from zhenxun.services.ai.core.stream_events import AgentStreamEvent
+from zhenxun.services.ai.flow.core.base import BaseRunnable
+from zhenxun.services.ai.run import RunContext
from zhenxun.services.ai.run.di import DependencyInjector
-from zhenxun.services.ai.run.models import AgentRunEnd
+from zhenxun.services.ai.run.models import AgentRunEnd, RunIntent
from zhenxun.services.ai.utils.logger import log_flow as logger
-from .base import BaseNode
+from .base import BaseNode, StreamCapturer
from .policies import BaseFailurePolicy
from .types import (
StepInput,
@@ -24,7 +26,7 @@ NodeSource = BaseNode | BaseRunnable | Callable
class Step(BaseNode):
"""
- 工作流中的最小执行单元门面 (Facade)。
+ 工作流中的最小执行单元门面
对外部隐藏了 AgentNode 和 FunctionNode 的具体实现。
当实例化 Step 时,底层会自动根据 executor 的类型返回专属的节点对象。
"""
@@ -35,7 +37,7 @@ class Step(BaseNode):
if executor is None and len(args) > 1:
executor = args[1]
- from zhenxun.services.ai.flow.base import BaseRunnable
+ from zhenxun.services.ai.flow.core.base import BaseRunnable
if isinstance(executor, BaseRunnable):
return object.__new__(RunnableNode)
@@ -75,7 +77,7 @@ class Step(BaseNode):
async def run_stream(
self, step_input: StepInput, context: RunContext
- ) -> AsyncIterator[Any]:
+ ) -> AsyncIterator[AgentStreamEvent | StepOutput]:
if False:
yield None
raise NotImplementedError("这是一个外观门面,实际的执行发生在子类中。")
@@ -86,32 +88,28 @@ class RunnableNode(Step):
async def run_stream(
self, step_input: StepInput, context: RunContext
- ) -> AsyncIterator[Any]:
+ ) -> AsyncIterator[AgentStreamEvent | StepOutput]:
executor = cast(BaseRunnable, self.executor)
- if isinstance(step_input.previous_step_content, AgentTask):
- prompt_data = step_input.previous_step_content
- prev_content_to_append = None
- else:
- prompt_data = self.prompt if self.prompt is not None else step_input.input
- prev_content_to_append = step_input.previous_step_content
+ base_prompt = self.prompt if self.prompt is not None else step_input.input
+ node_intent = RunIntent.from_input(base_prompt)
- if isinstance(prompt_data, AgentTask):
- prompt_data = copy.copy(prompt_data)
- if prev_content_to_append:
- prev_content = str(prev_content_to_append)
- prompt_data.description = (
+ prev_content = step_input.previous_step_content
+ prompt_data = base_prompt
+
+ if prev_content:
+ if node_intent.task_obj:
+ task_clone = copy.copy(node_intent.task_obj)
+ task_clone.description = (
f"### 🔙 [上游节点执行输出]\n{prev_content}\n\n"
- f"### 🎯 [当前需执行的任务]\n{prompt_data.description}"
+ f"### 🎯 [当前需执行的任务]\n{task_clone.description}"
)
- context.run.user_input = prompt_data.description
- else:
- if prev_content_to_append:
+ prompt_data = task_clone
+ else:
prompt_data = (
- f"[上游节点执行输出]:\n{prev_content_to_append}\n\n"
- f"[当前需执行的任务]:\n{prompt_data or ''}"
+ f"[上游节点执行输出]:\n{prev_content}\n\n"
+ f"[当前需执行的任务]:\n{node_intent.text}"
)
- context.run.user_input = str(prompt_data) if prompt_data else ""
final_result = None
sandbox_context = context.clone_for_member(self.name)
@@ -136,7 +134,7 @@ class FunctionNode(Step):
async def run_stream(
self, step_input: StepInput, context: RunContext
- ) -> AsyncIterator[Any]:
+ ) -> AsyncIterator[AgentStreamEvent | StepOutput]:
context.run.user_input = str(step_input.input) if step_input.input else ""
executor = cast(Callable, self.executor)
@@ -170,21 +168,20 @@ class Steps(BaseNode):
async def run_stream(
self, step_input: StepInput, context: RunContext
- ) -> AsyncIterator[Any]:
+ ) -> AsyncIterator[AgentStreamEvent | StepOutput]:
current_input = StepInput(
input=step_input.input,
+ intent=step_input.intent,
previous_step_content=step_input.previous_step_content,
additional_data=step_input.additional_data.copy(),
)
all_outputs: list[StepOutput] = []
for step_obj in self.steps:
- out_box: list[StepOutput] = []
- async for event in self._forward_stream(
- step_obj.aexecute_stream(current_input, context), out_box
- ):
+ capturer = StreamCapturer(step_obj.aexecute_stream(current_input, context))
+ async for event in capturer:
yield event
- step_out = out_box[0] if out_box else None
+ step_out = capturer.output
if step_out:
all_outputs.append(step_out)
@@ -199,12 +196,47 @@ class Steps(BaseNode):
)
-class Condition(BaseNode):
+class BranchingNode(BaseNode, ABC):
+ """
+ 处理基于条件或路由的单分支复合节点基类
+ """
+
+ @abstractmethod
+ async def select_branch(
+ self, step_input: StepInput, context: RunContext
+ ) -> tuple[list[BaseNode], str, str]:
+ """
+ 子类实现此方法进行路由决策。
+ 返回: (目标节点列表, 分支标识名, 未命中有效分支时的兜底提示信息)
+ """
+ pass
+
+ async def run_stream(
+ self, step_input: StepInput, context: RunContext
+ ) -> AsyncIterator[AgentStreamEvent | StepOutput]:
+ context.run.user_input = str(step_input.input) if step_input.input else ""
+ target_steps, branch_name, fallback_msg = await self.select_branch(
+ step_input, context
+ )
+
+ if not target_steps:
+ yield StepOutput(content=fallback_msg, success=True)
+ return
+
+ steps_container = Steps(steps=target_steps, name=f"{self.name}_{branch_name}")
+ capturer = StreamCapturer(steps_container.aexecute_stream(step_input, context))
+ async for event in capturer:
+ yield event
+ if capturer.output:
+ yield capturer.output
+
+
+class Condition(BranchingNode):
"""根据条件函数的返回结果,决定走向 steps 还是 else_steps"""
def __init__(
self,
- evaluator: Any,
+ evaluator: bool | Callable[..., bool | Awaitable[bool]],
steps: Sequence[NodeSource],
else_steps: Sequence[NodeSource] | None = None,
name: str = "ConditionGroup",
@@ -227,9 +259,9 @@ class Condition(BaseNode):
def node_type(self) -> StepType:
return StepType.CONDITION
- async def run_stream(
+ async def select_branch(
self, step_input: StepInput, context: RunContext
- ) -> AsyncIterator[Any]:
+ ) -> tuple[list[BaseNode], str, str]:
if callable(self.evaluator):
condition_result = await DependencyInjector.invoke(
self.evaluator, {"step_input": step_input}, context
@@ -238,32 +270,22 @@ class Condition(BaseNode):
condition_result = bool(self.evaluator)
target_steps = self.steps if condition_result else self.else_steps
- branch_name = "if" if condition_result else "else"
+ branch_name = "if_branch" if condition_result else "else_branch"
+ fallback_msg = f"条件求值为 {condition_result},无对应步骤需执行。"
- if not target_steps:
- yield StepOutput(
- content=f"条件求值为 {condition_result},无对应步骤需执行。",
- success=True,
- )
- return
-
- steps_container = Steps(
- steps=target_steps, name=f"{self.name}_{branch_name}_branch"
- )
- out_box: list[StepOutput] = []
- async for event in self._forward_stream(
- steps_container.aexecute_stream(step_input, context), out_box
- ):
- yield event
- if out_box:
- yield out_box[0]
+ return target_steps, branch_name, fallback_msg
-class Router(BaseNode):
+class Router(BranchingNode):
"""根据选择器函数的返回值(名称),从候选项中挑选步骤执行"""
def __init__(
- self, choices: Sequence[NodeSource], selector: Any, name: str = "RouterGroup"
+ self,
+ choices: Sequence[NodeSource],
+ selector: str
+ | list[str]
+ | Callable[..., str | list[str] | Awaitable[str | list[str]]],
+ name: str = "RouterGroup",
):
"""
初始化选择路由器节点。
@@ -285,9 +307,9 @@ class Router(BaseNode):
def node_type(self) -> StepType:
return StepType.ROUTER
- async def run_stream(
+ async def select_branch(
self, step_input: StepInput, context: RunContext
- ) -> AsyncIterator[Any]:
+ ) -> tuple[list[BaseNode], str, str]:
if callable(self.selector):
selected = await DependencyInjector.invoke(
self.selector, {"step_input": step_input}, context
@@ -308,18 +330,7 @@ class Router(BaseNode):
else:
target_steps.append(NodeFactory.build(s))
- if not target_steps:
- yield StepOutput(content="没有命中任何有效路由分支。", success=True)
- return
-
- steps_container = Steps(steps=target_steps, name=f"{self.name}_routed_steps")
- out_box: list[StepOutput] = []
- async for event in self._forward_stream(
- steps_container.aexecute_stream(step_input, context), out_box
- ):
- yield event
- if out_box:
- yield out_box[0]
+ return target_steps, "routed_steps", "没有命中任何有效路由分支。"
class Loop(BaseNode):
@@ -329,7 +340,7 @@ class Loop(BaseNode):
self,
steps: Sequence[NodeSource],
max_iterations: int = 3,
- end_condition: Any = None,
+ end_condition: bool | Callable[..., bool | Awaitable[bool]] | None = None,
name: str = "LoopGroup",
):
"""
@@ -353,7 +364,7 @@ class Loop(BaseNode):
async def run_stream(
self, step_input: StepInput, context: RunContext
- ) -> AsyncIterator[Any]:
+ ) -> AsyncIterator[AgentStreamEvent | StepOutput]:
logger.debug(
f" 🔁 开始循环: [Loop] `{self.name}` (最大 {self.max_iterations} 次)"
)
@@ -362,6 +373,7 @@ class Loop(BaseNode):
all_results: list[StepOutput] = []
current_input = StepInput(
input=step_input.input,
+ intent=step_input.intent,
previous_step_content=step_input.previous_step_content,
additional_data=step_input.additional_data.copy(),
)
@@ -372,12 +384,12 @@ class Loop(BaseNode):
steps_container = Steps(
steps=self.steps, name=f"{self.name}_iter_{iteration + 1}"
)
- out_box: list[StepOutput] = []
- async for event in self._forward_stream(
- steps_container.aexecute_stream(current_input, context), out_box
- ):
+ capturer = StreamCapturer(
+ steps_container.aexecute_stream(current_input, context)
+ )
+ async for event in capturer:
yield event
- iter_output = out_box[0] if out_box else None
+ iter_output = capturer.output
should_stop = False
if iter_output:
@@ -434,13 +446,13 @@ class Parallel(BaseNode):
async def run_stream(
self, step_input: StepInput, context: RunContext
- ) -> AsyncIterator[Any]:
+ ) -> AsyncIterator[AgentStreamEvent | StepOutput]:
logger.debug(f" 🔀 [并发] `{self.name}` 开启了 {len(self.steps)} 个并发任务")
queue = asyncio.Queue()
bg_tasks = []
- async def worker(idx: int, s_obj: Any, c_ctx: RunContext):
+ async def worker(idx: int, s_obj: BaseNode, c_ctx: RunContext):
try:
async for evt in s_obj.aexecute_stream(step_input, c_ctx):
await queue.put(("event", evt))
@@ -524,7 +536,7 @@ class NodeFactory:
cls,
executor: NodeSource,
name: str | None = None,
- failure_policy: Any = None,
+ failure_policy: BaseFailurePolicy | None = None,
) -> BaseNode:
"""底层物理实例化分发"""
kwargs = {
diff --git a/zhenxun/services/ai/flow/workflow/policies.py b/zhenxun/services/ai/flow/workflow/policies.py
index cf0c9864..cdb70677 100644
--- a/zhenxun/services/ai/flow/workflow/policies.py
+++ b/zhenxun/services/ai/flow/workflow/policies.py
@@ -1,10 +1,14 @@
import copy
from enum import Enum
-from typing import Any
+from typing import TYPE_CHECKING
+
+if TYPE_CHECKING:
+ from .base import BaseNode
from pydantic import BaseModel, Field
from zhenxun.services.ai.llm.api import generate_structured
+from zhenxun.services.ai.run.context import RunContext
from zhenxun.services.ai.utils.logger import log_flow as logger
from .types import StepInput
@@ -26,7 +30,7 @@ class PolicyResult(BaseModel):
"""执行延迟或重试前需要等待的缓冲秒数"""
new_input: StepInput | None = None
"""用于动态纠错自愈时替换传入的新参数结构"""
- fallback_node: Any | None = None
+ fallback_node: "BaseNode | None" = None
"""策略裁定降级时所指定的备用工作流节点"""
healer_agent_name: str | None = None
"""执行了高级自愈的大模型或修复者名称"""
@@ -36,7 +40,11 @@ class BaseFailurePolicy:
"""错误处理策略抽象基类"""
async def handle_failure(
- self, node: Any, exception: BaseException, step_input: StepInput, context: Any
+ self,
+ node: "BaseNode",
+ exception: BaseException,
+ step_input: StepInput,
+ context: RunContext,
) -> PolicyResult:
"""
处理节点执行失败的策略入口方法。
@@ -57,7 +65,11 @@ class AbortPolicy(BaseFailurePolicy):
"""直接中断策略"""
async def handle_failure(
- self, node: Any, exception: BaseException, step_input: StepInput, context: Any
+ self,
+ node: "BaseNode",
+ exception: BaseException,
+ step_input: StepInput,
+ context: RunContext,
) -> PolicyResult:
return PolicyResult(action=PolicyAction.ABORT)
@@ -66,7 +78,11 @@ class SkipPolicy(BaseFailurePolicy):
"""跳过并继续策略"""
async def handle_failure(
- self, node: Any, exception: BaseException, step_input: StepInput, context: Any
+ self,
+ node: "BaseNode",
+ exception: BaseException,
+ step_input: StepInput,
+ context: RunContext,
) -> PolicyResult:
return PolicyResult(action=PolicyAction.CONTINUE)
@@ -86,7 +102,11 @@ class RetryPolicy(BaseFailurePolicy):
self.delay = delay
async def handle_failure(
- self, node: Any, exception: BaseException, step_input: StepInput, context: Any
+ self,
+ node: "BaseNode",
+ exception: BaseException,
+ step_input: StepInput,
+ context: RunContext,
) -> PolicyResult:
counts = context.state.setdefault("__retry_counts__", {})
key = f"{node.name}_{id(self)}"
@@ -100,7 +120,7 @@ class RetryPolicy(BaseFailurePolicy):
class FallbackPolicy(BaseFailurePolicy):
"""降级路由策略"""
- def __init__(self, fallback_node: Any):
+ def __init__(self, fallback_node: "BaseNode"):
"""
初始化降级路由策略。
@@ -110,7 +130,11 @@ class FallbackPolicy(BaseFailurePolicy):
self.fallback_node = fallback_node
async def handle_failure(
- self, node: Any, exception: BaseException, step_input: StepInput, context: Any
+ self,
+ node: "BaseNode",
+ exception: BaseException,
+ step_input: StepInput,
+ context: RunContext,
) -> PolicyResult:
return PolicyResult(
action=PolicyAction.FALLBACK, fallback_node=self.fallback_node
@@ -132,7 +156,11 @@ class SelfHealingPolicy(BaseFailurePolicy):
self.max_retries = max_retries
async def handle_failure(
- self, node: Any, exception: BaseException, step_input: StepInput, context: Any
+ self,
+ node: "BaseNode",
+ exception: BaseException,
+ step_input: StepInput,
+ context: RunContext,
) -> PolicyResult:
counts = context.state.setdefault("__heal_counts__", {})
key = f"{node.name}_{id(self)}"
diff --git a/zhenxun/services/ai/flow/workflow/types.py b/zhenxun/services/ai/flow/workflow/types.py
index e891e445..f061aac6 100644
--- a/zhenxun/services/ai/flow/workflow/types.py
+++ b/zhenxun/services/ai/flow/workflow/types.py
@@ -3,6 +3,8 @@ from typing import Any
from pydantic import BaseModel, Field
+from zhenxun.services.ai.run.models import RunIntent
+
class StepType(str, Enum):
FUNCTION = "Function"
@@ -20,6 +22,9 @@ class StepInput(BaseModel):
input: Any = Field(default=None)
"""继承自 Workflow 的初始输入"""
+ intent: RunIntent | None = Field(default=None)
+ """归一化后的意图载体,避免下游节点猜测解析 input 的原始类型"""
+
previous_step_content: Any = Field(default=None)
"""上一个执行步骤产生的直接输出内容"""
diff --git a/zhenxun/services/ai/llm/api.py b/zhenxun/services/ai/llm/api.py
index 0f5fe3f5..e89bde9f 100644
--- a/zhenxun/services/ai/llm/api.py
+++ b/zhenxun/services/ai/llm/api.py
@@ -2,11 +2,13 @@
LLM 服务的高级 API 接口 - 便捷函数入口 (无状态)
"""
+import json
from pathlib import Path
from typing import Any, Literal, TypeVar, overload
from pydantic import BaseModel
+from zhenxun.services.ai.config import get_llm_config
from zhenxun.services.ai.core.exceptions import (
ControlFlowExit,
LLMException,
@@ -27,14 +29,18 @@ from zhenxun.services.ai.core.messages import (
RerankRequest,
RerankResult,
SpeechRequest,
+ UsageInfo,
)
from zhenxun.services.ai.core.models import ModelName
from zhenxun.services.ai.core.options import (
GenerationConfig,
LLMEmbeddingConfig,
+ OutputFormatConfig,
+ ResponseFormat,
+ StructuredOutputStrategy,
TTSConfig,
)
-from zhenxun.services.ai.guardrails import GuardrailSource
+from zhenxun.services.ai.guardrails import GuardrailSource, parse_guardrails
from zhenxun.services.ai.utils.logger import log_llm as logger
from .builder import IntentBuilder
@@ -109,7 +115,7 @@ async def embed(
@overload
async def embed(
- input_batch: list[Any],
+ input_batch: list[PromptInput],
*,
model: ModelName = None,
task: Literal[
@@ -122,7 +128,7 @@ async def embed(
async def embed(
- input_batch: PromptInput | list[Any],
+ input_batch: PromptInput | list[PromptInput],
*,
model: ModelName = None,
task: Literal[
@@ -158,8 +164,6 @@ async def embed(
)
if not batch.payloads:
- from zhenxun.services.ai.core.messages import UsageInfo
-
return EmbeddingResponse(
embeddings=[], usage=UsageInfo(), model_name=str(model)
)
@@ -256,21 +260,13 @@ async def generate_structured(
T: 解析验证通过后的 Pydantic 模型实例。
""" # noqa: E501
try:
- from zhenxun.services.ai.config import get_llm_config
from zhenxun.services.ai.core.engine.structured_parser import (
BaseOutputProcessor,
)
- from zhenxun.services.ai.core.options import (
- OutputFormatConfig,
- ResponseFormat,
- StructuredOutputStrategy,
- )
if max_retries is None:
max_retries = get_llm_config().client_settings.structured_retries
- from zhenxun.services.ai.guardrails import parse_guardrails
-
parsed_guardrails = parse_guardrails(guardrails)
output_processor = BaseOutputProcessor(
@@ -291,8 +287,6 @@ async def generate_structured(
if instruction:
prompt_parts.append(instruction)
- import json
-
schema_str = json.dumps(json_schema, ensure_ascii=False, indent=2)
prompt_parts.append(
"### ⚠️ [结构化输出要求]\n"
@@ -393,7 +387,9 @@ async def generate(
llm_context = LLMContext(request=request)
combined_cap = CombinedCapability(sys_caps)
- async def inner_handler(ctx: LLMContext[Any, Any]) -> ChatResponse:
+ async def inner_handler(
+ ctx: LLMContext[ChatRequest, ChatResponse],
+ ) -> ChatResponse:
return await LLMOrchestrator.invoke(
ctx.request,
model_name=model,
@@ -420,7 +416,7 @@ async def generate(
@overload
async def create_image(
- prompt: str | Any,
+ prompt: PromptInput,
*,
images: None = None,
model: ModelName = None,
@@ -432,7 +428,7 @@ async def create_image(
@overload
async def create_image(
- prompt: str | Any,
+ prompt: PromptInput,
*,
images: list[Path | bytes | str] | Path | bytes | str,
model: ModelName = None,
@@ -443,7 +439,7 @@ async def create_image(
async def create_image(
- prompt: str | Any,
+ prompt: PromptInput,
*,
images: list[Path | bytes | str] | Path | bytes | str | None = None,
model: ModelName = None,
diff --git a/zhenxun/services/ai/llm/builder.py b/zhenxun/services/ai/llm/builder.py
index 4593d7c6..79fabd9e 100644
--- a/zhenxun/services/ai/llm/builder.py
+++ b/zhenxun/services/ai/llm/builder.py
@@ -2,13 +2,18 @@
LLM 生成配置相关类和函数
"""
-from typing import Any, Literal
+import inspect
+from typing import Any, Literal, cast
from typing_extensions import Self
+from pydantic import BaseModel
+
+from zhenxun.services.ai.config import get_gemini_safety_threshold
from zhenxun.services.ai.core.exceptions import ConfigurationException
from zhenxun.services.ai.core.options import (
GenerationConfig,
ResponseFormat,
+ StructuredOutputStrategy,
)
from zhenxun.services.ai.utils.logger import log_llm as logger
from zhenxun.utils.pydantic_compat import model_json_schema, model_validate
@@ -94,20 +99,12 @@ class IntentBuilder:
强制要求结构化输出意图。
支持自动处理 Pydantic 模型并转换为厂商所需的 JSON Schema。
"""
- import inspect
-
- from pydantic import BaseModel
-
- from zhenxun.services.ai.core.options import StructuredOutputStrategy
-
self._config.output.response_format = ResponseFormat.JSON
self._config.output.response_mime_type = "application/json"
if schema:
if inspect.isclass(schema) and issubclass(schema, BaseModel):
self._config.output.response_schema = model_json_schema(schema)
else:
- from typing import cast
-
self._config.output.response_schema = cast(dict[str, Any], schema)
if strict:
self._config.output.structured_output_strategy = (
@@ -137,8 +134,6 @@ class IntentBuilder:
安全合规意图。
level 取值: 'strict' (最严格), 'moderate' (中等), 'none' (完全无限制)。
"""
- from zhenxun.services.ai.config import get_gemini_safety_threshold
-
if level == "strict":
self.gemini.set_safety_threshold("BLOCK_LOW_AND_ABOVE")
elif level == "none":
diff --git a/zhenxun/services/ai/llm/manager.py b/zhenxun/services/ai/llm/manager.py
index 4b5d8287..0d8199db 100644
--- a/zhenxun/services/ai/llm/manager.py
+++ b/zhenxun/services/ai/llm/manager.py
@@ -17,6 +17,7 @@ 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 .system.cache import clear_model_cache, get_or_create_model
from .system.capabilities import get_model_capabilities
from .system.network import health_manager
@@ -277,8 +278,6 @@ async def get_model_instance(
provider_config_found, model_detail_found = config_tuple_found
- from .system.cache import get_or_create_model
-
return await get_or_create_model(
provider_config_found, model_detail_found, override_config
)
@@ -288,8 +287,6 @@ def clear_all_cache() -> None:
"""
清空模型实例缓存与路由组解析缓存。
"""
- from .system.cache import clear_model_cache
-
clear_model_cache()
clear_resolved_group_cache()
logger.debug("已清空全局模型实例与路由组缓存")
@@ -300,11 +297,8 @@ async def _init_llm_config_on_startup():
"""启动时初始化 LLM 配置、密钥状态并预热工具提供者管理器。"""
logger.info("正在初始化 LLM 配置并加载遥测状态...")
try:
- from zhenxun.services.ai.config import get_llm_config
from zhenxun.services.ai.tools.engine.registry import tool_provider_manager
- from .system.network import health_manager
-
get_llm_config()
await health_manager.initialize()
await tool_provider_manager.initialize()
diff --git a/zhenxun/services/ai/message_builder.py b/zhenxun/services/ai/message_builder.py
index 5940b7c3..525cc5cb 100644
--- a/zhenxun/services/ai/message_builder.py
+++ b/zhenxun/services/ai/message_builder.py
@@ -187,7 +187,7 @@ class MessageBuilder:
allowed_modalities: set[str] | None = None,
) -> list[UserContentUnion]:
"""将 UniMessage 消息解析并转换为 LLM 内容部件列表"""
- namespace = namespace or infer_plugin_namespace(default="global")
+ namespace = namespace or infer_plugin_namespace()
parts: list[UserContentUnion] = []
for seg in message:
if allowed_modalities is not None:
@@ -237,7 +237,7 @@ class MessageBuilder:
allowed_modalities: set[str] | None = None,
) -> list[LLMContentPart] | None:
"""获取并解析引用消息的内容片段"""
- namespace = namespace or infer_plugin_namespace(default="global")
+ namespace = namespace or infer_plugin_namespace()
try:
orig_msg = await reply_fetch(event, bot)
if not orig_msg or not orig_msg.msg:
@@ -277,7 +277,7 @@ class MessageBuilder:
allowed_modalities: set[str] | None = None,
) -> list[LLMMessage]:
"""将任意类型的提示输入标准化为统一的 LLM 消息历史列表"""
- namespace = namespace or infer_plugin_namespace(default="global")
+ namespace = namespace or infer_plugin_namespace()
messages = []
if instruction:
messages.append(SystemMessage(content=[TextPart(text=instruction)]))
@@ -381,7 +381,7 @@ class MessageBuilder:
config: LLMEmbeddingConfig | None = None,
) -> list[LLMContentPart]:
"""为 Embed 向量化提取纯粹的内容片段,忽略杂项"""
- namespace = namespace or infer_plugin_namespace(default="global")
+ namespace = namespace or infer_plugin_namespace()
allowed_modalities = {"text"}
if config:
if config.multimodal is True:
@@ -418,7 +418,7 @@ class MessageBuilder:
config: LLMEmbeddingConfig | None = None,
) -> "EmbedBatch":
"""将任意输入标准化为嵌入向量批处理对象"""
- namespace = namespace or infer_plugin_namespace(default="global")
+ namespace = namespace or infer_plugin_namespace()
if isinstance(inputs, list) and not isinstance(inputs, UniMessage):
if not inputs:
diff --git a/zhenxun/services/ai/run/__init__.py b/zhenxun/services/ai/run/__init__.py
index f76d1ef7..462d8693 100644
--- a/zhenxun/services/ai/run/__init__.py
+++ b/zhenxun/services/ai/run/__init__.py
@@ -12,6 +12,7 @@ from .hooks import Hooks
from .models import (
AgentRunResult,
AgentTask,
+ RunIntent,
StreamedRunResult,
)
from .session import session_manager
@@ -28,6 +29,7 @@ __all__ = [
"Inject",
"NoneBotDeps",
"RunContext",
+ "RunIntent",
"StreamedRunResult",
"UIController",
"get_current_run_context",
diff --git a/zhenxun/services/ai/run/context.py b/zhenxun/services/ai/run/context.py
index c46ea584..1083c93c 100644
--- a/zhenxun/services/ai/run/context.py
+++ b/zhenxun/services/ai/run/context.py
@@ -8,6 +8,9 @@ from typing import TYPE_CHECKING, Any, Generic, cast, get_origin
from typing_extensions import TypeVar
import uuid
+if TYPE_CHECKING:
+ from zhenxun.services.ai.context.memory.types import SessionMetadata
+
from nonebot.adapters import Bot, Event
from nonebot.matcher import Matcher, current_bot, current_event, current_matcher
from pydantic import BaseModel, ConfigDict, Field
@@ -19,16 +22,12 @@ from zhenxun.services.ai.core.models import CancellationToken
from zhenxun.services.ai.core.protocols.tool import ToolExecutable
from zhenxun.services.ai.core.stream_events import AgentStreamEvent, EventBus
from zhenxun.services.ai.utils import ContextUtils
-from zhenxun.services.ai.utils.scope import ScopeSelector
from zhenxun.services.scheduler.types import ScheduleContext
from zhenxun.utils.platform import PlatformUtils
from zhenxun.utils.utils import infer_plugin_namespace
from .blackboard import BlackboardManager
-if TYPE_CHECKING:
- from zhenxun.services.ai.context.memory.facades import AgentSessionFacade
-
AgentDepsT = TypeVar("AgentDepsT", default=Any)
"""泛型类型变量:外部环境依赖对象 (Agent Dependencies)。"""
ProviderFunc = Callable[["RunContext"], Any | Awaitable[Any]]
@@ -114,31 +113,8 @@ class SessionContext(Generic[AgentDepsT]):
"""触发事件的插件命名空间"""
append_only_manager: Any = dataclasses.field(default=None)
"""用于大模型前缀缓存命中优化的追加写入管理器。"""
-
- @property
- def memory(self) -> "AgentSessionFacade":
- """
- 获取当前会话的持久化记忆访问门面 (AgentSessionFacade)。
- 提供极简的 history 和 slots 操作 API。
- """
- from zhenxun.services.ai.context.memory.facades import AgentSessionFacade
- from zhenxun.services.ai.context.memory.manager import memory_manager
- from zhenxun.services.ai.context.memory.types import SessionMetadata
-
- user_id = ContextUtils.extract_user_id(self.deps)
- group_id = ContextUtils.extract_group_id(self.deps)
- platform = ContextUtils.extract_platform(self.deps)
-
- meta = SessionMetadata(
- session_id=self.session_id,
- selector=ScopeSelector(
- user_id=user_id,
- group_id=group_id,
- platform=platform,
- namespace=self.namespace,
- ),
- )
- return AgentSessionFacade(memory_manager, meta)
+ session_meta: "SessionMetadata | None" = dataclasses.field(default=None)
+ """隔离会话的元信息(Session ID, 命名空间, 权限等),在状态流转中生成"""
@dataclasses.dataclass
@@ -324,7 +300,7 @@ class RunContext(Generic[AgentDepsT]):
self.deps.get("namespace") if isinstance(self.deps, dict) else None
)
if not ns:
- ns = infer_plugin_namespace(default="global")
+ ns = infer_plugin_namespace()
self.session = SessionContext(
session_id=self.session_id or "default_session",
diff --git a/zhenxun/services/ai/run/di.py b/zhenxun/services/ai/run/di.py
index 0e7965a6..c8552ff3 100644
--- a/zhenxun/services/ai/run/di.py
+++ b/zhenxun/services/ai/run/di.py
@@ -1,14 +1,27 @@
from collections import defaultdict
from collections.abc import Awaitable, Callable
+from functools import lru_cache
import inspect
from typing import Annotated, Any, ClassVar, cast
+
+@lru_cache(maxsize=2048)
+def _get_signature_cached(func: Callable) -> inspect.Signature:
+ return inspect.signature(func)
+
+
+def get_signature(func: Callable) -> inspect.Signature:
+ try:
+ return _get_signature_cached(func)
+ except TypeError:
+ return inspect.signature(func)
+
+
from nonebot.adapters import Bot, Event
from nonebot.matcher import Matcher
from nonebot.utils import is_coroutine_callable
from nonebot_plugin_session import EventSession, extract_session
-from zhenxun.services.ai.context.memory.facades import AgentSessionFacade
from zhenxun.utils.utils import infer_plugin_namespace
from .blackboard import BlackboardManager
@@ -59,7 +72,7 @@ CurrentEventPayload = Annotated[Any, Hidden(), _InjectMarker("stream_event")]
CurrentBlackboard = Annotated[
BlackboardManager | None, Hidden(), _InjectMarker("blackboard")
]
-CurrentMemory = Annotated[AgentSessionFacade, Hidden(), _InjectMarker("memory")]
+
CurrentSandbox = Annotated[Any, Hidden(), _InjectMarker("sandbox")]
@@ -160,9 +173,6 @@ class Inject:
Blackboard = CurrentBlackboard
"""自动注入:当前工作流/团队挂载的强类型黑板 (BlackboardManager) 实例"""
- Memory = CurrentMemory
- """自动注入:当前会话的持久化记忆存取门面 (AgentSessionFacade) 实例"""
-
Sandbox = CurrentSandbox
"""自动注入:当前沙箱环境管理器实例"""
@@ -314,7 +324,7 @@ class DependencyInjector:
context: RunContext,
) -> Any:
"""统一执行带有依赖注入的函数 (支持同步/异步)"""
- sig = inspect.signature(func)
+ sig = get_signature(func)
resolved_kwargs = await cls.resolve_all(sig, call_kwargs, context)
filtered_kwargs = {
k: v for k, v in resolved_kwargs.items() if k in sig.parameters
@@ -354,7 +364,6 @@ Inject.register_provider(
Inject.register_provider(
"shared_state", lambda ctx: ctx.session.shared_state, scope="global"
)
-Inject.register_provider("memory", lambda ctx: ctx.session.memory, scope="global")
def _resolve_blackboard(ctx: RunContext):
diff --git a/zhenxun/services/ai/run/models.py b/zhenxun/services/ai/run/models.py
index 6e86afa2..70e834b4 100644
--- a/zhenxun/services/ai/run/models.py
+++ b/zhenxun/services/ai/run/models.py
@@ -4,16 +4,15 @@ from __future__ import annotations
运行时(Run)相关核心类型定义
"""
-from collections.abc import AsyncIterator, Callable
+from collections.abc import AsyncIterator
import json
from typing import Any, Generic, cast
-from typing_extensions import TypeVar
from pydantic import BaseModel, ConfigDict, Field, PrivateAttr
from zhenxun.services.ai.core.messages import AgentMessage, LLMMessage, UsageInfo
+from zhenxun.services.ai.core.messages.types import OutputDataT
from zhenxun.services.ai.core.options import BaseOutputDefinition
-from zhenxun.services.ai.core.protocols.tool import ToolResolvable
from zhenxun.services.ai.core.stream_events import (
AgentStreamEvent,
EventBus,
@@ -21,7 +20,6 @@ from zhenxun.services.ai.core.stream_events import (
ToolStreamChunkEvent,
)
from zhenxun.services.ai.guardrails import BaseGuardrail, GuardrailSource
-from zhenxun.services.ai.tools.core.tool import BaseTool
from zhenxun.utils.pydantic_compat import model_dump, model_validator
@@ -71,9 +69,6 @@ class HandoffPayload(BaseModel):
"""随移交传递的上下文数据"""
-OutputDataT = TypeVar("OutputDataT", default=str)
-
-
class AgentRunResult(BaseModel, Generic[OutputDataT]):
"""Agent 单次无状态运行的结果"""
@@ -238,9 +233,7 @@ class AgentTask(BaseModel):
"""强制要求返回的强类型结构 (Pydantic Model) 或
OutputDefinition,为空则返回普通文本"""
- tools: list[str | Callable | dict[str, Any] | BaseTool | ToolResolvable] | None = (
- None
- )
+ tools: list[Any] | None = None
"""针对此特定任务动态追加或覆盖的工具列表"""
guardrails: list[GuardrailSource] | None = None
@@ -259,9 +252,74 @@ class AgentTask(BaseModel):
return self
+class RunIntent(BaseModel):
+ """标准化且归一化的运行时意图载体"""
+
+ text: str = ""
+ """提取出的纯文本指令(用于路由、日志和并发控制判断)"""
+ original_input: Any = None
+ """用户最原始的输入对象(如 UniMessage 等,用于多模态图像/音频数据提取)"""
+ payload_to_render: Any = None
+ """将要被压入 MessageBuilder 渲染为大模型 Prompt 的实际载体"""
+ task_obj: AgentTask | None = None
+ """如果是强类型任务契约,存储其原始对象引用"""
+ response_model: type[BaseModel] | BaseOutputDefinition | None = None
+ """提取出的强类型输出约束"""
+ extra_tools: list[Any] = Field(default_factory=list)
+ """提取出的附加工具集"""
+ guardrails: list[BaseGuardrail] = Field(default_factory=list)
+ """提取出的安全护栏集"""
+
+ model_config = ConfigDict(arbitrary_types_allowed=True)
+
+ @classmethod
+ def from_input(cls, prompt: Any) -> "RunIntent":
+ task_obj = None
+ text_content = ""
+ extra_tools = []
+ response_model = None
+ guardrails = []
+ payload_to_render = prompt
+
+ if isinstance(prompt, AgentTask):
+ task_obj = prompt
+ response_model = task_obj.response_model
+ if task_obj.tools:
+ extra_tools.extend(task_obj.tools)
+ if hasattr(task_obj, "_parsed_guardrails"):
+ guardrails.extend(task_obj._parsed_guardrails)
+
+ prompt_parts = [
+ f"### 📋 [任务指令]\n{task_obj.description}",
+ f"### 🎯 [预期产出要求]\n{task_obj.expected_output}",
+ ]
+ text_content = "\n\n".join(prompt_parts)
+ payload_to_render = text_content
+ elif prompt is not None:
+ text_content = getattr(prompt, "description", None) or (
+ getattr(prompt, "extract_plain_text", lambda: str(prompt))()
+ if prompt
+ else str(prompt)
+ )
+ else:
+ text_content = ""
+ payload_to_render = None
+
+ return cls(
+ text=text_content,
+ original_input=prompt,
+ payload_to_render=payload_to_render,
+ task_obj=task_obj,
+ response_model=response_model,
+ extra_tools=extra_tools,
+ guardrails=guardrails,
+ )
+
+
__all__ = [
"AgentRunResult",
"AgentTask",
"OutputDataT",
+ "RunIntent",
"StreamedRunResult",
]
diff --git a/zhenxun/services/ai/tools/bridges/delegate.py b/zhenxun/services/ai/tools/bridges/delegate.py
index 9e66208d..666c133c 100644
--- a/zhenxun/services/ai/tools/bridges/delegate.py
+++ b/zhenxun/services/ai/tools/bridges/delegate.py
@@ -7,7 +7,7 @@ from pydantic import BaseModel, Field
from zhenxun.services.ai.core.exceptions import AbortException, ControlFlowExit
from zhenxun.services.ai.core.stream_events import ToolStreamChunkEvent
-from zhenxun.services.ai.flow.base import BaseRunnable
+from zhenxun.services.ai.flow.core.base import BaseRunnable
from zhenxun.services.ai.run import RunContext
from zhenxun.services.ai.tools.core.tool import BaseTool
from zhenxun.services.ai.tools.models import ToolResult
@@ -49,7 +49,9 @@ class DelegateTool(BaseTool):
"""
resolved_name = name or getattr(runnable, "name", "SubRunnable")
resolved_desc = description or getattr(
- runnable, "description", f"将子任务委派给 {resolved_name} 执行"
+ runnable,
+ "profile_summary",
+ getattr(runnable, "description", f"将子任务委派给 {resolved_name} 执行"),
)
final_name = (
f"delegate_to_{resolved_name}"
diff --git a/zhenxun/services/ai/tools/bridges/matcher_bridge.py b/zhenxun/services/ai/tools/bridges/matcher_bridge.py
index 9b6df31c..fad6d111 100644
--- a/zhenxun/services/ai/tools/bridges/matcher_bridge.py
+++ b/zhenxun/services/ai/tools/bridges/matcher_bridge.py
@@ -450,7 +450,7 @@ def bind_matcher(
"""
if require_prefix:
- ns = infer_plugin_namespace(default="global")
+ ns = infer_plugin_namespace()
if ns and ns not in ("global", "unknown"):
if not name.startswith(f"{ns}_"):
name = f"{ns}_{name}"
diff --git a/zhenxun/services/ai/tools/core/decorators.py b/zhenxun/services/ai/tools/core/decorators.py
index 67ca2ade..dab15997 100644
--- a/zhenxun/services/ai/tools/core/decorators.py
+++ b/zhenxun/services/ai/tools/core/decorators.py
@@ -218,7 +218,7 @@ def tool(
if require_prefix:
from zhenxun.utils.utils import infer_plugin_namespace
- ns = infer_plugin_namespace(default="global")
+ ns = infer_plugin_namespace()
if ns and ns not in ("global", "unknown"):
if not tool_name.startswith(f"{ns}_"):
tool_name = f"{ns}_{tool_name}"
diff --git a/zhenxun/services/ai/tools/core/toolkit.py b/zhenxun/services/ai/tools/core/toolkit.py
index c4ae69c3..675f779a 100644
--- a/zhenxun/services/ai/tools/core/toolkit.py
+++ b/zhenxun/services/ai/tools/core/toolkit.py
@@ -217,6 +217,14 @@ class BaseToolkit:
)
return f"<{tag_name}>\n{text}\n{tag_name}>"
+ def clone_with(self, **kwargs: Any) -> "BaseToolkit":
+ """克隆当前工具箱原型,并透明注入新的运行时属性。"""
+ new_tk = copy.copy(self)
+ for k, v in kwargs.items():
+ setattr(new_tk, k, v)
+ new_tk._cached_tools = None
+ return new_tk
+
def prefixed(self, prefix: str) -> "BaseToolkit":
"""克隆工具箱并为其中所有工具追加统一的前缀"""
new_tk = copy.copy(self)
diff --git a/zhenxun/services/ai/tools/engine/registry.py b/zhenxun/services/ai/tools/engine/registry.py
index e7103a8e..8f178306 100644
--- a/zhenxun/services/ai/tools/engine/registry.py
+++ b/zhenxun/services/ai/tools/engine/registry.py
@@ -365,7 +365,7 @@ class ToolProviderManager:
def register_tool(self, tool: ToolExecutable):
"""注册由 @tool 生成的单一工具"""
- ns = infer_plugin_namespace(default="global")
+ ns = infer_plugin_namespace()
self.local_provider.register_tool(tool, ns)
tags = getattr(getattr(tool, "settings", None), "tags", [])
tag_str = f" | Tags: {tags}" if tags else ""
@@ -376,7 +376,7 @@ class ToolProviderManager:
"""
注册一个完整的 Toolkit 实例,使其可通过智能字符串路由(Tag或Name)被动态发现。
"""
- ns = infer_plugin_namespace(default="global")
+ ns = infer_plugin_namespace()
self.local_provider.register_toolkit(toolkit, ns)
tk_name = getattr(toolkit, "__class__", type).__name__
diff --git a/zhenxun/services/ai/tools/models.py b/zhenxun/services/ai/tools/models.py
index 37b576d8..49c45c25 100644
--- a/zhenxun/services/ai/tools/models.py
+++ b/zhenxun/services/ai/tools/models.py
@@ -4,7 +4,7 @@
from dataclasses import dataclass, field
import fnmatch
-from typing import TYPE_CHECKING, Any, Literal
+from typing import Any, Literal
from pydantic import BaseModel, ConfigDict, Field
@@ -13,9 +13,6 @@ from zhenxun.services.ai.run.context import RunContext
from zhenxun.services.ai.utils.logger import log_tool as logger
from zhenxun.utils.pydantic_compat import model_dump, model_validate
-if TYPE_CHECKING:
- from zhenxun.services.ai.tools.core.tool import BaseTool
-
class DirectivePayload(BaseModel):
"""工具执行产生的副作用控制流指令载荷"""
@@ -272,7 +269,7 @@ class Query(BaseModel):
metadata_filter: dict[str, Any] | None = Field(default=None)
"""如果提供,则工具的 metadata 必须包含这里列出的所有键值对。"""
- def match(self, tool: "BaseTool") -> bool:
+ def match(self, tool: Any) -> bool:
"""判断某个工具或工具箱是否符合当前 Query 的筛选条件"""
def _match_pattern(val: str, pattern: str | list[str]) -> bool:
diff --git a/zhenxun/services/ai/tools/providers/builtin/memory.py b/zhenxun/services/ai/tools/providers/builtin/memory.py
index c3118e65..d78b4283 100644
--- a/zhenxun/services/ai/tools/providers/builtin/memory.py
+++ b/zhenxun/services/ai/tools/providers/builtin/memory.py
@@ -2,15 +2,17 @@ from typing import Any, Literal, Optional
from pydantic import Field, create_model
-from zhenxun.services.ai.context.memory.manager import memory_manager
-from zhenxun.services.ai.context.memory.models import MemoryConfig
+from zhenxun.services.ai.context.memory.storage.backends import MemoryScope
from zhenxun.services.ai.context.memory.types import SessionMetadata
+from zhenxun.services.ai.context.rag.engine import ScopedRAGClient
from zhenxun.services.ai.run.context import RunContext
from zhenxun.services.ai.tools.core.decorators import tool
from zhenxun.services.ai.tools.core.tool import BaseTool
from zhenxun.services.ai.tools.core.toolkit import BaseToolkit
from zhenxun.services.ai.tools.models import ToolOptions, ToolResult
from zhenxun.services.ai.utils.logger import log_tool as logger
+from zhenxun.services.ai.utils.runtime import ContextUtils
+from zhenxun.services.ai.utils.scope import ScopeBuilder
class MemoryManagementToolkit(BaseToolkit):
@@ -26,23 +28,45 @@ class MemoryManagementToolkit(BaseToolkit):
shared_options = ToolOptions(silent=True)
- default_instructions = """\
-## 🧠 长期记忆管理系统 (Long-Term Memory)
-该系统是你的「无限档案馆」。⚠️ **注意:系统默认不会主动向你提供所有历史信息,你必须通过主动搜索来回忆。**
+ _INTRO_TEXT = (
+ "## 🧠 长期记忆管理系统 (Long-Term Memory)\n"
+ "该系统是你的「无限档案馆」。系统默认不会主动向你提供所有历史信息,你必须通过主动搜索来回忆。\n\n"
+ "### 📝 职责说明\n"
+ )
+ _READ_GUIDE = (
+ "- **寻找历史线索**:当遇到未知情况,或用户提及过去的事情、"
+ "特定设定时,必须主动检索历史库(使用 `search_memory`)。\n"
+ )
+ _WRITE_GUIDE = (
+ "- **记录离散事实与经验**:当需要记录某个独立事件、历史经验、"
+ "问题解决方案或具体事实时(使用 `save_memory`)。\n"
+ "- **隐式记录**:当接收到值得记忆的重要信息时,请静默记录。"
+ "除非用户主动提问,否则无需向用户显式汇报'我已记住'。\n"
+ "- **按需更新**:如果发现某项历史记录已过时或状态发生扭转,"
+ "请先检索出它的 ID,再修改(`update_memory`)或废弃(`delete_memory`)。\n"
+ "- **精准提炼**:保存记忆时请提炼核心价值,避免保存无意义的闲聊。\n"
+ )
-### 📝 何时使用长期记忆?
-- **记录离散事实与经验**:当需要记录某个独立事件、历史经验、问题解决方案或具体事实时(使用 `save_memory`)。
-- **寻找历史线索**:当遇到未知情况,或用户提及过去的事情、特定设定时,必须主动检索历史库(使用 `search_memory`)。
+ default_instructions = _INTRO_TEXT + _READ_GUIDE + _WRITE_GUIDE
-### ⚙️ 操作规范
-1. **隐式记录**:当接收到值得记忆的重要信息时,请静默记录。除非用户主动提问,否则无需向用户显式汇报"我已记住"。
-2. **按需更新**:如果发现某项历史记录已过时或状态发生扭转,请先检索出它的 ID,再进行修改(使用 `update_memory`)或废弃(使用 `delete_memory`)。
-3. **精准提炼**:保存记忆时请提炼核心价值,避免保存无意义的闲聊。\
-""" # noqa: E501
+ @classmethod
+ def read_only(cls, **kwargs) -> "MemoryManagementToolkit":
+ """[工厂方法] 创建一个只读模式的长期记忆工具箱。"""
+ kwargs["include"] = ["search_memory"]
+ kwargs.setdefault("instructions", cls._INTRO_TEXT + cls._READ_GUIDE)
+ return cls(**kwargs)
+
+ @classmethod
+ def write_only(cls, **kwargs) -> "MemoryManagementToolkit":
+ """[工厂方法] 创建一个仅写入模式的长期记忆工具箱。"""
+ kwargs["exclude"] = ["search_memory"]
+ kwargs.setdefault("instructions", cls._INTRO_TEXT + cls._WRITE_GUIDE)
+ return cls(**kwargs)
def __init__(
self,
- memory_config: MemoryConfig | None = None,
+ rag_client: ScopedRAGClient | None = None,
+ scopes: dict[str, ScopeBuilder] | None = None,
namespace: str | None = None,
**kwargs: Any,
):
@@ -50,85 +74,43 @@ class MemoryManagementToolkit(BaseToolkit):
初始化主动记忆管理工具箱。
参数:
- memory_config: 记忆系统的全局配置对象,为空则使用全局默认。
+ rag_client: 底层 RAG 检索引擎客户端实例。
+ scopes: 作用域构建器映射字典,用于动态限定存储的分区。
namespace: 当前隔离环境的命名空间。
kwargs: 其他透传给 BaseToolkit 的参数。
"""
super().__init__(**kwargs)
- self.memory_config = memory_config
+ self.rag_client = rag_client
+ from zhenxun.services.ai.context.memory.types import Isolation
+
+ self.scopes = scopes or {"私有": Isolation.AGENT_USER()}
self._namespace = namespace
def _get_runtime_meta_and_scope(
self, context: RunContext, scope_name: str | None = None
) -> tuple[Any, SessionMetadata]:
"""动态获取当前运行时的数据库实例与会话元信息,实现无状态化"""
- ns = self._namespace or getattr(context.session, "namespace", "global")
- scope = memory_manager.get_long_term_memory(self.memory_config, namespace=ns)
+ scope = MemoryScope(rag_client=self.rag_client) if self.rag_client else None
- scope_builder = None
- if (
- self.memory_config
- and self.memory_config.long_term
- and self.memory_config.long_term.scopes
- ):
- scopes_dict = self.memory_config.long_term.scopes
- if not scope_name:
- scope_builder = next(iter(scopes_dict.values()))
- else:
- scope_builder = scopes_dict.get(scope_name)
-
- if not scope_builder:
- scope_builder = getattr(self.memory_config, "base_isolation", None)
- if not scope_builder:
- from zhenxun.services.ai.context.memory.types import Isolation
-
- scope_builder = Isolation.AGENT_USER()
-
- selector = scope_builder.resolve(
- deps=context.deps,
- prefix="",
- default_namespace=ns,
- default_agent=context.run.agent_name,
+ scope_builder = (
+ self.scopes.get(scope_name)
+ if scope_name
+ else next(iter(self.scopes.values()), None)
)
- parts = selector.get_scope_parts()
- all_scopes = {"/"}
- current_path = ""
- for part in parts:
- current_path += f"/{part}"
- all_scopes.add(current_path)
- accessible_scopes = list(all_scopes)
- accessible_scopes.sort(key=lambda x: len(x.split("/")))
-
- scope_name_mapping = {}
- if (
- self.memory_config
- and self.memory_config.long_term
- and self.memory_config.long_term.scopes
- ):
- for name, builder in self.memory_config.long_term.scopes.items():
- sel = builder.resolve(
- deps=context.deps,
- prefix="",
- default_namespace=ns,
- default_agent=context.run.agent_name,
- )
- scope_name_mapping[sel.scope_prefix] = name
-
- session_meta = SessionMetadata(
- session_id=context.session_id or "default_session",
- selector=selector,
- scope_prefix=selector.scope_prefix,
- accessible_scopes=accessible_scopes,
- scope_name_mapping=scope_name_mapping,
+ session_meta = ContextUtils.build_session_meta(
+ context=context,
+ target_builder=scope_builder,
+ extra_scopes=self.scopes,
+ custom_namespace=self._namespace,
)
return scope, session_meta
async def get_tools(self, context: RunContext | None = None) -> dict[str, BaseTool]:
tools = await super().get_tools(context)
- if not self.memory_config or not self.memory_config.long_term.enable:
+ if not getattr(self, "rag_client", None):
return tools
- scopes_dict = self.memory_config.long_term.scopes
+ scopes_dict = getattr(self, "scopes", {})
if not scopes_dict:
return tools
diff --git a/zhenxun/services/ai/tools/providers/builtin/slots.py b/zhenxun/services/ai/tools/providers/builtin/slots.py
index daf48646..948eb3ad 100644
--- a/zhenxun/services/ai/tools/providers/builtin/slots.py
+++ b/zhenxun/services/ai/tools/providers/builtin/slots.py
@@ -6,7 +6,6 @@ from typing import Any, Literal
from pydantic import Field, create_model
from zhenxun.services.ai.context.memory.manager import memory_manager
-from zhenxun.services.ai.context.memory.models import MemoryConfig
from zhenxun.services.ai.context.memory.types import (
MemorySlot,
SessionMetadata,
@@ -16,6 +15,7 @@ from zhenxun.services.ai.tools.core.decorators import tool
from zhenxun.services.ai.tools.core.tool import BaseTool
from zhenxun.services.ai.tools.core.toolkit import BaseToolkit
from zhenxun.services.ai.tools.models import ToolOptions, ToolResult
+from zhenxun.services.ai.utils.runtime import ContextUtils
_SLOT_LOCKS: dict[str, asyncio.Lock] = {}
_GLOBAL_LOCK = asyncio.Lock()
@@ -45,119 +45,78 @@ class MemorySlotToolkit(BaseToolkit):
shared_options = ToolOptions(silent=True)
- default_instructions = """\
-## 📋 状态与规则面板 (Memory Slots / 中期记忆)
-该系统是你的「桌面便利贴」或「共享黑板」,用于保存你当前需要随时查阅的核心状态与全局规范。
+ _INTRO_TEXT = (
+ "## 📋 状态与规则面板 (Memory Slots / 中期记忆)\n"
+ "该系统是你的「桌面便利贴」或「共享黑板」,用于保存你当前需要随时查阅的核心状态与全局规范。\n\n"
+ "### 💡 核心机制\n"
+ "- 记忆槽内容会在每次对话时**直接注入提示词中**,你无需搜索即可看见。\n"
+ "- 槽位容量极其有限,仅用于维持最新的运行状态。\n\n"
+ "### 📝 职责说明\n"
+ )
+ _READ_GUIDE = (
+ "- **探索可用面板**:接手新任务时,可使用 `list_slots` "
+ "宏观查看当前存在哪些面板。\n"
+ )
+ _WRITE_GUIDE = (
+ "- **维护规范与进度**:如设定'沟通口吻'等规范(`update_slot`),"
+ "或记录'待办清单'(`append_slot`)。\n"
+ "- **保持精简**:内容过长时,主动将其归档到长期记忆,"
+ "再重新提炼或调用 `delete_slot` 删除。\n"
+ )
-### 💡 核心机制
-- 被保存在记忆槽中的内容(如果已置顶),会在每次对话时**直接注入到你的上下文提示词中**,你无需任何搜索即可看见。
-- 槽位容量极其有限,仅用于维持当前最新的运行状态。
+ default_instructions = _INTRO_TEXT + _READ_GUIDE + _WRITE_GUIDE
-### 📝 何时使用记忆槽?
-- **维护全局规则**:例如设定"用户整体偏好"、"沟通口吻"、"全局指导原则"等需要时刻遵守的规范(使用 `update_slot`)。
-- **追踪当前进度**:例如记录"待办事项清单"、"当前任务进度"、"上下文摘要"(使用 `append_slot` 列表或 `update_slot` 覆盖)。
+ @classmethod
+ def read_only(cls, **kwargs) -> "MemorySlotToolkit":
+ """[工厂方法] 创建一个只读模式的记忆槽工具箱。"""
+ kwargs["include"] = ["list_slots", "read_slot"]
+ kwargs.setdefault("instructions", cls._INTRO_TEXT + cls._READ_GUIDE)
+ return cls(**kwargs)
-### ⚙️ 操作规范
-1. **探索可用面板**:接手新任务时,可使用 `list_slots` 宏观查看当前存在哪些状态面板。
-2. **保持精简**:槽位有严格的字符数限制。当内容过长时,请主动将其归档到长期记忆后,重新提炼并覆盖槽位,或直接调用 `delete_slot` 删除不再需要的槽位。\
-""" # noqa: E501
+ @classmethod
+ def write_only(cls, **kwargs) -> "MemorySlotToolkit":
+ """[工厂方法] 创建一个仅写入模式的记忆槽工具箱。"""
+ kwargs["exclude"] = ["list_slots", "read_slot"]
+ kwargs.setdefault("instructions", cls._INTRO_TEXT + cls._WRITE_GUIDE)
+ return cls(**kwargs)
def __init__(
self,
- memory_config: MemoryConfig | None = None,
+ scopes: dict[str, Any] | None = None,
+ backend: Any = None,
namespace: str | None = None,
**kwargs: Any,
):
- """
- 初始化中期记忆槽工具箱。
-
- 参数:
- memory_config: 记忆系统的全局配置对象,为空则使用全局默认。
- namespace: 当前隔离环境的命名空间。
- kwargs: 其他透传给 BaseToolkit 的参数。
- """
super().__init__(**kwargs)
- self.memory_config = memory_config
+ self.scopes = scopes or {}
+ self.backend = backend
self._namespace = namespace
def _get_runtime_meta_and_ctx(
self, context: RunContext, scope_name: str | None = None
) -> tuple[Any, SessionMetadata]:
ns = self._namespace or getattr(context.session, "namespace", "global")
- slot_ctx = memory_manager.get_slot_context(self.memory_config, namespace=ns)
+ slot_ctx = self.backend or memory_manager.get_backend("slots", namespace=ns)
- scope_builder = None
- if (
- self.memory_config
- and self.memory_config.slots
- and self.memory_config.slots.scopes
- ):
- scopes_dict = self.memory_config.slots.scopes
- if not scope_name:
- scope_builder = next(iter(scopes_dict.values()))
- else:
- scope_builder = scopes_dict.get(scope_name)
-
- if not scope_builder:
- scope_builder = getattr(self.memory_config, "base_isolation", None)
- if not scope_builder:
- from zhenxun.services.ai.context.memory.types import Isolation
-
- scope_builder = Isolation.AGENT_USER()
-
- selector = scope_builder.resolve(
- deps=context.deps,
- prefix="",
- default_namespace=ns,
- default_agent=context.run.agent_name,
+ scope_builder = (
+ self.scopes.get(scope_name)
+ if scope_name
+ else next(iter(self.scopes.values()), None)
)
- parts = selector.get_scope_parts()
- all_scopes = {"/"}
- current_path = ""
- for part in parts:
- current_path += f"/{part}"
- all_scopes.add(current_path)
- accessible_scopes = list(all_scopes)
- accessible_scopes.sort(key=lambda x: len(x.split("/")))
-
- scope_name_mapping = {}
- if (
- self.memory_config
- and self.memory_config.slots
- and self.memory_config.slots.scopes
- ):
- for name, builder in self.memory_config.slots.scopes.items():
- sel = builder.resolve(
- deps=context.deps,
- prefix="",
- default_namespace=ns,
- default_agent=context.run.agent_name,
- )
- scope_name_mapping[sel.scope_prefix] = name
-
- session_meta = SessionMetadata(
- session_id=context.session_id or "default_session",
- selector=selector,
- scope_prefix=selector.scope_prefix,
- accessible_scopes=accessible_scopes,
- scope_name_mapping=scope_name_mapping,
+ session_meta = ContextUtils.build_session_meta(
+ context=context,
+ target_builder=scope_builder,
+ extra_scopes=self.scopes,
+ custom_namespace=self._namespace,
)
return slot_ctx, session_meta
async def get_tools(self, context: RunContext | None = None) -> dict[str, BaseTool]:
tools = await super().get_tools(context)
- if (
- not self.memory_config
- or not self.memory_config.slots
- or not self.memory_config.slots.enable
- ):
+ if not self.scopes:
return tools
-
- scopes_dict = self.memory_config.slots.scopes
- if not scopes_dict:
- return tools
- scope_keys = tuple(scopes_dict.keys())
+ scope_keys = tuple(self.scopes.keys())
if len(scope_keys) > 1:
ScopeType = Literal[scope_keys]
@@ -237,12 +196,7 @@ class MemorySlotToolkit(BaseToolkit):
res = ["已创建的记忆槽列表:"]
show_scope = False
- if (
- self.memory_config
- and self.memory_config.slots
- and self.memory_config.slots.scopes
- and len(self.memory_config.slots.scopes) > 1
- ):
+ if len(self.scopes) > 1:
show_scope = True
for s in slots:
diff --git a/zhenxun/services/ai/tools/providers/skills/toolkit.py b/zhenxun/services/ai/tools/providers/skills/toolkit.py
index ec81b42d..a35d9e6e 100644
--- a/zhenxun/services/ai/tools/providers/skills/toolkit.py
+++ b/zhenxun/services/ai/tools/providers/skills/toolkit.py
@@ -295,6 +295,9 @@ class SkillMetaToolkit(BaseToolkit, SkillSandboxExecutionMixin):
if sandbox is None and context is not None:
sandbox = Inject._providers["sandbox"]["global"](context)
+ if sandbox is None:
+ return ToolResult(output="❌ 缺少沙箱环境或执行上下文").as_error()
+
executor = await sandbox.get_or_create_session(session_id, blueprint=bp)
fs_executor = cast(SupportsFileSystem, executor)
diff --git a/zhenxun/services/ai/utils/runtime.py b/zhenxun/services/ai/utils/runtime.py
index b573e22b..22c881b5 100644
--- a/zhenxun/services/ai/utils/runtime.py
+++ b/zhenxun/services/ai/utils/runtime.py
@@ -60,7 +60,7 @@ class ContextUtils:
context: Any, scope: Any, default_session_id: str
) -> str:
"""根据并发隔离范围 scope 动态计算并返回当前会话的并发锁 ID"""
- from zhenxun.services.ai.flow.base import ConcurrencyScope
+ from zhenxun.services.ai.flow.core.models import ConcurrencyScope
scope = scope or ConcurrencyScope.GROUP
@@ -135,6 +135,66 @@ class ContextUtils:
isolation_level=scope_builder,
)
+ @staticmethod
+ def build_session_meta(
+ context: Any,
+ target_builder: Any | None = None,
+ extra_scopes: dict[str, Any] | None = None,
+ custom_namespace: str | None = None,
+ ) -> Any:
+ """基于 RunContext 动态提取并生成 SessionMetadata"""
+ from zhenxun.services.ai.context.memory.types import Isolation, SessionMetadata
+
+ ns = custom_namespace or getattr(
+ getattr(context, "session", None), "namespace", "global"
+ )
+ agent_name = getattr(getattr(context, "run", None), "agent_name", None)
+ deps = getattr(context, "deps", None)
+
+ if target_builder is None:
+ target_builder = Isolation.AGENT_USER()
+
+ selector = target_builder.resolve(
+ deps=deps,
+ prefix="",
+ default_namespace=ns,
+ default_agent=agent_name,
+ )
+
+ all_scopes = {"/"}
+ parts = selector.get_scope_parts()
+ current_path = ""
+ for part in parts:
+ current_path += f"/{part}"
+ all_scopes.add(current_path)
+
+ scope_name_mapping = {}
+ if extra_scopes:
+ for name, builder in extra_scopes.items():
+ sel = builder.resolve(
+ deps=deps, prefix="", default_namespace=ns, default_agent=agent_name
+ )
+ all_scopes.add(sel.scope_prefix)
+ scope_name_mapping[sel.scope_prefix] = name
+
+ accessible_scopes = list(all_scopes)
+ accessible_scopes.sort(key=lambda x: len(x.split("/")))
+
+ session_id = (
+ getattr(context, "session_id", None)
+ or getattr(getattr(context, "session", None), "session_id", None)
+ or selector.scope_prefix
+ )
+
+ return SessionMetadata(
+ session_id=session_id,
+ selector=selector,
+ scope_prefix=selector.scope_prefix,
+ accessible_scopes=accessible_scopes,
+ scope_name_mapping=scope_name_mapping,
+ isolation_level=target_builder,
+ )
+
class PermissionUtils:
"""运行时权限校验通用工具类"""
diff --git a/zhenxun/services/ai/utils/scope.py b/zhenxun/services/ai/utils/scope.py
index be564316..f88e6b06 100644
--- a/zhenxun/services/ai/utils/scope.py
+++ b/zhenxun/services/ai/utils/scope.py
@@ -214,7 +214,7 @@ class ScopeBuilder:
selector.namespace = (
default_namespace
or getattr(deps, "namespace", None)
- or infer_plugin_namespace(default="global")
+ or infer_plugin_namespace()
)
if "agent" in self._dims:
selector.agent_name = default_agent