mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-09 13:50:00 +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
@@ -1,17 +1,15 @@
|
||||
"""
|
||||
Zhenxun AI - 上下文、记忆与知识管理子系统门面 (Context, Memory & Knowledge Facade)
|
||||
Zhenxun AI - 上下文、记忆与知识管理子系统门面
|
||||
"""
|
||||
|
||||
from .knowledge import FileSystemKnowledge, VectorKnowledge
|
||||
from .memory import (
|
||||
AgentSessionFacade,
|
||||
MemoryBuilder,
|
||||
memory_manager,
|
||||
)
|
||||
from .rag import RAGBuilder
|
||||
|
||||
__all__ = [
|
||||
"AgentSessionFacade",
|
||||
"FileSystemKnowledge",
|
||||
"MemoryBuilder",
|
||||
"RAGBuilder",
|
||||
|
||||
@@ -5,6 +5,7 @@ import anyio
|
||||
from nonebot.adapters import Bot, Event
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from zhenxun.services.ai.context.rag.backends import StorageBackend
|
||||
from zhenxun.services.ai.context.rag.engine import ScopedRAGClient
|
||||
from zhenxun.services.ai.context.rag.models import BaseRecord
|
||||
from zhenxun.services.ai.core.messages import LLMMessage
|
||||
@@ -55,7 +56,7 @@ class VectorKnowledge(BaseKnowledge):
|
||||
"{knowledge_text}"
|
||||
)
|
||||
|
||||
_global_storage: Any = None
|
||||
_global_storage: StorageBackend | None = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -194,7 +195,7 @@ class VectorKnowledge(BaseKnowledge):
|
||||
return super().get_instructions()
|
||||
|
||||
async def before_llm_request(
|
||||
self, context: RunContext, messages: list[Any]
|
||||
self, context: RunContext, messages: list[LLMMessage]
|
||||
) -> None:
|
||||
"""
|
||||
生命周期钩子:在向底层 LLM 发起请求前触发。
|
||||
@@ -285,13 +286,27 @@ class VectorKnowledge(BaseKnowledge):
|
||||
return await self.rag_client.ingest([doc])
|
||||
|
||||
async def add_directory(self, dir_path: str | Path) -> int:
|
||||
"""扫描目录并注入所有支持的文件"""
|
||||
total_chunks = 0
|
||||
"""扫描目录并批量注入所有支持的文件"""
|
||||
aio_path = anyio.Path(dir_path)
|
||||
docs_to_ingest = []
|
||||
|
||||
async for p in aio_path.rglob("*"):
|
||||
if await p.is_file():
|
||||
total_chunks += await self.add_file(Path(p))
|
||||
return total_chunks
|
||||
std_path = Path(p)
|
||||
ext = std_path.suffix.lower()
|
||||
reader = self.readers.get(ext)
|
||||
if not reader:
|
||||
logger.warning(f"当前知识库未配置支持解析文件后缀: {ext}")
|
||||
continue
|
||||
|
||||
doc = await reader.read(std_path)
|
||||
if doc:
|
||||
docs_to_ingest.append(doc)
|
||||
|
||||
if not docs_to_ingest:
|
||||
return 0
|
||||
|
||||
return await self.rag_client.ingest(docs_to_ingest)
|
||||
|
||||
@tool(
|
||||
name="search_knowledge",
|
||||
|
||||
@@ -1,9 +1,7 @@
|
||||
from .builder import MemoryBuilder
|
||||
from .compression import MemoryPolicy
|
||||
from .facades import AgentSessionFacade
|
||||
from .manager import memory_manager
|
||||
from .models import (
|
||||
BaseMemoryIngestionMiddleware,
|
||||
MemoryConfig,
|
||||
)
|
||||
from .types import (
|
||||
@@ -12,8 +10,6 @@ from .types import (
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"AgentSessionFacade",
|
||||
"BaseMemoryIngestionMiddleware",
|
||||
"Isolation",
|
||||
"MemoryBuilder",
|
||||
"MemoryConfig",
|
||||
|
||||
@@ -6,27 +6,19 @@ from typing_extensions import Self
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from zhenxun.services.ai.context.memory.compression import MemoryPolicy
|
||||
from zhenxun.services.ai.context.memory.models import (
|
||||
from zhenxun.services.ai.utils.scope import ScopeBuilder
|
||||
|
||||
from .compression import MemoryPolicy
|
||||
from .models import (
|
||||
ContextCompressionConfig,
|
||||
IngestionConfig,
|
||||
LongTermConfig,
|
||||
MemoryConfig,
|
||||
MemorySlot,
|
||||
ShortTermConfig,
|
||||
SlotMemoryConfig,
|
||||
)
|
||||
from zhenxun.services.ai.context.memory.storage.interfaces import (
|
||||
from .storage.interfaces import (
|
||||
BaseChatContext,
|
||||
BaseMemoryIngestionMiddleware,
|
||||
BaseSlotContext,
|
||||
)
|
||||
from zhenxun.services.ai.context.memory.types import (
|
||||
AutoRecallPolicy,
|
||||
)
|
||||
from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend
|
||||
from zhenxun.services.ai.context.rag.engine import ScopedRAGClient
|
||||
from zhenxun.services.ai.utils.scope import ScopeBuilder
|
||||
|
||||
|
||||
class MemoryBuilder:
|
||||
@@ -42,8 +34,6 @@ class MemoryBuilder:
|
||||
"""
|
||||
self._config = MemoryConfig(
|
||||
short_term=ShortTermConfig(enable=False),
|
||||
slots=SlotMemoryConfig(enable=False),
|
||||
long_term=LongTermConfig(enable=False),
|
||||
compression=ContextCompressionConfig(),
|
||||
ingestion=IngestionConfig(),
|
||||
)
|
||||
@@ -68,17 +58,10 @@ class MemoryBuilder:
|
||||
if isinstance(memory, bool):
|
||||
return MemoryConfig(
|
||||
short_term=ShortTermConfig(enable=memory),
|
||||
long_term=LongTermConfig(enable=memory),
|
||||
)
|
||||
|
||||
return MemoryConfig(short_term=ShortTermConfig(enable=False))
|
||||
|
||||
def with_base_isolation(self, isolation: ScopeBuilder) -> Self:
|
||||
"""设置顶层基准隔离级别,短期/中期/长期记忆将默认继承此级别"""
|
||||
self._config.base_isolation = isolation
|
||||
self._config.short_term.isolation = isolation
|
||||
return self
|
||||
|
||||
def with_short_term(
|
||||
self,
|
||||
enable: bool = True,
|
||||
@@ -95,77 +78,11 @@ class MemoryBuilder:
|
||||
"""
|
||||
self._config.short_term.enable = enable
|
||||
if isolation is not None:
|
||||
self._config.base_isolation = isolation
|
||||
self._config.short_term.isolation = isolation
|
||||
if backend is not None:
|
||||
self._config.short_term.backend = backend
|
||||
return self
|
||||
|
||||
def with_slots(
|
||||
self,
|
||||
enable: bool = True,
|
||||
scopes: dict[str, ScopeBuilder] | None = None,
|
||||
default_slots: list[MemorySlot] | None = None,
|
||||
backend: str | BaseSlotContext | None = None,
|
||||
instructions: str | None = None,
|
||||
) -> Self:
|
||||
"""
|
||||
配置核心槽位记忆 (Memory Slots)。
|
||||
|
||||
参数:
|
||||
enable: 是否启用槽位记忆。
|
||||
scopes: 语义化作用域映射字典。如果只有一个键值对,则大模型不可见该参数。
|
||||
default_slots: 首次初始化时自动写入的默认槽位列表。
|
||||
backend: 槽位记忆存储后端,如果为 None 则使用全局默认后端。
|
||||
instructions: 覆写内置槽位管理工具箱的默认系统提示词规则。
|
||||
"""
|
||||
self._config.slots.enable = enable
|
||||
if scopes is not None:
|
||||
self._config.slots.scopes = scopes
|
||||
if default_slots is not None:
|
||||
self._config.slots.default_slots = default_slots
|
||||
if backend is not None:
|
||||
self._config.slots.backend = backend
|
||||
if instructions is not None:
|
||||
self._config.slots.instructions = instructions
|
||||
return self
|
||||
|
||||
def with_long_term(
|
||||
self,
|
||||
enable: bool = True,
|
||||
scopes: dict[str, ScopeBuilder] | None = None,
|
||||
engine: ScopedRAGClient | None = None,
|
||||
backend: str | StorageBackend | None = None,
|
||||
embedder: Embedder | str | None = None,
|
||||
agentic: bool = True,
|
||||
auto_recall: AutoRecallPolicy = False,
|
||||
instructions: str | None = None,
|
||||
) -> Self:
|
||||
"""
|
||||
配置长期向量记忆与 RAG 设定。
|
||||
|
||||
参数:
|
||||
enable: 是否启用长期记忆。
|
||||
scopes: 语义化作用域映射字典。如果只有一个键值对,则大模型不可见该参数。
|
||||
engine: 高级 RAG 检索引擎实例 (推荐)。若提供,将接管记忆的底层检索、混合与重排。
|
||||
backend: 长期记忆存储后端。
|
||||
embedder: 用于向量化的文本嵌入模型实例。
|
||||
agentic: 是否开启主动智能体记忆管理 (增删改查工具自动注入)。
|
||||
auto_recall: 长期记忆的自动召回策略,支持 bool 或 Callable 函数。
|
||||
instructions: 覆写内置长期记忆工具箱的默认系统提示词规则。
|
||||
""" # noqa: E501
|
||||
self._config.long_term.enable = enable
|
||||
self._config.long_term.engine = engine
|
||||
if scopes is not None:
|
||||
self._config.long_term.scopes = scopes
|
||||
self._config.long_term.backend = backend
|
||||
self._config.long_term.embedder = embedder
|
||||
self._config.long_term.agentic = agentic
|
||||
self._config.long_term.auto_recall = auto_recall
|
||||
if instructions is not None:
|
||||
self._config.long_term.instructions = instructions
|
||||
return self
|
||||
|
||||
def with_multimodal_window(self, window_size: int = 5) -> Self:
|
||||
"""
|
||||
配置多模态历史视窗大小。
|
||||
@@ -261,8 +178,4 @@ class MemoryBuilder:
|
||||
"""
|
||||
生成最终构建好的 MemoryConfig 配置对象。
|
||||
"""
|
||||
if not self._config.slots.scopes:
|
||||
self._config.slots.scopes = {"私有": self._config.base_isolation}
|
||||
if not self._config.long_term.scopes:
|
||||
self._config.long_term.scopes = {"私有": self._config.base_isolation}
|
||||
return self._config
|
||||
|
||||
@@ -1,52 +1,270 @@
|
||||
import inspect
|
||||
from typing import Any
|
||||
|
||||
from zhenxun.services.ai.capabilities.base import AbstractCapability
|
||||
from zhenxun.services.ai.context.memory.models import MemoryConfig
|
||||
from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend
|
||||
from zhenxun.services.ai.context.rag.engine import ScopedRAGClient
|
||||
from zhenxun.services.ai.run.context import RunContext
|
||||
from zhenxun.services.ai.tools.core.toolkit import BaseToolkit
|
||||
from zhenxun.services.ai.tools.providers.builtin.memory import MemoryManagementToolkit
|
||||
from zhenxun.services.ai.utils.logger import log_memory as logger
|
||||
from zhenxun.services.ai.utils.runtime import ContextUtils
|
||||
from zhenxun.services.ai.utils.scope import ScopeBuilder
|
||||
|
||||
from .manager import memory_manager
|
||||
from .models import MemoryScoringConfig
|
||||
from .storage.backends import MemoryScope
|
||||
from .storage.interfaces import BaseSlotContext
|
||||
from .types import (
|
||||
AutoRecallPolicy,
|
||||
Isolation,
|
||||
MemorySlot,
|
||||
SessionMetadata,
|
||||
)
|
||||
|
||||
|
||||
class AgenticMemoryCapability(AbstractCapability):
|
||||
class LongTermMemoryCapability(AbstractCapability):
|
||||
"""
|
||||
智能体主动记忆管理能力 (Agentic Memory Management)。
|
||||
当 `MemoryConfig.long_term.enable == True` 且 `agentic == True` 时隐式挂载,
|
||||
在运行时动态组装并向大模型提供 `MemoryManagementToolkit` 工具箱。
|
||||
长期向量记忆 (RAG) 核心能力组件。
|
||||
负责静默执行自动召回 (Auto Recall),并在必要时提供读写 RAG 数据库的工具链。
|
||||
"""
|
||||
|
||||
def __init__(self, memory_config: MemoryConfig, namespace: str):
|
||||
self.memory_config = memory_config
|
||||
def __init__(
|
||||
self,
|
||||
engine: ScopedRAGClient | None = None,
|
||||
storage_backend: StorageBackend | None = None,
|
||||
embedder: Embedder | str | None = None,
|
||||
scopes: dict[str, ScopeBuilder] | ScopeBuilder | None = None,
|
||||
toolkit: bool | BaseToolkit = True,
|
||||
scoring_config: MemoryScoringConfig | None = None,
|
||||
auto_recall: AutoRecallPolicy = False,
|
||||
recall_limit: int = 5,
|
||||
recall_threshold: float = 0.5,
|
||||
namespace: str | None = None,
|
||||
):
|
||||
"""
|
||||
初始化长期记忆能力组件。
|
||||
|
||||
参数:
|
||||
engine: RAG 客户端引擎实例。若为 None 则在运行时按需构建。
|
||||
storage_backend: RAG 向量存储后端。若为 None 则从管理器按命名空间提取。
|
||||
embedder: 嵌入模型实例或模型名称。
|
||||
scopes: 控制长期记忆的隔离级别。支持单 ScopeBuilder 或 映射字典。
|
||||
toolkit: 是否启用记忆管理工具箱,或传入自定义的工具箱实例。
|
||||
scoring_config: 记忆检索打分配置(包含时间衰减等参数)。
|
||||
auto_recall: 自动召回策略,可以是布尔值或自定义的回调函数。
|
||||
recall_limit: 自动召回的记忆条数限制。
|
||||
recall_threshold: 自动召回的相似度分数阈值。
|
||||
namespace: 指定的命名空间,用于自动路由存储后端及构建器。
|
||||
"""
|
||||
self.engine = engine
|
||||
self.storage_backend = storage_backend
|
||||
self.embedder = embedder
|
||||
if isinstance(scopes, ScopeBuilder):
|
||||
self.scopes = {"默认": scopes}
|
||||
else:
|
||||
self.scopes = scopes or {"私有": Isolation.AGENT_USER()}
|
||||
self.toolkit = toolkit
|
||||
self.scoring_config = scoring_config or MemoryScoringConfig()
|
||||
self.auto_recall = auto_recall
|
||||
self.recall_limit = recall_limit
|
||||
self.recall_threshold = recall_threshold
|
||||
self.namespace = namespace
|
||||
|
||||
async def get_tools(self, context: RunContext) -> list[Any]:
|
||||
kwargs = self.memory_config.long_term.toolkit_kwargs.copy()
|
||||
kwargs["memory_config"] = self.memory_config
|
||||
kwargs["namespace"] = self.namespace
|
||||
if self.memory_config.long_term.instructions is not None:
|
||||
kwargs["instructions"] = self.memory_config.long_term.instructions
|
||||
def _build_session_meta(self, context: RunContext) -> SessionMetadata:
|
||||
scope_builder = next(iter(self.scopes.values())) if self.scopes else None
|
||||
return ContextUtils.build_session_meta(
|
||||
context=context,
|
||||
target_builder=scope_builder,
|
||||
extra_scopes=self.scopes,
|
||||
custom_namespace=self.namespace,
|
||||
)
|
||||
|
||||
toolkit = MemoryManagementToolkit(**kwargs)
|
||||
return [toolkit]
|
||||
def _ensure_engine(self, context: RunContext) -> ScopedRAGClient | None:
|
||||
if self.engine is not None:
|
||||
return self.engine
|
||||
|
||||
ns = self.namespace or getattr(context.session, "namespace", "global")
|
||||
|
||||
storage_instance = self.storage_backend
|
||||
if not storage_instance:
|
||||
factory = memory_manager._storage_factories.get(
|
||||
ns
|
||||
) or memory_manager._storage_factories.get("global")
|
||||
if factory:
|
||||
storage_instance = factory()
|
||||
|
||||
if not storage_instance:
|
||||
return None
|
||||
|
||||
embedder_instance = self.embedder
|
||||
if isinstance(embedder_instance, str):
|
||||
from zhenxun.services.ai.context.rag.backends.embedders import (
|
||||
DefaultEmbedder,
|
||||
)
|
||||
|
||||
embedder_instance = DefaultEmbedder(model_name=embedder_instance)
|
||||
|
||||
from zhenxun.services.ai.context.rag.builder import RAGBuilder
|
||||
|
||||
builder = RAGBuilder(storage_instance).with_scope("/")
|
||||
if embedder_instance:
|
||||
builder.with_embedder(embedder_instance)
|
||||
|
||||
builder.enable_lifecycle_scoring(
|
||||
half_life_days=self.scoring_config.recency_half_life_days,
|
||||
decay_weight=self.scoring_config.recency_weight,
|
||||
semantic_weight=self.scoring_config.semantic_weight,
|
||||
importance_weight=self.scoring_config.importance_weight,
|
||||
reinforcement_weight=self.scoring_config.reinforcement_weight,
|
||||
)
|
||||
|
||||
self.engine = builder.build()
|
||||
return self.engine
|
||||
|
||||
async def get_system_prompts(self, context: RunContext) -> list[str]:
|
||||
engine = self._ensure_engine(context)
|
||||
user_input = context.run.user_input
|
||||
if not user_input or not engine:
|
||||
return []
|
||||
|
||||
should_recall = False
|
||||
session_meta = self._build_session_meta(context)
|
||||
|
||||
if isinstance(self.auto_recall, bool):
|
||||
should_recall = self.auto_recall
|
||||
elif callable(self.auto_recall):
|
||||
try:
|
||||
res = self.auto_recall(user_input, session_meta)
|
||||
if inspect.isawaitable(res):
|
||||
should_recall = await res
|
||||
else:
|
||||
should_recall = bool(res)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"[LongTermMemoryCapability] auto_recall 函数执行失败: {e}"
|
||||
)
|
||||
should_recall = False
|
||||
|
||||
if not should_recall:
|
||||
return []
|
||||
|
||||
scope = MemoryScope(rag_client=engine)
|
||||
matches = await scope.recall(
|
||||
session=session_meta,
|
||||
query=user_input,
|
||||
limit=self.recall_limit,
|
||||
)
|
||||
if matches:
|
||||
valid_matches = [m for m in matches if m.score >= self.recall_threshold]
|
||||
if valid_matches:
|
||||
fact_str = "\n".join(f"- {m.record.content}" for m in valid_matches)
|
||||
return [f"[系统补充:有关用户的长期记忆设定]\n{fact_str}"]
|
||||
return []
|
||||
|
||||
async def get_tools(self, context: RunContext) -> list[Any]:
|
||||
if self.toolkit is False:
|
||||
return []
|
||||
|
||||
engine = self._ensure_engine(context)
|
||||
if not engine:
|
||||
return []
|
||||
|
||||
if isinstance(self.toolkit, BaseToolkit):
|
||||
tk = self.toolkit.clone_with(
|
||||
rag_client=engine, scopes=self.scopes, _namespace=self.namespace
|
||||
)
|
||||
return [tk]
|
||||
|
||||
return [
|
||||
MemoryManagementToolkit(
|
||||
rag_client=engine, scopes=self.scopes, namespace=self.namespace
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
class SlotMemoryCapability(AbstractCapability):
|
||||
"""
|
||||
槽位记忆能力组件。
|
||||
当 `MemoryConfig.slots.enable == True` 时隐式挂载,
|
||||
在运行时动态组装并向大模型提供 `MemorySlotToolkit` 工具箱。
|
||||
独立的槽位记忆 (Memory Slots) 能力组件。
|
||||
直接作为插件挂载至 Agent 的 capabilities 列表中。
|
||||
"""
|
||||
|
||||
def __init__(self, memory_config: MemoryConfig, namespace: str):
|
||||
self.memory_config = memory_config
|
||||
def __init__(
|
||||
self,
|
||||
scopes: dict[str, ScopeBuilder] | ScopeBuilder | None = None,
|
||||
default_slots: list[MemorySlot] | None = None,
|
||||
backend: BaseSlotContext | None = None,
|
||||
toolkit: bool | BaseToolkit = True,
|
||||
namespace: str | None = None,
|
||||
):
|
||||
"""
|
||||
初始化槽位记忆能力组件。
|
||||
|
||||
参数:
|
||||
scopes: 控制槽位记忆的隔离级别。支持单 ScopeBuilder 或 映射字典。
|
||||
default_slots: 默认记忆槽列表,初始时自动创建未存在的槽位。
|
||||
backend: 中期记忆槽持久化后端。若为 None 则从管理器按命名空间提取。
|
||||
toolkit: 是否启用记忆槽管理工具箱,或传入自定义的工具箱实例。
|
||||
namespace: 指定的命名空间,用于自动路由后端及工具箱。
|
||||
"""
|
||||
if isinstance(scopes, ScopeBuilder):
|
||||
self.scopes = {"默认": scopes}
|
||||
else:
|
||||
self.scopes = scopes or {"私有": Isolation.AGENT_USER()}
|
||||
self.default_slots = default_slots or []
|
||||
self.backend = backend
|
||||
self.toolkit = toolkit
|
||||
self.namespace = namespace
|
||||
|
||||
async def _get_slot_ctx_and_meta(self, context: RunContext):
|
||||
ns = self.namespace or getattr(context.session, "namespace", "global")
|
||||
slot_ctx = self.backend or memory_manager.get_backend("slots", namespace=ns)
|
||||
|
||||
target_builder = next(iter(self.scopes.values())) if self.scopes else None
|
||||
session_meta = ContextUtils.build_session_meta(
|
||||
context=context,
|
||||
target_builder=target_builder,
|
||||
extra_scopes=self.scopes,
|
||||
custom_namespace=self.namespace,
|
||||
)
|
||||
return slot_ctx, session_meta
|
||||
|
||||
async def get_system_prompts(self, context: RunContext) -> list[str]:
|
||||
slot_ctx, session_meta = await self._get_slot_ctx_and_meta(context)
|
||||
if not slot_ctx:
|
||||
return []
|
||||
|
||||
for default_slot in self.default_slots:
|
||||
if not await slot_ctx.get_slot(session_meta, default_slot.label):
|
||||
await slot_ctx.set_slot(session_meta, default_slot)
|
||||
|
||||
slots = await slot_ctx.list_pinned_slots(session_meta)
|
||||
if not slots:
|
||||
return []
|
||||
|
||||
xml_parts = ["<memory_slots>"]
|
||||
for slot in slots:
|
||||
semantic_name = session_meta.scope_name_mapping.get(slot.scope, "未知")
|
||||
xml_parts.append(
|
||||
f' <slot name="{slot.label}" scope="{semantic_name}">\n'
|
||||
f" {slot.content}\n"
|
||||
" </slot>"
|
||||
)
|
||||
xml_parts.append("</memory_slots>")
|
||||
return ["\n".join(xml_parts)]
|
||||
|
||||
async def get_tools(self, context: RunContext) -> list[Any]:
|
||||
from zhenxun.services.ai.tools.providers.builtin.slots import MemorySlotToolkit
|
||||
|
||||
kwargs = self.memory_config.slots.toolkit_kwargs.copy()
|
||||
kwargs["memory_config"] = self.memory_config
|
||||
kwargs["namespace"] = self.namespace
|
||||
if self.memory_config.slots.instructions is not None:
|
||||
kwargs["instructions"] = self.memory_config.slots.instructions
|
||||
if self.toolkit is False:
|
||||
return []
|
||||
|
||||
toolkit = MemorySlotToolkit(**kwargs)
|
||||
if isinstance(self.toolkit, BaseToolkit):
|
||||
toolkit = self.toolkit.clone_with(
|
||||
scopes=self.scopes, backend=self.backend, _namespace=self.namespace
|
||||
)
|
||||
else:
|
||||
toolkit = MemorySlotToolkit(
|
||||
scopes=self.scopes, backend=self.backend, namespace=self.namespace
|
||||
)
|
||||
return [toolkit]
|
||||
|
||||
@@ -4,6 +4,7 @@ from typing import Any, Generic, TypeVar
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from zhenxun.services.ai.context.memory.models import MemoryConfig
|
||||
from zhenxun.services.ai.core.engine.token_counter import token_counter
|
||||
from zhenxun.services.ai.core.messages import (
|
||||
AudioPart,
|
||||
@@ -386,7 +387,7 @@ class LLMSummarizerReducer(AbstractSummarizerReducer):
|
||||
) -> LLMMessage | None:
|
||||
prompt_text = f"### 📋 [对话摘要任务]\n{self.summarization_prompt}\n\n"
|
||||
if prev_summary:
|
||||
prompt_text += "#### önceki_summary (参考先前的快照):\n"
|
||||
prompt_text += "#### prev_summary (参考先前的快照):\n"
|
||||
prompt_text += f"> {prev_summary}\n\n"
|
||||
prompt_text += "#### 待处理的历史消息流:\n"
|
||||
for m in to_summarize:
|
||||
@@ -512,7 +513,7 @@ class CondenserPipeline:
|
||||
|
||||
@classmethod
|
||||
def create_from_configs(
|
||||
cls, memory_config: Any, model_name: str
|
||||
cls, memory_config: MemoryConfig | None, model_name: str
|
||||
) -> "CondenserPipeline":
|
||||
"""基于全局和局部配置组装压缩管线工厂方法"""
|
||||
from zhenxun.services.ai.config import get_llm_config
|
||||
@@ -602,7 +603,7 @@ class CondenserPipeline:
|
||||
|
||||
class MemoryPolicy:
|
||||
"""
|
||||
记忆策略工厂 (Strategy Factory Facade)。
|
||||
记忆策略工厂。
|
||||
为开发者提供开箱即用的上下文压缩管线组装方案。
|
||||
"""
|
||||
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
from collections.abc import Sequence
|
||||
from typing import Any, cast
|
||||
from typing import cast
|
||||
|
||||
from zhenxun.services.ai.core.engine.context_renderer import ContextConverter
|
||||
from zhenxun.services.ai.core.messages import AgentMessage
|
||||
from zhenxun.services.ai.run.context import RunContext
|
||||
from zhenxun.services.ai.utils.logger import log_memory as logger
|
||||
from zhenxun.utils.pydantic_compat import model_copy
|
||||
|
||||
@@ -14,144 +15,37 @@ from .models import MemoryConfig
|
||||
from .types import SessionMetadata
|
||||
|
||||
|
||||
class MemoryReader:
|
||||
class SessionMemoryContext:
|
||||
"""
|
||||
记忆读取器 (Memory Reader)。
|
||||
负责从数据库中提取短期上下文历史,召回长期的背景知识,并执行自动压缩。
|
||||
会话记忆门面。
|
||||
封装了当前会话记忆的读写操作、上下文压缩管线以及入库清洗中间件。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, session_meta: SessionMetadata, memory_config: MemoryConfig | None
|
||||
self,
|
||||
session_meta: SessionMetadata,
|
||||
memory_config: MemoryConfig | None,
|
||||
context: RunContext,
|
||||
):
|
||||
"""
|
||||
初始化记忆读取器。
|
||||
初始化会话记忆门面。
|
||||
|
||||
参数:
|
||||
session_meta: 会话元数据,包含 Namespace 与作用域映射等上下文信息。
|
||||
memory_config: 记忆系统的配置对象,控制长期、短期及槽位记忆的启用与逻辑。
|
||||
memory_config: 记忆系统的配置对象,控制长期、短期记忆的启用与逻辑。
|
||||
context: 必填的运行时上下文环境 (RunContext),供中间件进行依赖注入。
|
||||
"""
|
||||
self.session_meta = session_meta
|
||||
self.memory_config = memory_config
|
||||
self.context = context
|
||||
|
||||
async def get_long_term_context(self, user_input: str) -> str:
|
||||
"""
|
||||
基于用户输入召回长期记忆(RAG),返回格式化后的背景提示词。
|
||||
"""
|
||||
if (
|
||||
not self.memory_config
|
||||
or not self.memory_config.long_term.enable
|
||||
or not user_input
|
||||
):
|
||||
return ""
|
||||
|
||||
policy = self.memory_config.long_term.auto_recall
|
||||
should_recall = False
|
||||
|
||||
if isinstance(policy, bool):
|
||||
should_recall = policy
|
||||
elif callable(policy):
|
||||
import inspect
|
||||
|
||||
try:
|
||||
res = policy(user_input, self.session_meta)
|
||||
if inspect.isawaitable(res):
|
||||
should_recall = await res
|
||||
else:
|
||||
should_recall = bool(res)
|
||||
except Exception as e:
|
||||
logger.error(f"[MemoryReader] 自定义 auto_recall 函数执行失败: {e}")
|
||||
should_recall = False
|
||||
|
||||
if not should_recall:
|
||||
return ""
|
||||
|
||||
ltm_scope = memory_manager.get_long_term_memory(
|
||||
self.memory_config,
|
||||
namespace=self.session_meta.selector.namespace or "global",
|
||||
)
|
||||
if not ltm_scope:
|
||||
return ""
|
||||
|
||||
matches = await ltm_scope.recall(session=self.session_meta, query=user_input)
|
||||
if matches:
|
||||
logger.debug(f"🧠 [MemoryReader] 长期记忆召回详情 (Query: '{user_input}'):")
|
||||
for i, m in enumerate(matches):
|
||||
logger.debug(
|
||||
f" [{i + 1}] 得分: {m.score:.4f} | 内容: {m.record.content}"
|
||||
)
|
||||
|
||||
threshold = self.memory_config.long_term.recall_threshold
|
||||
valid_matches = [m for m in matches if m.score >= threshold]
|
||||
if not valid_matches:
|
||||
logger.debug("🧠 [MemoryReader] 召回的记忆均未达到相关性阈值,已丢弃。")
|
||||
return ""
|
||||
|
||||
fact_str = "\n".join(f"- {m.record.content}" for m in valid_matches)
|
||||
logger.debug(
|
||||
f"🧠 [MemoryReader]"
|
||||
f"成功截取并注入 {len(valid_matches)} 条高价值长期记忆。"
|
||||
)
|
||||
return f"[系统补充:有关用户的长期记忆设定]\n{fact_str}"
|
||||
return ""
|
||||
|
||||
async def get_slots_context(self) -> str:
|
||||
"""
|
||||
读取并组装核心槽位记忆 (Memory Slots),返回 XML 格式字符串供大模型使用。
|
||||
"""
|
||||
if not self.memory_config or not self.memory_config.slots.enable:
|
||||
return ""
|
||||
slot_ctx = memory_manager.get_slot_context(
|
||||
self.memory_config,
|
||||
namespace=self.session_meta.selector.namespace or "global",
|
||||
)
|
||||
if not slot_ctx:
|
||||
return ""
|
||||
|
||||
if self.memory_config.slots.default_slots:
|
||||
for default_slot in self.memory_config.slots.default_slots:
|
||||
existing = await slot_ctx.get_slot(
|
||||
self.session_meta, default_slot.label
|
||||
)
|
||||
if not existing:
|
||||
await slot_ctx.set_slot(self.session_meta, default_slot)
|
||||
|
||||
slots = await slot_ctx.list_pinned_slots(self.session_meta)
|
||||
if not slots:
|
||||
return ""
|
||||
|
||||
show_scope = False
|
||||
if (
|
||||
self.memory_config
|
||||
and self.memory_config.slots.scopes
|
||||
and len(self.memory_config.slots.scopes) > 1
|
||||
):
|
||||
show_scope = True
|
||||
|
||||
xml_parts = ["<memory_slots>"]
|
||||
for slot in slots:
|
||||
if show_scope:
|
||||
semantic_name = self.session_meta.scope_name_mapping.get(
|
||||
slot.scope, "未知"
|
||||
)
|
||||
xml_parts.append(
|
||||
f' <slot name="{slot.label}" scope="{semantic_name}">\n'
|
||||
f" {slot.content}\n"
|
||||
" </slot>"
|
||||
)
|
||||
else:
|
||||
xml_parts.append(
|
||||
f' <slot name="{slot.label}">\n {slot.content}\n </slot>'
|
||||
)
|
||||
xml_parts.append("</memory_slots>")
|
||||
return "\n".join(xml_parts)
|
||||
|
||||
async def get_short_term_context(
|
||||
async def read(
|
||||
self,
|
||||
model_name: str,
|
||||
override_history: Sequence[AgentMessage] | None = None,
|
||||
) -> list[AgentMessage]:
|
||||
"""
|
||||
拉取短期对话历史,并执行 Token 压缩。
|
||||
拉取短期对话历史,并执行 Token 压缩与管线修剪。
|
||||
"""
|
||||
current_history: list[AgentMessage] = []
|
||||
if override_history is not None:
|
||||
@@ -187,43 +81,18 @@ class MemoryReader:
|
||||
if changed:
|
||||
await chat_context.set_messages(self.session_meta, new_history)
|
||||
logger.info(
|
||||
"💾 [MemoryReader] 压缩截断完毕,已同步覆写数据库。"
|
||||
"💾 [SessionMemory] 压缩截断完毕,已同步覆写数据库。"
|
||||
f"压缩后条数: {len(new_history)}"
|
||||
)
|
||||
current_history = cast(list[AgentMessage], new_history)
|
||||
|
||||
return current_history
|
||||
|
||||
|
||||
class MemoryWriter:
|
||||
"""
|
||||
记忆写入器 (Memory Writer)。
|
||||
负责将对话增量安全地写入数据库。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
session_meta: SessionMetadata,
|
||||
memory_config: MemoryConfig | None,
|
||||
context: Any = None,
|
||||
):
|
||||
"""
|
||||
初始化记忆写入器。
|
||||
|
||||
参数:
|
||||
session_meta: 会话元数据,包含 Namespace 与作用域映射等上下文信息。
|
||||
memory_config: 记忆系统的配置对象,控制记忆存入的逻辑。
|
||||
context: 运行时上下文环境,作为可选参数传入,供中间件使用,默认 None。
|
||||
"""
|
||||
self.session_meta = session_meta
|
||||
self.memory_config = memory_config
|
||||
self.context = context
|
||||
|
||||
async def save_new_messages(
|
||||
async def write(
|
||||
self,
|
||||
new_messages: Sequence[AgentMessage],
|
||||
):
|
||||
"""将新产生的对话增量保存到数据库"""
|
||||
) -> None:
|
||||
"""将新产生的对话增量,经过入库中间件清洗后保存到数据库"""
|
||||
if not new_messages:
|
||||
return
|
||||
|
||||
|
||||
@@ -1,129 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from typing import Literal
|
||||
|
||||
from zhenxun.services.ai.core.messages import AgentMessage, LLMMessage
|
||||
|
||||
from .manager import GlobalMemoryManager
|
||||
from .storage.interfaces import (
|
||||
BaseChatContext,
|
||||
BaseSlotContext,
|
||||
)
|
||||
from .types import (
|
||||
MemorySlot,
|
||||
SessionMetadata,
|
||||
)
|
||||
|
||||
|
||||
class ChatHistoryFacade:
|
||||
"""短期对话历史门面"""
|
||||
|
||||
def __init__(self, manager: GlobalMemoryManager, session_meta: SessionMetadata):
|
||||
self.manager = manager
|
||||
self.session_meta = session_meta
|
||||
|
||||
@property
|
||||
def _backend(self) -> BaseChatContext | None:
|
||||
return self.manager.get_chat_context(
|
||||
None, self.session_meta.namespace or "global"
|
||||
)
|
||||
|
||||
async def get(self, limit: int | None = None) -> list[LLMMessage]:
|
||||
"""获取当前会话的历史消息"""
|
||||
if not self._backend:
|
||||
return []
|
||||
msgs = await self._backend.get_messages(self.session_meta)
|
||||
return msgs[-limit:] if limit else msgs
|
||||
|
||||
async def add(self, messages: Sequence[AgentMessage] | AgentMessage) -> None:
|
||||
"""向当前会话追加一条或多条历史消息"""
|
||||
if not self._backend:
|
||||
return
|
||||
from zhenxun.services.ai.core.engine.context_renderer import ContextConverter
|
||||
|
||||
msgs = messages if isinstance(messages, Sequence) else [messages]
|
||||
flattened = ContextConverter.flatten_to_llm_messages(msgs)
|
||||
if flattened:
|
||||
await self._backend.add_messages(self.session_meta, flattened)
|
||||
|
||||
async def clear(self) -> None:
|
||||
"""清空当前会话的短期对话历史"""
|
||||
if not self._backend:
|
||||
return
|
||||
await self._backend.clear(self.session_meta)
|
||||
|
||||
|
||||
class SlotFacade:
|
||||
"""中期记忆槽门面"""
|
||||
|
||||
def __init__(self, manager: GlobalMemoryManager, session_meta: SessionMetadata):
|
||||
self.manager = manager
|
||||
self.session_meta = session_meta
|
||||
|
||||
@property
|
||||
def _backend(self) -> BaseSlotContext | None:
|
||||
"""获取底层槽位存储后端"""
|
||||
return self.manager.get_slot_context(
|
||||
None, self.session_meta.namespace or "global"
|
||||
)
|
||||
|
||||
async def get(self, label: str) -> str | None:
|
||||
"""获取指定标识的槽位记忆内容"""
|
||||
if not self._backend:
|
||||
return None
|
||||
slot = await self._backend.get_slot(self.session_meta, label)
|
||||
return slot.content if slot else None
|
||||
|
||||
async def set(
|
||||
self,
|
||||
label: str,
|
||||
content: str,
|
||||
scope: Literal["session", "global"] = "session",
|
||||
size_limit: int = 2000,
|
||||
pinned: bool = True,
|
||||
) -> None:
|
||||
"""设置或更新指定的槽位记忆"""
|
||||
if not self._backend:
|
||||
return
|
||||
slot = MemorySlot(
|
||||
label=label,
|
||||
content=content,
|
||||
scope=scope,
|
||||
size_limit=size_limit,
|
||||
pinned=pinned,
|
||||
)
|
||||
await self._backend.set_slot(self.session_meta, slot)
|
||||
|
||||
async def delete(
|
||||
self, label: str, scope: Literal["session", "global"] = "session"
|
||||
) -> None:
|
||||
"""删除指定的槽位记忆"""
|
||||
if not self._backend:
|
||||
return
|
||||
await self._backend.delete_slot(self.session_meta, label, scope)
|
||||
|
||||
async def list_all(self) -> dict[str, str]:
|
||||
"""获取当前会话所有被置顶的槽位记忆"""
|
||||
if not self._backend:
|
||||
return {}
|
||||
slots = await self._backend.list_pinned_slots(self.session_meta)
|
||||
return {s.label: s.content for s in slots}
|
||||
|
||||
|
||||
class AgentSessionFacade:
|
||||
"""
|
||||
提供给第三方开发者的会话记忆访问聚合门面 (Facade)。
|
||||
"""
|
||||
|
||||
def __init__(self, manager: GlobalMemoryManager, session_meta: SessionMetadata):
|
||||
self.manager = manager
|
||||
self.session_meta = session_meta
|
||||
self.history = ChatHistoryFacade(manager, session_meta)
|
||||
self.slots = SlotFacade(manager, session_meta)
|
||||
|
||||
async def clear_all(self) -> None:
|
||||
"""一键清空当前会话下的短期对话历史与记忆槽"""
|
||||
cleaner = self.manager.cleaner().session(self.session_meta.session_id)
|
||||
await cleaner.clear_short_term()
|
||||
await cleaner.clear_slots()
|
||||
@@ -1,7 +1,8 @@
|
||||
from collections import defaultdict
|
||||
from collections.abc import Callable
|
||||
from typing import Any, cast
|
||||
|
||||
from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend
|
||||
from zhenxun.services.ai.context.rag.backends import StorageBackend
|
||||
from zhenxun.services.ai.utils.logger import log_memory as logger
|
||||
from zhenxun.services.ai.utils.scope import BaseScopeBuilder
|
||||
from zhenxun.utils.utils import infer_plugin_namespace
|
||||
@@ -9,11 +10,10 @@ from zhenxun.utils.utils import infer_plugin_namespace
|
||||
from .models import MemoryConfig
|
||||
from .storage.backends import (
|
||||
InMemoryChatContext,
|
||||
MemoryScope,
|
||||
)
|
||||
from .storage.interfaces import (
|
||||
BaseChatContext,
|
||||
BaseSlotContext,
|
||||
IClearableBackend,
|
||||
)
|
||||
|
||||
|
||||
@@ -33,41 +33,46 @@ class MemoryCleaner(BaseScopeBuilder["MemoryCleaner"]):
|
||||
self._config = cfg.build() if hasattr(cfg, "build") else cfg
|
||||
return self
|
||||
|
||||
async def clear_target(self, target_name: str):
|
||||
"""底层派发器:定向清理指定注册名称的泛型扩展后端数据"""
|
||||
ns_dict = self.manager._backends.get(target_name, {})
|
||||
for backend in ns_dict.values():
|
||||
if isinstance(backend, IClearableBackend):
|
||||
await backend.clear_by_query(self._selector)
|
||||
else:
|
||||
logger.warning(
|
||||
f"后端 {backend.__class__.__name__}"
|
||||
"未实现 IClearableBackend 协议,已跳过清理。"
|
||||
)
|
||||
|
||||
async def clear_short_term(self):
|
||||
"""一键清理目标范围下的短期对话历史记忆"""
|
||||
if self._config and self._config.short_term.backend:
|
||||
if isinstance(self._config, MemoryConfig) and isinstance(
|
||||
self._config.short_term.backend, IClearableBackend
|
||||
):
|
||||
await self._config.short_term.backend.clear_by_query(self._selector)
|
||||
else:
|
||||
for backend in self.manager._chat_backends.values():
|
||||
await backend.clear_by_query(self._selector)
|
||||
await self.clear_target("chat")
|
||||
|
||||
async def clear_slots(self):
|
||||
"""一键清理目标范围下的中期记忆槽 (Memory Slots)"""
|
||||
if self._config and self._config.slots.backend:
|
||||
await self._config.slots.backend.clear_by_query(self._selector)
|
||||
else:
|
||||
for backend in self.manager._slot_backends.values():
|
||||
await backend.clear_by_query(self._selector)
|
||||
"""一键清理目标范围下的记忆槽数据"""
|
||||
await self.clear_target("slots")
|
||||
|
||||
async def clear_long_term(self):
|
||||
"""一键清理目标范围下的长期向量记忆 (RAG Vector Database)"""
|
||||
if self._config and self._config.long_term.backend:
|
||||
from zhenxun.services.ai.context.rag.backends import StorageBackend
|
||||
|
||||
storage = cast(StorageBackend, self._config.long_term.backend)
|
||||
await storage.clear_by_query(self._selector)
|
||||
else:
|
||||
for factory in self.manager._storage_factories.values():
|
||||
storage = factory()
|
||||
if hasattr(storage, "clear_by_query"):
|
||||
await storage.clear_by_query(self._selector)
|
||||
else:
|
||||
await storage.delete(scope_prefix=self._selector.scope_prefix)
|
||||
for factory in self.manager._storage_factories.values():
|
||||
storage = factory()
|
||||
if isinstance(storage, IClearableBackend):
|
||||
await storage.clear_by_query(self._selector)
|
||||
|
||||
async def clear_all(self):
|
||||
"""一键清理指定范围下的所有生命周期记忆(对话、槽位、RAG)"""
|
||||
await self.clear_short_term()
|
||||
await self.clear_slots()
|
||||
"""一键清理指定范围下的所有生命周期记忆(对话、记忆槽、RAG、及其他泛型扩展后端)"""
|
||||
if isinstance(self._config, MemoryConfig) and isinstance(
|
||||
self._config.short_term.backend, IClearableBackend
|
||||
):
|
||||
await self._config.short_term.backend.clear_by_query(self._selector)
|
||||
for target_name in self.manager._backends:
|
||||
await self.clear_target(target_name)
|
||||
await self.clear_long_term()
|
||||
logger.info(
|
||||
f"🧹 成功清理作用域 '{self._selector.scope_prefix}'下的所有记忆痕迹!"
|
||||
@@ -81,10 +86,8 @@ class GlobalMemoryManager:
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._chat_backends: dict[str, BaseChatContext] = {
|
||||
"global": InMemoryChatContext()
|
||||
}
|
||||
self._slot_backends: dict[str, BaseSlotContext] = {}
|
||||
self._backends: dict[str, dict[str, Any]] = defaultdict(dict)
|
||||
self.register_backend("chat", InMemoryChatContext(), "global")
|
||||
|
||||
from zhenxun.services.ai.context.rag.backends import DictStorageBackend
|
||||
|
||||
@@ -92,19 +95,23 @@ class GlobalMemoryManager:
|
||||
"global": lambda: DictStorageBackend()
|
||||
}
|
||||
|
||||
def register_backend(
|
||||
self, backend_type: str, backend: Any, scope: str | None = None
|
||||
) -> None:
|
||||
"""泛型注册:注册任意类型的存储后端"""
|
||||
ns = scope if scope is not None else infer_plugin_namespace()
|
||||
self._backends[backend_type][ns] = backend
|
||||
|
||||
def get_backend(self, backend_type: str, namespace: str = "global") -> Any | None:
|
||||
"""泛型获取:获取任意类型的存储后端"""
|
||||
backends = self._backends.get(backend_type, {})
|
||||
return backends.get(namespace) or backends.get("global")
|
||||
|
||||
def register_chat_backend(
|
||||
self, backend: BaseChatContext, scope: str | None = None
|
||||
) -> None:
|
||||
"""注册特定命名空间的短期记忆存储后端。"""
|
||||
ns = scope if scope is not None else infer_plugin_namespace()
|
||||
self._chat_backends[ns] = backend
|
||||
|
||||
def register_slot_backend(
|
||||
self, backend: BaseSlotContext, scope: str | None = None
|
||||
) -> None:
|
||||
"""注册特定命名空间的中期记忆槽存储后端。"""
|
||||
ns = scope if scope is not None else infer_plugin_namespace()
|
||||
self._slot_backends[ns] = backend
|
||||
self.register_backend("chat", backend, scope)
|
||||
|
||||
def register_storage_factory(
|
||||
self, factory: Callable[[], StorageBackend], scope: str | None = None
|
||||
@@ -117,20 +124,6 @@ class GlobalMemoryManager:
|
||||
"""获取声明式记忆清理构建器,供第三方开发者极速清理指定记忆"""
|
||||
return MemoryCleaner(self)
|
||||
|
||||
def get_embedder(self, embedder_val: "Embedder | str | None") -> Embedder | None:
|
||||
"""获取向量化引擎实例。如果传入的是字符串,则视为 API 模型名称。"""
|
||||
if not embedder_val:
|
||||
return None
|
||||
|
||||
if isinstance(embedder_val, str):
|
||||
from zhenxun.services.ai.context.rag.backends.embedders import (
|
||||
DefaultEmbedder,
|
||||
)
|
||||
|
||||
return DefaultEmbedder(model_name=embedder_val)
|
||||
|
||||
return embedder_val
|
||||
|
||||
def get_chat_context(
|
||||
self, config: MemoryConfig | None, namespace: str = "global"
|
||||
) -> BaseChatContext | None:
|
||||
@@ -142,69 +135,7 @@ class GlobalMemoryManager:
|
||||
if backend_cfg is not None:
|
||||
return cast(BaseChatContext, backend_cfg)
|
||||
|
||||
return self._chat_backends.get(namespace) or self._chat_backends["global"]
|
||||
|
||||
def get_slot_context(
|
||||
self, config: MemoryConfig | None, namespace: str = "global"
|
||||
) -> BaseSlotContext | None:
|
||||
"""根据配置分配对应的槽位记忆实例"""
|
||||
if not config or not config.slots.enable:
|
||||
return None
|
||||
|
||||
backend_cfg = config.slots.backend
|
||||
if backend_cfg is not None:
|
||||
return cast(BaseSlotContext, backend_cfg)
|
||||
|
||||
return self._slot_backends.get(namespace) or self._slot_backends["global"]
|
||||
|
||||
def get_long_term_memory(
|
||||
self, config: MemoryConfig | None, namespace: str = "global"
|
||||
) -> MemoryScope | None:
|
||||
"""根据声明式配置动态组装长期向量记忆实例"""
|
||||
if not config or not config.long_term.enable:
|
||||
return None
|
||||
|
||||
if config.long_term.engine is not None:
|
||||
return MemoryScope(
|
||||
rag_client=config.long_term.engine,
|
||||
)
|
||||
|
||||
storage_instance = None
|
||||
backend_cfg = config.long_term.backend
|
||||
if backend_cfg is not None:
|
||||
storage_instance = cast(StorageBackend, backend_cfg)
|
||||
else:
|
||||
factory = (
|
||||
self._storage_factories.get(namespace)
|
||||
or self._storage_factories["global"]
|
||||
)
|
||||
storage_instance = factory()
|
||||
|
||||
embedder = self.get_embedder(config.long_term.embedder)
|
||||
|
||||
from zhenxun.services.ai.context.rag.builder import RAGBuilder
|
||||
|
||||
builder = RAGBuilder(storage_instance).with_scope("/")
|
||||
if embedder:
|
||||
builder.with_embedder(embedder)
|
||||
|
||||
from .models import MemoryScoringConfig
|
||||
|
||||
scoring_cfg = MemoryScoringConfig()
|
||||
|
||||
builder.enable_lifecycle_scoring(
|
||||
half_life_days=scoring_cfg.recency_half_life_days,
|
||||
decay_weight=scoring_cfg.recency_weight,
|
||||
semantic_weight=scoring_cfg.semantic_weight,
|
||||
importance_weight=scoring_cfg.importance_weight,
|
||||
reinforcement_weight=scoring_cfg.reinforcement_weight,
|
||||
)
|
||||
|
||||
client = builder.build()
|
||||
|
||||
return MemoryScope(
|
||||
rag_client=client,
|
||||
)
|
||||
return self.get_backend("chat", namespace)
|
||||
|
||||
|
||||
memory_manager = GlobalMemoryManager()
|
||||
|
||||
@@ -2,50 +2,21 @@
|
||||
记忆域类型定义
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend
|
||||
from zhenxun.services.ai.context.rag.engine import ScopedRAGClient
|
||||
from zhenxun.services.ai.utils.scope import ScopeBuilder
|
||||
|
||||
from .storage.interfaces import (
|
||||
BaseChatContext,
|
||||
BaseMemoryIngestionMiddleware,
|
||||
BaseMemoryReducer,
|
||||
BaseSlotContext,
|
||||
)
|
||||
from .types import (
|
||||
AutoRecallPolicy,
|
||||
Isolation,
|
||||
MemorySlot,
|
||||
SessionMetadata,
|
||||
)
|
||||
|
||||
|
||||
class SlotMemoryConfig(BaseModel):
|
||||
"""槽位记忆 (Memory Slots) 配置"""
|
||||
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
enable: bool = Field(default=False)
|
||||
"""是否启用中期记忆槽"""
|
||||
scopes: dict[str, ScopeBuilder] | None = Field(default=None)
|
||||
"""语义化作用域映射字典,供大模型作为 Literal 选择。如果只有一个,则自动隐藏参数"""
|
||||
default_slots: list[MemorySlot] = Field(default_factory=list)
|
||||
"""首次初始化时自动写入的默认槽位列表"""
|
||||
backend: str | BaseSlotContext | None = Field(default=None)
|
||||
"""
|
||||
指定底层槽位记忆数据库注册名称,或直接传入 BaseSlotContext 实例。
|
||||
为空则使用全局默认
|
||||
"""
|
||||
instructions: str | None = Field(default=None)
|
||||
"""覆写内置槽位管理工具箱的系统提示词"""
|
||||
toolkit_kwargs: dict[str, Any] = Field(default_factory=dict)
|
||||
"""透传给底层 MemorySlotToolkit 的高级参数 (如 prefix, exclude, shared_options)"""
|
||||
|
||||
|
||||
class MemoryScoringConfig(BaseModel):
|
||||
"""长期记忆的复合打分与检索配置"""
|
||||
|
||||
@@ -57,7 +28,6 @@ class MemoryScoringConfig(BaseModel):
|
||||
"""重要性权重"""
|
||||
recency_half_life_days: int = Field(default=30)
|
||||
"""时间衰减的半衰期(天)"""
|
||||
|
||||
reinforcement_weight: float = Field(default=0.2)
|
||||
"""访问强化的加权权重 (被检索越多得分越高)"""
|
||||
|
||||
@@ -78,44 +48,6 @@ class ShortTermConfig(BaseModel):
|
||||
"""单一的记忆隔离级别 (ScopeBuilder),决定短期记忆存储边界"""
|
||||
|
||||
|
||||
class LongTermConfig(BaseModel):
|
||||
"""长期向量记忆配置"""
|
||||
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
enable: bool = Field(default=False)
|
||||
"""是否启用长期记忆(开启后自动赋予 Agent 存取记忆的工具,并附加 RAG 召回能力)"""
|
||||
engine: ScopedRAGClient | None = Field(default=None)
|
||||
"""
|
||||
[推荐] 指定底层的高级 RAG 检索引擎实例。若传入此项,将覆盖默认的 backend
|
||||
和 embedder 配置。
|
||||
"""
|
||||
backend: str | StorageBackend | None = Field(default=None)
|
||||
"""
|
||||
指定底层长期向量数据库 (Storage) 注册名称,或直接传入 StorageBackend 实例。
|
||||
为空则使用全局默认
|
||||
"""
|
||||
scopes: dict[str, ScopeBuilder] | None = Field(default=None)
|
||||
"""语义化作用域映射字典,决定长期记忆存储边界。如果只有一个,则自动隐藏参数"""
|
||||
embedder: str | Embedder | None = Field(default=None)
|
||||
"""
|
||||
指定底层向量化引擎 (Embedder) 实例,若为字符串则视为 API 模型名称。
|
||||
为空则使用全局默认
|
||||
"""
|
||||
|
||||
agentic: bool = Field(default=True)
|
||||
"""是否赋予大模型主动管理记忆的能力 (Agentic Memory)"""
|
||||
auto_recall: AutoRecallPolicy = Field(default=False)
|
||||
"""长期记忆的自动召回策略,默认 False (从不自动召回),由大模型自主
|
||||
决定调用搜索工具"""
|
||||
recall_threshold: float = Field(default=0.5)
|
||||
"""长期记忆召回的最低余弦相似度要求"""
|
||||
instructions: str | None = Field(default=None)
|
||||
"""覆写内置长期记忆管理工具箱的系统提示词"""
|
||||
toolkit_kwargs: dict[str, Any] = Field(default_factory=dict)
|
||||
"""透传给底层 MemoryManagementToolkit 的高级参数"""
|
||||
|
||||
|
||||
class ContextCompressionConfig(BaseModel):
|
||||
"""上下文压缩与管理配置"""
|
||||
|
||||
@@ -144,14 +76,8 @@ class MemoryConfig(BaseModel):
|
||||
"""统一的记忆配置项声明"""
|
||||
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
base_isolation: ScopeBuilder = Field(default_factory=Isolation.AGENT_USER)
|
||||
"""顶层基准隔离级别,短期/中期/长期记忆将默认继承此级别"""
|
||||
short_term: ShortTermConfig = Field(default_factory=ShortTermConfig)
|
||||
"""短期对话记忆配置"""
|
||||
slots: SlotMemoryConfig = Field(default_factory=SlotMemoryConfig)
|
||||
"""槽位记忆配置"""
|
||||
long_term: LongTermConfig = Field(default_factory=LongTermConfig)
|
||||
"""长期向量记忆配置"""
|
||||
compression: ContextCompressionConfig = Field(
|
||||
default_factory=ContextCompressionConfig
|
||||
)
|
||||
@@ -161,12 +87,10 @@ class MemoryConfig(BaseModel):
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AutoRecallPolicy",
|
||||
"BaseMemoryIngestionMiddleware",
|
||||
"ContextCompressionConfig",
|
||||
"IngestionConfig",
|
||||
"Isolation",
|
||||
"LongTermConfig",
|
||||
"MemoryConfig",
|
||||
"MemoryScoringConfig",
|
||||
"SessionMetadata",
|
||||
|
||||
@@ -95,7 +95,9 @@ class DBMessageSerializer:
|
||||
return content_parts
|
||||
|
||||
@staticmethod
|
||||
def serialize_content(content_payload: Any) -> list[dict[str, Any]]:
|
||||
def serialize_content(
|
||||
content_payload: list[LLMContentPart] | str,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""将 LLMMessage 消息内容序列化为可存储于数据库的 JSON 格式"""
|
||||
if isinstance(content_payload, str):
|
||||
return [{"type": "text", "text": content_payload}]
|
||||
@@ -134,7 +136,7 @@ class MemoryScope:
|
||||
):
|
||||
"""初始化长期记忆作用域与 RAG 客户端"""
|
||||
self.rag_client = rag_client
|
||||
self._background_tasks: set[Any] = set()
|
||||
self._background_tasks: set[asyncio.Task[Any]] = set()
|
||||
|
||||
async def remember(
|
||||
self,
|
||||
@@ -492,3 +494,13 @@ def get_orm_slot_context(model_class: type[AbstractSlotRecord]) -> TortoiseSlotC
|
||||
[工厂方法] 供第三方开发者调用,将 Tortoise ORM 表直接包装为记忆槽存储系统。
|
||||
"""
|
||||
return TortoiseSlotContext(model_class=model_class)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AbstractMemoryRecord",
|
||||
"AbstractSlotRecord",
|
||||
"InMemoryChatContext",
|
||||
"MemoryScope",
|
||||
"TortoiseChatContext",
|
||||
"TortoiseSlotContext",
|
||||
]
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Sequence
|
||||
from typing import Protocol, runtime_checkable
|
||||
|
||||
from zhenxun.services.ai.context.memory.types import (
|
||||
MemorySlot,
|
||||
@@ -10,7 +11,14 @@ from zhenxun.services.ai.run.context import RunContext
|
||||
from zhenxun.services.ai.utils.scope import ScopeSelector
|
||||
|
||||
|
||||
class BaseChatContext(ABC):
|
||||
@runtime_checkable
|
||||
class IClearableBackend(Protocol):
|
||||
"""支持声明式清理的作用域后端协议"""
|
||||
|
||||
async def clear_by_query(self, query: ScopeSelector) -> int | None: ...
|
||||
|
||||
|
||||
class BaseChatContext(IClearableBackend, ABC):
|
||||
"""短期对话历史记忆接口"""
|
||||
|
||||
@abstractmethod
|
||||
@@ -44,13 +52,8 @@ class BaseChatContext(ABC):
|
||||
"""清空当前会话的历史消息。"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def clear_by_query(self, query: ScopeSelector) -> None:
|
||||
"""根据条件领域查询对象清理对话历史。"""
|
||||
...
|
||||
|
||||
|
||||
class BaseSlotContext(ABC):
|
||||
class BaseSlotContext(IClearableBackend, ABC):
|
||||
"""中期记忆槽持久化接口"""
|
||||
|
||||
@abstractmethod
|
||||
@@ -80,11 +83,6 @@ class BaseSlotContext(ABC):
|
||||
"""列出当前会话的所有记忆槽(包括未置顶的)。"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def clear_by_query(self, query: ScopeSelector) -> None:
|
||||
"""根据条件领域查询对象清理记忆槽。"""
|
||||
...
|
||||
|
||||
|
||||
class BaseMemoryReducer(ABC):
|
||||
"""记忆压缩器基类"""
|
||||
@@ -97,7 +95,19 @@ class BaseMemoryReducer(ABC):
|
||||
model_name: str,
|
||||
base_overhead: int = 0,
|
||||
) -> tuple[list[LLMMessage], bool, int]:
|
||||
"""对消息列表进行压缩处理。"""
|
||||
"""
|
||||
执行记忆压缩处理,精简或提炼对话上下文以降低 Token 消耗。
|
||||
|
||||
参数:
|
||||
messages: 需要进行压缩的原始 LLM 消息历史列表。
|
||||
current_tokens: 压缩前消息列表的当前 Token 总数。
|
||||
model_name: 用于判定压缩阈值或计算 Token 的底层大模型名称。
|
||||
base_overhead: 基础系统提示词等静态开销的 Token 计数。
|
||||
|
||||
返回:
|
||||
tuple[list[LLMMessage], bool, int]: 包含压缩后的新消息历史列表、
|
||||
本次是否实际触发了压缩的布尔标记、以及压缩后的新 Token 总数。
|
||||
"""
|
||||
...
|
||||
|
||||
|
||||
|
||||
@@ -4,6 +4,7 @@ import threading
|
||||
from typing import Any, Literal, Protocol, runtime_checkable
|
||||
|
||||
from zhenxun.services.ai.core.messages import EmbedBatch
|
||||
from zhenxun.services.ai.core.options import LLMEmbeddingConfig
|
||||
from zhenxun.services.ai.llm.api import embed as api_embed
|
||||
from zhenxun.services.ai.message_builder import MessageBuilder
|
||||
from zhenxun.services.ai.utils.logger import log_rag as logger
|
||||
@@ -16,7 +17,7 @@ EmbedTaskType = Literal[
|
||||
@runtime_checkable
|
||||
class Embedder(Protocol):
|
||||
"""
|
||||
向量化引擎协议 (Callable Protocol)。
|
||||
向量化引擎协议。
|
||||
任何实现了异步 __call__ 的对象或闭包函数均可作为 Embedder。
|
||||
"""
|
||||
|
||||
@@ -32,7 +33,9 @@ class Embedder(Protocol):
|
||||
class DefaultEmbedder(Embedder):
|
||||
"""系统默认的向量化引擎,调用大模型底座 API"""
|
||||
|
||||
def __init__(self, model_name: str | None = None, config: Any = None):
|
||||
def __init__(
|
||||
self, model_name: str | None = None, config: LLMEmbeddingConfig | None = None
|
||||
):
|
||||
self.model_name = model_name
|
||||
self.config = config
|
||||
|
||||
@@ -57,10 +60,28 @@ class BaseLocalEmbedder(Embedder, ABC):
|
||||
def __init__(self, model_name: str):
|
||||
self.model_name = model_name
|
||||
self._model_lock = threading.Lock()
|
||||
self._model: Any | None = None
|
||||
|
||||
def _ensure_model_loaded(self) -> None:
|
||||
"""线程安全的懒加载机制"""
|
||||
if self._model is None:
|
||||
with self._model_lock:
|
||||
if self._model is None:
|
||||
logger.info(
|
||||
f"正在后台加载本地向量模型: {self.model_name} ... "
|
||||
"(首次加载可能需要较长时间下载)"
|
||||
)
|
||||
self._model = self._load_model_impl()
|
||||
logger.info(f"本地向量模型 {self.model_name} 加载完毕!")
|
||||
|
||||
@abstractmethod
|
||||
def _load_model_impl(self) -> Any:
|
||||
"""子类实现:执行具体的依赖导入与模型实例化,并返回模型对象。"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def _encode_texts(self, texts: list[str]) -> list[list[float]]:
|
||||
"""子类只需实现此同步的批量文本向量化方法即可。"""
|
||||
"""子类实现:执行同步的批量文本向量化方法。"""
|
||||
pass
|
||||
|
||||
async def __call__(
|
||||
@@ -93,7 +114,6 @@ class FastEmbedder(BaseLocalEmbedder):
|
||||
|
||||
def __init__(self, model_name: str | None = None):
|
||||
super().__init__(model_name or "BAAI/bge-small-zh-v1.5")
|
||||
self.model = None
|
||||
|
||||
import importlib.util
|
||||
|
||||
@@ -102,29 +122,19 @@ class FastEmbedder(BaseLocalEmbedder):
|
||||
"⚠️ 使用 FastEmbed 需要额外依赖,请在终端执行: pip install fastembed"
|
||||
)
|
||||
|
||||
def _ensure_model_loaded(self):
|
||||
"""线程安全的懒加载机制"""
|
||||
if self.model is None:
|
||||
with self._model_lock:
|
||||
if self.model is None:
|
||||
try:
|
||||
from fastembed import TextEmbedding
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"⚠️ 使用 FastEmbed 需要额外依赖,"
|
||||
"请在终端执行: pip install fastembed"
|
||||
)
|
||||
logger.info(
|
||||
f"正在后台加载 FastEmbed 本地模型: {self.model_name} ... "
|
||||
"(首次加载可能需要极长时间下载)"
|
||||
)
|
||||
self.model = TextEmbedding(model_name=self.model_name)
|
||||
logger.info(f"FastEmbed 模型 {self.model_name} 加载完毕!")
|
||||
def _load_model_impl(self) -> Any:
|
||||
try:
|
||||
from fastembed import TextEmbedding
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"⚠️ 使用 FastEmbed 需要额外依赖,请在终端执行: pip install fastembed"
|
||||
)
|
||||
return TextEmbedding(model_name=self.model_name)
|
||||
|
||||
def _encode_texts(self, texts: list[str]) -> list[list[float]]:
|
||||
self._ensure_model_loaded()
|
||||
assert self.model is not None
|
||||
return [vec.tolist() for vec in self.model.embed(texts)]
|
||||
assert self._model is not None
|
||||
return [vec.tolist() for vec in self._model.embed(texts)]
|
||||
|
||||
|
||||
class SentenceTransformerEmbedder(BaseLocalEmbedder):
|
||||
@@ -135,7 +145,6 @@ class SentenceTransformerEmbedder(BaseLocalEmbedder):
|
||||
|
||||
def __init__(self, model_name: str | None = None):
|
||||
super().__init__(model_name or "BAAI/bge-small-zh-v1.5")
|
||||
self.model = None
|
||||
|
||||
import importlib.util
|
||||
|
||||
@@ -145,32 +154,18 @@ class SentenceTransformerEmbedder(BaseLocalEmbedder):
|
||||
"请在终端执行: pip install sentence-transformers"
|
||||
)
|
||||
|
||||
def _ensure_model_loaded(self):
|
||||
"""线程安全的懒加载机制"""
|
||||
if self.model is None:
|
||||
with self._model_lock:
|
||||
if self.model is None:
|
||||
try:
|
||||
from sentence_transformers import (
|
||||
SentenceTransformer,
|
||||
)
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"⚠️ 使用 SentenceTransformers 需要额外依赖,"
|
||||
"请在终端执行: pip install sentence-transformers"
|
||||
)
|
||||
logger.info(
|
||||
"正在后台加载 SentenceTransformer "
|
||||
f"本地模型: {self.model_name} ... "
|
||||
"(首次加载可能需要极长时间下载)"
|
||||
)
|
||||
self.model = SentenceTransformer(self.model_name)
|
||||
logger.info(
|
||||
f"SentenceTransformer 模型 {self.model_name} 加载完毕!"
|
||||
)
|
||||
def _load_model_impl(self) -> Any:
|
||||
try:
|
||||
from sentence_transformers import SentenceTransformer
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"⚠️ 使用 SentenceTransformers 需要额外依赖,"
|
||||
"请在终端执行: pip install sentence-transformers"
|
||||
)
|
||||
return SentenceTransformer(self.model_name)
|
||||
|
||||
def _encode_texts(self, texts: list[str]) -> list[list[float]]:
|
||||
self._ensure_model_loaded()
|
||||
assert self.model is not None
|
||||
embeddings = self.model.encode(texts)
|
||||
assert self._model is not None
|
||||
embeddings = self._model.encode(texts)
|
||||
return embeddings.tolist()
|
||||
|
||||
@@ -11,6 +11,10 @@ from zhenxun.services.ai.context.rag.models import (
|
||||
SearchResult,
|
||||
)
|
||||
from zhenxun.services.ai.context.rag.retrieval import FilterEvaluator
|
||||
from zhenxun.services.ai.context.rag.utils import (
|
||||
InMemoryScorer,
|
||||
normalize_vector,
|
||||
)
|
||||
from zhenxun.services.ai.utils.logger import log_rag as logger
|
||||
from zhenxun.services.ai.utils.scope import ScopeSelector
|
||||
from zhenxun.services.db_context import Model
|
||||
@@ -24,9 +28,7 @@ class StorageBackend(Protocol):
|
||||
"""保存或更新数据块"""
|
||||
...
|
||||
|
||||
async def search(
|
||||
self, query: QueryRequest, scopes: list[str] | None = None
|
||||
) -> list[SearchResult]:
|
||||
async def search(self, request: QueryRequest) -> list[SearchResult]:
|
||||
"""按向量和前缀检索数据块"""
|
||||
...
|
||||
|
||||
@@ -49,22 +51,6 @@ class StorageBackend(Protocol):
|
||||
...
|
||||
|
||||
|
||||
def normalize_vector(vec: list[float] | np.ndarray) -> np.ndarray:
|
||||
"""将一维向量转化为 float32 数组并进行 L2 归一化"""
|
||||
v = np.array(vec, dtype=np.float32)
|
||||
norm = np.linalg.norm(v)
|
||||
if norm == 0:
|
||||
return v
|
||||
return v / norm
|
||||
|
||||
|
||||
def normalize_matrix(mat: np.ndarray) -> np.ndarray:
|
||||
"""将二维矩阵的每一行进行 L2 归一化"""
|
||||
norms = np.linalg.norm(mat, axis=1, keepdims=True)
|
||||
norms[norms == 0] = 1.0
|
||||
return mat / norms
|
||||
|
||||
|
||||
class DictStorageBackend(StorageBackend):
|
||||
"""基于内存字典的轻量级纯净 RAG 存储实现"""
|
||||
|
||||
@@ -83,62 +69,36 @@ class DictStorageBackend(StorageBackend):
|
||||
else:
|
||||
self._vectors.pop(r.id, None)
|
||||
|
||||
async def search(
|
||||
self, query: QueryRequest, scopes: list[str] | None = None
|
||||
) -> list[SearchResult]:
|
||||
async def search(self, request: QueryRequest) -> list[SearchResult]:
|
||||
candidate_ids = []
|
||||
for record in self._records.values():
|
||||
if scopes is not None:
|
||||
if record.metadata.get("scope", "/") not in scopes:
|
||||
if request.scopes is not None:
|
||||
if record.metadata.get("scope", "/") not in request.scopes:
|
||||
continue
|
||||
if not FilterEvaluator.evaluate(record.metadata, query.metadata_filters):
|
||||
if not FilterEvaluator.evaluate(record.metadata, request.metadata_filters):
|
||||
continue
|
||||
if not query.embedding and query.text and query.text not in record.content:
|
||||
if (
|
||||
not request.embedding
|
||||
and request.text
|
||||
and request.text not in record.content
|
||||
):
|
||||
continue
|
||||
candidate_ids.append(record.id)
|
||||
|
||||
if not candidate_ids:
|
||||
return []
|
||||
|
||||
results = []
|
||||
if query.search_type == "sparse":
|
||||
import jieba
|
||||
records = [self._records[r_id] for r_id in candidate_ids]
|
||||
|
||||
tokens = set(jieba.lcut_for_search(query.text.lower()))
|
||||
for r_id in candidate_ids:
|
||||
record = self._records[r_id]
|
||||
content = record.content.lower()
|
||||
matched_count = sum(1 for t in tokens if t in content)
|
||||
if matched_count > 0:
|
||||
score = matched_count / len(tokens)
|
||||
results.append(SearchResult(record=record, score=score))
|
||||
elif query.search_type == "dense" and query.embedding:
|
||||
q_vec = normalize_vector(query.embedding)
|
||||
valid_ids = [r_id for r_id in candidate_ids if r_id in self._vectors]
|
||||
|
||||
if valid_ids:
|
||||
try:
|
||||
mat = np.array([self._vectors[r_id] for r_id in valid_ids])
|
||||
scores = mat @ q_vec
|
||||
for r_id, score in zip(valid_ids, scores):
|
||||
results.append(
|
||||
SearchResult(record=self._records[r_id], score=float(score))
|
||||
)
|
||||
except ValueError as e:
|
||||
logger.warning(
|
||||
"⚠️ DictStorage 中缓存的向量维度与当前查询维度不匹配,"
|
||||
f"跳过向量检索。原因: {e}"
|
||||
)
|
||||
|
||||
missing_ids = [r_id for r_id in candidate_ids if r_id not in self._vectors]
|
||||
for r_id in missing_ids:
|
||||
results.append(SearchResult(record=self._records[r_id], score=0.1))
|
||||
if request.search_type == "sparse":
|
||||
results = InMemoryScorer.calculate_sparse_scores(request.text, records)
|
||||
elif request.search_type == "dense" and request.embedding:
|
||||
results = InMemoryScorer.calculate_dense_scores(request.embedding, records)
|
||||
else:
|
||||
for r_id in candidate_ids:
|
||||
results.append(SearchResult(record=self._records[r_id], score=0.1))
|
||||
results = [SearchResult(record=r, score=0.1) for r in records]
|
||||
|
||||
results.sort(key=lambda x: x.score, reverse=True)
|
||||
return results[: query.limit]
|
||||
return results[: request.limit]
|
||||
|
||||
async def update(self, record: BaseRecord) -> None:
|
||||
if record.id in self._records:
|
||||
@@ -212,94 +172,48 @@ class TortoiseStorageBackend(StorageBackend):
|
||||
},
|
||||
)
|
||||
|
||||
async def search(
|
||||
self, query: QueryRequest, scopes: list[str] | None = None
|
||||
) -> list[SearchResult]:
|
||||
async def search(self, request: QueryRequest) -> list[SearchResult]:
|
||||
query_orm = self.model_class.all()
|
||||
if scopes is not None:
|
||||
query_orm = query_orm.filter(scope__in=scopes)
|
||||
if request.scopes is not None:
|
||||
query_orm = query_orm.filter(scope__in=request.scopes)
|
||||
|
||||
if query.search_type == "sparse" and query.text:
|
||||
if request.search_type == "sparse" and request.text:
|
||||
import jieba
|
||||
from tortoise.expressions import Q
|
||||
|
||||
tokens = [
|
||||
t for t in jieba.lcut_for_search(query.text) if len(t.strip()) > 1
|
||||
] or [query.text]
|
||||
t for t in jieba.lcut_for_search(request.text) if len(t.strip()) > 1
|
||||
] or [request.text]
|
||||
q_expr = Q()
|
||||
for token in tokens:
|
||||
q_expr |= Q(content__icontains=token)
|
||||
query_orm = query_orm.filter(q_expr)
|
||||
elif query.search_type == "dense" and not query.embedding and query.text:
|
||||
query_orm = query_orm.filter(content__icontains=query.text)
|
||||
elif request.search_type == "dense" and not request.embedding and request.text:
|
||||
query_orm = query_orm.filter(content__icontains=request.text)
|
||||
|
||||
rows = await query_orm
|
||||
|
||||
valid_rows = []
|
||||
for row in rows:
|
||||
row_meta = row.meta_data if isinstance(row.meta_data, dict) else {}
|
||||
if not FilterEvaluator.evaluate(row_meta, query.metadata_filters):
|
||||
if not FilterEvaluator.evaluate(row_meta, request.metadata_filters):
|
||||
continue
|
||||
valid_rows.append(row)
|
||||
|
||||
if not valid_rows:
|
||||
return []
|
||||
|
||||
results = []
|
||||
if query.search_type == "sparse":
|
||||
import jieba
|
||||
records = [self._to_base_record(row) for row in valid_rows]
|
||||
|
||||
tokens = set(jieba.lcut_for_search(query.text.lower()))
|
||||
for row in valid_rows:
|
||||
content = row.content.lower()
|
||||
matched_count = sum(1 for t in tokens if t in content)
|
||||
score = matched_count / len(tokens) if tokens else 0.1
|
||||
results.append(
|
||||
SearchResult(record=self._to_base_record(row), score=score)
|
||||
)
|
||||
elif query.search_type == "dense" and query.embedding:
|
||||
q_vec = normalize_vector(query.embedding)
|
||||
vec_rows = []
|
||||
missing_rows = []
|
||||
|
||||
for row in valid_rows:
|
||||
if isinstance(row.embedding, list):
|
||||
vec_rows.append(row)
|
||||
else:
|
||||
missing_rows.append(row)
|
||||
|
||||
if vec_rows:
|
||||
try:
|
||||
raw_mat = np.array(
|
||||
[r.embedding for r in vec_rows], dtype=np.float32
|
||||
)
|
||||
norm_mat = normalize_matrix(raw_mat)
|
||||
scores = norm_mat @ q_vec
|
||||
|
||||
for row, score in zip(vec_rows, scores):
|
||||
results.append(
|
||||
SearchResult(
|
||||
record=self._to_base_record(row), score=float(score)
|
||||
)
|
||||
)
|
||||
except ValueError as e:
|
||||
logger.warning(
|
||||
"⚠️ 数据库中缓存的向量维度与当前模型查询维度不匹配,"
|
||||
f"已安全跳过向量检索(降级为稀疏匹配)。原因: {e}"
|
||||
)
|
||||
|
||||
for row in missing_rows:
|
||||
results.append(
|
||||
SearchResult(record=self._to_base_record(row), score=0.1)
|
||||
)
|
||||
if request.search_type == "sparse":
|
||||
results = InMemoryScorer.calculate_sparse_scores(request.text, records)
|
||||
elif request.search_type == "dense" and request.embedding:
|
||||
results = InMemoryScorer.calculate_dense_scores(request.embedding, records)
|
||||
else:
|
||||
for row in valid_rows:
|
||||
results.append(
|
||||
SearchResult(record=self._to_base_record(row), score=0.1)
|
||||
)
|
||||
results = [SearchResult(record=r, score=0.1) for r in records]
|
||||
|
||||
results.sort(key=lambda x: x.score, reverse=True)
|
||||
return results[: query.limit]
|
||||
return results[: request.limit]
|
||||
|
||||
async def update(self, record: BaseRecord) -> None:
|
||||
await self.model_class.filter(id=record.id).update(
|
||||
@@ -386,50 +300,50 @@ class QdrantStorageBackend(StorageBackend):
|
||||
)
|
||||
await self.client.upsert(collection_name=self.collection_name, points=points)
|
||||
|
||||
async def search(
|
||||
self, query: QueryRequest, scopes: list[str] | None = None
|
||||
) -> list[SearchResult]:
|
||||
if query.search_type == "dense" and not query.embedding:
|
||||
async def search(self, request: QueryRequest) -> list[SearchResult]:
|
||||
if request.search_type == "dense" and not request.embedding:
|
||||
return []
|
||||
if query.embedding:
|
||||
await self._ensure_collection(len(query.embedding))
|
||||
if request.embedding:
|
||||
await self._ensure_collection(len(request.embedding))
|
||||
|
||||
from qdrant_client.models import FieldCondition, Filter, MatchText, MatchValue
|
||||
|
||||
must_conditions = []
|
||||
|
||||
if scopes is not None:
|
||||
if request.scopes is not None:
|
||||
try:
|
||||
from qdrant_client.models import MatchAny
|
||||
|
||||
must_conditions.append(
|
||||
FieldCondition(key="metadata.scope", match=MatchAny(any=scopes))
|
||||
FieldCondition(
|
||||
key="metadata.scope", match=MatchAny(any=request.scopes)
|
||||
)
|
||||
)
|
||||
except ImportError:
|
||||
scope_conditions = [
|
||||
FieldCondition(key="metadata.scope", match=MatchValue(value=s))
|
||||
for s in scopes
|
||||
for s in request.scopes
|
||||
]
|
||||
must_conditions.append(Filter(should=scope_conditions))
|
||||
|
||||
if query.metadata_filters:
|
||||
for k, v in query.metadata_filters.items():
|
||||
if request.metadata_filters:
|
||||
for k, v in request.metadata_filters.items():
|
||||
must_conditions.append(
|
||||
FieldCondition(key=f"metadata.{k}", match=MatchValue(value=v))
|
||||
)
|
||||
|
||||
if query.search_type == "sparse":
|
||||
if request.search_type == "sparse":
|
||||
must_conditions.append(
|
||||
FieldCondition(key="content", match=MatchText(text=query.text))
|
||||
FieldCondition(key="content", match=MatchText(text=request.text))
|
||||
)
|
||||
|
||||
query_filter = Filter(must=must_conditions) if must_conditions else None
|
||||
|
||||
if query.search_type == "sparse":
|
||||
if request.search_type == "sparse":
|
||||
results = await self.client.scroll(
|
||||
collection_name=self.collection_name,
|
||||
scroll_filter=query_filter,
|
||||
limit=query.limit,
|
||||
limit=request.limit,
|
||||
with_payload=True,
|
||||
)
|
||||
return [
|
||||
@@ -446,8 +360,8 @@ class QdrantStorageBackend(StorageBackend):
|
||||
|
||||
results = await self.client.search( # type: ignore
|
||||
collection_name=self.collection_name,
|
||||
query_vector=query.embedding,
|
||||
limit=query.limit,
|
||||
query_vector=request.embedding,
|
||||
limit=request.limit,
|
||||
query_filter=query_filter,
|
||||
)
|
||||
|
||||
@@ -559,27 +473,25 @@ class LanceDBStorageBackend(StorageBackend):
|
||||
else:
|
||||
self.db.open_table(self.table_name).add(data)
|
||||
|
||||
async def search(
|
||||
self, query: QueryRequest, scopes: list[str] | None = None
|
||||
) -> list[SearchResult]:
|
||||
async def search(self, request: QueryRequest) -> list[SearchResult]:
|
||||
if self.table_name not in self.db.table_names():
|
||||
return []
|
||||
if query.search_type == "dense" and not query.embedding:
|
||||
if request.search_type == "dense" and not request.embedding:
|
||||
return []
|
||||
|
||||
tbl = self.db.open_table(self.table_name)
|
||||
if query.search_type == "sparse":
|
||||
if request.search_type == "sparse":
|
||||
try:
|
||||
results = (
|
||||
tbl.search(query.text, query_type="fts")
|
||||
.limit(query.limit)
|
||||
tbl.search(request.text, query_type="fts")
|
||||
.limit(request.limit)
|
||||
.to_list()
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"LanceDB FTS 检索失败(可能是由于尚未创建FTS索引): {e}")
|
||||
return []
|
||||
else:
|
||||
results = tbl.search(query.embedding).limit(query.limit).to_list()
|
||||
results = tbl.search(request.embedding).limit(request.limit).to_list()
|
||||
|
||||
import ast
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ from typing import Any
|
||||
|
||||
from zhenxun.services.ai.utils.logger import log_rag as logger
|
||||
|
||||
from .backends import StorageBackend
|
||||
from .backends import Embedder, StorageBackend
|
||||
from .configs import RAGConfig
|
||||
from .engine import ScopedRAGClient
|
||||
from .ingestion import (
|
||||
@@ -29,18 +29,25 @@ from .retrieval import (
|
||||
|
||||
class RAGBuilder:
|
||||
"""
|
||||
RAG 管线组装构建器 (Fluent Builder Pattern)。
|
||||
使用内部状态驱动模式 (The Memory Pattern),内部维护私有的 RAGConfig 实例。
|
||||
RAG 管线组装构建器。
|
||||
使用内部状态驱动模式,内部维护私有的 RAGConfig 实例。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, storage: StorageBackend | None = None, config: RAGConfig | None = None
|
||||
):
|
||||
"""
|
||||
初始化 RAG 管线组装构建器。
|
||||
|
||||
参数:
|
||||
storage: (可选) 底层存储后端实例。
|
||||
config: (可选) RAG 配置对象。若不提供,将新建默认的 RAGConfig 实例。
|
||||
"""
|
||||
self._config = config or RAGConfig()
|
||||
if storage is not None:
|
||||
self._config.storage = storage
|
||||
|
||||
def with_embedder(self, embedder: Any) -> "RAGBuilder":
|
||||
def with_embedder(self, embedder: Embedder) -> "RAGBuilder":
|
||||
"""
|
||||
设置向量化引擎。
|
||||
|
||||
|
||||
@@ -11,11 +11,13 @@ from .ingestion import (
|
||||
)
|
||||
from .models import (
|
||||
BaseRecord,
|
||||
QueryRequest,
|
||||
SearchResult,
|
||||
)
|
||||
from .retrieval import (
|
||||
BaseRetriever,
|
||||
)
|
||||
from .utils import normalize_query_text
|
||||
|
||||
|
||||
class ScopedRAGClient:
|
||||
@@ -55,7 +57,15 @@ class ScopedRAGClient:
|
||||
self.scope_prefix = self.scopes[0] if self.scopes else "/"
|
||||
|
||||
async def ingest(self, records: list[BaseRecord]) -> int:
|
||||
"""通过 RAG Ingestion Pipeline 处理并入库数据"""
|
||||
"""
|
||||
通过 RAG Ingestion Pipeline 处理并入库数据。
|
||||
|
||||
参数:
|
||||
records: 待导入的数据记录列表。
|
||||
|
||||
返回:
|
||||
int: 成功导入并存入的记录数量。
|
||||
"""
|
||||
for r in records:
|
||||
if "scope" not in r.metadata:
|
||||
r.metadata["scope"] = self.scope_prefix
|
||||
@@ -73,7 +83,16 @@ class ScopedRAGClient:
|
||||
"""
|
||||
多作用域联合切片视图检索 (Union Search)。
|
||||
并发向多个独立的作用域发起检索,并对结果进行合并、去重和重排。
|
||||
"""
|
||||
|
||||
参数:
|
||||
query: 检索的查询对象,可以是文本字符串、向量或高级查询对象。
|
||||
limit: 最大返回结果条数。
|
||||
scopes: (可选) 自定义检索的作用域。若不提供,则使用客户端初始化时设定的 scopes。
|
||||
**kwargs: 传递给底层检索器的其他关键字参数。
|
||||
|
||||
返回:
|
||||
list[SearchResult]: 检索到的相似度排序后的结果列表。
|
||||
""" # noqa: E501
|
||||
target_scopes = self.scopes
|
||||
if scopes is not None:
|
||||
target_scopes = (
|
||||
@@ -85,13 +104,40 @@ class ScopedRAGClient:
|
||||
if not target_scopes:
|
||||
return []
|
||||
|
||||
kwargs["scopes"] = target_scopes
|
||||
return await self.retriever.retrieve(query, limit=limit, **kwargs)
|
||||
text_query = normalize_query_text(query)
|
||||
metadata_filters = kwargs.pop("metadata_filters", None)
|
||||
|
||||
request = QueryRequest(
|
||||
text=text_query,
|
||||
limit=limit,
|
||||
scopes=target_scopes,
|
||||
metadata_filters=metadata_filters,
|
||||
extra=kwargs,
|
||||
)
|
||||
|
||||
return await self.retriever.retrieve(request)
|
||||
|
||||
async def update(self, record: BaseRecord) -> None:
|
||||
"""
|
||||
更新指定的已有数据记录。
|
||||
会自动强行附加当前客户端的单作用域前缀 (scope_prefix)。
|
||||
|
||||
参数:
|
||||
record: 待更新的完整数据记录对象。
|
||||
"""
|
||||
record.metadata["scope"] = self.scope_prefix
|
||||
await self.storage.update(record)
|
||||
|
||||
async def delete(self, record_ids: list[str] | None = None, **kwargs: Any) -> int:
|
||||
"""
|
||||
删除指定的已有数据记录,操作限定在当前客户端的单作用域前缀下。
|
||||
|
||||
参数:
|
||||
record_ids: (可选) 待删除的记录 ID 列表。
|
||||
**kwargs: 传递给底层存储后端的其他过滤或删除参数。
|
||||
|
||||
返回:
|
||||
int: 成功删除的记录条数。
|
||||
"""
|
||||
kwargs["scope_prefix"] = self.scope_prefix
|
||||
return await self.storage.delete(record_ids=record_ids, **kwargs)
|
||||
|
||||
@@ -42,6 +42,10 @@ class QueryRequest(BaseModel):
|
||||
"""元数据精确匹配字典"""
|
||||
limit: int = Field(default=10)
|
||||
"""返回的最大条数"""
|
||||
scopes: list[str] | None = Field(default=None)
|
||||
"""检索的数据隔离作用域列表"""
|
||||
extra: dict[str, Any] = Field(default_factory=dict)
|
||||
"""透传参数逃生舱"""
|
||||
|
||||
|
||||
StorageConfigType = dict[str, Any]
|
||||
|
||||
@@ -4,22 +4,15 @@ import time
|
||||
from typing import TYPE_CHECKING, Any, Protocol, cast, runtime_checkable
|
||||
|
||||
from zhenxun.services.ai.utils.logger import log_rag as logger
|
||||
from zhenxun.utils.pydantic_compat import model_copy
|
||||
|
||||
from .backends.embedders import Embedder
|
||||
from .models import QueryRequest, SearchResult
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .backends.storages import StorageBackend
|
||||
|
||||
|
||||
def normalize_query_text(query: Any) -> str:
|
||||
"""辅助函数:提取各种输入形式(如字符串、平台Message对象)的纯文本用于检索"""
|
||||
if isinstance(query, str):
|
||||
return query
|
||||
if hasattr(query, "extract_plain_text"):
|
||||
return query.extract_plain_text()
|
||||
return str(query) if query is not None else ""
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class BaseRetriever(Protocol):
|
||||
"""
|
||||
@@ -29,9 +22,7 @@ class BaseRetriever(Protocol):
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def retrieve(
|
||||
self, query: Any, limit: int = 10, **kwargs: Any
|
||||
) -> list[SearchResult]: ...
|
||||
async def retrieve(self, request: QueryRequest) -> list[SearchResult]: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
@@ -72,7 +63,7 @@ class VectorDBRetriever(BaseRetriever):
|
||||
def __init__(
|
||||
self,
|
||||
storage: "StorageBackend",
|
||||
embedder: Any,
|
||||
embedder: Embedder,
|
||||
scope_prefix: str | None = None,
|
||||
score_threshold: float = 0.4,
|
||||
):
|
||||
@@ -90,29 +81,25 @@ class VectorDBRetriever(BaseRetriever):
|
||||
self.scope_prefix = scope_prefix
|
||||
self.score_threshold = score_threshold
|
||||
|
||||
async def retrieve(
|
||||
self, query: Any, limit: int = 10, **kwargs: Any
|
||||
) -> list[SearchResult]:
|
||||
text_query = normalize_query_text(query)
|
||||
async def retrieve(self, request: QueryRequest) -> list[SearchResult]:
|
||||
if not request.embedding and request.text:
|
||||
vecs = await self.embedder(request.text, task="query")
|
||||
request.embedding = vecs[0] if vecs else None
|
||||
|
||||
vecs = await self.embedder(query, task="query")
|
||||
query_vec = vecs[0] if vecs else None
|
||||
|
||||
if not text_query.strip() and not query_vec:
|
||||
if not request.text.strip() and not request.embedding:
|
||||
return []
|
||||
|
||||
req = QueryRequest(
|
||||
text=text_query,
|
||||
embedding=query_vec,
|
||||
limit=limit * 2,
|
||||
search_type="dense",
|
||||
metadata_filters=kwargs.get("metadata_filters"),
|
||||
)
|
||||
effective_scopes = kwargs.get(
|
||||
"scopes", [self.scope_prefix] if self.scope_prefix else None
|
||||
)
|
||||
results = await self.storage.search(req, scopes=effective_scopes)
|
||||
return [r for r in results if r.score >= self.score_threshold][:limit]
|
||||
request.search_type = "dense"
|
||||
if not request.scopes and self.scope_prefix:
|
||||
request.scopes = [self.scope_prefix]
|
||||
|
||||
original_limit = request.limit
|
||||
request.limit = original_limit * 2
|
||||
|
||||
results = await self.storage.search(request)
|
||||
|
||||
request.limit = original_limit
|
||||
return [r for r in results if r.score >= self.score_threshold][: request.limit]
|
||||
|
||||
|
||||
class DatabaseSparseRetriever(BaseRetriever):
|
||||
@@ -136,24 +123,20 @@ class DatabaseSparseRetriever(BaseRetriever):
|
||||
self.scope_prefix = scope_prefix
|
||||
self.score_threshold = score_threshold
|
||||
|
||||
async def retrieve(
|
||||
self, query: Any, limit: int = 10, **kwargs: Any
|
||||
) -> list[SearchResult]:
|
||||
text_query = normalize_query_text(query)
|
||||
if not text_query.strip():
|
||||
async def retrieve(self, request: QueryRequest) -> list[SearchResult]:
|
||||
if not request.text.strip():
|
||||
return []
|
||||
|
||||
req = QueryRequest(
|
||||
text=text_query,
|
||||
limit=limit * 2,
|
||||
search_type="sparse",
|
||||
metadata_filters=kwargs.get("metadata_filters"),
|
||||
)
|
||||
effective_scopes = kwargs.get(
|
||||
"scopes", [self.scope_prefix] if self.scope_prefix else None
|
||||
)
|
||||
results = await self.storage.search(req, scopes=effective_scopes)
|
||||
return [r for r in results if r.score > self.score_threshold][:limit]
|
||||
request.search_type = "sparse"
|
||||
if not request.scopes and self.scope_prefix:
|
||||
request.scopes = [self.scope_prefix]
|
||||
|
||||
original_limit = request.limit
|
||||
request.limit = original_limit * 2
|
||||
|
||||
results = await self.storage.search(request)
|
||||
request.limit = original_limit
|
||||
return [r for r in results if r.score > self.score_threshold][: request.limit]
|
||||
|
||||
|
||||
class RerankRetriever(BaseRetriever):
|
||||
@@ -183,19 +166,17 @@ class RerankRetriever(BaseRetriever):
|
||||
self.oversample_factor = oversample_factor
|
||||
self.min_oversample = min_oversample
|
||||
|
||||
async def retrieve(
|
||||
self, query: Any, limit: int = 10, **kwargs: Any
|
||||
) -> list[SearchResult]:
|
||||
oversample_limit = max(limit * self.oversample_factor, self.min_oversample)
|
||||
initial_results = await self.base_retriever.retrieve(
|
||||
query, limit=oversample_limit, **kwargs
|
||||
async def retrieve(self, request: QueryRequest) -> list[SearchResult]:
|
||||
req_clone = model_copy(request, deep=True)
|
||||
req_clone.limit = max(
|
||||
request.limit * self.oversample_factor, self.min_oversample
|
||||
)
|
||||
|
||||
initial_results = await self.base_retriever.retrieve(req_clone)
|
||||
|
||||
if not initial_results:
|
||||
return []
|
||||
|
||||
text_query = normalize_query_text(query)
|
||||
|
||||
docs: list[str | dict[str, str]] = [
|
||||
res.record.content for res in initial_results
|
||||
]
|
||||
@@ -204,14 +185,14 @@ class RerankRetriever(BaseRetriever):
|
||||
|
||||
try:
|
||||
reranked = await rerank(
|
||||
query=text_query,
|
||||
query=request.text,
|
||||
documents=docs,
|
||||
top_n=min(limit, self.top_n),
|
||||
top_n=min(request.limit, self.top_n),
|
||||
model=self.model_name,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"Rerank 重排请求失败,将降级返回初筛结果: {e}")
|
||||
return initial_results[:limit]
|
||||
return initial_results[: request.limit]
|
||||
|
||||
final_results = []
|
||||
for rr in reranked:
|
||||
@@ -243,15 +224,11 @@ class PipelineRetriever(BaseRetriever):
|
||||
self.post_processors = post_processors or []
|
||||
self.pre_processors = pre_processors or []
|
||||
|
||||
async def retrieve(
|
||||
self, query: Any, limit: int = 10, **kwargs: Any
|
||||
) -> list[SearchResult]:
|
||||
text_query = normalize_query_text(query)
|
||||
async def retrieve(self, request: QueryRequest) -> list[SearchResult]:
|
||||
requests_to_search = [request]
|
||||
|
||||
queries_to_search = [query]
|
||||
|
||||
if text_query.strip():
|
||||
processed_texts = [text_query]
|
||||
if request.text.strip():
|
||||
processed_texts = [request.text]
|
||||
for pp in self.pre_processors:
|
||||
new_texts = []
|
||||
for t in processed_texts:
|
||||
@@ -259,15 +236,23 @@ class PipelineRetriever(BaseRetriever):
|
||||
processed_texts = new_texts
|
||||
|
||||
if len(processed_texts) > 1 or (
|
||||
len(processed_texts) == 1 and processed_texts[0] != text_query
|
||||
len(processed_texts) == 1 and processed_texts[0] != request.text
|
||||
):
|
||||
queries_to_search.extend(processed_texts)
|
||||
requests_to_search = []
|
||||
for pt in processed_texts:
|
||||
new_req = model_copy(request, deep=True)
|
||||
new_req.text = pt
|
||||
requests_to_search.append(new_req)
|
||||
|
||||
all_results = []
|
||||
seen_ids = set()
|
||||
|
||||
for q in queries_to_search:
|
||||
res = await self.base_retriever.retrieve(q, limit=limit * 2, **kwargs)
|
||||
for req in requests_to_search:
|
||||
original_limit = req.limit
|
||||
req.limit = original_limit * 2
|
||||
res = await self.base_retriever.retrieve(req)
|
||||
req.limit = original_limit
|
||||
|
||||
for r in res:
|
||||
if r.record.id not in seen_ids:
|
||||
seen_ids.add(r.record.id)
|
||||
@@ -276,9 +261,9 @@ class PipelineRetriever(BaseRetriever):
|
||||
results = sorted(all_results, key=lambda x: x.score, reverse=True)
|
||||
|
||||
for pp in self.post_processors:
|
||||
results = await pp.process(results, query)
|
||||
results = await pp.process(results, request.text)
|
||||
|
||||
return results[:limit]
|
||||
return results[: request.limit]
|
||||
|
||||
|
||||
class LifecyclePostProcessor(PostProcessor):
|
||||
@@ -367,14 +352,22 @@ class HybridRetriever(BaseRetriever):
|
||||
self.oversample_factor = oversample_factor
|
||||
self.min_oversample = min_oversample
|
||||
|
||||
async def retrieve(
|
||||
self, query: Any, limit: int = 10, **kwargs: Any
|
||||
) -> list[SearchResult]:
|
||||
oversample_limit = max(limit * self.oversample_factor, self.min_oversample)
|
||||
async def retrieve(self, request: QueryRequest) -> list[SearchResult]:
|
||||
oversample_limit = max(
|
||||
request.limit * self.oversample_factor, self.min_oversample
|
||||
)
|
||||
|
||||
dense_req = model_copy(request, deep=True)
|
||||
dense_req.limit = oversample_limit
|
||||
dense_req.search_type = "dense"
|
||||
|
||||
sparse_req = model_copy(request, deep=True)
|
||||
sparse_req.limit = oversample_limit
|
||||
sparse_req.search_type = "sparse"
|
||||
|
||||
results = await asyncio.gather(
|
||||
self.dense_retriever.retrieve(query, limit=oversample_limit, **kwargs),
|
||||
self.sparse_retriever.retrieve(query, limit=oversample_limit, **kwargs),
|
||||
self.dense_retriever.retrieve(dense_req),
|
||||
self.sparse_retriever.retrieve(sparse_req),
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
@@ -435,6 +428,6 @@ class HybridRetriever(BaseRetriever):
|
||||
logger.debug(
|
||||
f"⚖️ [HybridSearch] 融合完成: "
|
||||
f"Dense({len(dense_res)}) + Sparse({len(sparse_res)}) "
|
||||
f"-> Merged({len(final_results)}), 截取 Top {limit}"
|
||||
f"-> Merged({len(final_results)}), 截取 Top {request.limit}"
|
||||
)
|
||||
return final_results[:limit]
|
||||
return final_results[: request.limit]
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from zhenxun.services.ai.context.rag.models import BaseRecord, SearchResult
|
||||
from zhenxun.services.ai.utils.logger import log_rag as logger
|
||||
|
||||
|
||||
def cosine_similarity(vec1: list[float], vec2: list[float]) -> float:
|
||||
"""使用 numpy 计算两组向量的余弦相似度"""
|
||||
@@ -10,3 +15,79 @@ def cosine_similarity(vec1: list[float], vec2: list[float]) -> float:
|
||||
if norm1 == 0 or norm2 == 0:
|
||||
return 0.0
|
||||
return float(np.dot(v1, v2) / (norm1 * norm2))
|
||||
|
||||
|
||||
def normalize_vector(vec: list[float] | np.ndarray) -> np.ndarray:
|
||||
"""将一维向量转化为 float32 数组并进行 L2 归一化"""
|
||||
v = np.array(vec, dtype=np.float32)
|
||||
norm = np.linalg.norm(v)
|
||||
if norm == 0:
|
||||
return v
|
||||
return v / norm
|
||||
|
||||
|
||||
def normalize_matrix(mat: np.ndarray) -> np.ndarray:
|
||||
"""将二维矩阵的每一行进行 L2 归一化"""
|
||||
norms = np.linalg.norm(mat, axis=1, keepdims=True)
|
||||
norms[norms == 0] = 1.0
|
||||
return mat / norms
|
||||
|
||||
|
||||
def normalize_query_text(query: Any) -> str:
|
||||
"""辅助函数:提取各种输入形式(如字符串、平台Message对象)的纯文本用于检索"""
|
||||
if isinstance(query, str):
|
||||
return query
|
||||
if hasattr(query, "extract_plain_text"):
|
||||
return query.extract_plain_text()
|
||||
return str(query) if query is not None else ""
|
||||
|
||||
|
||||
class InMemoryScorer:
|
||||
"""提供内存级别的纯 Python 向量打分与 BM25 稀疏打分工具类"""
|
||||
|
||||
@staticmethod
|
||||
def calculate_sparse_scores(
|
||||
query: str, candidate_records: list[BaseRecord]
|
||||
) -> list[SearchResult]:
|
||||
import jieba
|
||||
|
||||
tokens = set(jieba.lcut_for_search(query.lower()))
|
||||
if not tokens:
|
||||
return [SearchResult(record=r, score=0.1) for r in candidate_records]
|
||||
|
||||
results = []
|
||||
for r in candidate_records:
|
||||
content_lower = r.content.lower()
|
||||
matched_count = sum(1 for t in tokens if t in content_lower)
|
||||
score = matched_count / len(tokens) if tokens else 0.1
|
||||
results.append(SearchResult(record=r, score=score))
|
||||
return results
|
||||
|
||||
@staticmethod
|
||||
def calculate_dense_scores(
|
||||
query_embedding: list[float], candidate_records: list[BaseRecord]
|
||||
) -> list[SearchResult]:
|
||||
if not candidate_records or not query_embedding:
|
||||
return []
|
||||
|
||||
q_vec = normalize_vector(query_embedding)
|
||||
vec_records = [r for r in candidate_records if r.embedding]
|
||||
missing_records = [r for r in candidate_records if not r.embedding]
|
||||
|
||||
results = []
|
||||
if vec_records:
|
||||
try:
|
||||
raw_mat = np.array([r.embedding for r in vec_records], dtype=np.float32)
|
||||
norm_mat = normalize_matrix(raw_mat)
|
||||
scores = norm_mat @ q_vec
|
||||
for r, score in zip(vec_records, scores):
|
||||
results.append(SearchResult(record=r, score=float(score)))
|
||||
except ValueError as e:
|
||||
logger.warning(
|
||||
"⚠️ 维度不匹配,已安全跳过向量检索(降级)。原因: " + str(e)
|
||||
)
|
||||
|
||||
for r in missing_records:
|
||||
results.append(SearchResult(record=r, score=0.1))
|
||||
|
||||
return results
|
||||
|
||||
Reference in New Issue
Block a user