- 后端 API(Flask + Gunicorn) - RAG 引擎(混合检索 + 云端 Reranker + 引用溯源) - 文档解析(MinerU + 多格式支持) - Docker 生产部署配置 - 排除前端项目、敏感配置、模型文件
112 lines
3.3 KiB
Python
112 lines
3.3 KiB
Python
"""
|
||
Agentic RAG - 上下文处理 Mixin
|
||
|
||
包含上下文压缩、去重、Token 控制等方法
|
||
"""
|
||
|
||
import logging
|
||
|
||
from .agentic_base import logger, MAX_CONTEXT_TOKENS, RERANK_THRESHOLD
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
class ContextMixin:
|
||
"""上下文处理方法"""
|
||
|
||
def _compress_contexts(self, query: str, contexts: list) -> list:
|
||
"""上下文压缩三步走:Rerank 过滤 → 去重 → Token 控制"""
|
||
if not contexts:
|
||
return contexts
|
||
|
||
# Step 1: Rerank 过滤
|
||
filtered = self._rerank_filter(contexts)
|
||
|
||
# Step 2: 去重
|
||
deduped = self._deduplicate_contexts(filtered)
|
||
|
||
# Step 3: Token 控制
|
||
result = self._truncate_to_tokens(deduped, self.MAX_CONTEXT_TOKENS)
|
||
|
||
return result
|
||
|
||
def _rerank_filter(self, contexts: list) -> list:
|
||
"""Rerank 过滤 - 保留相关性分数 >= 阈值的上下文"""
|
||
scored_contexts = [c for c in contexts if c.get('score') is not None]
|
||
|
||
if scored_contexts:
|
||
filtered = [c for c in contexts if c.get('score', 0) >= self.RERANK_THRESHOLD]
|
||
return filtered if filtered else contexts
|
||
|
||
return contexts
|
||
|
||
def _deduplicate_contexts(self, contexts: list, threshold: float = 0.9) -> list:
|
||
"""去重 - 基于内容相似度去重"""
|
||
if len(contexts) <= 1:
|
||
return contexts
|
||
|
||
result = []
|
||
seen_keys = set()
|
||
|
||
for c in contexts:
|
||
doc = c.get('doc', '')
|
||
key = doc[:100] if doc else ''
|
||
|
||
meta = c.get('meta', {})
|
||
source = meta.get('source', '')
|
||
page = meta.get('page', '')
|
||
|
||
composite_key = f"{source}|{page}|{key}"
|
||
|
||
if composite_key not in seen_keys:
|
||
seen_keys.add(composite_key)
|
||
result.append(c)
|
||
|
||
return result
|
||
|
||
def _truncate_to_tokens(self, contexts: list, max_tokens: int) -> list:
|
||
"""Token 控制 - 截断到最大 Token 数"""
|
||
result = []
|
||
total_tokens = 0
|
||
|
||
for c in contexts:
|
||
doc = c.get('doc', '')
|
||
# 简单估算:1 token ≈ 1.5 中文字符
|
||
tokens = len(doc) // 1.5
|
||
|
||
if total_tokens + tokens <= max_tokens:
|
||
result.append(c)
|
||
total_tokens += tokens
|
||
else:
|
||
break
|
||
|
||
return result
|
||
|
||
def _merge_and_deduplicate(self, old_contexts: list, new_contexts: list) -> list:
|
||
"""合并并去重两个上下文列表"""
|
||
result = list(old_contexts)
|
||
seen_keys = set()
|
||
|
||
# 记录已有上下文的 key
|
||
for c in old_contexts:
|
||
doc = c.get('doc', '')
|
||
key = doc[:100] if doc else ''
|
||
meta = c.get('meta', {})
|
||
source = meta.get('source', '')
|
||
composite_key = f"{source}|{key}"
|
||
seen_keys.add(composite_key)
|
||
|
||
# 添加新上下文(去重)
|
||
for c in new_contexts:
|
||
doc = c.get('doc', '')
|
||
key = doc[:100] if doc else ''
|
||
meta = c.get('meta', {})
|
||
source = meta.get('source', '')
|
||
composite_key = f"{source}|{key}"
|
||
|
||
if composite_key not in seen_keys:
|
||
seen_keys.add(composite_key)
|
||
result.append(c)
|
||
|
||
return result
|