mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-09 22:00:01 +08:00
✨ feat!(llm): 重构并升级大语言模型服务为全新 AI 智能体框架 (#2146)
* ✨ feat!(llm): 重构并升级大语言模型服务为全新 AI 智能体框架 - 【重构】将原 services/llm 重构并迁移至全新的 services/ai 架构,提供向下兼容垫片 - 【新增】引入 Agent、Team、Workflow 三大智能体与工作流编排范式 - 【新增】引入基于 RAG 的长期向量记忆与中期槽位记忆系统 - 【新增】引入基于 Docker 的安全代码执行沙箱环境 - 【新增】支持 MCP 协议,允许动态管理和调用 MCP 服务 - 【新增】引入输入输出安全合规护栏与自愈反思机制 - 【优化】重构并优化多厂商 API 适配器 (Gemini, OpenAI, DeepSeek, GLM 等) - 【优化】优化日志脱敏与 Token 预估机制 - 【移除】移除旧版 llm default 和 llm reset-key 命令,新增 llm mcp 管理命令 * 🔧 chore(deps): 更新项目依赖与配置 - 添加 mcp、jieba 和 aiodocker 依赖到配置文件及 requirements.txt - 在 pyright 配置中设置 reportMissingImports 为 none - 调整 .gitignore 中 resources 目录的忽略规则 * ♻️ refactor(tools): 重构工具终止机制并清理知识库日志输出 - 统一使用 `context.state["__end_run__"]` 替代 `EndRunResult` 控制任务结束 - 移除文件系统和向量知识库检索工具中 `ToolResult` 的 `.with_log` 调用 - 调整指令处理器(Directive)的返回值为 `tool_res.output` - 修复部分类型检查警告并优化联合类型判断语法 * ♻️ refactor(tools): 重构工具副作用指令与控制流熔断机制 - 引入 `DirectivePayload` 及 `ToolResult` 的子类以结构化表达工具副作用 - 移除通过 `context.state` 传递魔术变量的隐式控制流设计 - 重构 `DirectiveManager` 处理器接口,直接在处理器中修改 `AgentState` 并构建 `AgentRunResult` - 在 `StandardAgentExecutor` 中统一通过 `directive_manager` 调度工具返回的副作用指令 - 补全 `MessageBuilder` 中部分核心方法的文档注释 * 🐛 fix(sandbox): 修复 Docker 沙箱容器状态检测与会话清理逻辑 -【修复】修正 `is_alive` 中直接读取私有属性的问题,改用 `show()` 返回值 -【修复】解决 `execute_code` 中缓存的执行器与当前会话不一致的问题 -【优化】在清理工作区前增加容器存活检测,避免向已死容器发送请求 -【优化】创建容器时增加运行状态校验,若已停止则自动从缓存中移除并重建 -【优化】优化容器销毁和清理逻辑,静默处理容器不存在 (404) 的异常 * 📝 docs(core): 补充核心模块初始化方法的文档注释 * 🚨 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
bdc1374848
commit
80fc5b86a7
@@ -0,0 +1,23 @@
|
||||
from .builder import MemoryBuilder
|
||||
from .compression import MemoryPolicy
|
||||
from .facades import AgentSessionFacade
|
||||
from .manager import memory_manager
|
||||
from .models import (
|
||||
BaseMemoryIngestionMiddleware,
|
||||
MemoryConfig,
|
||||
)
|
||||
from .types import (
|
||||
Isolation,
|
||||
SessionMetadata,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"AgentSessionFacade",
|
||||
"BaseMemoryIngestionMiddleware",
|
||||
"Isolation",
|
||||
"MemoryBuilder",
|
||||
"MemoryConfig",
|
||||
"MemoryPolicy",
|
||||
"SessionMetadata",
|
||||
"memory_manager",
|
||||
]
|
||||
@@ -0,0 +1,270 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing_extensions import Self
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from zhenxun.services.ai.context.memory.models import MemorySlot
|
||||
from zhenxun.services.ai.context.memory.storage.interfaces import (
|
||||
BaseChatContext,
|
||||
BaseMemoryIngestionMiddleware,
|
||||
BaseSlotContext,
|
||||
)
|
||||
from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend
|
||||
from zhenxun.services.ai.context.rag.engine import ScopedRAGClient
|
||||
|
||||
from zhenxun.services.ai.context.memory.compression import MemoryPolicy
|
||||
from zhenxun.services.ai.context.memory.models import (
|
||||
ContextCompressionConfig,
|
||||
IngestionConfig,
|
||||
LongTermConfig,
|
||||
MemoryConfig,
|
||||
ShortTermConfig,
|
||||
SlotMemoryConfig,
|
||||
)
|
||||
from zhenxun.services.ai.context.memory.types import (
|
||||
AutoRecallPolicy,
|
||||
)
|
||||
from zhenxun.services.ai.utils.scope import ScopeBuilder
|
||||
|
||||
|
||||
class MemoryBuilder:
|
||||
"""
|
||||
记忆配置的链式构建器 (Fluent Builder)。
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
"""
|
||||
初始化 MemoryBuilder 实例。
|
||||
|
||||
创建一个默认关闭短期和长期记忆,并包含默认上下文压缩配置的构建器。
|
||||
"""
|
||||
self._config = MemoryConfig(
|
||||
short_term=ShortTermConfig(enable=False),
|
||||
slots=SlotMemoryConfig(enable=False),
|
||||
long_term=LongTermConfig(enable=False),
|
||||
compression=ContextCompressionConfig(),
|
||||
ingestion=IngestionConfig(),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def auto(cls) -> "MemoryBuilder":
|
||||
"""
|
||||
创建一个开箱即用的默认记忆配置构建器。
|
||||
|
||||
默认开启隔离的短期记忆,并使用 LLM 对话摘要进行上下文压缩。
|
||||
"""
|
||||
return cls().with_short_term(enable=True).with_llm_summary()
|
||||
|
||||
@classmethod
|
||||
def resolve(
|
||||
cls, memory: bool | MemoryConfig | "MemoryBuilder" | None
|
||||
) -> MemoryConfig:
|
||||
if isinstance(memory, MemoryConfig):
|
||||
return memory
|
||||
if isinstance(memory, cls):
|
||||
return memory.build()
|
||||
if isinstance(memory, bool):
|
||||
return MemoryConfig(
|
||||
short_term=ShortTermConfig(enable=memory),
|
||||
long_term=LongTermConfig(enable=memory),
|
||||
)
|
||||
|
||||
return MemoryConfig(short_term=ShortTermConfig(enable=False))
|
||||
|
||||
def with_base_isolation(self, isolation: ScopeBuilder) -> Self:
|
||||
"""设置顶层基准隔离级别,短期/中期/长期记忆将默认继承此级别"""
|
||||
self._config.base_isolation = isolation
|
||||
self._config.short_term.isolation = isolation
|
||||
return self
|
||||
|
||||
def with_short_term(
|
||||
self,
|
||||
enable: bool = True,
|
||||
isolation: ScopeBuilder | None = None,
|
||||
backend: "str | BaseChatContext | None" = None,
|
||||
) -> Self:
|
||||
"""
|
||||
配置短期对话历史记忆。
|
||||
|
||||
参数:
|
||||
enable: 是否开启短期记忆。
|
||||
isolation: 记忆隔离级别 (ScopeBuilder),决定会话历史记录的区分范围。
|
||||
backend: 短期记忆存储后端实例,如果为 None 则使用全局默认后端。
|
||||
"""
|
||||
self._config.short_term.enable = enable
|
||||
if isolation is not None:
|
||||
self._config.base_isolation = isolation
|
||||
self._config.short_term.isolation = isolation
|
||||
if backend is not None:
|
||||
self._config.short_term.backend = backend
|
||||
return self
|
||||
|
||||
def with_slots(
|
||||
self,
|
||||
enable: bool = True,
|
||||
scopes: dict[str, ScopeBuilder] | None = None,
|
||||
default_slots: list["MemorySlot"] | None = None,
|
||||
backend: "str | BaseSlotContext | None" = None,
|
||||
instructions: str | None = None,
|
||||
) -> Self:
|
||||
"""
|
||||
配置核心槽位记忆 (Memory Slots)。
|
||||
|
||||
参数:
|
||||
enable: 是否启用槽位记忆。
|
||||
scopes: 语义化作用域映射字典。如果只有一个键值对,则大模型不可见该参数。
|
||||
default_slots: 首次初始化时自动写入的默认槽位列表。
|
||||
backend: 槽位记忆存储后端,如果为 None 则使用全局默认后端。
|
||||
instructions: 覆写内置槽位管理工具箱的默认系统提示词规则。
|
||||
"""
|
||||
self._config.slots.enable = enable
|
||||
if scopes is not None:
|
||||
self._config.slots.scopes = scopes
|
||||
if default_slots is not None:
|
||||
self._config.slots.default_slots = default_slots
|
||||
if backend is not None:
|
||||
self._config.slots.backend = backend
|
||||
if instructions is not None:
|
||||
self._config.slots.instructions = instructions
|
||||
return self
|
||||
|
||||
def with_long_term(
|
||||
self,
|
||||
enable: bool = True,
|
||||
scopes: dict[str, ScopeBuilder] | None = None,
|
||||
engine: "ScopedRAGClient | None" = None,
|
||||
backend: "str | StorageBackend | None" = None,
|
||||
embedder: "Embedder | str | None" = None,
|
||||
agentic: bool = True,
|
||||
auto_recall: AutoRecallPolicy = False,
|
||||
instructions: str | None = None,
|
||||
) -> Self:
|
||||
"""
|
||||
配置长期向量记忆与 RAG 设定。
|
||||
|
||||
参数:
|
||||
enable: 是否启用长期记忆。
|
||||
scopes: 语义化作用域映射字典。如果只有一个键值对,则大模型不可见该参数。
|
||||
engine: 高级 RAG 检索引擎实例 (推荐)。若提供,将接管记忆的底层检索、混合与重排。
|
||||
backend: 长期记忆存储后端。
|
||||
embedder: 用于向量化的文本嵌入模型实例。
|
||||
agentic: 是否开启主动智能体记忆管理 (增删改查工具自动注入)。
|
||||
auto_recall: 长期记忆的自动召回策略,支持 bool 或 Callable 函数。
|
||||
instructions: 覆写内置长期记忆工具箱的默认系统提示词规则。
|
||||
""" # noqa: E501
|
||||
self._config.long_term.enable = enable
|
||||
self._config.long_term.engine = engine
|
||||
if scopes is not None:
|
||||
self._config.long_term.scopes = scopes
|
||||
self._config.long_term.backend = backend
|
||||
self._config.long_term.embedder = embedder
|
||||
self._config.long_term.agentic = agentic
|
||||
self._config.long_term.auto_recall = auto_recall
|
||||
if instructions is not None:
|
||||
self._config.long_term.instructions = instructions
|
||||
return self
|
||||
|
||||
def with_multimodal_window(self, window_size: int = 5) -> Self:
|
||||
"""
|
||||
配置多模态历史视窗大小。
|
||||
|
||||
超出此窗口的图片/视频等富媒体消息会自动转换为纯文本占位符,以节省 Token 预算。
|
||||
|
||||
参数:
|
||||
window_size: 允许保留多模态信息的最新的对话轮数。
|
||||
"""
|
||||
self._config.compression.vision_window = window_size
|
||||
return self
|
||||
|
||||
def with_llm_summary(
|
||||
self,
|
||||
trigger_tokens: int = 4000,
|
||||
max_turns: int = 0,
|
||||
keep_recent_turns: int = 0,
|
||||
summarization_model: str | None = None,
|
||||
summarization_prompt: str = "请概括以下对话内容,保留关键的约束条件、用户偏好、已完成的任务状态和未解决的问题。", # noqa: E501
|
||||
) -> Self:
|
||||
"""
|
||||
配置使用大模型自然语言总结作为上下文压缩策略。
|
||||
|
||||
参数:
|
||||
trigger_tokens: 触发压缩的 Token 门槛。
|
||||
max_turns: 压缩策略作用的最大历史对话轮数上限。
|
||||
keep_recent_turns: 在大模型总结之外,强制保留的最近原始对话轮数。
|
||||
summarization_model: 负责生成总结的大模型名称。
|
||||
summarization_prompt: 生成总结时所使用的系统提示词。
|
||||
"""
|
||||
self._config.compression.policy = MemoryPolicy.llm_summarize(
|
||||
trigger_tokens=trigger_tokens,
|
||||
max_turns=max_turns,
|
||||
keep_recent_turns=keep_recent_turns,
|
||||
summarization_model=summarization_model,
|
||||
summarization_prompt=summarization_prompt,
|
||||
)
|
||||
return self
|
||||
|
||||
def with_structured_summary(
|
||||
self,
|
||||
trigger_tokens: int = 4000,
|
||||
max_turns: int = 0,
|
||||
keep_recent_turns: int = 0,
|
||||
summarization_model: str | None = None,
|
||||
response_model: type[BaseModel] | None = None,
|
||||
prompt_template: str | None = None,
|
||||
format_callback: Callable[[Any], str] | None = None,
|
||||
) -> Self:
|
||||
"""
|
||||
配置使用自定义结构化 JSON 提取作为上下文压缩策略。
|
||||
|
||||
参数:
|
||||
trigger_tokens: 触发压缩的 Token 门槛。
|
||||
max_turns: 压缩策略作用的最大历史对话轮数上限。
|
||||
keep_recent_turns: 强制保留的最近原始对话轮数。
|
||||
summarization_model: 负责生成结构化总结的大模型名称。
|
||||
response_model: (可选) 自定义的 Pydantic 数据模型,用于指导提取的结构。
|
||||
prompt_template: (可选) 提取提示词模板,
|
||||
支持 {prev_summary} 和 {dialogue} 变量。
|
||||
format_callback: (可选) 将提取出的 Pydantic 实例格式化为字符串的回调函数。
|
||||
"""
|
||||
self._config.compression.policy = MemoryPolicy.structured_summarize(
|
||||
trigger_tokens=trigger_tokens,
|
||||
max_turns=max_turns,
|
||||
keep_recent_turns=keep_recent_turns,
|
||||
summarization_model=summarization_model,
|
||||
response_model=response_model,
|
||||
prompt_template=prompt_template,
|
||||
format_callback=format_callback,
|
||||
)
|
||||
return self
|
||||
|
||||
def unlimited(self) -> Self:
|
||||
"""
|
||||
配置为不进行任何截断和压缩的策略。
|
||||
适用于短程会话或者具备超长上下文窗口的底层语言模型。
|
||||
"""
|
||||
self._config.compression.policy = MemoryPolicy.unlimited()
|
||||
return self
|
||||
|
||||
def with_ingestion_middlewares(
|
||||
self, *middlewares: "BaseMemoryIngestionMiddleware"
|
||||
) -> Self:
|
||||
"""
|
||||
配置记忆入库管线中间件。
|
||||
用于在消息正式落盘前进行实体消解、隐私脱敏、自动打标签等操作。
|
||||
"""
|
||||
self._config.ingestion.middlewares.extend(middlewares)
|
||||
return self
|
||||
|
||||
def build(self) -> 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
|
||||
@@ -0,0 +1,52 @@
|
||||
from typing import Any
|
||||
|
||||
from zhenxun.services.ai.capabilities.base import AbstractCapability
|
||||
from zhenxun.services.ai.context.memory.models import MemoryConfig
|
||||
from zhenxun.services.ai.run.context import RunContext
|
||||
from zhenxun.services.ai.tools.providers.builtin.memory import MemoryManagementToolkit
|
||||
|
||||
|
||||
class AgenticMemoryCapability(AbstractCapability):
|
||||
"""
|
||||
智能体主动记忆管理能力 (Agentic Memory Management)。
|
||||
当 `MemoryConfig.long_term.enable == True` 且 `agentic == True` 时隐式挂载,
|
||||
在运行时动态组装并向大模型提供 `MemoryManagementToolkit` 工具箱。
|
||||
"""
|
||||
|
||||
def __init__(self, memory_config: MemoryConfig, namespace: str):
|
||||
self.memory_config = memory_config
|
||||
self.namespace = namespace
|
||||
|
||||
async def get_tools(self, context: RunContext) -> list[Any]:
|
||||
kwargs = self.memory_config.long_term.toolkit_kwargs.copy()
|
||||
kwargs["memory_config"] = self.memory_config
|
||||
kwargs["namespace"] = self.namespace
|
||||
if self.memory_config.long_term.instructions is not None:
|
||||
kwargs["instructions"] = self.memory_config.long_term.instructions
|
||||
|
||||
toolkit = MemoryManagementToolkit(**kwargs)
|
||||
return [toolkit]
|
||||
|
||||
|
||||
class SlotMemoryCapability(AbstractCapability):
|
||||
"""
|
||||
槽位记忆能力组件。
|
||||
当 `MemoryConfig.slots.enable == True` 时隐式挂载,
|
||||
在运行时动态组装并向大模型提供 `MemorySlotToolkit` 工具箱。
|
||||
"""
|
||||
|
||||
def __init__(self, memory_config: MemoryConfig, namespace: str):
|
||||
self.memory_config = memory_config
|
||||
self.namespace = namespace
|
||||
|
||||
async def get_tools(self, context: RunContext) -> list[Any]:
|
||||
from zhenxun.services.ai.tools.providers.builtin.slots import MemorySlotToolkit
|
||||
|
||||
kwargs = self.memory_config.slots.toolkit_kwargs.copy()
|
||||
kwargs["memory_config"] = self.memory_config
|
||||
kwargs["namespace"] = self.namespace
|
||||
if self.memory_config.slots.instructions is not None:
|
||||
kwargs["instructions"] = self.memory_config.slots.instructions
|
||||
|
||||
toolkit = MemorySlotToolkit(**kwargs)
|
||||
return [toolkit]
|
||||
@@ -0,0 +1,688 @@
|
||||
from abc import abstractmethod
|
||||
from collections.abc import Callable
|
||||
from typing import Any, Generic, TypeVar
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from zhenxun.services.ai.context.memory.storage.interfaces import (
|
||||
BaseMemoryReducer,
|
||||
)
|
||||
from zhenxun.services.ai.core.engine.token_counter import token_counter
|
||||
from zhenxun.services.ai.core.messages import (
|
||||
AudioPart,
|
||||
FilePart,
|
||||
ImagePart,
|
||||
LLMMessage,
|
||||
SystemMessage,
|
||||
TextPart,
|
||||
VideoPart,
|
||||
)
|
||||
from zhenxun.services.ai.llm.manager import get_default_model
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.pydantic_compat import model_copy
|
||||
|
||||
|
||||
class MultimodalPlaceholderReducer(BaseMemoryReducer):
|
||||
"""视觉媒体降级:将超过一定轮数的老图片/视频替换为 <图片> 占位符文本"""
|
||||
|
||||
def __init__(self, window_size: int = 5):
|
||||
"""
|
||||
初始化多模态占位符修剪器。
|
||||
|
||||
参数:
|
||||
window_size: 多模态视窗大小。在保留最近的指定数量的多模态消息对后,
|
||||
超出的旧消息中多模态内容(如图片、视频等)将被替换为占位文本。
|
||||
"""
|
||||
self.window_size = window_size
|
||||
|
||||
@staticmethod
|
||||
def apply_multimodal_placeholder(message: LLMMessage) -> LLMMessage:
|
||||
sanitized_message = model_copy(message, deep=False)
|
||||
new_content_parts = []
|
||||
|
||||
for part in sanitized_message.content:
|
||||
if isinstance(part, ImagePart):
|
||||
new_content_parts.append(TextPart(text="<图片>"))
|
||||
elif isinstance(part, AudioPart):
|
||||
new_content_parts.append(TextPart(text="<音频>"))
|
||||
elif isinstance(part, VideoPart):
|
||||
new_content_parts.append(TextPart(text="<视频>"))
|
||||
elif isinstance(part, FilePart):
|
||||
new_content_parts.append(TextPart(text="<文件>"))
|
||||
elif isinstance(part, TextPart) and "[多模态内容:" in part.text:
|
||||
new_content_parts.append(TextPart(text="<图片>"))
|
||||
else:
|
||||
new_content_parts.append(part)
|
||||
|
||||
merged_parts = []
|
||||
for part in new_content_parts:
|
||||
if (
|
||||
isinstance(part, TextPart)
|
||||
and merged_parts
|
||||
and isinstance(merged_parts[-1], TextPart)
|
||||
):
|
||||
new_text = (merged_parts[-1].text or "") + " " + (part.text or "")
|
||||
merged_parts[-1] = TextPart(text=new_text.strip())
|
||||
else:
|
||||
merged_parts.append(part)
|
||||
|
||||
sanitized_message.content = merged_parts
|
||||
sanitized_message.token_cost = None
|
||||
return sanitized_message
|
||||
|
||||
async def reduce(self, messages, current_tokens, model_name, base_overhead=0):
|
||||
if self.window_size <= 0:
|
||||
return messages, False, current_tokens
|
||||
|
||||
processed_messages = []
|
||||
user_multimodal_count = 0
|
||||
changed = False
|
||||
|
||||
for msg in reversed(messages):
|
||||
has_multimodal = False
|
||||
if isinstance(msg.content, list):
|
||||
has_multimodal = any(
|
||||
isinstance(p, ImagePart | AudioPart | VideoPart | FilePart)
|
||||
or (isinstance(p, TextPart) and "[多模态内容:" in p.text)
|
||||
for p in msg.content
|
||||
)
|
||||
|
||||
if has_multimodal:
|
||||
if msg.role == "user":
|
||||
user_multimodal_count += 1
|
||||
if user_multimodal_count > self.window_size:
|
||||
processed_messages.append(self.apply_multimodal_placeholder(msg))
|
||||
changed = True
|
||||
else:
|
||||
processed_messages.append(msg)
|
||||
else:
|
||||
processed_messages.append(msg)
|
||||
|
||||
if not changed:
|
||||
return messages, False, current_tokens
|
||||
|
||||
processed_messages.reverse()
|
||||
new_tokens = token_counter.count_context(
|
||||
processed_messages, model_name, base_overhead
|
||||
)
|
||||
return processed_messages, True, new_tokens
|
||||
|
||||
|
||||
class MessageDropper(BaseMemoryReducer):
|
||||
"""消息丢弃器:在 Token 超过阈值时丢弃最早的非置顶消息对。"""
|
||||
|
||||
def __init__(self, trigger_tokens: int = 4000):
|
||||
"""
|
||||
初始化消息丢弃器。
|
||||
|
||||
参数:
|
||||
trigger_tokens: 触发丢弃策略的 Token 阈值上限。
|
||||
当当前对话 Token 总数超过此值时,将触发硬截断。
|
||||
"""
|
||||
self.trigger_tokens = trigger_tokens
|
||||
|
||||
async def reduce(self, messages, current_tokens, model_name, base_overhead=0):
|
||||
if current_tokens <= self.trigger_tokens:
|
||||
return messages, False, current_tokens
|
||||
|
||||
logger.info(
|
||||
"✂️ [MemoryCompression] 触发硬截断丢弃策略 | 原因: "
|
||||
f"当前 Token 预估 ({current_tokens}) 仍超过硬性上限 ({self.trigger_tokens})," # noqa: E501
|
||||
"开始丢弃最旧的历史对话..."
|
||||
)
|
||||
new_messages = list(messages)
|
||||
changed = False
|
||||
|
||||
while current_tokens > self.trigger_tokens:
|
||||
user_indices = [
|
||||
i
|
||||
for i, m in enumerate(new_messages)
|
||||
if m.role == "user"
|
||||
and not (m.metadata and m.metadata.get("pinned", False))
|
||||
]
|
||||
|
||||
if len(user_indices) < 2:
|
||||
break
|
||||
|
||||
start_idx = user_indices[0]
|
||||
end_idx = user_indices[1]
|
||||
|
||||
del new_messages[start_idx:end_idx]
|
||||
changed = True
|
||||
|
||||
current_tokens = token_counter.count_context(
|
||||
new_messages, model_name, base_overhead
|
||||
)
|
||||
return new_messages, changed, current_tokens
|
||||
|
||||
|
||||
class ToolPrunerReducer(BaseMemoryReducer):
|
||||
"""工具结果修剪器:纯粹计算工具输出的 Token 和轮数,超标时剔除老旧工具返回结果"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
keep_recent_turns: int = 3,
|
||||
trigger_tokens: int = 4000,
|
||||
max_turns: int = 0,
|
||||
):
|
||||
"""
|
||||
初始化工具结果修剪器。
|
||||
|
||||
参数:
|
||||
keep_recent_turns: 保留最近的工具调用轮数(不进行内容截断的轮数)。
|
||||
trigger_tokens: 触发工具修剪策略的工具总 Token 阈值上限。
|
||||
max_turns: 触发工具修剪的最大工具调用轮数上限。若为 0,则不限制轮数。
|
||||
"""
|
||||
self.keep_recent_turns = keep_recent_turns
|
||||
self.trigger_tokens = trigger_tokens
|
||||
self.max_turns = max_turns
|
||||
|
||||
async def reduce(self, messages, current_tokens, model_name, base_overhead=0):
|
||||
tool_msgs = [m for m in messages if m.role == "tool"]
|
||||
tool_turns = len(tool_msgs)
|
||||
|
||||
if tool_turns == 0:
|
||||
return messages, False, current_tokens
|
||||
|
||||
tool_tokens = sum(token_counter.count_message(m, model_name) for m in tool_msgs)
|
||||
|
||||
is_token_exceeded = tool_tokens > self.trigger_tokens
|
||||
is_turn_exceeded = self.max_turns > 0 and tool_turns > self.max_turns
|
||||
|
||||
if not (is_token_exceeded or is_turn_exceeded):
|
||||
return messages, False, current_tokens
|
||||
|
||||
reasons = []
|
||||
if is_token_exceeded:
|
||||
reasons.append(f"工具Token超标 ({tool_tokens} > {self.trigger_tokens})")
|
||||
if is_turn_exceeded:
|
||||
reasons.append(f"工具调用轮数超限 ({tool_turns} > {self.max_turns})")
|
||||
|
||||
logger.info(
|
||||
f"✂️ [MemoryCompression] 触发工具结果修剪策略 | 原因: {' 且 '.join(reasons)}"
|
||||
)
|
||||
|
||||
from zhenxun.services.ai.core.messages import ToolReturnPart
|
||||
|
||||
new_messages = []
|
||||
tools_kept = 0
|
||||
changed = False
|
||||
|
||||
for msg in reversed(messages):
|
||||
if msg.role != "tool":
|
||||
new_messages.append(msg)
|
||||
continue
|
||||
|
||||
if tools_kept < self.keep_recent_turns:
|
||||
tools_kept += 1
|
||||
new_messages.append(msg)
|
||||
continue
|
||||
|
||||
new_content = []
|
||||
part_changed = False
|
||||
for p in msg.content:
|
||||
if isinstance(p, ToolReturnPart):
|
||||
old_len = len(str(p.output))
|
||||
new_p = model_copy(
|
||||
p,
|
||||
update={
|
||||
"output": f"[数据过载自动截断 - 原长度: {old_len} 字符]"
|
||||
},
|
||||
)
|
||||
new_content.append(new_p)
|
||||
part_changed = True
|
||||
changed = True
|
||||
else:
|
||||
new_content.append(p)
|
||||
|
||||
if part_changed:
|
||||
new_msg = model_copy(
|
||||
msg, update={"content": new_content, "token_cost": None}
|
||||
)
|
||||
new_messages.append(new_msg)
|
||||
else:
|
||||
new_messages.append(msg)
|
||||
|
||||
if not changed:
|
||||
return messages, False, current_tokens
|
||||
|
||||
new_messages.reverse()
|
||||
new_total = token_counter.count_context(new_messages, model_name, base_overhead)
|
||||
return new_messages, True, new_total
|
||||
|
||||
|
||||
class AbstractSummarizerReducer(BaseMemoryReducer):
|
||||
"""抽象总结压缩器:提取阈值判断与上下文分流的公共逻辑"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
strategy_name: str,
|
||||
keep_recent_turns: int = 0,
|
||||
trigger_tokens: int = 4000,
|
||||
max_turns: int | None = None,
|
||||
summarization_model: str | None = None,
|
||||
):
|
||||
"""
|
||||
初始化抽象总结压缩基类。
|
||||
|
||||
参数:
|
||||
strategy_name: 压缩策略的名称,用于日志输出和追踪。
|
||||
keep_recent_turns: 压缩时需要保留的最新的对话轮数(不参与总结的轮数)。
|
||||
trigger_tokens: 触发总结策略的 Token 阈值上限。
|
||||
max_turns: 触发总结策略的最大对话轮数上限。
|
||||
summarization_model: 用于执行总结压缩大模型请求的模型名称,若为 None 则使用默认模型。
|
||||
""" # noqa: E501
|
||||
self.strategy_name = strategy_name
|
||||
self.keep_recent_turns = keep_recent_turns
|
||||
self.trigger_tokens = trigger_tokens
|
||||
self.max_turns = max_turns
|
||||
self.summarization_model = summarization_model
|
||||
|
||||
@abstractmethod
|
||||
async def _execute_summarization(
|
||||
self, to_summarize: list[LLMMessage], prev_summary: str
|
||||
) -> LLMMessage | None:
|
||||
"""由子类实现具体的 LLM 调用逻辑,返回新的总结消息"""
|
||||
pass
|
||||
|
||||
async def reduce(self, messages, current_tokens, model_name, base_overhead=0):
|
||||
user_turns = sum(
|
||||
1
|
||||
for m in messages
|
||||
if m.role == "user"
|
||||
and not (m.metadata and m.metadata.get("is_summary", False))
|
||||
)
|
||||
is_token_exceeded = current_tokens > self.trigger_tokens
|
||||
is_turn_exceeded = (
|
||||
self.max_turns is not None
|
||||
and self.max_turns > 0
|
||||
and user_turns > self.max_turns
|
||||
)
|
||||
|
||||
if not (is_token_exceeded or is_turn_exceeded):
|
||||
return messages, False, current_tokens
|
||||
|
||||
reasons = []
|
||||
if is_token_exceeded:
|
||||
reasons.append(f"Token 预估超限 ({current_tokens} > {self.trigger_tokens})")
|
||||
if is_turn_exceeded:
|
||||
reasons.append(f"有效对话轮次超限 ({user_turns} > {self.max_turns})")
|
||||
logger.info(
|
||||
f"🔄 [MemoryCompression] 触发{self.strategy_name}策略 | 原因: "
|
||||
f"{' 且 '.join(reasons)}"
|
||||
)
|
||||
|
||||
pinned_msgs, working_msgs, prev_summary = [], [], ""
|
||||
for msg in messages:
|
||||
is_pinned = isinstance(msg, SystemMessage) or (
|
||||
msg.metadata and msg.metadata.get("pinned", False)
|
||||
)
|
||||
if msg.metadata and msg.metadata.get("is_summary", False):
|
||||
prev_summary = msg.extract_text
|
||||
elif is_pinned:
|
||||
pinned_msgs.append(msg)
|
||||
else:
|
||||
working_msgs.append(msg)
|
||||
|
||||
user_indices = [i for i, m in enumerate(working_msgs) if m.role == "user"]
|
||||
|
||||
if len(user_indices) <= self.keep_recent_turns:
|
||||
return messages, False, current_tokens
|
||||
|
||||
split_idx = (
|
||||
user_indices[-self.keep_recent_turns]
|
||||
if self.keep_recent_turns > 0
|
||||
else len(working_msgs)
|
||||
)
|
||||
to_summarize = working_msgs[:split_idx]
|
||||
to_keep = working_msgs[split_idx:]
|
||||
|
||||
new_summary_msg = await self._execute_summarization(to_summarize, prev_summary)
|
||||
if not new_summary_msg:
|
||||
return messages, False, current_tokens
|
||||
|
||||
new_messages = [*pinned_msgs, new_summary_msg, *to_keep]
|
||||
return (
|
||||
new_messages,
|
||||
True,
|
||||
token_counter.count_context(new_messages, model_name, base_overhead),
|
||||
)
|
||||
|
||||
|
||||
class LLMSummarizerReducer(AbstractSummarizerReducer):
|
||||
"""大模型总结压缩器:将较早的历史对话记录通过 LLM 压缩合并为一段文本摘要。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
keep_recent_turns: int = 0,
|
||||
trigger_tokens: int = 4000,
|
||||
max_turns: int | None = None,
|
||||
summarization_model: str | None = None,
|
||||
summarization_prompt: str = (
|
||||
"请概括以下对话内容,保留关键的约束条件、用户偏好、"
|
||||
"已完成的任务状态和未解决的问题。"
|
||||
),
|
||||
):
|
||||
"""
|
||||
初始化大模型总结压缩器。
|
||||
|
||||
参数:
|
||||
keep_recent_turns: 压缩时需要保留的最新的对话轮数。
|
||||
trigger_tokens: 触发总结策略的 Token 阈值上限。
|
||||
max_turns: 触发总结策略的最大对话轮数上限。
|
||||
summarization_model: 用于执行总结压缩的大模型名称。
|
||||
summarization_prompt: 发送给大模型的总结引导 Prompt 提示词。
|
||||
"""
|
||||
super().__init__(
|
||||
strategy_name="历史对话合并总结",
|
||||
keep_recent_turns=keep_recent_turns,
|
||||
trigger_tokens=trigger_tokens,
|
||||
max_turns=max_turns,
|
||||
summarization_model=summarization_model,
|
||||
)
|
||||
self.summarization_prompt = summarization_prompt
|
||||
|
||||
async def _execute_summarization(
|
||||
self, to_summarize: list[LLMMessage], prev_summary: str
|
||||
) -> LLMMessage | None:
|
||||
prompt_text = f"### 📋 [对话摘要任务]\n{self.summarization_prompt}\n\n"
|
||||
if prev_summary:
|
||||
prompt_text += "#### önceki_summary (参考先前的快照):\n"
|
||||
prompt_text += f"> {prev_summary}\n\n"
|
||||
prompt_text += "#### 待处理的历史消息流:\n"
|
||||
for m in to_summarize:
|
||||
c_str = m.extract_text[:1500]
|
||||
speaker = m.source_name if m.source_name else m.role.capitalize()
|
||||
prompt_text += f"[{speaker}]: {c_str}\n"
|
||||
prompt_text += "</需要合并的旧对话记录>\n"
|
||||
|
||||
from zhenxun.services.ai.llm.api import chat
|
||||
|
||||
try:
|
||||
model_to_use = self.summarization_model or get_default_model("chat")
|
||||
response = await chat(
|
||||
prompt_text,
|
||||
model=model_to_use,
|
||||
instruction="你是后台记忆整理引擎。请客观、简明输出当前对话全局摘要。",
|
||||
)
|
||||
new_summary_msg = LLMMessage.assistant_text_response(
|
||||
f"【历史对话摘要记忆】\n{response.text}"
|
||||
)
|
||||
new_summary_msg.metadata = {"is_summary": True, "pinned": True}
|
||||
return new_summary_msg
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"[{self.__class__.__name__}] 压缩总结调用失败,已跳过本次压缩: {e}"
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
_T_Summary = TypeVar("_T_Summary", bound=BaseModel)
|
||||
|
||||
|
||||
class StructuredSummaryReducer(AbstractSummarizerReducer, Generic[_T_Summary]):
|
||||
"""结构化总结压缩器:基于 JSON Schema 格式化抽取长上下文状态信息并合并"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
response_model: type[_T_Summary],
|
||||
prompt_template: str,
|
||||
format_callback: Callable[[_T_Summary], str],
|
||||
keep_recent_turns: int = 0,
|
||||
trigger_tokens: int = 4000,
|
||||
max_turns: int | None = None,
|
||||
summarization_model: str | None = None,
|
||||
instruction: str = (
|
||||
"请提取并合并先前的状态和最新的对话内容,保持精简,不要编造事实"
|
||||
),
|
||||
):
|
||||
"""
|
||||
初始化结构化总结压缩器。
|
||||
|
||||
参数:
|
||||
response_model: 接收结构化输出的 Pydantic 模型类,需继承自 BaseModel。
|
||||
prompt_template: 用于抽取合并状态的 Prompt 模板,包含 {prev_summary} 和 {dialogue} 占位符。
|
||||
format_callback: 格式化回调函数,用于将结构化 Pydantic 响应对象转换为便于大模型阅读的字符串。
|
||||
keep_recent_turns: 压缩时需要保留的最新的对话轮数。
|
||||
trigger_tokens: 触发总结策略的 Token 阈值上限。
|
||||
max_turns: 触发总结策略的最大对话轮数上限。
|
||||
summarization_model: 用于执行总结压缩的大模型名称。
|
||||
instruction: 指导大模型生成结构化数据时的系统指令说明。
|
||||
""" # noqa: E501
|
||||
super().__init__(
|
||||
strategy_name="结构化状态抽取压缩",
|
||||
keep_recent_turns=keep_recent_turns,
|
||||
trigger_tokens=trigger_tokens,
|
||||
max_turns=max_turns,
|
||||
summarization_model=summarization_model,
|
||||
)
|
||||
self.response_model = response_model
|
||||
self.prompt_template = prompt_template
|
||||
self.format_callback = format_callback
|
||||
self.instruction = instruction
|
||||
|
||||
async def _execute_summarization(
|
||||
self, to_summarize: list[LLMMessage], prev_summary: str
|
||||
) -> LLMMessage | None:
|
||||
dialogue_text = ""
|
||||
for m in to_summarize:
|
||||
c_str = m.extract_text[:1500]
|
||||
speaker = m.source_name if m.source_name else m.role.capitalize()
|
||||
dialogue_text += f"[{speaker}]: {c_str}\n"
|
||||
|
||||
prompt_text = self.prompt_template.format(
|
||||
prev_summary=prev_summary, dialogue=dialogue_text
|
||||
)
|
||||
|
||||
from zhenxun.services.ai.llm.api import generate_structured
|
||||
|
||||
try:
|
||||
model_to_use = self.summarization_model or get_default_model("chat")
|
||||
summary_obj = await generate_structured(
|
||||
prompt_text,
|
||||
response_model=self.response_model,
|
||||
model=model_to_use,
|
||||
instruction=self.instruction,
|
||||
)
|
||||
|
||||
summary_text = self.format_callback(summary_obj)
|
||||
|
||||
new_summary_msg = LLMMessage.assistant_text_response(
|
||||
f"【历史状态摘要记忆】\n{summary_text}"
|
||||
)
|
||||
new_summary_msg.metadata = {"is_summary": True, "pinned": True}
|
||||
return new_summary_msg
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"[{self.__class__.__name__}] 结构化总结失败,已跳过本次压缩: {e}"
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
class CondenserPipeline:
|
||||
"""上下文压缩流水线:按顺序依次执行各阶段的记忆压缩减项。"""
|
||||
|
||||
def __init__(self, reducers: list[BaseMemoryReducer]):
|
||||
"""
|
||||
初始化上下文压缩流水线。
|
||||
|
||||
参数:
|
||||
reducers: 压缩减项器列表,将按顺序对记忆进行多阶段修剪和压缩。
|
||||
"""
|
||||
self.reducers = reducers
|
||||
|
||||
@classmethod
|
||||
def create_from_configs(
|
||||
cls, memory_config: Any, model_name: str
|
||||
) -> "CondenserPipeline":
|
||||
"""基于全局和局部配置组装压缩管线工厂方法"""
|
||||
from zhenxun.services.ai.config import get_llm_config
|
||||
from zhenxun.services.ai.llm.system.capabilities import get_model_capabilities
|
||||
|
||||
config = get_llm_config().context_settings
|
||||
pipeline_reducers = []
|
||||
caps = get_model_capabilities(model_name)
|
||||
|
||||
vw = config.vision_window_size
|
||||
if memory_config and memory_config.compression.vision_window is not None:
|
||||
vw = memory_config.compression.vision_window
|
||||
if vw > 0:
|
||||
pipeline_reducers.append(MultimodalPlaceholderReducer(window_size=vw))
|
||||
|
||||
tp = config.tool_pruning
|
||||
if tp.enable:
|
||||
tp_limit = (
|
||||
int(caps.max_input_tokens * tp.trigger_threshold)
|
||||
if tp.trigger_threshold <= 1.0
|
||||
else int(tp.trigger_threshold)
|
||||
)
|
||||
pipeline_reducers.append(
|
||||
ToolPrunerReducer(
|
||||
keep_recent_turns=tp.keep_recent_turns,
|
||||
trigger_tokens=tp_limit,
|
||||
max_turns=tp.max_history_turns,
|
||||
)
|
||||
)
|
||||
|
||||
policy = memory_config.compression.policy if memory_config else None
|
||||
if policy is not None:
|
||||
pipeline_reducers.extend(policy)
|
||||
else:
|
||||
threshold = config.llm_summary.trigger_threshold
|
||||
if memory_config and memory_config.compression.threshold is not None:
|
||||
threshold = memory_config.compression.threshold
|
||||
|
||||
limit = (
|
||||
int(caps.max_input_tokens * threshold)
|
||||
if threshold <= 1.0
|
||||
else int(threshold)
|
||||
)
|
||||
|
||||
max_turns = config.llm_summary.max_history_turns
|
||||
if (
|
||||
memory_config
|
||||
and memory_config.compression.max_history_turns is not None
|
||||
):
|
||||
max_turns = memory_config.compression.max_history_turns
|
||||
|
||||
if config.llm_summary.enable:
|
||||
pipeline_reducers.extend(
|
||||
MemoryPolicy.llm_summarize(
|
||||
trigger_tokens=limit,
|
||||
max_turns=max_turns,
|
||||
keep_recent_turns=config.llm_summary.keep_recent_turns,
|
||||
summarization_model=config.llm_summary.summarization_model,
|
||||
summarization_prompt=config.llm_summary.summarization_prompt,
|
||||
)
|
||||
)
|
||||
else:
|
||||
pipeline_reducers.extend(MemoryPolicy.unlimited())
|
||||
|
||||
return cls(pipeline_reducers)
|
||||
|
||||
async def run(
|
||||
self, messages, model_name, base_overhead=0
|
||||
) -> tuple[list[LLMMessage], bool]:
|
||||
current_tokens = token_counter.count_context(
|
||||
messages, model_name, base_overhead
|
||||
)
|
||||
|
||||
current_messages = messages
|
||||
any_changed = False
|
||||
for reducer in self.reducers:
|
||||
current_messages, changed, current_tokens = await reducer.reduce(
|
||||
current_messages,
|
||||
current_tokens,
|
||||
model_name,
|
||||
base_overhead,
|
||||
)
|
||||
if changed:
|
||||
any_changed = True
|
||||
return current_messages, any_changed
|
||||
|
||||
|
||||
class MemoryPolicy:
|
||||
"""
|
||||
记忆策略工厂 (Strategy Factory Facade)。
|
||||
为开发者提供开箱即用的上下文压缩管线组装方案。
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def unlimited() -> list[BaseMemoryReducer]:
|
||||
"""无限制模式。不进行任何形式的截断和总结,适用于短对话或纯 Agent 内部流转。"""
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def llm_summarize(
|
||||
trigger_tokens: int = 4000,
|
||||
max_turns: int | None = None,
|
||||
keep_recent_turns: int = 0,
|
||||
summarization_model: str | None = None,
|
||||
summarization_prompt: str = (
|
||||
"请概括以下对话内容,保留关键的约束条件、用户偏好、"
|
||||
"已完成的任务状态和未解决的问题。"
|
||||
),
|
||||
) -> list[BaseMemoryReducer]:
|
||||
"""LLM 总结压缩模式。Token 达标后,自动将历史对话合并为一段 Summary。"""
|
||||
return [
|
||||
LLMSummarizerReducer(
|
||||
keep_recent_turns=keep_recent_turns,
|
||||
trigger_tokens=trigger_tokens,
|
||||
max_turns=max_turns,
|
||||
summarization_model=summarization_model,
|
||||
summarization_prompt=summarization_prompt,
|
||||
),
|
||||
MessageDropper(trigger_tokens=trigger_tokens),
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def structured_summarize(
|
||||
trigger_tokens: int = 4000,
|
||||
max_turns: int | None = None,
|
||||
keep_recent_turns: int = 0,
|
||||
summarization_model: str | None = None,
|
||||
response_model: type[BaseModel] | None = None,
|
||||
prompt_template: str | None = None,
|
||||
format_callback: Callable[[Any], str] | None = None,
|
||||
) -> list[BaseMemoryReducer]:
|
||||
"""结构化总结压缩模式。使用 JSON Schema 强制大模型提取核心状态。"""
|
||||
|
||||
class DefaultStateSummary(BaseModel):
|
||||
user_context: str = Field(
|
||||
description="用户的核心意图、诉求、人设或长期记忆规则。"
|
||||
)
|
||||
completed_tasks: str = Field(description="已完成的操作或已经确认的情节。")
|
||||
pending_tasks: str = Field(description="正在进行中的任务或尚未解答的问题。")
|
||||
current_state: str = Field(
|
||||
description="当前状态,如重要变量、玩家血量、关键物品坐标等。"
|
||||
)
|
||||
|
||||
def default_format(obj: DefaultStateSummary) -> str:
|
||||
return (
|
||||
f"👤 用户上下文: {obj.user_context}\n"
|
||||
f"✅ 已完成/确认: {obj.completed_tasks}\n"
|
||||
f"⏳ 待处理/疑问: {obj.pending_tasks}\n"
|
||||
f"📌 当前状态: {obj.current_state}"
|
||||
)
|
||||
|
||||
default_prompt = (
|
||||
"你是一个专门用于长上下文状态压缩的引擎。请阅读以下先前的总结和旧对话,"
|
||||
"提取核心状态信息,并合并它们。\n\n"
|
||||
"<之前的状态摘要>\n{prev_summary}\n</之前的状态摘要>\n\n"
|
||||
"<需要合并的旧对话记录>\n"
|
||||
"{dialogue}"
|
||||
"</需要合并的旧对话记录>\n"
|
||||
)
|
||||
|
||||
return [
|
||||
StructuredSummaryReducer(
|
||||
response_model=response_model or DefaultStateSummary,
|
||||
prompt_template=prompt_template or default_prompt,
|
||||
format_callback=format_callback or default_format,
|
||||
keep_recent_turns=keep_recent_turns,
|
||||
trigger_tokens=trigger_tokens,
|
||||
max_turns=max_turns,
|
||||
summarization_model=summarization_model,
|
||||
),
|
||||
MessageDropper(trigger_tokens=trigger_tokens),
|
||||
]
|
||||
@@ -0,0 +1,260 @@
|
||||
from collections.abc import Sequence
|
||||
from typing import Any, cast
|
||||
|
||||
from zhenxun.services.ai.context.memory.compression import (
|
||||
CondenserPipeline,
|
||||
)
|
||||
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 SessionMetadata
|
||||
from zhenxun.services.ai.core.engine.context_renderer import ContextConverter
|
||||
from zhenxun.services.ai.core.messages import AgentMessage
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.pydantic_compat import model_copy
|
||||
|
||||
|
||||
class MemoryReader:
|
||||
"""
|
||||
记忆读取器 (Memory Reader)。
|
||||
负责从数据库中提取短期上下文历史,召回长期的背景知识,并执行自动压缩。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, session_meta: SessionMetadata, memory_config: MemoryConfig | None
|
||||
):
|
||||
"""
|
||||
初始化记忆读取器。
|
||||
|
||||
参数:
|
||||
session_meta: 会话元数据,包含 Namespace 与作用域映射等上下文信息。
|
||||
memory_config: 记忆系统的配置对象,控制长期、短期及槽位记忆的启用与逻辑。
|
||||
"""
|
||||
self.session_meta = session_meta
|
||||
self.memory_config = memory_config
|
||||
|
||||
async def get_long_term_context(self, user_input: str) -> str:
|
||||
"""
|
||||
基于用户输入召回长期记忆(RAG),返回格式化后的背景提示词。
|
||||
"""
|
||||
if (
|
||||
not self.memory_config
|
||||
or not self.memory_config.long_term.enable
|
||||
or not user_input
|
||||
):
|
||||
return ""
|
||||
|
||||
policy = self.memory_config.long_term.auto_recall
|
||||
should_recall = False
|
||||
|
||||
if isinstance(policy, bool):
|
||||
should_recall = policy
|
||||
elif callable(policy):
|
||||
import inspect
|
||||
|
||||
try:
|
||||
res = policy(user_input, self.session_meta)
|
||||
if inspect.isawaitable(res):
|
||||
should_recall = await res
|
||||
else:
|
||||
should_recall = bool(res)
|
||||
except Exception as e:
|
||||
logger.error(f"[MemoryReader] 自定义 auto_recall 函数执行失败: {e}")
|
||||
should_recall = False
|
||||
|
||||
if not should_recall:
|
||||
return ""
|
||||
|
||||
ltm_scope = memory_manager.get_long_term_memory(
|
||||
self.memory_config,
|
||||
namespace=self.session_meta.selector.namespace or "global",
|
||||
)
|
||||
if not ltm_scope:
|
||||
return ""
|
||||
|
||||
matches = await ltm_scope.recall(session=self.session_meta, query=user_input)
|
||||
if matches:
|
||||
logger.debug(f"🧠 [MemoryReader] 长期记忆召回详情 (Query: '{user_input}'):")
|
||||
for i, m in enumerate(matches):
|
||||
logger.debug(
|
||||
f" [{i + 1}] 得分: {m.score:.4f} | 内容: {m.record.content}"
|
||||
)
|
||||
|
||||
threshold = self.memory_config.long_term.recall_threshold
|
||||
valid_matches = [m for m in matches if m.score >= threshold]
|
||||
if not valid_matches:
|
||||
logger.debug("🧠 [MemoryReader] 召回的记忆均未达到相关性阈值,已丢弃。")
|
||||
return ""
|
||||
|
||||
fact_str = "\n".join(f"- {m.record.content}" for m in valid_matches)
|
||||
logger.debug(
|
||||
f"🧠 [MemoryReader]"
|
||||
f"成功截取并注入 {len(valid_matches)} 条高价值长期记忆。"
|
||||
)
|
||||
return f"[系统补充:有关用户的长期记忆设定]\n{fact_str}"
|
||||
return ""
|
||||
|
||||
async def get_slots_context(self) -> str:
|
||||
"""
|
||||
读取并组装核心槽位记忆 (Memory Slots),返回 XML 格式字符串供大模型使用。
|
||||
"""
|
||||
if not self.memory_config or not self.memory_config.slots.enable:
|
||||
return ""
|
||||
slot_ctx = memory_manager.get_slot_context(
|
||||
self.memory_config,
|
||||
namespace=self.session_meta.selector.namespace or "global",
|
||||
)
|
||||
if not slot_ctx:
|
||||
return ""
|
||||
|
||||
if self.memory_config.slots.default_slots:
|
||||
for default_slot in self.memory_config.slots.default_slots:
|
||||
existing = await slot_ctx.get_slot(
|
||||
self.session_meta, default_slot.label
|
||||
)
|
||||
if not existing:
|
||||
await slot_ctx.set_slot(self.session_meta, default_slot)
|
||||
|
||||
slots = await slot_ctx.list_pinned_slots(self.session_meta)
|
||||
if not slots:
|
||||
return ""
|
||||
|
||||
show_scope = False
|
||||
if (
|
||||
self.memory_config
|
||||
and self.memory_config.slots.scopes
|
||||
and len(self.memory_config.slots.scopes) > 1
|
||||
):
|
||||
show_scope = True
|
||||
|
||||
xml_parts = ["<memory_slots>"]
|
||||
for slot in slots:
|
||||
if show_scope:
|
||||
semantic_name = self.session_meta.scope_name_mapping.get(
|
||||
slot.scope, "未知"
|
||||
)
|
||||
xml_parts.append(
|
||||
f' <slot name="{slot.label}" scope="{semantic_name}">\n'
|
||||
f" {slot.content}\n"
|
||||
" </slot>"
|
||||
)
|
||||
else:
|
||||
xml_parts.append(
|
||||
f' <slot name="{slot.label}">\n {slot.content}\n </slot>'
|
||||
)
|
||||
xml_parts.append("</memory_slots>")
|
||||
return "\n".join(xml_parts)
|
||||
|
||||
async def get_short_term_context(
|
||||
self,
|
||||
model_name: str,
|
||||
override_history: Sequence[AgentMessage] | None = None,
|
||||
) -> list[AgentMessage]:
|
||||
"""
|
||||
拉取短期对话历史,并执行 Token 压缩。
|
||||
"""
|
||||
current_history: list[AgentMessage] = []
|
||||
if override_history is not None:
|
||||
current_history = list(override_history)
|
||||
|
||||
chat_context = memory_manager.get_chat_context(
|
||||
self.memory_config,
|
||||
namespace=self.session_meta.selector.namespace or "global",
|
||||
)
|
||||
|
||||
if self.memory_config and self.memory_config.short_term.enable and chat_context:
|
||||
if override_history is not None:
|
||||
flattened_override = ContextConverter.flatten_to_llm_messages(
|
||||
override_history
|
||||
)
|
||||
await chat_context.set_messages(self.session_meta, flattened_override)
|
||||
else:
|
||||
current_history = cast(
|
||||
list[AgentMessage],
|
||||
await chat_context.get_messages(self.session_meta),
|
||||
)
|
||||
|
||||
pipeline = CondenserPipeline.create_from_configs(
|
||||
self.memory_config, model_name
|
||||
)
|
||||
if pipeline.reducers:
|
||||
flattened_to_reduce = ContextConverter.flatten_to_llm_messages(
|
||||
current_history
|
||||
)
|
||||
new_history, changed = await pipeline.run(
|
||||
flattened_to_reduce, model_name=model_name, base_overhead=0
|
||||
)
|
||||
if changed:
|
||||
await chat_context.set_messages(self.session_meta, new_history)
|
||||
logger.info(
|
||||
"💾 [MemoryReader] 压缩截断完毕,已同步覆写数据库。"
|
||||
f"压缩后条数: {len(new_history)}"
|
||||
)
|
||||
current_history = cast(list[AgentMessage], new_history)
|
||||
|
||||
return current_history
|
||||
|
||||
|
||||
class MemoryWriter:
|
||||
"""
|
||||
记忆写入器 (Memory Writer)。
|
||||
负责将对话增量安全地写入数据库。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
session_meta: SessionMetadata,
|
||||
memory_config: MemoryConfig | None,
|
||||
context: Any = None,
|
||||
):
|
||||
"""
|
||||
初始化记忆写入器。
|
||||
|
||||
参数:
|
||||
session_meta: 会话元数据,包含 Namespace 与作用域映射等上下文信息。
|
||||
memory_config: 记忆系统的配置对象,控制记忆存入的逻辑。
|
||||
context: 运行时上下文环境,作为可选参数传入,供中间件使用,默认 None。
|
||||
"""
|
||||
self.session_meta = session_meta
|
||||
self.memory_config = memory_config
|
||||
self.context = context
|
||||
|
||||
async def save_new_messages(
|
||||
self,
|
||||
new_messages: Sequence[AgentMessage],
|
||||
):
|
||||
"""将新产生的对话增量保存到数据库"""
|
||||
if not new_messages:
|
||||
return
|
||||
|
||||
messages_to_save = new_messages
|
||||
|
||||
if self.memory_config and self.memory_config.ingestion.middlewares:
|
||||
messages_to_save = [model_copy(m, deep=True) for m in new_messages]
|
||||
|
||||
for middleware in self.memory_config.ingestion.middlewares:
|
||||
try:
|
||||
messages_to_save = await middleware.process(
|
||||
messages_to_save, self.context
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"[MemoryIngestion] 中间件 {middleware.__class__.__name__} "
|
||||
f"执行失败: {e}",
|
||||
e=e,
|
||||
)
|
||||
|
||||
if not messages_to_save:
|
||||
return
|
||||
|
||||
chat_ctx = memory_manager.get_chat_context(
|
||||
self.memory_config,
|
||||
namespace=self.session_meta.selector.namespace or "global",
|
||||
)
|
||||
|
||||
flattened_msgs = ContextConverter.flatten_to_llm_messages(
|
||||
messages_to_save, self.context
|
||||
)
|
||||
|
||||
if chat_ctx and self.memory_config and self.memory_config.short_term.enable:
|
||||
if flattened_msgs:
|
||||
await chat_ctx.add_messages(self.session_meta, flattened_msgs)
|
||||
@@ -0,0 +1,128 @@
|
||||
from collections.abc import Sequence
|
||||
from typing import TYPE_CHECKING, Literal
|
||||
|
||||
from zhenxun.services.ai.context.memory.types import (
|
||||
MemorySlot,
|
||||
SessionMetadata,
|
||||
)
|
||||
from zhenxun.services.ai.core.messages import AgentMessage, LLMMessage
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from zhenxun.services.ai.context.memory.manager import GlobalMemoryManager
|
||||
from zhenxun.services.ai.context.memory.storage.interfaces import (
|
||||
BaseChatContext,
|
||||
BaseSlotContext,
|
||||
)
|
||||
|
||||
|
||||
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()
|
||||
@@ -0,0 +1,211 @@
|
||||
from collections.abc import Callable
|
||||
from typing import Any, cast
|
||||
|
||||
from zhenxun.services.ai.context.memory.models import MemoryConfig
|
||||
from zhenxun.services.ai.context.memory.storage.backends import (
|
||||
InMemoryChatContext,
|
||||
MemoryScope,
|
||||
)
|
||||
from zhenxun.services.ai.context.memory.storage.interfaces import (
|
||||
BaseChatContext,
|
||||
BaseSlotContext,
|
||||
)
|
||||
from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend
|
||||
from zhenxun.services.ai.utils.scope import BaseScopeBuilder
|
||||
from zhenxun.utils.utils import infer_plugin_namespace
|
||||
|
||||
|
||||
class MemoryCleaner(BaseScopeBuilder["MemoryCleaner"]):
|
||||
"""
|
||||
声明式记忆清理构建器 (Query Builder)。
|
||||
为第三方开发者提供极端友好的链式 API,彻底屏蔽底层前缀逻辑。
|
||||
"""
|
||||
|
||||
def __init__(self, manager: "GlobalMemoryManager"):
|
||||
super().__init__()
|
||||
self.manager = manager
|
||||
self._config: Any = None
|
||||
|
||||
def config(self, cfg: Any):
|
||||
"""指定私有记忆配置(自动识别未全局注册 of 第三方私有数据库实例)"""
|
||||
self._config = cfg.build() if hasattr(cfg, "build") else cfg
|
||||
return self
|
||||
|
||||
async def clear_short_term(self):
|
||||
"""一键清理目标范围下的短期对话历史记忆"""
|
||||
if self._config and self._config.short_term.backend:
|
||||
await self._config.short_term.backend.clear_by_query(self._selector)
|
||||
else:
|
||||
for backend in self.manager._chat_backends.values():
|
||||
await backend.clear_by_query(self._selector)
|
||||
|
||||
async def clear_slots(self):
|
||||
"""一键清理目标范围下的中期记忆槽 (Memory Slots)"""
|
||||
if self._config and self._config.slots.backend:
|
||||
await self._config.slots.backend.clear_by_query(self._selector)
|
||||
else:
|
||||
for backend in self.manager._slot_backends.values():
|
||||
await backend.clear_by_query(self._selector)
|
||||
|
||||
async def clear_long_term(self):
|
||||
"""一键清理目标范围下的长期向量记忆 (RAG Vector Database)"""
|
||||
if self._config and self._config.long_term.backend:
|
||||
from zhenxun.services.ai.context.rag.backends import StorageBackend
|
||||
|
||||
storage = cast(StorageBackend, self._config.long_term.backend)
|
||||
await storage.clear_by_query(self._selector)
|
||||
else:
|
||||
for factory in self.manager._storage_factories.values():
|
||||
storage = factory()
|
||||
if hasattr(storage, "clear_by_query"):
|
||||
await storage.clear_by_query(self._selector)
|
||||
else:
|
||||
await storage.delete(scope_prefix=self._selector.scope_prefix)
|
||||
|
||||
async def clear_all(self):
|
||||
"""一键清理指定范围下的所有生命周期记忆(对话、槽位、RAG)"""
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
await self.clear_short_term()
|
||||
await self.clear_slots()
|
||||
await self.clear_long_term()
|
||||
logger.info(
|
||||
f"🧹 [MemoryCleaner] 成功清理作用域 '{self._selector.scope_prefix}'"
|
||||
"下的所有记忆痕迹!"
|
||||
)
|
||||
|
||||
|
||||
class GlobalMemoryManager:
|
||||
"""
|
||||
全局记忆大管家 (IoC 容器)。
|
||||
使用现代化依赖注入机制管理短/长期记忆引擎的默认实例。
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._chat_backends: dict[str, BaseChatContext] = {
|
||||
"global": InMemoryChatContext()
|
||||
}
|
||||
self._slot_backends: dict[str, BaseSlotContext] = {}
|
||||
|
||||
from zhenxun.services.ai.context.rag.backends import DictStorageBackend
|
||||
|
||||
self._storage_factories: dict[str, Callable[[], StorageBackend]] = {
|
||||
"global": lambda: DictStorageBackend()
|
||||
}
|
||||
|
||||
def register_chat_backend(
|
||||
self, backend: BaseChatContext, scope: str | None = None
|
||||
) -> None:
|
||||
"""注册特定命名空间的短期记忆存储后端。"""
|
||||
ns = scope if scope is not None else infer_plugin_namespace()
|
||||
self._chat_backends[ns] = backend
|
||||
|
||||
def register_slot_backend(
|
||||
self, backend: BaseSlotContext, scope: str | None = None
|
||||
) -> None:
|
||||
"""注册特定命名空间的中期记忆槽存储后端。"""
|
||||
ns = scope if scope is not None else infer_plugin_namespace()
|
||||
self._slot_backends[ns] = backend
|
||||
|
||||
def register_storage_factory(
|
||||
self, factory: Callable[[], StorageBackend], scope: str | None = None
|
||||
) -> None:
|
||||
"""注册特定命名空间的长期记忆向量存储工厂。"""
|
||||
ns = scope if scope is not None else infer_plugin_namespace()
|
||||
self._storage_factories[ns] = factory
|
||||
|
||||
def cleaner(self) -> MemoryCleaner:
|
||||
"""获取声明式记忆清理构建器,供第三方开发者极速清理指定记忆"""
|
||||
return MemoryCleaner(self)
|
||||
|
||||
def get_embedder(self, embedder_val: "Embedder | str | None") -> Embedder | None:
|
||||
"""获取向量化引擎实例。如果传入的是字符串,则视为 API 模型名称。"""
|
||||
if not embedder_val:
|
||||
return None
|
||||
|
||||
if isinstance(embedder_val, str):
|
||||
from zhenxun.services.ai.context.rag.backends.embedders import (
|
||||
DefaultEmbedder,
|
||||
)
|
||||
|
||||
return DefaultEmbedder(model_name=embedder_val)
|
||||
|
||||
return embedder_val
|
||||
|
||||
def get_chat_context(
|
||||
self, config: MemoryConfig | None, namespace: str = "global"
|
||||
) -> BaseChatContext | None:
|
||||
"""根据配置分配对应的短期对话历史实例"""
|
||||
if not config or not config.short_term.enable:
|
||||
return None
|
||||
|
||||
backend_cfg = config.short_term.backend
|
||||
if backend_cfg is not None:
|
||||
return cast(BaseChatContext, backend_cfg)
|
||||
|
||||
return self._chat_backends.get(namespace) or self._chat_backends["global"]
|
||||
|
||||
def get_slot_context(
|
||||
self, config: MemoryConfig | None, namespace: str = "global"
|
||||
) -> BaseSlotContext | None:
|
||||
"""根据配置分配对应的槽位记忆实例"""
|
||||
if not config or not config.slots.enable:
|
||||
return None
|
||||
|
||||
backend_cfg = config.slots.backend
|
||||
if backend_cfg is not None:
|
||||
return cast(BaseSlotContext, backend_cfg)
|
||||
|
||||
return self._slot_backends.get(namespace) or self._slot_backends["global"]
|
||||
|
||||
def get_long_term_memory(
|
||||
self, config: MemoryConfig | None, namespace: str = "global"
|
||||
) -> MemoryScope | None:
|
||||
"""根据声明式配置动态组装长期向量记忆实例"""
|
||||
if not config or not config.long_term.enable:
|
||||
return None
|
||||
|
||||
if config.long_term.engine is not None:
|
||||
return MemoryScope(
|
||||
rag_client=config.long_term.engine,
|
||||
)
|
||||
|
||||
storage_instance = None
|
||||
backend_cfg = config.long_term.backend
|
||||
if backend_cfg is not None:
|
||||
storage_instance = cast(StorageBackend, backend_cfg)
|
||||
else:
|
||||
factory = (
|
||||
self._storage_factories.get(namespace)
|
||||
or self._storage_factories["global"]
|
||||
)
|
||||
storage_instance = factory()
|
||||
|
||||
embedder = self.get_embedder(config.long_term.embedder)
|
||||
|
||||
from zhenxun.services.ai.context.rag.builder import RAGBuilder
|
||||
|
||||
builder = RAGBuilder(storage_instance).with_scope("/")
|
||||
if embedder:
|
||||
builder.with_embedder(embedder)
|
||||
|
||||
from zhenxun.services.ai.context.memory.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()
|
||||
@@ -0,0 +1,173 @@
|
||||
"""
|
||||
记忆域类型定义
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from zhenxun.services.ai.context.memory.storage.interfaces import (
|
||||
BaseChatContext,
|
||||
BaseMemoryIngestionMiddleware,
|
||||
BaseMemoryReducer,
|
||||
BaseSlotContext,
|
||||
)
|
||||
from zhenxun.services.ai.context.memory.types import (
|
||||
AutoRecallPolicy,
|
||||
Isolation,
|
||||
MemorySlot,
|
||||
SessionMetadata,
|
||||
)
|
||||
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 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):
|
||||
"""长期记忆的复合打分与检索配置"""
|
||||
|
||||
recency_weight: float = Field(default=0.3)
|
||||
"""时间衰减权重"""
|
||||
semantic_weight: float = Field(default=0.5)
|
||||
"""语义相似度权重"""
|
||||
importance_weight: float = Field(default=0.2)
|
||||
"""重要性权重"""
|
||||
recency_half_life_days: int = Field(default=30)
|
||||
"""时间衰减的半衰期(天)"""
|
||||
|
||||
reinforcement_weight: float = Field(default=0.2)
|
||||
"""访问强化的加权权重 (被检索越多得分越高)"""
|
||||
|
||||
|
||||
class ShortTermConfig(BaseModel):
|
||||
"""短期对话记忆配置"""
|
||||
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
enable: bool = Field(default=True)
|
||||
"""是否启用短期对话记忆上下文"""
|
||||
backend: str | BaseChatContext | None = Field(default=None)
|
||||
"""
|
||||
指定底层短期记忆数据库注册名称,或直接传入 BaseChatContext 实例。
|
||||
为空则使用全局默认
|
||||
"""
|
||||
isolation: ScopeBuilder = Field(default_factory=Isolation.AGENT_USER)
|
||||
"""单一的记忆隔离级别 (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):
|
||||
"""上下文压缩与管理配置"""
|
||||
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
threshold: float | None = Field(default=None)
|
||||
"""(局部重写) 触发记忆压缩的 Token 阈值"""
|
||||
max_history_turns: int | None = Field(default=None)
|
||||
"""(局部重写) 触发记忆压缩的对话轮数上限。设为 0 表示不限制轮数。"""
|
||||
vision_window: int | None = Field(default=None)
|
||||
"""多模态滑动窗口大小。0表示关闭该功能,>0表示仅保留最近N轮包含多模态数据的消息,None表示跟随全局配置。"""
|
||||
policy: list[BaseMemoryReducer] | None = Field(default=None)
|
||||
"""核心记忆压缩策略管线 (List[BaseMemoryReducer])。为 None 时将应用全局默认策略。"""
|
||||
|
||||
|
||||
class IngestionConfig(BaseModel):
|
||||
"""记忆入库管线配置"""
|
||||
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
middlewares: list[BaseMemoryIngestionMiddleware] = Field(default_factory=list)
|
||||
"""入库中间件列表(按顺序依次执行清洗过滤)"""
|
||||
|
||||
|
||||
class MemoryConfig(BaseModel):
|
||||
"""统一的记忆配置项声明"""
|
||||
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
base_isolation: ScopeBuilder = Field(default_factory=Isolation.AGENT_USER)
|
||||
"""顶层基准隔离级别,短期/中期/长期记忆将默认继承此级别"""
|
||||
short_term: ShortTermConfig = Field(default_factory=ShortTermConfig)
|
||||
"""短期对话记忆配置"""
|
||||
slots: SlotMemoryConfig = Field(default_factory=SlotMemoryConfig)
|
||||
"""槽位记忆配置"""
|
||||
long_term: LongTermConfig = Field(default_factory=LongTermConfig)
|
||||
"""长期向量记忆配置"""
|
||||
compression: ContextCompressionConfig = Field(
|
||||
default_factory=ContextCompressionConfig
|
||||
)
|
||||
"""上下文压缩与管理配置"""
|
||||
ingestion: IngestionConfig = Field(default_factory=IngestionConfig)
|
||||
"""记忆入库前的清洗与过滤管线配置"""
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AutoRecallPolicy",
|
||||
"BaseMemoryIngestionMiddleware",
|
||||
"ContextCompressionConfig",
|
||||
"IngestionConfig",
|
||||
"Isolation",
|
||||
"LongTermConfig",
|
||||
"MemoryConfig",
|
||||
"MemoryScoringConfig",
|
||||
"SessionMetadata",
|
||||
"ShortTermConfig",
|
||||
]
|
||||
@@ -0,0 +1,21 @@
|
||||
from .backends import (
|
||||
AbstractMemoryRecord,
|
||||
AbstractSlotRecord,
|
||||
InMemoryChatContext,
|
||||
MemoryScope,
|
||||
TortoiseChatContext,
|
||||
TortoiseSlotContext,
|
||||
get_orm_chat_context,
|
||||
get_orm_slot_context,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"AbstractMemoryRecord",
|
||||
"AbstractSlotRecord",
|
||||
"InMemoryChatContext",
|
||||
"MemoryScope",
|
||||
"TortoiseChatContext",
|
||||
"TortoiseSlotContext",
|
||||
"get_orm_chat_context",
|
||||
"get_orm_slot_context",
|
||||
]
|
||||
@@ -0,0 +1,465 @@
|
||||
import asyncio
|
||||
import base64
|
||||
from collections.abc import Callable
|
||||
import datetime
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from zhenxun.services.ai.context.rag.engine import ScopedRAGClient
|
||||
|
||||
from nonebot.utils import is_coroutine_callable
|
||||
from tortoise import fields
|
||||
from tortoise.timezone import now
|
||||
|
||||
from zhenxun.services.ai.context.memory.storage.interfaces import (
|
||||
BaseChatContext,
|
||||
BaseSlotContext,
|
||||
)
|
||||
from zhenxun.services.ai.context.memory.types import (
|
||||
MemorySlot,
|
||||
SessionMetadata,
|
||||
)
|
||||
from zhenxun.services.ai.context.rag.models import BaseRecord, SearchResult
|
||||
from zhenxun.services.ai.core.messages import (
|
||||
AssistantMessage,
|
||||
LLMContentPart,
|
||||
LLMMessage,
|
||||
SystemMessage,
|
||||
ToolMessage,
|
||||
UserMessage,
|
||||
)
|
||||
from zhenxun.services.ai.utils.scope import ScopeSelector
|
||||
from zhenxun.services.db_context import Model
|
||||
from zhenxun.utils.pydantic_compat import TypeAdapter, model_dump
|
||||
|
||||
|
||||
class DBMessageSerializer:
|
||||
"""将 LLMMessage 与数据库 JSON 格式进行序列化/反序列化的帮助类"""
|
||||
|
||||
@staticmethod
|
||||
def deserialize_content(content_raw: Any) -> list[LLMContentPart]:
|
||||
from zhenxun.services.ai.core.messages import TextPart
|
||||
|
||||
content_parts: list[LLMContentPart] = []
|
||||
if isinstance(content_raw, list):
|
||||
adapter = TypeAdapter(LLMContentPart)
|
||||
for p in content_raw:
|
||||
if isinstance(p, dict):
|
||||
for k in list(p.keys()):
|
||||
if k.startswith("_is_b64_"):
|
||||
orig_k = k[8:]
|
||||
if orig_k in p and isinstance(p[orig_k], str):
|
||||
p[orig_k] = base64.b64decode(p[orig_k])
|
||||
p.pop(k, None)
|
||||
content_parts.append(adapter.validate_python(p))
|
||||
elif isinstance(content_raw, str):
|
||||
content_parts.append(TextPart(text=content_raw))
|
||||
return content_parts
|
||||
|
||||
@staticmethod
|
||||
def serialize_content(content_payload: Any) -> list[dict[str, Any]]:
|
||||
from pathlib import Path
|
||||
|
||||
if isinstance(content_payload, str):
|
||||
return [{"type": "text", "text": content_payload}]
|
||||
elif isinstance(content_payload, list):
|
||||
processed_content = []
|
||||
for p in content_payload:
|
||||
p_dump = (
|
||||
model_dump(p, exclude_none=True)
|
||||
if hasattr(p, "model_dump")
|
||||
else (p.copy() if isinstance(p, dict) else p)
|
||||
)
|
||||
if isinstance(p_dump, dict):
|
||||
for k, v in list(p_dump.items()):
|
||||
if isinstance(v, bytes):
|
||||
p_dump[k] = base64.b64encode(v).decode("utf-8")
|
||||
p_dump[f"_is_b64_{k}"] = True
|
||||
elif isinstance(v, Path):
|
||||
p_dump[k] = str(v)
|
||||
processed_content.append(p_dump)
|
||||
return (
|
||||
processed_content
|
||||
if processed_content
|
||||
else [
|
||||
{"type": "text", "text": "[仅包含思维链或工具调度,无实质文本输出]"}
|
||||
]
|
||||
)
|
||||
return []
|
||||
|
||||
|
||||
class MemoryScope:
|
||||
"""长期记忆的作用域视图与 RAG 管线。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
rag_client: "ScopedRAGClient",
|
||||
):
|
||||
self.rag_client = rag_client
|
||||
self._background_tasks: set[Any] = set()
|
||||
|
||||
async def remember(
|
||||
self,
|
||||
session: SessionMetadata,
|
||||
content: str,
|
||||
importance: float = 0.5,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
"""通过 RAG Ingestion Pipeline 完成记忆落盘"""
|
||||
meta = metadata.copy() if metadata else {}
|
||||
meta.update(
|
||||
{
|
||||
"scope": session.scope_prefix,
|
||||
"importance": importance,
|
||||
"created_at": time.time(),
|
||||
}
|
||||
)
|
||||
record = BaseRecord(content=content, metadata=meta)
|
||||
|
||||
await self.rag_client.ingest([record])
|
||||
|
||||
async def recall(
|
||||
self,
|
||||
session: SessionMetadata,
|
||||
query: str,
|
||||
limit: int = 10,
|
||||
metadata_filter: dict[str, Any] | None = None,
|
||||
) -> list[SearchResult]:
|
||||
"""委托至 Retriever 检索与重排,并触发读时惰性强化"""
|
||||
matches = await self.rag_client.search(
|
||||
query=query,
|
||||
limit=limit,
|
||||
scopes=session.accessible_scopes,
|
||||
metadata_filters=metadata_filter,
|
||||
)
|
||||
if matches:
|
||||
task = asyncio.create_task(
|
||||
self._reinforce_memories([m.record for m in matches])
|
||||
)
|
||||
self._background_tasks.add(task)
|
||||
task.add_done_callback(self._background_tasks.discard)
|
||||
return matches
|
||||
|
||||
async def update(
|
||||
self,
|
||||
session: SessionMetadata,
|
||||
record_id: str,
|
||||
new_content: str,
|
||||
importance: float = 0.5,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
) -> bool:
|
||||
"""原子更新:通过先删后插,确保底层向量(Embedding)能根据新文本被正确刷新"""
|
||||
deleted_count = await self.forget(session, record_ids=[record_id])
|
||||
if deleted_count > 0:
|
||||
await self.remember(
|
||||
session=session,
|
||||
content=new_content,
|
||||
importance=importance,
|
||||
metadata=metadata,
|
||||
)
|
||||
return True
|
||||
return False
|
||||
|
||||
async def forget(
|
||||
self, session: SessionMetadata, record_ids: list[str] | None = None
|
||||
) -> int:
|
||||
return await self.rag_client.delete(
|
||||
record_ids=record_ids,
|
||||
)
|
||||
|
||||
async def _reinforce_memories(self, records: list[BaseRecord]):
|
||||
import time
|
||||
|
||||
now = time.time()
|
||||
for r in records:
|
||||
r.metadata["access_count"] = r.metadata.get("access_count", 0) + 1
|
||||
r.metadata["last_accessed_at"] = now
|
||||
await self.rag_client.storage.update(r)
|
||||
|
||||
|
||||
class InMemoryChatContext(BaseChatContext):
|
||||
def __init__(self):
|
||||
self._messages: dict[str, list[LLMMessage]] = {}
|
||||
|
||||
async def get_messages(self, session: SessionMetadata) -> list[LLMMessage]:
|
||||
return list(self._messages.get(session.session_id, []))
|
||||
|
||||
async def search(
|
||||
self, query: str, session: SessionMetadata, limit: int = 10
|
||||
) -> list[LLMMessage]:
|
||||
results = []
|
||||
for msg in self._messages.get(session.session_id, []):
|
||||
if query in msg.extract_text:
|
||||
results.append(msg)
|
||||
if len(results) >= limit:
|
||||
break
|
||||
return results
|
||||
|
||||
async def add_messages(
|
||||
self, session: SessionMetadata, messages: list[LLMMessage]
|
||||
) -> None:
|
||||
if session.session_id not in self._messages:
|
||||
self._messages[session.session_id] = []
|
||||
self._messages[session.session_id].extend(messages)
|
||||
|
||||
async def set_messages(
|
||||
self, session: SessionMetadata, messages: list[LLMMessage]
|
||||
) -> None:
|
||||
self._messages[session.session_id] = list(messages)
|
||||
|
||||
async def clear(self, session: SessionMetadata) -> None:
|
||||
self._messages.pop(session.session_id, None)
|
||||
|
||||
async def clear_by_query(self, query: ScopeSelector) -> None:
|
||||
"""内存级:前缀匹配清理所有符合要求的短期会话"""
|
||||
scope_prefix = query.scope_prefix
|
||||
keys_to_delete = [
|
||||
sid for sid in self._messages.keys() if sid.startswith(scope_prefix)
|
||||
]
|
||||
for sid in keys_to_delete:
|
||||
self._messages.pop(sid, None)
|
||||
|
||||
|
||||
class AbstractMemoryRecord(Model):
|
||||
"""Tortoise ORM 短期记忆持久化基类 (Mixin)。"""
|
||||
|
||||
id = fields.UUIDField(pk=True, description="主键")
|
||||
session_id = fields.CharField(max_length=255, index=True)
|
||||
role = fields.CharField(max_length=32)
|
||||
content = fields.JSONField()
|
||||
api_context = fields.JSONField(null=True)
|
||||
created_at = fields.DatetimeField(auto_now_add=True)
|
||||
metadata = fields.JSONField(null=True)
|
||||
|
||||
class Meta: # type: ignore
|
||||
abstract = True
|
||||
|
||||
|
||||
class TortoiseChatContext(BaseChatContext):
|
||||
def __init__(
|
||||
self,
|
||||
model_class: type[AbstractMemoryRecord],
|
||||
custom_save_hook: Callable[
|
||||
[AbstractMemoryRecord, LLMMessage, SessionMetadata], Any
|
||||
]
|
||||
| None = None,
|
||||
):
|
||||
self.model_class = model_class
|
||||
self.custom_save_hook = custom_save_hook
|
||||
|
||||
def _row_to_message(self, row: AbstractMemoryRecord) -> LLMMessage:
|
||||
content_parts = DBMessageSerializer.deserialize_content(row.content)
|
||||
metadata: dict[str, Any] | None = (
|
||||
row.metadata if isinstance(row.metadata, dict) else None
|
||||
)
|
||||
kwargs = {
|
||||
"content": content_parts,
|
||||
"metadata": metadata,
|
||||
"created_at": row.created_at.timestamp() if row.created_at else time.time(),
|
||||
}
|
||||
role = row.role
|
||||
if role == "system":
|
||||
return cast(LLMMessage, SystemMessage(**kwargs))
|
||||
elif role == "user":
|
||||
return cast(LLMMessage, UserMessage(**kwargs))
|
||||
elif role == "assistant":
|
||||
return cast(LLMMessage, AssistantMessage(**kwargs))
|
||||
elif role == "tool":
|
||||
return cast(LLMMessage, ToolMessage(**kwargs))
|
||||
return cast(LLMMessage, LLMMessage(role=role, **kwargs))
|
||||
|
||||
async def get_messages(self, session: SessionMetadata) -> list[LLMMessage]:
|
||||
rows = (
|
||||
await self.model_class.filter(session_id=session.session_id)
|
||||
.order_by("created_at")
|
||||
.all()
|
||||
)
|
||||
return [self._row_to_message(row) for row in rows]
|
||||
|
||||
async def search(
|
||||
self, query: str, session: SessionMetadata, limit: int = 10
|
||||
) -> list[LLMMessage]:
|
||||
rows = (
|
||||
await self.model_class.filter(
|
||||
session_id=session.session_id, content__icontains=query
|
||||
)
|
||||
.order_by("-created_at")
|
||||
.limit(limit)
|
||||
.all()
|
||||
)
|
||||
return [self._row_to_message(row) for row in reversed(rows)]
|
||||
|
||||
async def add_messages(
|
||||
self, session: SessionMetadata, messages: list[LLMMessage]
|
||||
) -> None:
|
||||
if not messages:
|
||||
return
|
||||
|
||||
base_time = now()
|
||||
|
||||
last_msg = (
|
||||
await self.model_class.filter(session_id=session.session_id)
|
||||
.order_by("-created_at")
|
||||
.first()
|
||||
)
|
||||
if last_msg and last_msg.created_at and last_msg.created_at >= base_time:
|
||||
base_time = last_msg.created_at + datetime.timedelta(milliseconds=10)
|
||||
|
||||
orm_objects = []
|
||||
for i, msg in enumerate(messages):
|
||||
content_payload = DBMessageSerializer.serialize_content(msg.content)
|
||||
|
||||
msg_time = base_time + datetime.timedelta(milliseconds=i * 10)
|
||||
orm_obj = self.model_class(
|
||||
session_id=session.session_id,
|
||||
role=msg.role,
|
||||
content=content_payload,
|
||||
api_context=None,
|
||||
metadata=msg.metadata,
|
||||
created_at=msg_time,
|
||||
)
|
||||
if self.custom_save_hook:
|
||||
if is_coroutine_callable(self.custom_save_hook):
|
||||
await self.custom_save_hook(orm_obj, msg, session)
|
||||
else:
|
||||
self.custom_save_hook(orm_obj, msg, session)
|
||||
orm_objects.append(orm_obj)
|
||||
if orm_objects:
|
||||
await self.model_class.bulk_create(orm_objects)
|
||||
|
||||
async def set_messages(
|
||||
self, session: SessionMetadata, messages: list[LLMMessage]
|
||||
) -> None:
|
||||
await self.clear(session)
|
||||
await self.add_messages(session, messages)
|
||||
|
||||
async def clear(self, session: SessionMetadata) -> None:
|
||||
await self.model_class.filter(session_id=session.session_id).delete()
|
||||
|
||||
async def clear_by_query(self, query: ScopeSelector) -> None:
|
||||
"""ORM 级:利用数据库 startswith 原生语法批量级联删除短期记忆"""
|
||||
scope_prefix = query.scope_prefix
|
||||
await self.model_class.filter(session_id__startswith=scope_prefix).delete()
|
||||
|
||||
|
||||
def get_orm_chat_context(
|
||||
model_class: type[AbstractMemoryRecord],
|
||||
custom_save_hook: Callable[[AbstractMemoryRecord, LLMMessage, SessionMetadata], Any]
|
||||
| None = None,
|
||||
) -> TortoiseChatContext:
|
||||
"""
|
||||
[工厂方法] 供第三方开发者调用,
|
||||
将 Tortoise ORM 表直接包装为对话历史记录系统。
|
||||
"""
|
||||
return TortoiseChatContext(
|
||||
model_class=model_class, custom_save_hook=custom_save_hook
|
||||
)
|
||||
|
||||
|
||||
class AbstractSlotRecord(Model):
|
||||
"""Tortoise ORM 记忆槽持久化基类 (Mixin)。"""
|
||||
|
||||
id = fields.CharField(
|
||||
pk=True, max_length=128, description="复合主键: session_id + label"
|
||||
)
|
||||
session_id = fields.CharField(max_length=255, index=True)
|
||||
label = fields.CharField(max_length=64, index=True)
|
||||
content = fields.TextField()
|
||||
size_limit = fields.IntField(default=2000)
|
||||
pinned = fields.BooleanField(default=True)
|
||||
scope = fields.CharField(max_length=255)
|
||||
description = fields.CharField(max_length=255, default="")
|
||||
created_at = fields.FloatField()
|
||||
updated_at = fields.FloatField()
|
||||
|
||||
class Meta: # type: ignore
|
||||
abstract = True
|
||||
|
||||
|
||||
class TortoiseSlotContext(BaseSlotContext):
|
||||
def __init__(self, model_class: type[AbstractSlotRecord]):
|
||||
self.model_class = model_class
|
||||
|
||||
def _row_to_slot(self, row: AbstractSlotRecord) -> MemorySlot:
|
||||
return MemorySlot(
|
||||
label=row.label,
|
||||
content=row.content,
|
||||
size_limit=row.size_limit,
|
||||
pinned=row.pinned,
|
||||
scope=row.scope,
|
||||
description=row.description,
|
||||
created_at=row.created_at,
|
||||
updated_at=row.updated_at,
|
||||
)
|
||||
|
||||
async def get_slot(self, session: SessionMetadata, label: str) -> MemorySlot | None:
|
||||
rows = await self.model_class.filter(
|
||||
session_id__in=session.accessible_scopes, label=label
|
||||
).all()
|
||||
|
||||
row_map = {r.session_id: r for r in rows}
|
||||
for scope in reversed(session.accessible_scopes):
|
||||
if scope in row_map:
|
||||
return self._row_to_slot(row_map[scope])
|
||||
return None
|
||||
|
||||
async def set_slot(self, session: SessionMetadata, slot: MemorySlot) -> None:
|
||||
composite_id = f"{slot.scope}_{slot.label}"
|
||||
|
||||
await self.model_class.update_or_create(
|
||||
id=composite_id,
|
||||
defaults={
|
||||
"session_id": slot.scope,
|
||||
"label": slot.label,
|
||||
"content": slot.content,
|
||||
"size_limit": slot.size_limit,
|
||||
"pinned": slot.pinned,
|
||||
"scope": slot.scope,
|
||||
"description": slot.description,
|
||||
"created_at": slot.created_at,
|
||||
"updated_at": slot.updated_at,
|
||||
},
|
||||
)
|
||||
|
||||
async def delete_slot(
|
||||
self, session: SessionMetadata, label: str, scope: str
|
||||
) -> None:
|
||||
composite_id = f"{scope}_{label}"
|
||||
await self.model_class.filter(id=composite_id).delete()
|
||||
|
||||
async def list_pinned_slots(self, session: SessionMetadata) -> list[MemorySlot]:
|
||||
rows = await self.model_class.filter(
|
||||
session_id__in=session.accessible_scopes, pinned=True
|
||||
).all()
|
||||
|
||||
merged = {}
|
||||
for scope in session.accessible_scopes:
|
||||
for row in rows:
|
||||
if row.session_id == scope:
|
||||
merged[row.label] = self._row_to_slot(row)
|
||||
|
||||
return [s for s in merged.values() if s.content.strip()]
|
||||
|
||||
async def list_all_slots(self, session: SessionMetadata) -> list[MemorySlot]:
|
||||
rows = await self.model_class.filter(
|
||||
session_id__in=session.accessible_scopes
|
||||
).all()
|
||||
|
||||
merged = {}
|
||||
for scope in session.accessible_scopes:
|
||||
for row in rows:
|
||||
if row.session_id == scope:
|
||||
merged[row.label] = self._row_to_slot(row)
|
||||
return list(merged.values())
|
||||
|
||||
async def clear_by_query(self, query: ScopeSelector) -> None:
|
||||
scope_prefix = query.scope_prefix
|
||||
await self.model_class.filter(session_id__startswith=scope_prefix).delete()
|
||||
|
||||
|
||||
def get_orm_slot_context(model_class: type[AbstractSlotRecord]) -> TortoiseSlotContext:
|
||||
"""
|
||||
[工厂方法] 供第三方开发者调用,将 Tortoise ORM 表直接包装为记忆槽存储系统。
|
||||
"""
|
||||
return TortoiseSlotContext(model_class=model_class)
|
||||
@@ -0,0 +1,110 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Sequence
|
||||
|
||||
from zhenxun.services.ai.context.memory.types import (
|
||||
MemorySlot,
|
||||
SessionMetadata,
|
||||
)
|
||||
from zhenxun.services.ai.core.messages import AgentMessage, LLMMessage
|
||||
from zhenxun.services.ai.run.context import RunContext
|
||||
from zhenxun.services.ai.utils.scope import ScopeSelector
|
||||
|
||||
|
||||
class BaseChatContext(ABC):
|
||||
"""短期对话历史记忆接口"""
|
||||
|
||||
@abstractmethod
|
||||
async def get_messages(self, session: SessionMetadata) -> list[LLMMessage]:
|
||||
"""获取当前会话的所有历史消息。"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def search(
|
||||
self, query: str, session: SessionMetadata, limit: int = 10
|
||||
) -> list[LLMMessage]:
|
||||
"""根据查询词搜索当前会话的历史消息。"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def add_messages(
|
||||
self, session: SessionMetadata, messages: list[LLMMessage]
|
||||
) -> None:
|
||||
"""向当前会话追加消息。"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def set_messages(
|
||||
self, session: SessionMetadata, messages: list[LLMMessage]
|
||||
) -> None:
|
||||
"""重置并设置当前会话的消息列表。"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def clear(self, session: SessionMetadata) -> None:
|
||||
"""清空当前会话的历史消息。"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def clear_by_query(self, query: ScopeSelector) -> None:
|
||||
"""根据条件领域查询对象清理对话历史。"""
|
||||
...
|
||||
|
||||
|
||||
class BaseSlotContext(ABC):
|
||||
"""中期记忆槽持久化接口"""
|
||||
|
||||
@abstractmethod
|
||||
async def get_slot(self, session: SessionMetadata, label: str) -> MemorySlot | None:
|
||||
"""获取指定会话下特定标签的记忆槽。"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def set_slot(self, session: SessionMetadata, slot: MemorySlot) -> None:
|
||||
"""设置或更新指定会话下的记忆槽。"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def delete_slot(
|
||||
self, session: SessionMetadata, label: str, scope: str
|
||||
) -> None:
|
||||
"""删除指定会话下特定标签和作用域的记忆槽。"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def list_pinned_slots(self, session: SessionMetadata) -> list[MemorySlot]:
|
||||
"""列出当前会话所有固定的记忆槽。"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def list_all_slots(self, session: SessionMetadata) -> list[MemorySlot]:
|
||||
"""列出当前会话的所有记忆槽(包括未置顶的)。"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def clear_by_query(self, query: ScopeSelector) -> None:
|
||||
"""根据条件领域查询对象清理记忆槽。"""
|
||||
...
|
||||
|
||||
|
||||
class BaseMemoryReducer(ABC):
|
||||
"""记忆压缩器基类"""
|
||||
|
||||
@abstractmethod
|
||||
async def reduce(
|
||||
self,
|
||||
messages: list[LLMMessage],
|
||||
current_tokens: int,
|
||||
model_name: str,
|
||||
base_overhead: int = 0,
|
||||
) -> tuple[list[LLMMessage], bool, int]:
|
||||
"""对消息列表进行压缩处理。"""
|
||||
...
|
||||
|
||||
|
||||
class BaseMemoryIngestionMiddleware(ABC):
|
||||
"""记忆入库中间件基类,在写入数据库前拦截并修改/清洗消息"""
|
||||
|
||||
@abstractmethod
|
||||
async def process(
|
||||
self, messages: Sequence[AgentMessage], context: RunContext
|
||||
) -> list[AgentMessage]: ...
|
||||
@@ -0,0 +1,103 @@
|
||||
from collections.abc import Awaitable, Callable
|
||||
import time
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from zhenxun.services.ai.utils.scope import ScopeBuilder, ScopeSelector
|
||||
|
||||
AutoRecallPolicy = bool | Callable[[str, "SessionMetadata"], Awaitable[bool] | bool]
|
||||
"""长期记忆的自动召回策略"""
|
||||
|
||||
|
||||
class MemorySlot(BaseModel):
|
||||
"""可编辑的持久化记忆槽 (Mid-Term Memory)"""
|
||||
|
||||
label: str = Field(...)
|
||||
"""槽位唯一标签标识 (如 persona, preferences)"""
|
||||
content: str = Field(default="")
|
||||
"""槽位存储的具体文本内容"""
|
||||
size_limit: int = Field(default=2000)
|
||||
"""槽位内容的最大字符数限制"""
|
||||
pinned: bool = Field(default=True)
|
||||
"""是否固定注入到大模型的每次系统提示词中"""
|
||||
scope: str = Field(...)
|
||||
"""作用域:表示隔离的路径前缀 (scope_prefix)"""
|
||||
description: str = Field(default="")
|
||||
"""该记忆槽的用途说明,便于大模型理解"""
|
||||
created_at: float = Field(default_factory=time.time)
|
||||
"""创建时间戳"""
|
||||
updated_at: float = Field(default_factory=time.time)
|
||||
"""最近更新时间戳"""
|
||||
|
||||
|
||||
class Isolation:
|
||||
"""预设策略工厂,提供友好的隔离级别声明式 API"""
|
||||
|
||||
@staticmethod
|
||||
def _base() -> ScopeBuilder:
|
||||
"""获取底座通用隔离级别(包含Bot、平台、命名空间、智能体)。"""
|
||||
return ScopeBuilder().bot().platform().namespace().agent()
|
||||
|
||||
@classmethod
|
||||
def GROUP_SHARED(cls) -> ScopeBuilder:
|
||||
"""群组共享隔离:同群内共享金库与记忆。"""
|
||||
return cls._base().group()
|
||||
|
||||
@classmethod
|
||||
def USER_GLOBAL(cls) -> ScopeBuilder:
|
||||
"""用户全局隔离:跨群、跨插件共享用户记忆。"""
|
||||
return cls._base().user()
|
||||
|
||||
@classmethod
|
||||
def GROUP_USER(cls) -> ScopeBuilder:
|
||||
"""群组用户隔离:单群内单用户独立隔离。"""
|
||||
return cls._base().group().user()
|
||||
|
||||
@classmethod
|
||||
def AGENT_USER(cls) -> ScopeBuilder:
|
||||
"""智能体用户隔离:单智能体单用户物理隔离。"""
|
||||
return cls.GROUP_USER()
|
||||
|
||||
|
||||
class SessionMetadata(BaseModel):
|
||||
"""结构化会话元数据"""
|
||||
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
session_id: str = Field(...)
|
||||
"""核心会话标识符。"""
|
||||
selector: ScopeSelector = Field(default_factory=ScopeSelector)
|
||||
"""统一的作用域与实体资源选择器。"""
|
||||
isolation_level: ScopeBuilder | None = Field(default=None)
|
||||
"""生成此会话时的隔离级别。"""
|
||||
scope_prefix: str = Field(default="/")
|
||||
"""基于隔离级别生成的路径作用域,用于长期记忆 (RAG) 的向量检索前缀过滤。"""
|
||||
accessible_scopes: list[str] = Field(default_factory=lambda: ["/"])
|
||||
"""
|
||||
当前会话有权访问的作用域列表,用于 Slice 联合检索。
|
||||
"""
|
||||
scope_name_mapping: dict[str, str] = Field(default_factory=dict)
|
||||
"""物理路径到语义化名称的逆向映射字典,供大模型友好阅读"""
|
||||
|
||||
@property
|
||||
def platform(self) -> str | None:
|
||||
return self.selector.platform
|
||||
|
||||
@property
|
||||
def group_id(self) -> str | None:
|
||||
return self.selector.group_id
|
||||
|
||||
@property
|
||||
def user_id(self) -> str | None:
|
||||
return self.selector.user_id
|
||||
|
||||
@property
|
||||
def namespace(self) -> str | None:
|
||||
return self.selector.namespace
|
||||
|
||||
@property
|
||||
def agent_name(self) -> str | None:
|
||||
return self.selector.agent_name
|
||||
|
||||
def __str__(self) -> str:
|
||||
return self.session_id
|
||||
Reference in New Issue
Block a user