mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-08 13:20:01 +08:00
* ♻️ refactor(core): 重构 AI 能力与定时任务调度系统 - 【AI 能力与工具】重构 Capability 注册与管理机制,引入 CapabilityManager 统一管理 - 移除全局能力注册表,改用声明式装饰器 `@capability` 进行解耦注册 - 重构工具解析器链,使用统一的 BaseToolResolver 代替原有的多个特定解析器 - 增强工具查询过滤,支持通配符匹配、工具箱过滤和排除标签 - 【定时任务调度】重构定时任务管理器,引入 SchedulerRegistry 统一管理任务元数据 - 引入 JobConfig 聚合定时任务配置,支持用户维度的定时任务调度 - 重构执行分发器,支持并发限制、串行间隔和随机延迟打散 - 【运行上下文】引入 ScheduledDeps 以支持后台和定时任务环境下的依赖注入 - 优化 RunContext,支持从定时任务上下文快速构造,并提供 emit 辅助方法 - 【日志与监控】引入 AILoggerProxy,实现 AI 各模块的专属日志输出 - 将各模块的全局 logger 替换为对应的模块专属日志代理 - 【其他优化】修复 Pydantic V1 兼容层中 model_validator 的装饰器兼容性问题 - 在非交互式环境(如定时任务)中自动隐藏 HITL 交互工具以节省 Token * ♻️ refactor(core): 优化内部导入路径并提升 Pydantic 兼容性 - 【重构】将 `services/ai` 模块内的绝对导入重构为相对导入,优化包结构 - 【重构】移除不必要的 `if TYPE_CHECKING` 保护,通过 `from __future__ import annotations` 直接导入类型 - 【清理】清理 `core/messages/types.py` 中未使用的 `AssistantContentUnion` 等联合类型定义 - 【优化】在 `utils/pydantic_compat.py` 中新增 `model_rebuild` 兼容函数,统一 Pydantic V1/V2 的模型重建逻辑 - 【优化】将部分函数内部的延迟导入提升至模块顶部,规范代码结构 * ♻️ refactor(imports): 优化导入路径为相对导入并清理冗余导入 - 【重构】将 AI 服务相关模块中的绝对导入路径修改为相对导入,提升模块内聚性与可移植性 - 【清理】移除多处函数内部或类方法中未使用的冗余导入,避免循环引用和资源浪费 - 【格式化】微调部分工具装饰器和返回语句的格式与尾随逗号 * 🚨 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>
441 lines
15 KiB
Python
441 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.utils.logger import log_rag as logger
|
|
|
|
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):
|
|
"""
|
|
检索器核心协议。
|
|
任何实现了 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]
|