Files
zhenxun_bot/zhenxun/services/ai/context/rag/retrieval.py
T
80fc5b86a7 ✨ 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>
2026-07-03 08:53:56 +08:00

440 lines
15 KiB
Python

from abc import abstractmethod
import asyncio
import time
from typing import TYPE_CHECKING, Any, Protocol, cast, runtime_checkable
from zhenxun.services.ai.context.rag.models import QueryRequest, SearchResult
from zhenxun.services.log import logger
if TYPE_CHECKING:
from zhenxun.services.ai.context.rag.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):
"""
检索器核心协议。
任何实现了 retrieve 方法的对象均可作为检索器
(不仅限于向量检索,也可包含 BM25、SQL 搜索等)。
"""
@abstractmethod
async def retrieve(
self, query: Any, limit: int = 10, **kwargs: Any
) -> list[SearchResult]: ...
@runtime_checkable
class PostProcessor(Protocol):
"""后处理器协议(如重排、时间衰减打分等)。"""
@abstractmethod
async def process(
self, results: list[SearchResult], query: str
) -> list[SearchResult]: ...
@runtime_checkable
class PreProcessor(Protocol):
"""预处理器协议(如 LLM Query 改写、意图提取等)。"""
@abstractmethod
async def process(self, query: str) -> list[str]:
"""接收原始查询,返回一个或多个处理/改写后的查询词"""
...
class FilterEvaluator:
"""纯 Python 内存求值器,用于为轻量级 Storage 提供字典精确匹配过滤"""
@classmethod
def evaluate(
cls, metadata: dict[str, Any], filter_dict: dict[str, Any] | None
) -> bool:
if filter_dict is None:
return True
return all(metadata.get(k) == v for k, v in filter_dict.items())
class VectorDBRetriever(BaseRetriever):
"""基于向量数据库的标准检索器"""
def __init__(
self,
storage: "StorageBackend",
embedder: Any,
scope_prefix: str | None = None,
score_threshold: float = 0.4,
):
"""
初始化向量数据库检索器。
参数:
storage: 存储后端,用于执行向量相似度搜索。
embedder: 向量嵌入模型/函数,用于将文本转换为向量。
scope_prefix: 作用域前缀,用于限制检索范围,默认 None。
score_threshold: 分数阈值,过滤掉相似度低于该值的检索结果,默认 0.4。
"""
self.storage = storage
self.embedder = embedder
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)
vecs = await self.embedder(query, task="query")
query_vec = vecs[0] if vecs else None
if not text_query.strip() and not query_vec:
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]
class DatabaseSparseRetriever(BaseRetriever):
"""纯数据库下沉的稀疏检索器 (Keyword/FTS)"""
def __init__(
self,
storage: "StorageBackend",
scope_prefix: str | None = None,
score_threshold: float = 0.0,
):
"""
初始化数据库稀疏检索器。
参数:
storage: 存储后端,用于执行全文检索/关键词检索。
scope_prefix: 作用域前缀,用于限制检索范围,默认 None。
score_threshold: 分数阈值,过滤掉相关度低于该值的检索结果,默认 0.0。
"""
self.storage = storage
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():
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]
class RerankRetriever(BaseRetriever):
"""带大模型交叉注意力重排的高阶检索器 (Decorator Pattern)"""
def __init__(
self,
base_retriever: BaseRetriever,
model_name: str | None = None,
top_n: int = 5,
oversample_factor: int = 2,
min_oversample: int = 20,
):
"""
初始化重排检索器。
参数:
base_retriever: 基础检索器,用于初筛。
model_name: 重排模型的名称,默认 None。
top_n: 重排后保留的前 N 个文档数,默认 5。
oversample_factor: 过采样系数,决定初筛检索的文档数量倍数,默认 2。
min_oversample: 最小过采样文档数,默认 20。
"""
self.base_retriever = base_retriever
self.model_name = model_name
self.top_n = top_n
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
)
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
]
from zhenxun.services.ai.llm.api import rerank
try:
reranked = await rerank(
query=text_query,
documents=docs,
top_n=min(limit, self.top_n),
model=self.model_name,
)
except Exception as e:
logger.warning(f"Rerank 重排请求失败,将降级返回初筛结果: {e}")
return initial_results[:limit]
final_results = []
for rr in reranked:
original_res = initial_results[rr.index]
original_res.score = rr.relevance_score
final_results.append(original_res)
return final_results
class PipelineRetriever(BaseRetriever):
"""支持挂载多个后处理器的流水线检索器"""
def __init__(
self,
base_retriever: BaseRetriever,
post_processors: list[PostProcessor] | None = None,
pre_processors: list[PreProcessor] | None = None,
):
"""
初始化流水线检索器。
参数:
base_retriever: 基础检索器,执行最初的检索过程。
post_processors: 后处理器列表,用于对检索到的结果进行重排、过滤等后处理,默认 None。
pre_processors: 预处理器列表,用于对查询词进行改写、扩展等预处理,默认 None。
""" # noqa: E501
self.base_retriever = base_retriever
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)
queries_to_search = [query]
if text_query.strip():
processed_texts = [text_query]
for pp in self.pre_processors:
new_texts = []
for t in processed_texts:
new_texts.extend(await pp.process(t))
processed_texts = new_texts
if len(processed_texts) > 1 or (
len(processed_texts) == 1 and processed_texts[0] != text_query
):
queries_to_search.extend(processed_texts)
all_results = []
seen_ids = set()
for q in queries_to_search:
res = await self.base_retriever.retrieve(q, limit=limit * 2, **kwargs)
for r in res:
if r.record.id not in seen_ids:
seen_ids.add(r.record.id)
all_results.append(r)
results = sorted(all_results, key=lambda x: x.score, reverse=True)
for pp in self.post_processors:
results = await pp.process(results, query)
return results[:limit]
class LifecyclePostProcessor(PostProcessor):
"""生命周期后处理器(融合时间衰减与惰性访问强化)"""
def __init__(
self,
half_life_days: int = 30,
decay_weight: float = 0.3,
semantic_weight: float = 0.7,
importance_weight: float = 0.0,
reinforcement_weight: float = 0.2,
):
"""
初始化生命周期后处理器。
参数:
half_life_days: 记忆衰减半衰期天数,控制信息随时间的降权速度,默认 30。
decay_weight: 时间衰减得分的权重,默认 0.3。
semantic_weight: 语义相关度得分的权重,默认 0.7。
importance_weight: 信息重要性得分的权重,默认 0.0。
reinforcement_weight: 惰性访问强化(如访问次数得分)的权重,默认 0.2。
"""
self.half_life_days = half_life_days
self.decay_weight = decay_weight
self.semantic_weight = semantic_weight
self.importance_weight = importance_weight
self.reinforcement_weight = reinforcement_weight
async def process(
self, results: list[SearchResult], query: str
) -> list[SearchResult]:
now = time.time()
import math
for res in results:
created_at = res.record.metadata.get("created_at", now)
importance = res.record.metadata.get("importance", 0.5)
access_count = res.record.metadata.get("access_count", 0)
last_accessed_at = res.record.metadata.get("last_accessed_at", created_at)
age_days = max(0.0, (now - last_accessed_at) / 86400.0)
decay = 0.5 ** (age_days / self.half_life_days)
access_score = min(1.0, math.log1p(access_count) / 5.0)
res.score = (
(self.semantic_weight * res.score)
+ (self.decay_weight * decay)
+ (self.importance_weight * importance)
+ (self.reinforcement_weight * access_score)
)
results.sort(key=lambda x: x.score, reverse=True)
return results
class HybridRetriever(BaseRetriever):
"""
双轨混合检索器 (Hybrid Search Engine)。
并发调用 Dense (VectorDB) 和 Sparse (BM25),并使用倒数秩融合 (RRF) 算法合并结果。
"""
def __init__(
self,
dense_retriever: BaseRetriever,
sparse_retriever: BaseRetriever,
dense_weight: float = 0.7,
sparse_weight: float = 0.3,
rrf_k: int = 60,
oversample_factor: int = 2,
min_oversample: int = 20,
):
"""
初始化双轨混合检索器。
参数:
dense_retriever: 稠密向量检索器,用于语义召回。
sparse_retriever: 稀疏文本检索器,用于关键词召回(如 BM25)。
dense_weight: 稠密向量检索的加权权重,默认 0.7。
sparse_weight: 稀疏文本检索的加权权重,默认 0.3。
rrf_k: 倒数秩融合(RRF)算法中的常数参数,默认 60。
"""
self.dense_retriever = dense_retriever
self.sparse_retriever = sparse_retriever
self.dense_weight = dense_weight
self.sparse_weight = sparse_weight
self.rrf_k = rrf_k
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)
results = await asyncio.gather(
self.dense_retriever.retrieve(query, limit=oversample_limit, **kwargs),
self.sparse_retriever.retrieve(query, limit=oversample_limit, **kwargs),
return_exceptions=True,
)
for res in results:
if isinstance(res, ImportError):
raise res
dense_res = (
cast(list[SearchResult], results[0])
if not isinstance(results[0], BaseException)
else []
)
sparse_res = (
cast(list[SearchResult], results[1])
if not isinstance(results[1], BaseException)
else []
)
if isinstance(results[0], BaseException):
logger.error(f"[HybridSearch] 向量检索异常: {results[0]}")
if isinstance(results[1], BaseException):
logger.error(f"[HybridSearch] BM25 检索异常: {results[1]}")
rrf_scores: dict[str, float] = {}
merged_records = {}
for rank, res in enumerate(dense_res):
record_id = res.record.id
merged_records[record_id] = res.record
rrf_score = 1.0 / (self.rrf_k + rank + 1)
rrf_scores[record_id] = rrf_scores.get(record_id, 0.0) + (
self.dense_weight * rrf_score
)
for rank, res in enumerate(sparse_res):
record_id = res.record.id
merged_records[record_id] = res.record
rrf_score = 1.0 / (self.rrf_k + rank + 1)
rrf_scores[record_id] = rrf_scores.get(record_id, 0.0) + (
self.sparse_weight * rrf_score
)
max_possible_score = (self.dense_weight * (1.0 / (self.rrf_k + 1))) + (
self.sparse_weight * (1.0 / (self.rrf_k + 1))
)
final_results = []
for record_id, score in sorted(
rrf_scores.items(), key=lambda x: x[1], reverse=True
):
normalized_score = (
score / max_possible_score if max_possible_score > 0 else 0.0
)
final_results.append(
SearchResult(record=merged_records[record_id], score=normalized_score)
)
logger.debug(
f"⚖️ [HybridSearch] 融合完成: "
f"Dense({len(dense_res)}) + Sparse({len(sparse_res)}) "
f"-> Merged({len(final_results)}), 截取 Top {limit}"
)
return final_results[:limit]