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