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

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

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

* 🚨 auto fix by pre-commit hooks

---------

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