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