Compare commits
7 Commits
c6a17ad1e4
...
server-bas
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
279e2bf47c | ||
|
|
acb84b804d | ||
|
|
0b2ef8c161 | ||
|
|
8e3e9832ff | ||
|
|
43261e9aff | ||
|
|
a340eaaeee | ||
|
|
8af8d38c01 |
@@ -1,23 +0,0 @@
|
||||
# 大体积数据目录(通过 volume 挂载,不需要打进镜像)
|
||||
models/
|
||||
knowledge/vector_store/
|
||||
documents/
|
||||
.data/
|
||||
data/
|
||||
|
||||
# Python 虚拟环境
|
||||
venv/
|
||||
__pycache__/
|
||||
*.pyc
|
||||
|
||||
# Git
|
||||
.git/
|
||||
|
||||
# IDE 和编辑器
|
||||
.vscode/
|
||||
.idea/
|
||||
*.swp
|
||||
|
||||
# 其他
|
||||
*.log
|
||||
.env*
|
||||
5
.gitignore
vendored
5
.gitignore
vendored
@@ -120,8 +120,9 @@ test_*.json
|
||||
rag_response.json
|
||||
nul
|
||||
|
||||
# 临时调试脚本(下划线开头)
|
||||
scripts/_*.py
|
||||
# 调试脚本和临时计划(仅本地使用)
|
||||
scripts/
|
||||
plans/
|
||||
|
||||
# Qoder 工具目录
|
||||
.qoder/
|
||||
|
||||
@@ -3,7 +3,7 @@ API 路由层 — Flask 应用工厂
|
||||
|
||||
本模块实现 Flask 应用工厂模式,负责:
|
||||
- 创建和配置 Flask 应用实例
|
||||
- 初始化核心服务(AgenticRAG、同步服务)
|
||||
- 初始化核心服务(同步服务)
|
||||
- 注册所有 API Blueprint
|
||||
- 配置前端静态文件路由
|
||||
|
||||
@@ -47,7 +47,7 @@ def create_app() -> 'Flask':
|
||||
|
||||
1. 创建 Flask 应用,配置 CORS
|
||||
2. 初始化 Repository(会话存储)
|
||||
3. 初始化核心服务(AgenticRAG、同步服务)
|
||||
3. 初始化核心服务(同步服务)
|
||||
4. 注册所有 API Blueprint
|
||||
5. 配置前端静态文件路由
|
||||
6. 执行生产环境配置校验
|
||||
@@ -93,19 +93,6 @@ def create_app() -> 'Flask':
|
||||
|
||||
# ==================== 核心服务初始化 ====================
|
||||
|
||||
# Agentic RAG 引擎
|
||||
try:
|
||||
from core.agentic import AgenticRAG
|
||||
from config import ENABLE_WEB_SEARCH
|
||||
agentic_rag = AgenticRAG(
|
||||
enable_web_search=ENABLE_WEB_SEARCH,
|
||||
)
|
||||
app.config['AGENTIC_RAG'] = agentic_rag
|
||||
logger.info(f"Agentic RAG 引擎已初始化(网络搜索={'启用' if ENABLE_WEB_SEARCH else '关闭'})")
|
||||
except Exception as e:
|
||||
app.config['AGENTIC_RAG'] = None
|
||||
logger.warning(f"Agentic RAG 初始化失败: {e}")
|
||||
|
||||
# 同步服务
|
||||
try:
|
||||
from knowledge.sync import KnowledgeSyncService
|
||||
|
||||
@@ -11,6 +11,7 @@
|
||||
from flask import Blueprint, request, jsonify
|
||||
from auth.gateway import require_gateway_auth, require_role, get_user_permissions, MOCK_USERS
|
||||
import os
|
||||
import time
|
||||
from pathlib import Path
|
||||
from dotenv import load_dotenv
|
||||
|
||||
@@ -21,6 +22,37 @@ load_dotenv(env_path)
|
||||
auth_bp = Blueprint('auth', __name__)
|
||||
|
||||
|
||||
# ==================== 登录速率限制 ====================
|
||||
# 简单的内存级速率限制:每个 IP 在时间窗口内最多允许 N 次登录尝试
|
||||
_login_attempts = {} # {ip: [(timestamp, ...), ...]}
|
||||
_RATE_LIMIT_WINDOW = 300 # 5 分钟窗口
|
||||
_RATE_LIMIT_MAX = 10 # 窗口内最多 10 次尝试
|
||||
|
||||
|
||||
def _check_rate_limit(client_ip: str) -> bool:
|
||||
"""检查是否超出登录速率限制,返回 True 表示允许"""
|
||||
now = time.time()
|
||||
if client_ip not in _login_attempts:
|
||||
_login_attempts[client_ip] = []
|
||||
|
||||
# 清理过期记录
|
||||
_login_attempts[client_ip] = [
|
||||
t for t in _login_attempts[client_ip]
|
||||
if now - t < _RATE_LIMIT_WINDOW
|
||||
]
|
||||
|
||||
if len(_login_attempts[client_ip]) >= _RATE_LIMIT_MAX:
|
||||
return False
|
||||
|
||||
_login_attempts[client_ip].append(now)
|
||||
return True
|
||||
|
||||
|
||||
def _is_dev_mode() -> bool:
|
||||
"""统一的开发模式判断"""
|
||||
return os.environ.get('DEV_MODE', 'true').lower() != 'false'
|
||||
|
||||
|
||||
@auth_bp.route('/auth/login', methods=['POST'])
|
||||
def mock_login():
|
||||
"""
|
||||
@@ -50,10 +82,14 @@ def mock_login():
|
||||
- manager / manager123 (经理,财务部)
|
||||
- user / test123 (普通用户,技术部)
|
||||
"""
|
||||
# 默认开启开发模式(生产环境需设置 DEV_MODE=false)
|
||||
if os.environ.get('DEV_MODE', 'true').lower() == 'false':
|
||||
if not _is_dev_mode():
|
||||
return jsonify({"error": "仅开发环境可用,请设置 DEV_MODE=true"}), 403
|
||||
|
||||
# 速率限制检查
|
||||
client_ip = request.remote_addr or 'unknown'
|
||||
if not _check_rate_limit(client_ip):
|
||||
return jsonify({"error": f"登录尝试过于频繁,请 {_RATE_LIMIT_WINDOW // 60} 分钟后再试"}), 429
|
||||
|
||||
data = request.json or {}
|
||||
username = data.get('username')
|
||||
password = data.get('password')
|
||||
@@ -79,7 +115,9 @@ def mock_login():
|
||||
def get_stats():
|
||||
"""获取系统统计信息(仅管理员)"""
|
||||
from flask import current_app
|
||||
session_manager = current_app.config['SESSION_MANAGER']
|
||||
session_manager = current_app.config.get('SESSION_MANAGER')
|
||||
if not session_manager:
|
||||
return jsonify({"error": "会话管理器未启用"}), 503
|
||||
return jsonify(session_manager.get_stats())
|
||||
|
||||
|
||||
@@ -131,8 +169,7 @@ def get_users():
|
||||
]
|
||||
}
|
||||
"""
|
||||
dev_mode = os.environ.get('DEV_MODE', 'true').lower() != 'false'
|
||||
if not dev_mode:
|
||||
if not _is_dev_mode():
|
||||
return jsonify({"error": "仅开发环境可用"}), 403
|
||||
|
||||
users = []
|
||||
@@ -159,12 +196,26 @@ def update_user(user_id):
|
||||
"is_active": false
|
||||
}
|
||||
"""
|
||||
dev_mode = os.environ.get('DEV_MODE', 'true').lower() != 'false'
|
||||
if not dev_mode:
|
||||
if not _is_dev_mode():
|
||||
return jsonify({"error": "仅开发环境可用"}), 403
|
||||
|
||||
# 模拟用户不支持真正的状态切换,直接返回成功
|
||||
return jsonify({"message": "操作成功(模拟)", "user_id": user_id})
|
||||
# 验证目标用户是否存在
|
||||
target_user = None
|
||||
for username, info in MOCK_USERS.items():
|
||||
if info['user_id'] == user_id:
|
||||
target_user = info
|
||||
break
|
||||
|
||||
if not target_user:
|
||||
return jsonify({"error": f"用户 {user_id} 不存在"}), 404
|
||||
|
||||
data = request.json or {}
|
||||
# 模拟操作:记录请求但不实际执行(mock 用户数据是静态的)
|
||||
return jsonify({
|
||||
"message": "操作成功(模拟)",
|
||||
"user_id": user_id,
|
||||
"applied_changes": data
|
||||
})
|
||||
|
||||
|
||||
@auth_bp.route('/auth/change-password', methods=['POST'])
|
||||
@@ -179,8 +230,7 @@ def change_password():
|
||||
"new_password": "xxx"
|
||||
}
|
||||
"""
|
||||
dev_mode = os.environ.get('DEV_MODE', 'true').lower() != 'false'
|
||||
if not dev_mode:
|
||||
if not _is_dev_mode():
|
||||
return jsonify({"error": "仅开发环境可用"}), 403
|
||||
|
||||
data = request.json or {}
|
||||
@@ -193,5 +243,12 @@ def change_password():
|
||||
if len(new_password) < 6:
|
||||
return jsonify({"error": "新密码至少6位"}), 400
|
||||
|
||||
# 模拟环境直接返回成功
|
||||
# 验证当前用户的旧密码
|
||||
user = request.current_user
|
||||
username = user.get('username', '')
|
||||
mock_user = MOCK_USERS.get(username)
|
||||
if mock_user and mock_user['password'] != old_password:
|
||||
return jsonify({"error": "旧密码错误"}), 401
|
||||
|
||||
# 模拟环境返回成功(不实际修改密码,mock 数据是静态的)
|
||||
return jsonify({"message": "密码修改成功(模拟)"})
|
||||
|
||||
1065
api/chat_routes.py
1065
api/chat_routes.py
File diff suppressed because it is too large
Load Diff
@@ -45,6 +45,7 @@ import logging
|
||||
logger = logging.getLogger(__name__)
|
||||
from werkzeug.utils import secure_filename
|
||||
from auth.gateway import require_gateway_auth
|
||||
from config import DEV_MODE
|
||||
from core.status_codes import (
|
||||
UPLOAD_SUCCESS, BATCH_UPLOAD_SUCCESS, BAD_REQUEST,
|
||||
NO_FILE, NO_FILE_SELECTED, NO_COLLECTION,
|
||||
@@ -169,9 +170,9 @@ def serve_document_file(doc_path: str) -> Tuple[Any, int]:
|
||||
文件内容或错误响应
|
||||
|
||||
Note:
|
||||
仅在 DEV_MODE=true 时可用
|
||||
仅在 DEV_MODE=true 时可用(需在 .env 中显式设置)
|
||||
"""
|
||||
if os.environ.get('DEV_MODE', 'true').lower() == 'false':
|
||||
if not DEV_MODE:
|
||||
return jsonify({"error": "仅开发环境可用"}), 403
|
||||
|
||||
from config import DOCUMENTS_PATH
|
||||
@@ -775,22 +776,7 @@ def delete_document(doc_path: str) -> Tuple[Any, int]:
|
||||
if kb_manager:
|
||||
kb_manager.delete_document(collection, filename)
|
||||
|
||||
# 2. 缓存失效
|
||||
try:
|
||||
from core.cache import get_cache_manager
|
||||
_cm = get_cache_manager()
|
||||
_cm.increment_kb_version(collection)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
from core.semantic_cache import get_semantic_cache
|
||||
_sc = get_semantic_cache()
|
||||
if _sc:
|
||||
_sc.clear()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 3. 删除文件
|
||||
# 2. 删除文件
|
||||
os.remove(filepath)
|
||||
|
||||
return jsonify({
|
||||
|
||||
@@ -260,21 +260,6 @@ def delete_collection(kb_name: str) -> Tuple[Any, int]:
|
||||
success, message = kb_manager.delete_collection(kb_name, delete_documents)
|
||||
|
||||
if success:
|
||||
# 缓存失效
|
||||
try:
|
||||
from core.cache import get_cache_manager
|
||||
_cm = get_cache_manager()
|
||||
_cm.increment_kb_version(kb_name)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
from core.semantic_cache import get_semantic_cache
|
||||
_sc = get_semantic_cache()
|
||||
if _sc:
|
||||
_sc.clear()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
"message": message,
|
||||
|
||||
@@ -69,7 +69,7 @@ def get_history(session_id):
|
||||
user_id = request.current_user["user_id"]
|
||||
|
||||
# 验证会话归属
|
||||
sessions = session_manager.get_user_sessions(user_id)
|
||||
sessions = session_manager.get_user_sessions(user_id, limit=20)
|
||||
session_ids = [s["session_id"] for s in sessions]
|
||||
|
||||
if session_id not in session_ids:
|
||||
@@ -90,7 +90,7 @@ def delete_session(session_id):
|
||||
user_id = request.current_user["user_id"]
|
||||
|
||||
# 验证会话归属
|
||||
sessions = session_manager.get_user_sessions(user_id)
|
||||
sessions = session_manager.get_user_sessions(user_id, limit=20)
|
||||
session_ids = [s["session_id"] for s in sessions]
|
||||
|
||||
if session_id not in session_ids:
|
||||
@@ -111,7 +111,7 @@ def clear_history(session_id):
|
||||
user_id = request.current_user["user_id"]
|
||||
|
||||
# 验证会话归属
|
||||
sessions = session_manager.get_user_sessions(user_id)
|
||||
sessions = session_manager.get_user_sessions(user_id, limit=20)
|
||||
session_ids = [s["session_id"] for s in sessions]
|
||||
|
||||
if session_id not in session_ids:
|
||||
|
||||
@@ -20,12 +20,12 @@
|
||||
|
||||
## 模式说明
|
||||
|
||||
开发模式 (DEV_MODE=true,默认):
|
||||
开发模式 (DEV_MODE=true):
|
||||
- 支持 mock token 模拟用户:Authorization: Bearer mock-token-admin
|
||||
- 无 Header 时自动使用开发测试用户
|
||||
- 适用于前端测试和开发调试
|
||||
|
||||
生产模式 (DEV_MODE=false):
|
||||
生产模式 (DEV_MODE=false,默认):
|
||||
- 不需要 Header,直接放行
|
||||
- 权限由后端完全控制,通过 collections 参数传入
|
||||
- RAG 服务完全无状态,只负责问答检索
|
||||
@@ -34,17 +34,11 @@
|
||||
from functools import wraps
|
||||
from flask import request, jsonify
|
||||
from typing import Dict, Optional
|
||||
import os
|
||||
from pathlib import Path
|
||||
from dotenv import load_dotenv
|
||||
|
||||
# 加载 .env 文件(从项目根目录)
|
||||
env_path = Path(__file__).parent.parent / '.env'
|
||||
load_dotenv(env_path)
|
||||
from config import DEV_MODE
|
||||
|
||||
|
||||
# ==================== 模拟用户数据(开发环境)====================
|
||||
# 用于前端模拟登录测试,仅 DEV_MODE=true 时生效
|
||||
# 用于前端模拟登录测试,仅 DEV_MODE=true 时生效(需在 .env 中显式开启)
|
||||
MOCK_USERS = {
|
||||
'admin': {
|
||||
'user_id': 'admin001',
|
||||
@@ -83,19 +77,19 @@ def require_gateway_auth(f):
|
||||
"""
|
||||
网关认证装饰器 - 从 Header 读取用户信息
|
||||
|
||||
开发模式 (DEV_MODE=true,默认):
|
||||
开发模式 (DEV_MODE=true):
|
||||
- 支持 mock token: Authorization: Bearer mock-token-admin
|
||||
- 无 Header 时自动使用开发测试用户(admin 角色)
|
||||
|
||||
生产模式 (DEV_MODE=false):
|
||||
生产模式 (DEV_MODE=false,默认):
|
||||
- 不需要 Header,直接放行
|
||||
- 用户信息设为默认值
|
||||
- 权限由后端通过 collections 参数控制
|
||||
"""
|
||||
@wraps(f)
|
||||
def decorated(*args, **kwargs):
|
||||
# 开发模式开关(默认开启,生产环境设置 DEV_MODE=false)
|
||||
dev_mode = os.environ.get('DEV_MODE', 'true').lower() != 'false'
|
||||
# 开发模式开关(统一由 config.py 管理)
|
||||
dev_mode = DEV_MODE
|
||||
|
||||
# 开发模式:支持 mock token
|
||||
if dev_mode:
|
||||
@@ -202,11 +196,24 @@ def can_delete_collection(role: str) -> bool:
|
||||
|
||||
def require_role(*roles):
|
||||
"""
|
||||
兼容旧代码 - 权限由后端管理,此装饰器不再执行权限验证
|
||||
角色验证装饰器(开发和生产环境均生效)
|
||||
|
||||
需搭配 @require_gateway_auth 使用(先设置 current_user,再验证角色)。
|
||||
|
||||
开发模式: 检查 mock token 对应用户的角色
|
||||
生产模式: 检查网关注入的 X-User-Role Header
|
||||
"""
|
||||
def decorator(f):
|
||||
@wraps(f)
|
||||
def decorated(*args, **kwargs):
|
||||
if roles:
|
||||
user = get_current_user()
|
||||
if user is None or user.get('role') not in roles:
|
||||
from flask import jsonify
|
||||
return jsonify({
|
||||
"error": "权限不足,需要角色: {}".format(', '.join(roles)),
|
||||
"status": "FORBIDDEN"
|
||||
}), 403
|
||||
return f(*args, **kwargs)
|
||||
return decorated
|
||||
return decorator
|
||||
@@ -214,11 +221,12 @@ def require_role(*roles):
|
||||
|
||||
def require_collection_permission(operation: str):
|
||||
"""
|
||||
兼容旧代码 - 权限由后端管理,此装饰器不再执行权限验证
|
||||
集合权限验证装饰器(开发模式下为占位实现,生产环境权限由网关控制)
|
||||
"""
|
||||
def decorator(f):
|
||||
@wraps(f)
|
||||
def decorated(*args, **kwargs):
|
||||
# 生产环境下权限由网关/后端统一管控,此处放行
|
||||
return f(*args, **kwargs)
|
||||
return decorated
|
||||
return decorator
|
||||
|
||||
@@ -1,207 +0,0 @@
|
||||
"""
|
||||
孤儿文件清理工具
|
||||
|
||||
清理 .data/images/ 和 .data/cache/vlm/ 中不再被任何 ChromaDB 切片引用的文件。
|
||||
|
||||
用法:
|
||||
# 预览模式(不删除,只显示孤儿文件)
|
||||
python cleanup_orphans.py
|
||||
|
||||
# 实际删除
|
||||
python cleanup_orphans.py --force
|
||||
|
||||
# 仅清理特定知识库
|
||||
python cleanup_orphans.py --collections public_kb test_kb
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import hashlib
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format='%(levelname)s: %(message)s')
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
IMAGES_DIR = Path(".data/images")
|
||||
VLM_CACHE_DIR = Path(".data/cache/vlm")
|
||||
|
||||
|
||||
def compute_file_hash(file_path: str) -> str:
|
||||
"""计算文件 MD5"""
|
||||
with open(file_path, 'rb') as f:
|
||||
return hashlib.md5(f.read()).hexdigest()
|
||||
|
||||
|
||||
def collect_referenced_images(manager, collections=None) -> dict:
|
||||
"""
|
||||
从 ChromaDB 收集所有被引用的图片路径。
|
||||
|
||||
Returns:
|
||||
{image_filename: set of chunk_ids referencing it}
|
||||
"""
|
||||
referenced = {}
|
||||
|
||||
if collections:
|
||||
kb_names = collections
|
||||
else:
|
||||
kb_names = [c.name if hasattr(c, 'name') else str(c)
|
||||
for c in manager.list_collections()]
|
||||
|
||||
for kb_name in kb_names:
|
||||
try:
|
||||
col = manager.get_collection(kb_name)
|
||||
except Exception as e:
|
||||
logger.warning(f"无法获取 {kb_name}: {e}")
|
||||
continue
|
||||
|
||||
# 查找所有带 image_path 的切片
|
||||
results = col.get(include=['metadatas'])
|
||||
if not results['ids']:
|
||||
continue
|
||||
|
||||
for chunk_id, meta in zip(results['ids'], results['metadatas']):
|
||||
image_path = meta.get('image_path', '')
|
||||
if not image_path:
|
||||
continue
|
||||
|
||||
# image_path 可能是: "185a7a75d246.png" 或相对路径
|
||||
filename = os.path.basename(image_path)
|
||||
if filename not in referenced:
|
||||
referenced[filename] = set()
|
||||
referenced[filename].add(f"{kb_name}/{chunk_id}")
|
||||
|
||||
return referenced
|
||||
|
||||
|
||||
def find_orphan_images(referenced: dict) -> list:
|
||||
"""
|
||||
查找 .data/images/ 中不再被引用的图片文件。
|
||||
|
||||
Returns:
|
||||
[(filepath, filename, size_bytes)]
|
||||
"""
|
||||
orphans = []
|
||||
if not IMAGES_DIR.exists():
|
||||
return orphans
|
||||
|
||||
for f in IMAGES_DIR.iterdir():
|
||||
if not f.is_file():
|
||||
continue
|
||||
if f.name not in referenced:
|
||||
orphans.append((str(f), f.name, f.stat().st_size))
|
||||
|
||||
return orphans
|
||||
|
||||
|
||||
def find_orphan_vlm_caches(referenced: dict) -> list:
|
||||
"""
|
||||
查找 .data/cache/vlm/ 中对应的图片已不存在的缓存文件。
|
||||
|
||||
缓存文件以图片 MD5 命名,如果图片被删了,缓存也应该是孤儿。
|
||||
|
||||
Returns:
|
||||
[(filepath, filename, size_bytes)]
|
||||
"""
|
||||
orphans = []
|
||||
if not VLM_CACHE_DIR.exists():
|
||||
return orphans
|
||||
|
||||
# 构建 referenced 中所有图片的 MD5 集合
|
||||
referenced_hashes = set()
|
||||
for filename in referenced:
|
||||
full_path = IMAGES_DIR / filename
|
||||
if full_path.exists():
|
||||
try:
|
||||
img_hash = compute_file_hash(str(full_path))
|
||||
referenced_hashes.add(img_hash)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
for f in VLM_CACHE_DIR.iterdir():
|
||||
if not f.is_file() or not f.suffix == '.txt':
|
||||
continue
|
||||
cache_hash = f.stem # 文件名就是 MD5
|
||||
if cache_hash not in referenced_hashes:
|
||||
orphans.append((str(f), f.name, f.stat().st_size))
|
||||
|
||||
return orphans
|
||||
|
||||
|
||||
def delete_files(file_list: list, dry_run: bool = True) -> int:
|
||||
"""删除文件列表,返回删除数量"""
|
||||
deleted = 0
|
||||
for filepath, filename, size in file_list:
|
||||
if dry_run:
|
||||
logger.info(f" [DRY-RUN] 将删除: {filename} ({size/1024:.1f} KB)")
|
||||
else:
|
||||
try:
|
||||
os.remove(filepath)
|
||||
logger.info(f" 已删除: {filename} ({size/1024:.1f} KB)")
|
||||
deleted += 1
|
||||
except OSError as e:
|
||||
logger.warning(f" 删除失败: {filename} - {e}")
|
||||
return deleted
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="孤儿文件清理工具")
|
||||
parser.add_argument("--force", action="store_true",
|
||||
help="实际删除文件(默认仅预览)")
|
||||
parser.add_argument("--collections", nargs="+",
|
||||
help="仅检查指定知识库(默认检查全部)")
|
||||
args = parser.parse_args()
|
||||
|
||||
os.chdir(os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
mode = "删除" if args.force else "预览(不删除)"
|
||||
print("=" * 60)
|
||||
print(f"孤儿文件清理 - {mode}")
|
||||
print("=" * 60)
|
||||
|
||||
# 1. 收集引用
|
||||
print("\n[1/4] 扫描 ChromaDB 引用...")
|
||||
sys.path.insert(0, '.')
|
||||
from knowledge.manager import get_kb_manager
|
||||
manager = get_kb_manager()
|
||||
referenced = collect_referenced_images(manager, args.collections)
|
||||
print(f" 被引用的图片: {len(referenced)} 个")
|
||||
|
||||
# 2. 查找孤儿图片
|
||||
print("\n[2/4] 查找孤儿图片...")
|
||||
orphan_images = find_orphan_images(referenced)
|
||||
total_img_size = sum(s for _, _, s in orphan_images)
|
||||
print(f" 孤儿图片: {len(orphan_images)} 个 ({total_img_size/1024:.1f} KB)")
|
||||
|
||||
# 3. 查找孤儿 VLM 缓存
|
||||
print("\n[3/4] 查找孤儿 VLM 缓存...")
|
||||
orphan_caches = find_orphan_vlm_caches(referenced)
|
||||
total_cache_size = sum(s for _, _, s in orphan_caches)
|
||||
print(f" 孤儿缓存: {len(orphan_caches)} 个 ({total_cache_size/1024:.1f} KB)")
|
||||
|
||||
# 4. 清理
|
||||
print("\n[4/4] 清理...")
|
||||
if not orphan_images and not orphan_caches:
|
||||
print(" 没有需要清理的文件")
|
||||
else:
|
||||
if not args.force:
|
||||
print(" 预览模式,以下文件将被删除:")
|
||||
|
||||
img_deleted = delete_files(orphan_images, dry_run=not args.force)
|
||||
cache_deleted = delete_files(orphan_caches, dry_run=not args.force)
|
||||
|
||||
if args.force:
|
||||
print(f"\n 删除了 {img_deleted} 个图片 + {cache_deleted} 个缓存")
|
||||
freed = total_img_size + total_cache_size
|
||||
print(f" 释放空间: {freed/1024:.1f} KB")
|
||||
else:
|
||||
print(f"\n 共 {len(orphan_images) + len(orphan_caches)} 个文件待清理")
|
||||
print(f" 释放空间: {(total_img_size + total_cache_size)/1024:.1f} KB")
|
||||
print(" 使用 --force 参数执行实际删除")
|
||||
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
@@ -109,11 +109,6 @@ USE_RERANK = True
|
||||
RERANK_CANDIDATES = 20 # 送入重排序的候选数
|
||||
RERANK_TOP_K = 15 # 重排序后保留数
|
||||
RERANK_USE_ONNX = os.getenv("RERANK_USE_ONNX", "true").lower() == "true"
|
||||
RERANK_BACKEND = os.getenv("RERANK_BACKEND", "local") # local / cloud / fallback
|
||||
RERANK_CLOUD_MODEL = os.getenv("RERANK_CLOUD_MODEL", "qwen3-rerank")
|
||||
RERANK_CLOUD_API_KEY = os.getenv("RERANK_CLOUD_API_KEY", "")
|
||||
RERANK_CLOUD_BASE_URL = os.getenv("RERANK_CLOUD_BASE_URL", "https://dashscope.aliyuncs.com/compatible-api/v1/reranks")
|
||||
RERANK_CLOUD_TIMEOUT = int(os.getenv("RERANK_CLOUD_TIMEOUT", "15"))
|
||||
RERANK_CONTEXT_MIN_SCORE = 0.05 # Rerank 分数低于此值的切片不送入 LLM
|
||||
|
||||
# ----- RRF 融合 -----
|
||||
@@ -186,7 +181,10 @@ SEMANTIC_CACHE_ENABLED = True
|
||||
SEMANTIC_CACHE_THRESHOLD = 0.92 # 相似度阈值
|
||||
|
||||
# 缓存写入最低置信度
|
||||
CACHE_MIN_SCORE = 0.3
|
||||
# 注意:ChromaDB cosine distance 范围 [0,2],score = 1 - dist
|
||||
# 当前 embedding 模型的 cosine similarity 普遍在 0.03-0.06 之间
|
||||
# 搜索管线已通过 rerank 过滤低质量结果,此处不再额外限制
|
||||
CACHE_MIN_SCORE = 0.0
|
||||
|
||||
# LLM 调用预算
|
||||
MAX_LLM_CALLS_PER_QUERY = 2
|
||||
@@ -203,6 +201,18 @@ MINERU_DEVICE_MODE = os.getenv("MINERU_DEVICE_MODE", "cpu") # cpu / cuda
|
||||
MINERU_API_TOKEN = os.getenv("MINERU_API_TOKEN", "") # 在 https://mineru.net/apiManage/token 申请
|
||||
MINERU_API_URL = os.getenv("MINERU_API_URL", "https://mineru.net/api/v4/extract/task")
|
||||
MINERU_PREFER_ONLINE = os.getenv("MINERU_PREFER_ONLINE", "true").lower() == "true"
|
||||
MINERU_PREFER_V2 = os.getenv("MINERU_PREFER_V2", "true").lower() == "true" # 优先使用 v2 格式(含 style 信息)
|
||||
|
||||
# 标题识别规则引擎
|
||||
# 规则定义见 parsers/heading_rules.py,支持通过 config.py 覆盖
|
||||
HEADING_RULES_CONFIG = None # None=使用内置默认规则; dict 列表=自定义规则覆盖
|
||||
HEADING_SHORT_TEXT_ENABLED = True # 是否启用短中文文本标题识别(最易误判的规则,可单独关闭)
|
||||
HEADING_SHORT_TEXT_MAX_LENGTH = 20 # 短文本最大字符数阈值
|
||||
|
||||
# 表单类型二次校正
|
||||
# MinerU 解析 Word 文档时,某些表单被标记为 text,根据内容特征修正为 table
|
||||
FORM_RECLASSIFY_ENABLED = True # 是否启用 text->table 表单检测
|
||||
FORM_RECLASSIFY_MIN_INDICATORS = 2 # 最少命中几个表单特征指标才校正
|
||||
|
||||
# 分块参数
|
||||
CHUNK_SIZE = 1000
|
||||
|
||||
@@ -3,7 +3,6 @@ RAG 核心引擎模块
|
||||
|
||||
包含:
|
||||
- engine: RAGEngine 单例类,管理模型和共享资源
|
||||
- agentic: AgenticRAG 智能问答
|
||||
- bm25_index: BM25 关键词检索索引
|
||||
- chunker: 文本分块器
|
||||
"""
|
||||
|
||||
360
core/agentic.py
360
core/agentic.py
@@ -1,360 +0,0 @@
|
||||
"""
|
||||
Agentic RAG - 知识库智能问答系统
|
||||
|
||||
核心能力:
|
||||
1. 知识库检索 - 向量检索 + BM25 + Rerank
|
||||
2. 网络搜索 - 当知识库不足时自动搜索(需配置SERPER_API_KEY)
|
||||
3. 图谱检索 - 实体关系推理(需配置Neo4j)
|
||||
4. 多源融合 - 智能处理知识库和网络内容
|
||||
5. Agent决策 - 动态决定检索、改写、分解等操作
|
||||
|
||||
使用方式:
|
||||
from core.agentic import AgenticRAG
|
||||
|
||||
rag = AgenticRAG()
|
||||
result = rag.process("你的问题")
|
||||
print(result["answer"])
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from openai import OpenAI
|
||||
|
||||
# 配置日志
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 导入基础模块
|
||||
from core.engine import get_engine
|
||||
from core.llm_utils import call_llm, quick_yes_no, parse_json_from_response
|
||||
|
||||
# 导入基础常量和配置
|
||||
from .agentic_base import (
|
||||
API_KEY, BASE_URL, MODEL,
|
||||
HAS_SERPER,
|
||||
HAS_BUDGET,
|
||||
MAX_CONTEXT_TOKENS, MAX_CONTEXT_COUNT, RERANK_THRESHOLD,
|
||||
SOURCE_KB, SOURCE_WEB,
|
||||
)
|
||||
|
||||
# 导入 Mixin 类
|
||||
from .agentic_query import QueryRewriteMixin
|
||||
from .agentic_search import SearchMixin
|
||||
from .agentic_answer import AnswerMixin
|
||||
from .agentic_citation import CitationMixin
|
||||
from .agentic_media import RichMediaMixin
|
||||
from .agentic_quality import QualityMixin
|
||||
from .agentic_context import ContextMixin
|
||||
from .agentic_meta import MetaQuestionMixin
|
||||
|
||||
|
||||
class AgenticRAG(
|
||||
QueryRewriteMixin,
|
||||
SearchMixin,
|
||||
AnswerMixin,
|
||||
CitationMixin,
|
||||
RichMediaMixin,
|
||||
QualityMixin,
|
||||
ContextMixin,
|
||||
MetaQuestionMixin
|
||||
):
|
||||
"""
|
||||
Agentic RAG - 知识库智能问答
|
||||
|
||||
通过 Mixin 模式组合功能:
|
||||
- QueryRewriteMixin: 查询重写
|
||||
- SearchMixin: 检索功能
|
||||
- AnswerMixin: 答案生成
|
||||
- CitationMixin: 引用处理
|
||||
- RichMediaMixin: 富媒体处理
|
||||
- QualityMixin: 质量评估
|
||||
- ContextMixin: 上下文处理
|
||||
- MetaQuestionMixin: 元问题处理
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
max_iterations: int = 3,
|
||||
enable_web_search: bool = True,
|
||||
**kwargs
|
||||
):
|
||||
"""初始化"""
|
||||
self.max_iterations = max_iterations
|
||||
self.enable_web_search = enable_web_search and HAS_SERPER
|
||||
self.client = OpenAI(api_key=API_KEY, base_url=BASE_URL)
|
||||
|
||||
# 信息来源标记
|
||||
self.SOURCE_KB = SOURCE_KB
|
||||
self.SOURCE_WEB = SOURCE_WEB
|
||||
|
||||
# 初始化置信度门控
|
||||
try:
|
||||
from core.confidence_gate import create_gate
|
||||
self.confidence_gate = create_gate()
|
||||
except ImportError:
|
||||
self.confidence_gate = None
|
||||
|
||||
# 初始化多维质量评估器
|
||||
try:
|
||||
from core.quality_assessor import create_assessor
|
||||
self.quality_assessor = create_assessor()
|
||||
except ImportError:
|
||||
self.quality_assessor = None
|
||||
|
||||
# 初始化推理反思器
|
||||
try:
|
||||
from core.reasoning_reflector import create_reflector
|
||||
self.reasoning_reflector = create_reflector()
|
||||
except ImportError:
|
||||
self.reasoning_reflector = None
|
||||
|
||||
# 初始化循环防护器
|
||||
try:
|
||||
from core.loop_guard import create_guard
|
||||
self.loop_guard = create_guard(max_iterations=max_iterations)
|
||||
except ImportError:
|
||||
self.loop_guard = None
|
||||
|
||||
# Context Compression 配置
|
||||
self.MAX_CONTEXT_TOKENS = MAX_CONTEXT_TOKENS
|
||||
self.MAX_CONTEXT_COUNT = MAX_CONTEXT_COUNT
|
||||
self.RERANK_THRESHOLD = RERANK_THRESHOLD
|
||||
|
||||
# Answer Grounding 配置
|
||||
self.MAX_GROUNDING_RETRY = 1
|
||||
self.grounding_retry_count = 0
|
||||
|
||||
def should_rewrite(self, query: str, history: list = None) -> bool:
|
||||
"""判断是否需要重写查询"""
|
||||
# 口语化表达模式
|
||||
colloquial_patterns = [
|
||||
"这个", "那个", "它", "这", "那",
|
||||
"上面", "下面", "刚才", "之前",
|
||||
"能不能", "可以吗", "行不行",
|
||||
"怎么办", "怎么弄", "咋整"
|
||||
]
|
||||
|
||||
for pattern in colloquial_patterns:
|
||||
if pattern in query:
|
||||
return True
|
||||
|
||||
# 查询太短
|
||||
if len(query) < 5:
|
||||
return True
|
||||
|
||||
# 有对话历史时,可能需要实体补全
|
||||
if history:
|
||||
for msg in reversed(history[-3:]):
|
||||
if msg.get("role") == "user":
|
||||
prev_query = msg.get("content", "")
|
||||
# 如果当前查询缺少主语,可能需要补全
|
||||
if any(kw in query for kw in ["标准", "规定", "流程", "制度"]):
|
||||
if not any(kw in query for kw in ["报销", "出差", "请假", "工资", "合同"]):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
def process(self, query: str, verbose: bool = True, history: list = None,
|
||||
allowed_levels: list = None, role: str = None, department: str = None,
|
||||
emit_log=None) -> dict:
|
||||
"""
|
||||
主流程:智能问答
|
||||
|
||||
Args:
|
||||
query: 用户问题
|
||||
verbose: 是否打印详细过程
|
||||
history: 对话历史
|
||||
allowed_levels: 允许访问的安全级别
|
||||
role: 用户角色
|
||||
department: 用户部门
|
||||
emit_log: 日志发射函数(流式输出)
|
||||
|
||||
Returns:
|
||||
{
|
||||
"answer": str,
|
||||
"sources": list,
|
||||
"images": list,
|
||||
"tables": list,
|
||||
"citations": list,
|
||||
"log_trace": list
|
||||
}
|
||||
"""
|
||||
log_trace = []
|
||||
|
||||
# 1. 检查元问题
|
||||
if self._is_meta_question(query):
|
||||
answer = self._answer_meta_question(query, allowed_levels, role, department)
|
||||
return {
|
||||
"answer": answer,
|
||||
"sources": [],
|
||||
"images": [],
|
||||
"tables": [],
|
||||
"citations": [],
|
||||
"log_trace": [{"phase": "meta_question", "query": query}]
|
||||
}
|
||||
|
||||
# 2. 查询重写
|
||||
current_query = query
|
||||
if self.should_rewrite(query, history):
|
||||
current_query = self._rewrite_query(query, history)
|
||||
log_trace.append({"phase": "rewrite", "original": query, "rewritten": current_query})
|
||||
if emit_log:
|
||||
emit_log(f"📝 查询重写: {query} → {current_query}")
|
||||
|
||||
# 3. 知识库检索
|
||||
contexts = []
|
||||
try:
|
||||
engine = get_engine()
|
||||
if not engine._initialized:
|
||||
engine.initialize()
|
||||
|
||||
# 获取用户可访问的向量库
|
||||
from knowledge.manager import get_kb_manager
|
||||
kb_mgr = get_kb_manager()
|
||||
accessible = kb_mgr.get_accessible_collections(role or 'user', department or '', 'read')
|
||||
|
||||
# 统一使用 search_knowledge() — 生产路径的同一 API
|
||||
# search_knowledge() 返回 dict: {ids, documents, metadatas, distances},每项为 list[list]
|
||||
# top_k 与生产路径对齐(30),确保 Rerank 后仍有足够结果
|
||||
results = engine.search_knowledge(
|
||||
query=current_query,
|
||||
top_k=30,
|
||||
collections=accessible if accessible else None,
|
||||
)
|
||||
|
||||
docs = results.get('documents', [[]])[0]
|
||||
metas = results.get('metadatas', [[]])[0]
|
||||
dists = results.get('distances', [[]])[0]
|
||||
|
||||
for doc, meta, score in zip(docs, metas, dists):
|
||||
contexts.append({
|
||||
'doc': doc,
|
||||
'meta': meta,
|
||||
'score': 1 - score if score <= 1 else 1 / (1 + score), # 距离→相似度
|
||||
'source_type': self.SOURCE_KB,
|
||||
'query': current_query
|
||||
})
|
||||
|
||||
log_trace.append({"phase": "kb_search", "query": current_query, "results": len(contexts)})
|
||||
if emit_log:
|
||||
emit_log(f"🔍 知识库检索: {len(contexts)} 条结果")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"知识库检索失败: {e}")
|
||||
log_trace.append({"phase": "kb_search", "error": str(e)})
|
||||
|
||||
# 3.5 图片独立检索 + 打分选择(与生产路径对齐)
|
||||
selected_images = []
|
||||
try:
|
||||
from api.chat_routes import select_images
|
||||
selected_images = select_images(contexts, current_query)
|
||||
if selected_images and emit_log:
|
||||
emit_log(f"🖼️ 图片选择: {len(selected_images)} 张相关图片")
|
||||
log_trace.append({"phase": "image_selection", "count": len(selected_images)})
|
||||
except Exception as e:
|
||||
logger.debug(f"图片选择失败: {e}")
|
||||
|
||||
# 4. 上下文压缩
|
||||
contexts = self._compress_contexts(current_query, contexts)
|
||||
|
||||
# 5. 网络搜索(如果需要)
|
||||
web_contexts = []
|
||||
if self.enable_web_search and (
|
||||
not contexts or
|
||||
not self._is_kb_result_sufficient(current_query, [c['doc'] for c in contexts]) or
|
||||
self._should_web_search(current_query)
|
||||
):
|
||||
web_contexts = self._web_search_flow(current_query, log_trace, emit_log, verbose, allowed_levels)
|
||||
contexts.extend(web_contexts)
|
||||
|
||||
# 7. 生成答案(注入图片描述到上下文)
|
||||
if contexts:
|
||||
# 将选中图片的描述注入上下文,让 LLM 能"看到"图片内容
|
||||
if selected_images:
|
||||
image_contexts = []
|
||||
for i, img in enumerate(selected_images, 1):
|
||||
full_desc = img.get('full_description', '') or img.get('description', '')
|
||||
if full_desc:
|
||||
img_source = img.get('source', '')
|
||||
img_page = img.get('page', '')
|
||||
source_info = f"(来源:{img_source} 第{img_page}页)" if img_source and img_page else ""
|
||||
image_contexts.append({
|
||||
'doc': f"【图片{i}】{full_desc}{source_info}",
|
||||
'meta': {'source': img_source, 'page': img_page, 'chunk_type': img.get('type', 'image')},
|
||||
'score': img.get('score', 0),
|
||||
'source_type': self.SOURCE_KB,
|
||||
'query': current_query
|
||||
})
|
||||
# 图片上下文追加到知识库上下文前面(让 LLM 优先看到图片信息)
|
||||
contexts = image_contexts + contexts
|
||||
|
||||
answer = self._generate_fused_answer(current_query, contexts, allowed_levels)
|
||||
|
||||
# 答案验证(防止幻觉)
|
||||
if self.grounding_retry_count < self.MAX_GROUNDING_RETRY:
|
||||
answer = self._verify_and_refine_answer(current_query, answer, contexts)
|
||||
else:
|
||||
answer = self._generate_no_context_answer(current_query, allowed_levels)
|
||||
|
||||
# 8. 构建引用
|
||||
citations = self._attach_citations(answer, contexts)
|
||||
|
||||
# 9. 图片结果:优先使用 select_images 的结构化结果(含 URL + 打分)
|
||||
# 回退到 _extract_rich_media(从 metadata 提取)
|
||||
if selected_images:
|
||||
images_result = selected_images
|
||||
else:
|
||||
rich_media = self._extract_rich_media(contexts)
|
||||
images_result = rich_media.get("images", [])
|
||||
|
||||
return {
|
||||
"answer": answer,
|
||||
"sources": citations.get("sources", []),
|
||||
"images": images_result,
|
||||
"tables": citations.get("tables", []) if isinstance(citations, dict) else [],
|
||||
"citations": citations.get("citations", []),
|
||||
"log_trace": log_trace
|
||||
}
|
||||
|
||||
def chat_search(self, query: str, history: list = None, role: str = None,
|
||||
department: str = None, allowed_levels: list = None) -> dict:
|
||||
"""聊天式检索接口"""
|
||||
return self.process(
|
||||
query,
|
||||
verbose=False,
|
||||
history=history,
|
||||
role=role,
|
||||
department=department,
|
||||
allowed_levels=allowed_levels
|
||||
)
|
||||
|
||||
def chat(self):
|
||||
"""命令行交互模式"""
|
||||
print("🤖 Agentic RAG 已启动,输入 'quit' 退出")
|
||||
print("-" * 50)
|
||||
|
||||
history = []
|
||||
while True:
|
||||
try:
|
||||
query = input("\n👤 你: ").strip()
|
||||
if not query:
|
||||
continue
|
||||
if query.lower() in ['quit', 'exit', 'q']:
|
||||
print("👋 再见!")
|
||||
break
|
||||
|
||||
result = self.process(query, verbose=True, history=history)
|
||||
print(f"\n🤖 AI: {result['answer']}")
|
||||
|
||||
if result['sources']:
|
||||
print("\n📚 来源:")
|
||||
for src in result['sources'][:3]:
|
||||
print(f" - {src['source']}")
|
||||
|
||||
history.append({"role": "user", "content": query})
|
||||
history.append({"role": "assistant", "content": result['answer']})
|
||||
|
||||
except KeyboardInterrupt:
|
||||
print("\n👋 再见!")
|
||||
break
|
||||
except Exception as e:
|
||||
print(f"❌ 错误: {e}")
|
||||
@@ -1,214 +0,0 @@
|
||||
"""
|
||||
Agentic RAG - 答案生成 Mixin
|
||||
|
||||
包含答案生成、上下文构建、融合回答等方法
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
from .agentic_base import logger, MODEL, SOURCE_KB, SOURCE_WEB
|
||||
from core.llm_utils import call_llm
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class AnswerMixin:
|
||||
"""答案生成方法"""
|
||||
|
||||
def _generate_fused_answer(self, query: str, contexts: list, allowed_levels: list = None) -> str:
|
||||
"""生成融合答案 - 智能处理多源信息"""
|
||||
# 分离不同来源
|
||||
kb_contexts = [c for c in contexts if c.get('source_type') == self.SOURCE_KB]
|
||||
web_contexts = [c for c in contexts if c.get('source_type') == self.SOURCE_WEB]
|
||||
|
||||
# 如果没有任何上下文,检测是否因权限限制
|
||||
if not contexts:
|
||||
return self._generate_no_context_answer(query, allowed_levels)
|
||||
|
||||
# 正常生成答案
|
||||
context_str = self._build_context_string(kb_contexts, web_contexts)
|
||||
prompt = self._build_normal_answer_prompt(query, context_str, kb_contexts, web_contexts)
|
||||
|
||||
result = call_llm(
|
||||
self.client, prompt, MODEL,
|
||||
temperature=0.7,
|
||||
max_tokens=2000
|
||||
)
|
||||
return result or f"生成答案失败"
|
||||
|
||||
def _build_context_string(self, kb_contexts, web_contexts):
|
||||
"""构建上下文字符串 - FAQ 优先策略,按分数排序"""
|
||||
# 分离 FAQ 和普通知识库内容
|
||||
faq_contexts = [c for c in kb_contexts if c.get('meta', {}).get('chunk_type') == 'faq']
|
||||
regular_contexts = [c for c in kb_contexts if c.get('meta', {}).get('chunk_type') != 'faq']
|
||||
|
||||
# 按分数降序排列,确保最相关的内容优先展示
|
||||
regular_contexts.sort(key=lambda c: c.get('score', 0), reverse=True)
|
||||
|
||||
# FAQ 部分(优先展示)
|
||||
faq_parts = []
|
||||
for i, c in enumerate(faq_contexts[:3], 1):
|
||||
meta = c['meta']
|
||||
answer = meta.get('faq_answer', c['doc'])
|
||||
faq_parts.append(f"[FAQ-{i}] 常见问题\n问题:{c['doc']}\n标准答案:{answer}")
|
||||
|
||||
# 普通知识库部分(用 12 条,提升覆盖率)
|
||||
kb_parts = []
|
||||
for i, c in enumerate(regular_contexts[:12], 1):
|
||||
meta = c['meta']
|
||||
source_str = meta.get('source', '未知')
|
||||
section = meta.get('section', '')
|
||||
source_info = f"{source_str}"
|
||||
if section:
|
||||
source_info += f" > {section[:60]}"
|
||||
kb_parts.append(f"[知识库-{i}] {source_info}\n{c['doc']}")
|
||||
|
||||
web_parts = []
|
||||
for i, c in enumerate(web_contexts[:5], 1):
|
||||
meta = c['meta']
|
||||
web_parts.append(f"[网络-{i}] {meta.get('title', '')}\n来源:{meta.get('source', '')}\n{c['doc']}")
|
||||
|
||||
return "\n\n".join(faq_parts + kb_parts + web_parts)
|
||||
|
||||
def _build_normal_answer_prompt(self, query, context_str, kb_contexts, web_contexts):
|
||||
"""构建正常回答的提示词(与生产路径 generate_answer_stream 对齐)"""
|
||||
# 检测是否有图片上下文
|
||||
has_images = any(c.get('meta', {}).get('chunk_type') in ('image', 'chart', 'table')
|
||||
for c in kb_contexts)
|
||||
|
||||
image_instruction = ""
|
||||
if has_images:
|
||||
image_instruction = "\n5. 如果参考资料中包含【图片N】信息,请在回答中简要介绍每张图片的内容和用途"
|
||||
|
||||
return f"""你是一个严谨的知识库问答助手。你必须且只能根据用户提供的【参考资料】回答问题。
|
||||
|
||||
【参考资料】
|
||||
{context_str}
|
||||
|
||||
【用户问题】
|
||||
{query}
|
||||
|
||||
【回答要求】
|
||||
1. 如果参考资料中有答案,必须引用对应内容回答,并在回答末尾标注引用编号(如[1]、[2])
|
||||
2. 如果参考资料中确实没有相关信息,简短说明"参考资料中没有相关信息"即可,不要编造或补充资料外的内容
|
||||
3. 禁止使用参考资料以外的知识进行补充或推测
|
||||
4. 分点列举,条理清晰,语言简洁{image_instruction}
|
||||
|
||||
请仔细阅读以上全部参考资料后回答:"""
|
||||
|
||||
def _build_answer_prompt_with_permission(self, query, context_str, levels_str, sources_str, kb_contexts, web_contexts):
|
||||
"""构建带权限提示的回答提示词"""
|
||||
return f"""你是一个严谨的智能助手。
|
||||
|
||||
【用户问题】
|
||||
{query}
|
||||
|
||||
【重要提示】
|
||||
检测到与用户问题更相关的信息可能存在于「{levels_str}」级别的文档中,但用户当前的权限级别无法访问。
|
||||
|
||||
【可访问的信息来源】
|
||||
{context_str}
|
||||
|
||||
【回答要求】
|
||||
1. 首先明确告知用户:当前回答基于您有权限访问的文档,可能不完整
|
||||
2. 基于现有信息如实回答
|
||||
3. 建议用户如需完整信息,请联系管理员申请相应权限
|
||||
|
||||
请回答:"""
|
||||
|
||||
def _generate_no_context_answer(self, query: str, allowed_levels: list = None) -> str:
|
||||
"""无上下文时的回答 — 诚实告知,不编造"""
|
||||
return "参考资料中没有找到与该问题相关的信息,无法根据现有知识库内容回答您的问题。"
|
||||
|
||||
def _verify_and_refine_answer(self, query: str, answer: str, contexts: list) -> str:
|
||||
"""验证并精炼答案 - 防止幻觉
|
||||
|
||||
返回值始终是干净的答案文本,不包含验证推理过程。
|
||||
"""
|
||||
prompt = f"""请检查以下回答是否存在"幻觉"(与参考信息不符的内容)。
|
||||
|
||||
【用户问题】
|
||||
{query}
|
||||
|
||||
【参考信息】
|
||||
{chr(10).join([f"[{i+1}] {c['doc'][:200]}" for i, c in enumerate(contexts[:8])])}
|
||||
|
||||
【AI回答】
|
||||
{answer}
|
||||
|
||||
【检查规则】
|
||||
1. 逐条核对回答中的事实是否能在参考信息中找到依据
|
||||
2. 如果没有幻觉,只回复一个英文单词:PASS
|
||||
3. 如果有幻觉,只输出修正后的完整回答(不要输出检查过程、不要加标题、不要加"检查结果"等前缀)
|
||||
|
||||
修正后的回答:"""
|
||||
|
||||
try:
|
||||
result = call_llm(
|
||||
self.client, prompt, MODEL,
|
||||
temperature=0.1,
|
||||
max_tokens=2000
|
||||
)
|
||||
if not result:
|
||||
return answer
|
||||
# 如果返回 PASS 或很短的确认,说明无幻觉
|
||||
cleaned = result.strip()
|
||||
if cleaned.upper() == "PASS" or len(cleaned) < 10:
|
||||
return answer
|
||||
# 有幻觉时,返回修正后的干净答案(去掉可能的前缀)
|
||||
for prefix in ["修正后的回答:", "修正后回答:", "修正回答:", "修正后:"]:
|
||||
if cleaned.startswith(prefix):
|
||||
cleaned = cleaned[len(prefix):].strip()
|
||||
return cleaned
|
||||
except Exception as e:
|
||||
logger.warning(f"答案验证失败: {e}")
|
||||
return answer
|
||||
|
||||
def _generate_uncertain_answer(self, query: str, contexts: list) -> str:
|
||||
"""生成不确定性回答"""
|
||||
context_str = "\n".join([c['doc'][:200] for c in contexts[:3]])
|
||||
|
||||
prompt = f"""用户问题:{query}
|
||||
|
||||
找到的信息可能不够完整或相关性不高:
|
||||
{context_str}
|
||||
|
||||
请基于这些信息给出一个谨慎的回答,明确说明哪些部分是有依据的,哪些部分可能需要更多验证。
|
||||
|
||||
回答:"""
|
||||
|
||||
try:
|
||||
result = call_llm(
|
||||
self.client, prompt, MODEL,
|
||||
temperature=0.7,
|
||||
max_tokens=1000
|
||||
)
|
||||
return result or "根据现有信息无法确定答案。"
|
||||
except Exception as e:
|
||||
logger.error(f"生成不确定性回答失败: {e}")
|
||||
return "根据现有信息无法确定答案。"
|
||||
|
||||
def _direct_answer(self, query: str, history: list = None) -> str:
|
||||
"""直接使用 LLM 回答(无知识库检索)"""
|
||||
messages = [
|
||||
{"role": "system", "content": "你是一个专业的助手,请用中文回答用户的问题。"}
|
||||
]
|
||||
|
||||
if history:
|
||||
for h in history[-4:]:
|
||||
if h.get("role") in ["user", "assistant"]:
|
||||
messages.append({"role": h["role"], "content": h.get("content", "")})
|
||||
|
||||
messages.append({"role": "user", "content": query})
|
||||
|
||||
try:
|
||||
result = call_llm(
|
||||
self.client, "", MODEL,
|
||||
temperature=0.7,
|
||||
max_tokens=1500,
|
||||
messages=messages
|
||||
)
|
||||
return result or "抱歉,我无法回答这个问题。"
|
||||
except Exception as e:
|
||||
logger.error(f"直接回答失败: {e}")
|
||||
return f"回答生成失败:{str(e)}"
|
||||
@@ -1,46 +0,0 @@
|
||||
"""
|
||||
Agentic RAG - 基础模块
|
||||
|
||||
包含常量、导入和共享配置
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
# 配置日志
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 尝试导入搜索API配置
|
||||
try:
|
||||
from config import SERPER_API_KEY
|
||||
HAS_SERPER = True
|
||||
except ImportError:
|
||||
HAS_SERPER = False
|
||||
SERPER_API_KEY = None
|
||||
|
||||
# LLM 预算控制
|
||||
try:
|
||||
from core.llm_budget import get_budget_controller, should_use_agent, CallType
|
||||
HAS_BUDGET = True
|
||||
except ImportError:
|
||||
HAS_BUDGET = False
|
||||
CallType = None
|
||||
|
||||
# LLM 配置
|
||||
try:
|
||||
from config import API_KEY, BASE_URL, MODEL
|
||||
except ImportError:
|
||||
API_KEY = None
|
||||
BASE_URL = None
|
||||
MODEL = None
|
||||
|
||||
# 来源标记
|
||||
SOURCE_KB = "知识库"
|
||||
SOURCE_WEB = "网络搜索"
|
||||
|
||||
# Context Compression 配置
|
||||
MAX_CONTEXT_TOKENS = 8000 # 最大上下文 token 数(与生产路径对齐)
|
||||
MAX_CONTEXT_COUNT = 20 # 最大上下文数量
|
||||
RERANK_THRESHOLD = 0.3 # Rerank 过滤阈值
|
||||
|
||||
# Answer Grounding 配置
|
||||
MAX_GROUNDING_RETRY = 1 # 幻觉修正最多重试次数
|
||||
@@ -1,237 +0,0 @@
|
||||
"""
|
||||
Agentic RAG - 引用处理 Mixin
|
||||
|
||||
包含来源提取、引用构建、引用附加等方法
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
|
||||
from .agentic_base import logger, SOURCE_KB
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class CitationMixin:
|
||||
"""引用处理方法"""
|
||||
|
||||
def _extract_sources(self, contexts: list) -> list:
|
||||
"""提取来源列表,返回结构化定位信息"""
|
||||
source_map = {}
|
||||
|
||||
for c in contexts:
|
||||
meta = c.get('meta', {})
|
||||
source_type = c.get('source_type', '未知')
|
||||
|
||||
if source_type == self.SOURCE_KB:
|
||||
source_key = meta.get('source', '未知')
|
||||
page = meta.get('page')
|
||||
page_end = meta.get('page_end', page)
|
||||
section = meta.get('section', '')
|
||||
doc_type = meta.get('doc_type', 'other')
|
||||
preview = meta.get('preview', '')
|
||||
section_chunk_id = meta.get('section_chunk_id')
|
||||
else:
|
||||
source_key = meta.get('title', meta.get('source', '未知'))
|
||||
page = None
|
||||
page_end = None
|
||||
section = ''
|
||||
doc_type = 'other'
|
||||
preview = ''
|
||||
section_chunk_id = None
|
||||
|
||||
if source_key not in source_map:
|
||||
source_map[source_key] = {
|
||||
"source": source_key,
|
||||
"type": source_type,
|
||||
"doc_type": doc_type,
|
||||
"count": 0,
|
||||
"pages": [],
|
||||
"sections": set(),
|
||||
"previews": [],
|
||||
"section_chunk_ids": set()
|
||||
}
|
||||
|
||||
source_map[source_key]["count"] += 1
|
||||
|
||||
if page:
|
||||
page_range = (page, page_end if page_end else page)
|
||||
if page_range not in source_map[source_key]["pages"]:
|
||||
source_map[source_key]["pages"].append(page_range)
|
||||
|
||||
if section:
|
||||
source_map[source_key]["sections"].add(section)
|
||||
|
||||
if preview and len(source_map[source_key]["previews"]) < 3:
|
||||
if preview not in source_map[source_key]["previews"]:
|
||||
source_map[source_key]["previews"].append(preview)
|
||||
|
||||
if section_chunk_id:
|
||||
source_map[source_key]["section_chunk_ids"].add(section_chunk_id)
|
||||
|
||||
sources = []
|
||||
for key, info in source_map.items():
|
||||
source_str = info["source"]
|
||||
doc_type = info.get("doc_type", "other")
|
||||
location_parts = []
|
||||
|
||||
if doc_type == 'pdf':
|
||||
if info["pages"]:
|
||||
valid_pages = [(s, e) for s, e in info["pages"] if s > 1 or e > 1]
|
||||
if valid_pages or not info["sections"]:
|
||||
page_strs = []
|
||||
for start, end in sorted(info["pages"], key=lambda x: x[0]):
|
||||
if start == end:
|
||||
page_strs.append(f"第{start}页")
|
||||
else:
|
||||
page_strs.append(f"第{start}-{end}页")
|
||||
location_parts.append(", ".join(page_strs))
|
||||
|
||||
if info["sections"]:
|
||||
sections_list = sorted(info["sections"])[:3]
|
||||
sections_str = "、".join(sections_list)
|
||||
if len(info["sections"]) > 3:
|
||||
sections_str += f"等{len(info['sections'])}个章节"
|
||||
location_parts.append(sections_str)
|
||||
|
||||
elif doc_type == 'word':
|
||||
if info["sections"]:
|
||||
sections_list = sorted(info["sections"])[:3]
|
||||
sections_str = "、".join(sections_list)
|
||||
if len(info["sections"]) > 3:
|
||||
sections_str += f"等{len(info['sections'])}个章节"
|
||||
location_parts.append(sections_str)
|
||||
|
||||
if info.get("section_chunk_ids"):
|
||||
chunk_ids = sorted(info["section_chunk_ids"])[:5]
|
||||
if chunk_ids:
|
||||
chunk_str = f"第{chunk_ids[0]}"
|
||||
if len(chunk_ids) > 1:
|
||||
chunk_str = f"第{chunk_ids[0]}-{chunk_ids[-1]}段"
|
||||
location_parts.append(chunk_str)
|
||||
|
||||
elif doc_type == 'excel':
|
||||
if info["sections"]:
|
||||
sections_list = sorted(info["sections"])[:3]
|
||||
sections_str = "、".join(sections_list)
|
||||
location_parts.append(sections_str)
|
||||
|
||||
else:
|
||||
if info["pages"]:
|
||||
valid_pages = [(s, e) for s, e in info["pages"] if s > 1 or e > 1]
|
||||
if valid_pages or not info["sections"]:
|
||||
page_strs = []
|
||||
for start, end in sorted(info["pages"], key=lambda x: x[0]):
|
||||
if start == end:
|
||||
page_strs.append(f"第{start}页")
|
||||
else:
|
||||
page_strs.append(f"第{start}-{end}页")
|
||||
location_parts.append(", ".join(page_strs))
|
||||
|
||||
if info["sections"]:
|
||||
sections_list = sorted(info["sections"])[:3]
|
||||
sections_str = "、".join(sections_list)
|
||||
if len(info["sections"]) > 3:
|
||||
sections_str += f"等{len(info['sections'])}个章节"
|
||||
location_parts.append(sections_str)
|
||||
|
||||
if location_parts:
|
||||
source_str = f"{source_str} ({' | '.join(location_parts)})"
|
||||
|
||||
sources.append({
|
||||
"source": source_str,
|
||||
"type": info["type"],
|
||||
"count": info["count"],
|
||||
"doc_type": doc_type,
|
||||
"previews": info.get("previews", []),
|
||||
"section_chunk_ids": sorted(info.get("section_chunk_ids", []))[:5]
|
||||
})
|
||||
|
||||
return sources
|
||||
|
||||
def _build_citation(self, meta: dict) -> dict:
|
||||
"""根据文档类型构建定位信息"""
|
||||
# 从 chunk_id 中提取全局切片序号(格式: "filename_N")
|
||||
chunk_id_raw = meta.get('chunk_id', '')
|
||||
chunk_index = None
|
||||
if chunk_id_raw and '_' in str(chunk_id_raw):
|
||||
try:
|
||||
chunk_index = int(str(chunk_id_raw).rsplit('_', 1)[-1])
|
||||
except (ValueError, IndexError):
|
||||
chunk_index = meta.get('chunk_index')
|
||||
else:
|
||||
chunk_index = meta.get('chunk_index')
|
||||
|
||||
citation = {
|
||||
"chunk_id": chunk_id_raw,
|
||||
"chunk_index": chunk_index, # 全局切片序号,用于前端文档预览跳转
|
||||
"source": meta.get('source', ''),
|
||||
"collection": meta.get('_collection', ''), # 所属向量库,用于前端文档预览跳转
|
||||
"doc_type": meta.get('doc_type', 'other'),
|
||||
"section": meta.get('section', ''),
|
||||
"preview": meta.get('preview', ''),
|
||||
"content": meta.get('preview', ''), # 初始用 preview,_attach_citations 中会用完整内容覆盖
|
||||
"chunk_type": meta.get('chunk_type', 'text'),
|
||||
}
|
||||
|
||||
doc_type = meta.get('doc_type', 'other')
|
||||
|
||||
if doc_type == 'pdf':
|
||||
bbox_raw = meta.get('bbox')
|
||||
bbox = None
|
||||
if bbox_raw:
|
||||
try:
|
||||
bbox = json.loads(bbox_raw) if isinstance(bbox_raw, str) else bbox_raw
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
bbox = bbox_raw
|
||||
|
||||
citation.update({
|
||||
"page": meta.get('page'),
|
||||
"page_end": meta.get('page_end'),
|
||||
"bbox": bbox,
|
||||
"bbox_mode": meta.get('bbox_mode'),
|
||||
})
|
||||
elif doc_type == 'word':
|
||||
citation.update({
|
||||
"section_chunk_id": meta.get('section_chunk_id'),
|
||||
})
|
||||
elif doc_type == 'excel':
|
||||
citation.update({
|
||||
"page": meta.get('page'),
|
||||
})
|
||||
else:
|
||||
bbox_raw = meta.get('bbox')
|
||||
bbox = None
|
||||
if bbox_raw:
|
||||
try:
|
||||
bbox = json.loads(bbox_raw) if isinstance(bbox_raw, str) else bbox_raw
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
bbox = bbox_raw
|
||||
|
||||
citation.update({
|
||||
"page": meta.get('page'),
|
||||
"page_end": meta.get('page_end'),
|
||||
"bbox": bbox,
|
||||
"bbox_mode": meta.get('bbox_mode'),
|
||||
})
|
||||
|
||||
return citation
|
||||
|
||||
def _attach_citations(self, answer: str, contexts: list) -> dict:
|
||||
"""将引用信息附加到答案"""
|
||||
citations = []
|
||||
|
||||
for c in contexts:
|
||||
meta = c.get('meta', {})
|
||||
full_content = c.get('doc', '')
|
||||
citation = self._build_citation(meta)
|
||||
# 用上下文中的完整文档内容覆盖 content 字段
|
||||
if full_content:
|
||||
citation['content'] = full_content[:300]
|
||||
citations.append(citation)
|
||||
|
||||
return {
|
||||
"answer": answer,
|
||||
"citations": citations,
|
||||
"sources": self._extract_sources(contexts)
|
||||
}
|
||||
@@ -1,111 +0,0 @@
|
||||
"""
|
||||
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
|
||||
@@ -1,200 +0,0 @@
|
||||
"""
|
||||
Agentic RAG - 富媒体处理 Mixin
|
||||
|
||||
包含图表查找、图片提取、富媒体附加等方法
|
||||
"""
|
||||
|
||||
import re
|
||||
import json
|
||||
import logging
|
||||
|
||||
from .agentic_base import logger
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class RichMediaMixin:
|
||||
"""富媒体处理方法"""
|
||||
|
||||
def _find_figure(self, query: str, contexts: list, source: str = None) -> dict:
|
||||
"""精确查找图表,带 fallback"""
|
||||
patterns = [
|
||||
r'图\s*(\d+[\.\-]\d+)',
|
||||
r'Fig\.?\s*(\d+[\.\-]\d+)',
|
||||
r'Figure\s*(\d+[\.\-]\d+)',
|
||||
]
|
||||
|
||||
target_figure = None
|
||||
for pattern in patterns:
|
||||
match = re.search(pattern, query, re.IGNORECASE)
|
||||
if match:
|
||||
target_figure = match.group(1).replace('-', '.')
|
||||
break
|
||||
|
||||
if not target_figure:
|
||||
return {"found": False}
|
||||
|
||||
# 从 contexts 中查找
|
||||
for ctx in contexts:
|
||||
meta = ctx.get('meta', {})
|
||||
fig_num = meta.get('figure_number', '')
|
||||
if fig_num == target_figure:
|
||||
if not source or meta.get('source') == source:
|
||||
return {
|
||||
"found": True,
|
||||
"chunk_id": meta.get('chunk_id'),
|
||||
"source": meta.get('source'),
|
||||
"page": meta.get('page'),
|
||||
"caption": meta.get('caption'),
|
||||
"image_path": meta.get('image_path'),
|
||||
}
|
||||
|
||||
# Fallback: 直接查向量库
|
||||
try:
|
||||
from knowledge.manager import get_kb_manager
|
||||
kb_mgr = get_kb_manager()
|
||||
coll = kb_mgr.get_collection('public_kb')
|
||||
|
||||
if coll:
|
||||
where_conditions = [{'chunk_type': {'$in': ['image', 'chart']}}]
|
||||
if source:
|
||||
where_conditions.append({'source': source})
|
||||
|
||||
result = coll.get(
|
||||
where={'$and': where_conditions} if len(where_conditions) > 1 else where_conditions[0],
|
||||
include=['metadatas', 'documents']
|
||||
)
|
||||
|
||||
for meta, doc in zip(result.get('metadatas', []), result.get('documents', [])):
|
||||
if meta.get('figure_number') == target_figure:
|
||||
return {
|
||||
"found": True,
|
||||
"chunk_id": meta.get('chunk_id'),
|
||||
"source": meta.get('source'),
|
||||
"page": meta.get('page'),
|
||||
"caption": meta.get('caption'),
|
||||
"image_path": meta.get('image_path'),
|
||||
}
|
||||
caption = meta.get('caption', '') or (doc if doc else '')
|
||||
if f"图{target_figure}" in caption or f"图 {target_figure}" in caption:
|
||||
return {
|
||||
"found": True,
|
||||
"chunk_id": meta.get('chunk_id'),
|
||||
"source": meta.get('source'),
|
||||
"page": meta.get('page'),
|
||||
"caption": meta.get('caption'),
|
||||
"image_path": meta.get('image_path'),
|
||||
}
|
||||
except Exception as e:
|
||||
logger.warning(f"_find_figure fallback 查询失败: {e}")
|
||||
|
||||
return {"found": False}
|
||||
|
||||
def _get_images_for_source(self, source: str, collections: list = None) -> list:
|
||||
"""直接从向量库获取指定文件的所有图片"""
|
||||
try:
|
||||
from knowledge.manager import get_kb_manager
|
||||
kb_mgr = get_kb_manager()
|
||||
except ImportError:
|
||||
return []
|
||||
|
||||
images = []
|
||||
seen_ids = set()
|
||||
|
||||
target_collections = collections or ['public_kb']
|
||||
|
||||
for kb_name in target_collections:
|
||||
try:
|
||||
coll = kb_mgr.get_collection(kb_name)
|
||||
if not coll:
|
||||
continue
|
||||
|
||||
result = coll.get(
|
||||
where={'source': source},
|
||||
include=['metadatas']
|
||||
)
|
||||
|
||||
for meta in result.get('metadatas', []):
|
||||
images_json = meta.get('images_json')
|
||||
if images_json:
|
||||
try:
|
||||
imgs = json.loads(images_json)
|
||||
for img in imgs:
|
||||
img_id = img.get('id')
|
||||
if img_id and img_id not in seen_ids:
|
||||
seen_ids.add(img_id)
|
||||
images.append({
|
||||
"id": img_id,
|
||||
"caption": img.get("caption", ""),
|
||||
"url": f"/images/{img_id}",
|
||||
"page": img.get("page") or meta.get("page"),
|
||||
"source": source,
|
||||
"width": img.get("width"),
|
||||
"height": img.get("height")
|
||||
})
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
except Exception as e:
|
||||
logger.warning(f"从 {kb_name} 获取图片失败: {e}")
|
||||
continue
|
||||
|
||||
return images
|
||||
|
||||
def _extract_rich_media(self, contexts: list, sources_filter: list = None, max_images: int = 10,
|
||||
max_tables: int = 5) -> dict:
|
||||
"""从检索结果中提取富媒体(图片、表格)"""
|
||||
images = []
|
||||
tables = []
|
||||
seen_image_ids = set()
|
||||
seen_table_ids = set()
|
||||
|
||||
for ctx in contexts:
|
||||
meta = ctx.get('meta', {})
|
||||
source = meta.get('source', '')
|
||||
|
||||
# 过滤来源
|
||||
if sources_filter and source not in sources_filter:
|
||||
continue
|
||||
|
||||
# 提取图片
|
||||
images_json = meta.get('images_json')
|
||||
if images_json:
|
||||
try:
|
||||
imgs = json.loads(images_json)
|
||||
for img in imgs:
|
||||
img_id = img.get('id')
|
||||
if img_id and img_id not in seen_image_ids:
|
||||
seen_image_ids.add(img_id)
|
||||
images.append({
|
||||
"id": img_id,
|
||||
"caption": img.get("caption", ""),
|
||||
"url": f"/images/{img_id}",
|
||||
"page": img.get("page") or meta.get("page"),
|
||||
"source": source,
|
||||
"type": img.get("type", "image")
|
||||
})
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
|
||||
# 提取表格
|
||||
table_json = meta.get('table_json')
|
||||
if table_json:
|
||||
try:
|
||||
tbl = json.loads(table_json)
|
||||
tbl_id = tbl.get('id') or meta.get('chunk_id')
|
||||
if tbl_id and tbl_id not in seen_table_ids:
|
||||
seen_table_ids.add(tbl_id)
|
||||
tables.append({
|
||||
"id": tbl_id,
|
||||
"caption": tbl.get("caption", ""),
|
||||
"markdown": tbl.get("markdown", ""),
|
||||
"page": meta.get("page"),
|
||||
"source": source
|
||||
})
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
|
||||
return {
|
||||
"images": images[:max_images],
|
||||
"tables": tables[:max_tables]
|
||||
}
|
||||
@@ -1,133 +0,0 @@
|
||||
"""
|
||||
Agentic RAG - 元问题处理 Mixin
|
||||
|
||||
包含元问题判断和知识库元数据回答方法
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
from .agentic_base import logger
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class MetaQuestionMixin:
|
||||
"""元问题处理方法"""
|
||||
|
||||
def _is_meta_question(self, query: str) -> bool:
|
||||
"""判断是否为元问题(关于知识库本身的问题)"""
|
||||
meta_patterns = [
|
||||
"有哪些文件", "什么文件", "哪些文件", "文件列表", "文件目录",
|
||||
"可以查看", "能查看", "有权限查看", "权限查看",
|
||||
"能访问", "可以访问", "有权限访问",
|
||||
"我的权限", "用户权限", "查看权限", "访问权限",
|
||||
"权限能", "权限可以", "有什么权限", "有哪些权限",
|
||||
"我能看", "我可以看", "我能查", "我可以查",
|
||||
"能看到什么", "能查到什么", "可以看什么", "可以查什么",
|
||||
"知识库有哪些", "库里有", "文档有哪些", "有哪些文档",
|
||||
"有什么文档", "有什么文件", "包含什么", "包含哪些",
|
||||
"你知道什么", "你都知道", "你能回答什么",
|
||||
"系统里有什么", "库里有什么",
|
||||
"public_kb", "dept_tech", "dept_hr", "dept_finance", "dept_operation",
|
||||
"kb里", "向量库", "有哪些库", "库列表", "kb有哪些"
|
||||
]
|
||||
query_lower = query.lower()
|
||||
return any(kw in query_lower for kw in meta_patterns)
|
||||
|
||||
def _answer_meta_question(self, query: str, allowed_levels: list = None,
|
||||
role: str = None, department: str = None) -> str:
|
||||
"""回答元问题(关于知识库本身的问题)"""
|
||||
try:
|
||||
source_map = {}
|
||||
|
||||
try:
|
||||
from knowledge.manager import get_kb_manager
|
||||
from auth.gateway import get_accessible_collections as _get_accessible
|
||||
|
||||
kb_mgr = get_kb_manager()
|
||||
accessible = _get_accessible(role or 'user', department or '', 'read')
|
||||
|
||||
for kb_name in accessible:
|
||||
coll = kb_mgr.get_collection(kb_name)
|
||||
if not coll:
|
||||
continue
|
||||
try:
|
||||
result = coll.get(include=['metadatas'])
|
||||
except Exception as e:
|
||||
logger.debug(f"获取{kb_name}元数据失败: {e}")
|
||||
continue
|
||||
|
||||
for meta in result.get('metadatas', []):
|
||||
source = meta.get('source', '未知')
|
||||
level = meta.get('security_level', 'public')
|
||||
page = meta.get('page')
|
||||
|
||||
if source not in source_map:
|
||||
source_map[source] = {
|
||||
'count': 0, 'levels': set(),
|
||||
'pages': set(), 'collections': set()
|
||||
}
|
||||
|
||||
source_map[source]['count'] += 1
|
||||
source_map[source]['levels'].add(level)
|
||||
source_map[source]['collections'].add(kb_name)
|
||||
if page:
|
||||
source_map[source]['pages'].add(page)
|
||||
|
||||
except ImportError:
|
||||
from core.engine import get_engine
|
||||
all_docs = get_engine().collection.get(include=['metadatas'])
|
||||
for meta in all_docs.get('metadatas', []):
|
||||
source = meta.get('source', '未知')
|
||||
level = meta.get('security_level', 'public')
|
||||
page = meta.get('page')
|
||||
|
||||
if source not in source_map:
|
||||
source_map[source] = {
|
||||
'count': 0, 'levels': set(),
|
||||
'pages': set(), 'collections': set()
|
||||
}
|
||||
|
||||
source_map[source]['count'] += 1
|
||||
source_map[source]['levels'].add(level)
|
||||
if page:
|
||||
source_map[source]['pages'].add(page)
|
||||
|
||||
# 根据安全级别过滤
|
||||
if allowed_levels:
|
||||
allowed_set = set(allowed_levels)
|
||||
filtered_sources = {}
|
||||
for source, info in source_map.items():
|
||||
if info['levels'] & allowed_set:
|
||||
filtered_sources[source] = info
|
||||
source_map = filtered_sources
|
||||
|
||||
if not source_map:
|
||||
return "抱歉,您当前没有权限查看任何文档,或者知识库为空。"
|
||||
|
||||
sorted_sources = sorted(source_map.items(), key=lambda x: x[1]['count'], reverse=True)
|
||||
|
||||
answer_parts = [f"📚 **知识库文档列表**(共 {len(sorted_sources)} 个文档)\n"]
|
||||
|
||||
for i, (source, info) in enumerate(sorted_sources, 1):
|
||||
colls = info.get('collections', set())
|
||||
coll_str = f",所属: {', '.join(sorted(colls))}" if colls else ""
|
||||
pages_str = ''
|
||||
if info['pages']:
|
||||
pages_list = sorted(info['pages'])
|
||||
if len(pages_list) <= 5:
|
||||
pages_str = f",页码: {', '.join(map(str, pages_list))}"
|
||||
else:
|
||||
pages_str = f",共 {len(info['pages'])} 页"
|
||||
|
||||
answer_parts.append(f"{i}. **{source}** ({info['count']} 条片段{coll_str}{pages_str})")
|
||||
|
||||
answer_parts.append(f"\n**总计**: {sum(s[1]['count'] for s in sorted_sources)} 条知识片段")
|
||||
answer_parts.append(f"\n**您的权限级别**: {', '.join(allowed_levels) if allowed_levels else '全部'}")
|
||||
|
||||
answer_parts.append("\n\n💡 **提示**: 您可以直接提问关于这些文档内容的问题。")
|
||||
|
||||
return '\n'.join(answer_parts)
|
||||
|
||||
except Exception as e:
|
||||
return f"获取文档列表时出错: {str(e)}\n\n您可以直接提问,我会尝试从知识库中检索相关信息。"
|
||||
@@ -1,137 +0,0 @@
|
||||
"""
|
||||
Agentic RAG - 质量评估 Mixin
|
||||
|
||||
包含置信度门控、质量评估、推理反思等方法
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
from .agentic_base import logger
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class QualityMixin:
|
||||
"""质量评估方法"""
|
||||
|
||||
def _check_confidence_gate(self, query: str, docs: list, verbose: bool = True,
|
||||
precomputed_scores: list = None):
|
||||
"""检查置信度门控
|
||||
|
||||
Args:
|
||||
query: 用户查询
|
||||
docs: 文档列表
|
||||
verbose: 是否详细输出
|
||||
precomputed_scores: 预计算的 Rerank 分数(可选,避免重复推理)
|
||||
"""
|
||||
if not self.confidence_gate:
|
||||
return {"passed": True, "reason": "no_gate"}
|
||||
|
||||
try:
|
||||
result = self.confidence_gate.evaluate(query, docs,
|
||||
precomputed_scores=precomputed_scores)
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.warning(f"置信度门控检查失败: {e}")
|
||||
return {"passed": True, "reason": "error"}
|
||||
|
||||
def _assess_quality(self, query: str, docs: list, metas: list = None,
|
||||
verbose: bool = True) -> dict:
|
||||
"""多维质量评估"""
|
||||
if not self.quality_assessor:
|
||||
return {"overall_score": 0.5, "dimensions": {}}
|
||||
|
||||
try:
|
||||
result = self.quality_assessor.assess(query, docs, metas)
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.warning(f"质量评估失败: {e}")
|
||||
return {"overall_score": 0.5, "dimensions": {}}
|
||||
|
||||
def _reflect_on_answer(self, query: str, answer: str, contexts: list,
|
||||
verbose: bool = True) -> dict:
|
||||
"""推理反思"""
|
||||
if not self.reasoning_reflector:
|
||||
return {"needs_reflection": False, "issues": []}
|
||||
|
||||
try:
|
||||
result = self.reasoning_reflector.reflect(query, answer, contexts)
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.warning(f"推理反思失败: {e}")
|
||||
return {"needs_reflection": False, "issues": []}
|
||||
|
||||
def _think(self, original_query: str, current_query: str,
|
||||
iteration: int, contexts: list, verbose: bool = True) -> dict:
|
||||
"""
|
||||
Agent 思考:决定下一步行动
|
||||
|
||||
Returns:
|
||||
{
|
||||
"action": "answer" | "rewrite" | "search_web" | "decompose",
|
||||
"reason": "...",
|
||||
"rewrite_query": "..." # 如果 action == "rewrite"
|
||||
}
|
||||
"""
|
||||
from core.llm_utils import call_llm, parse_json_from_response
|
||||
from .agentic_base import MODEL
|
||||
|
||||
# 构建思考提示
|
||||
context_summary = ""
|
||||
if contexts:
|
||||
for i, ctx in enumerate(contexts[:3], 1):
|
||||
meta = ctx.get('meta', {})
|
||||
source = meta.get('source', '未知')
|
||||
doc_preview = ctx.get('doc', '')[:100]
|
||||
context_summary += f"{i}. [{source}] {doc_preview}...\n"
|
||||
|
||||
prompt = f"""你是一个 RAG 系统的决策 Agent,需要判断下一步行动。
|
||||
|
||||
【原始问题】
|
||||
{original_query}
|
||||
|
||||
【当前问题】
|
||||
{current_query}
|
||||
|
||||
【迭代轮次】
|
||||
{iteration} / {self.max_iterations}
|
||||
|
||||
【已检索到的上下文】
|
||||
{context_summary if context_summary else "(无)"}
|
||||
|
||||
【可选行动】
|
||||
1. answer - 已有足够信息,可以回答
|
||||
2. rewrite - 查询不够清晰,需要重写
|
||||
3. search_web - 知识库信息不足,需要网络搜索
|
||||
4. decompose - 问题太复杂,需要分解
|
||||
|
||||
【决策要求】
|
||||
- 如果上下文足够回答问题,选择 answer
|
||||
- 如果上下文不足且迭代未超限,选择 search_web 或 rewrite
|
||||
- 返回 JSON 格式
|
||||
|
||||
请决策:"""
|
||||
|
||||
try:
|
||||
result = call_llm(
|
||||
self.client, prompt, MODEL,
|
||||
temperature=0.3,
|
||||
max_tokens=200
|
||||
)
|
||||
|
||||
decision = parse_json_from_response(result) if result else {}
|
||||
|
||||
# 默认决策
|
||||
if not decision or "action" not in decision:
|
||||
if contexts and len(contexts) >= 2:
|
||||
decision = {"action": "answer", "reason": "有足够上下文"}
|
||||
else:
|
||||
decision = {"action": "rewrite", "reason": "上下文不足"}
|
||||
|
||||
return decision
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Agent 思考失败: {e}")
|
||||
if contexts:
|
||||
return {"action": "answer", "reason": "默认回答"}
|
||||
return {"action": "rewrite", "reason": "默认重写"}
|
||||
@@ -1,271 +0,0 @@
|
||||
"""
|
||||
Agentic RAG - 查询重写 Mixin
|
||||
|
||||
包含查询改写、实体补全、专业术语映射等方法
|
||||
"""
|
||||
|
||||
import re
|
||||
import logging
|
||||
|
||||
from .agentic_base import logger, MODEL
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class QueryRewriteMixin:
|
||||
"""查询重写方法"""
|
||||
|
||||
def _rewrite_query(self, query: str, history: list = None,
|
||||
strategy: str = "professional") -> str:
|
||||
"""
|
||||
增强版查询重写:将口语化表达转为专业术语
|
||||
|
||||
Args:
|
||||
query: 原始查询
|
||||
history: 对话历史(用于实体补全)
|
||||
strategy: 重写策略
|
||||
- professional: 口语化→专业术语
|
||||
- expand: 扩展关键词
|
||||
- clarify: 消歧义
|
||||
- entity: 实体补全
|
||||
|
||||
Returns:
|
||||
str: 重写后的查询
|
||||
"""
|
||||
# 尝试多种策略组合
|
||||
rewritten = query
|
||||
|
||||
# 策略1: 口语化→专业术语映射
|
||||
if strategy in ["professional", "all"]:
|
||||
rewritten = self._apply_professional_mapping(rewritten)
|
||||
|
||||
# 策略2: 实体补全(利用对话历史)
|
||||
if strategy in ["entity", "all"] and history:
|
||||
rewritten = self._complete_entities(rewritten, history)
|
||||
|
||||
# 策略3: LLM 深度重写(仅在需要时调用)
|
||||
if strategy in ["professional", "all"]:
|
||||
llm_rewritten = self._llm_rewrite(rewritten)
|
||||
if llm_rewritten and len(llm_rewritten) > len(rewritten) * 0.5:
|
||||
rewritten = llm_rewritten
|
||||
|
||||
return rewritten
|
||||
|
||||
def _apply_professional_mapping(self, query: str) -> str:
|
||||
"""应用口语化→专业术语映射"""
|
||||
TERM_MAPPING = {
|
||||
"报销": "差旅报销 费用报销 报销审批",
|
||||
"请假": "休假申请 请假审批 考勤管理",
|
||||
"加班": "加班申请 工时管理 加班审批",
|
||||
"工资": "薪酬管理 工资发放 薪资结构",
|
||||
"合同": "合同管理 合同签署 合同审批",
|
||||
"流程": "审批流程 业务流程 工作流",
|
||||
"制度": "管理制度 规章制度 企业规范",
|
||||
"规定": "管理规定 制度规定 政策要求",
|
||||
"几天": "时限 审批时限 办理时限",
|
||||
"多久": "处理时效 审批周期 办理周期",
|
||||
"多少": "标准 额度 限额 标准",
|
||||
"能不能": "是否允许 是否可以 权限",
|
||||
"人事": "人力资源 HR 人力部门",
|
||||
"财务": "财务部 财务部门 财务管理",
|
||||
"技术": "技术部 研发部 IT部门",
|
||||
}
|
||||
|
||||
result = query
|
||||
for colloquial, professional in TERM_MAPPING.items():
|
||||
if colloquial in query:
|
||||
result = result.replace(colloquial, f"{colloquial} {professional.split()[0]}")
|
||||
|
||||
return result
|
||||
|
||||
def _complete_entities(self, query: str, history: list) -> str:
|
||||
"""实体补全:利用对话历史补充缺失的实体"""
|
||||
if not history:
|
||||
return query
|
||||
|
||||
# 图片指代识别
|
||||
image_reference = self._detect_image_reference(query, history)
|
||||
if image_reference:
|
||||
return image_reference
|
||||
|
||||
# 获取最近用户消息
|
||||
last_user_msg = None
|
||||
for msg in reversed(history):
|
||||
if msg.get("role") == "user":
|
||||
last_user_msg = msg.get("content", "")
|
||||
break
|
||||
|
||||
if not last_user_msg:
|
||||
return query
|
||||
|
||||
# 检查当前查询是否缺少主语
|
||||
BUSINESS_KEYWORDS = ["报销", "出差", "请假", "工资", "合同", "审批", "流程",
|
||||
"制度", "规定", "标准", "金额", "时间"]
|
||||
|
||||
has_subject = any(kw in query for kw in BUSINESS_KEYWORDS)
|
||||
|
||||
if not has_subject:
|
||||
try:
|
||||
import jieba
|
||||
entities = []
|
||||
for word in jieba.cut(last_user_msg):
|
||||
word = word.strip()
|
||||
if len(word) >= 2 and any(kw in word for kw in BUSINESS_KEYWORDS):
|
||||
entities.append(word)
|
||||
|
||||
if entities:
|
||||
return f"{entities[0]} {query}"
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
return query
|
||||
|
||||
def _detect_image_reference(self, query: str, history: list) -> str:
|
||||
"""检测图片指代查询并重写"""
|
||||
IMAGE_REFERENCE_PATTERNS = [
|
||||
r'这[张些]图片', r'那[张些]图片', r'上面的图片', r'刚才的图片',
|
||||
r'这[张些]图', r'那[张些]图', r'上面的图', r'刚才的图',
|
||||
r'解释一下这[张些]图', r'说明一下这[张些]图',
|
||||
r'这[张些]是什么图', r'图[里内]是什么', r'图片[里内]是什么',
|
||||
]
|
||||
|
||||
is_image_reference = False
|
||||
for pattern in IMAGE_REFERENCE_PATTERNS:
|
||||
if re.search(pattern, query):
|
||||
is_image_reference = True
|
||||
break
|
||||
|
||||
if not is_image_reference:
|
||||
return ""
|
||||
|
||||
last_images = []
|
||||
for msg in reversed(history):
|
||||
if msg.get("role") == "assistant":
|
||||
metadata = msg.get("metadata", {})
|
||||
if isinstance(metadata, dict):
|
||||
images = metadata.get("images", [])
|
||||
if images:
|
||||
for img in images[:5]:
|
||||
if isinstance(img, dict):
|
||||
desc = img.get("description", "")
|
||||
img_type = img.get("type", "图片")
|
||||
if desc:
|
||||
last_images.append(f"{img_type}:{desc}")
|
||||
elif isinstance(img, str):
|
||||
last_images.append(f"图片:{img}")
|
||||
|
||||
if not last_images:
|
||||
content = msg.get("content", "")
|
||||
if "图片" in content or "图表" in content or "图" in content:
|
||||
sentences = content.split("。")
|
||||
for sentence in sentences:
|
||||
if "图片" in sentence or "图表" in sentence:
|
||||
last_images.append(sentence.strip())
|
||||
|
||||
if last_images:
|
||||
break
|
||||
|
||||
if last_images:
|
||||
image_context = " ".join(last_images[:3])
|
||||
question_intent = re.sub(
|
||||
r'这[张些]图片?|那[张些]图片?|上面的图片?|刚才的图片?|解释一下|说明一下',
|
||||
'', query
|
||||
).strip()
|
||||
|
||||
if question_intent:
|
||||
return f"{image_context} {question_intent}"
|
||||
else:
|
||||
return f"详细解释:{image_context}"
|
||||
|
||||
return query
|
||||
|
||||
def _extract_image_context_from_history(self, history: list) -> str:
|
||||
"""从对话历史中提取图片上下文"""
|
||||
if not history:
|
||||
return ""
|
||||
|
||||
for msg in reversed(history):
|
||||
if msg.get("role") == "assistant":
|
||||
metadata = msg.get("metadata", {})
|
||||
images = metadata.get("images", [])
|
||||
content = msg.get("content", "")
|
||||
|
||||
image_descriptions = []
|
||||
|
||||
if images:
|
||||
for i, img in enumerate(images[:5], 1):
|
||||
if isinstance(img, dict):
|
||||
desc = img.get("description", "")
|
||||
img_type = img.get("type", "图片")
|
||||
source = img.get("source", "")
|
||||
page = img.get("page", "")
|
||||
|
||||
img_info = f"图片{i}:{img_type}"
|
||||
if desc:
|
||||
img_info += f",描述:{desc}"
|
||||
if source:
|
||||
img_info += f",来源:{source}"
|
||||
if page:
|
||||
img_info += f",第{page}页"
|
||||
image_descriptions.append(img_info)
|
||||
|
||||
if not image_descriptions:
|
||||
if "图片" in content or "图表" in content:
|
||||
sentences = content.split("。")
|
||||
for sentence in sentences:
|
||||
if "图片" in sentence or "图表" in sentence:
|
||||
image_descriptions.append(sentence.strip())
|
||||
if len(image_descriptions) >= 3:
|
||||
break
|
||||
|
||||
if image_descriptions:
|
||||
return "\n".join(image_descriptions)
|
||||
|
||||
return ""
|
||||
|
||||
def _answer_image_reference(self, enhanced_query: str, history: list) -> str:
|
||||
"""回答图片引用问题"""
|
||||
from core.llm_utils import call_llm
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": "你是一个专业的助手,请根据提供的图片信息回答用户的问题。"}
|
||||
]
|
||||
|
||||
for h in history[-4:]:
|
||||
if h.get("role") in ["user", "assistant"]:
|
||||
messages.append({"role": h["role"], "content": h.get("content", "")})
|
||||
|
||||
messages.append({"role": "user", "content": enhanced_query})
|
||||
|
||||
try:
|
||||
result = call_llm(
|
||||
self.client, "", MODEL,
|
||||
temperature=0.3,
|
||||
max_tokens=1000,
|
||||
messages=messages
|
||||
)
|
||||
return result or ""
|
||||
except Exception as e:
|
||||
logger.error(f"图片引用回答失败: {e}")
|
||||
return f"抱歉,回答图片问题时出现错误:{str(e)}"
|
||||
|
||||
def _llm_rewrite(self, query: str) -> str:
|
||||
"""LLM 深度重写查询"""
|
||||
from core.llm_utils import call_llm
|
||||
|
||||
prompt = f"""请将以下用户问题改写为更专业、更清晰的表达,保持原意不变。
|
||||
|
||||
原问题:{query}
|
||||
|
||||
改写后的问题:"""
|
||||
|
||||
try:
|
||||
rewritten = call_llm(
|
||||
self.client, prompt, MODEL,
|
||||
temperature=0.3,
|
||||
max_tokens=100
|
||||
)
|
||||
return rewritten.strip() if rewritten else query
|
||||
except Exception as e:
|
||||
logger.warning(f"LLM 重写失败: {e}")
|
||||
return query
|
||||
@@ -1,152 +0,0 @@
|
||||
"""
|
||||
Agentic RAG - 检索 Mixin
|
||||
|
||||
包含知识库检索、网络搜索等方法
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import requests
|
||||
|
||||
from .agentic_base import (
|
||||
logger, HAS_SERPER, SERPER_API_KEY,
|
||||
SOURCE_KB, SOURCE_WEB
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SearchMixin:
|
||||
"""检索功能方法"""
|
||||
|
||||
def _web_search(self, query: str, top_k: int = 5) -> list:
|
||||
"""网络搜索(使用Serper API)"""
|
||||
if not HAS_SERPER:
|
||||
return []
|
||||
|
||||
try:
|
||||
url = "https://google.serper.dev/search"
|
||||
payload = json.dumps({
|
||||
"q": query,
|
||||
"gl": "cn",
|
||||
"hl": "zh-cn",
|
||||
"num": top_k
|
||||
})
|
||||
headers = {
|
||||
'X-API-KEY': SERPER_API_KEY,
|
||||
'Content-Type': 'application/json'
|
||||
}
|
||||
|
||||
response = requests.post(url, headers=headers, data=payload, timeout=10)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
results = []
|
||||
for item in data.get('organic', [])[:top_k]:
|
||||
results.append({
|
||||
'title': item.get('title', ''),
|
||||
'link': item.get('link', ''),
|
||||
'snippet': item.get('snippet', ''),
|
||||
'date': item.get('date', '')
|
||||
})
|
||||
|
||||
return results
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"网络搜索失败: {e}")
|
||||
return []
|
||||
|
||||
def _should_web_search(self, query: str) -> bool:
|
||||
"""判断是否需要网络搜索"""
|
||||
realtime_keywords = [
|
||||
"今天", "最新", "今日", "当前", "现在",
|
||||
"天气", "新闻", "股价", "行情", "汇率",
|
||||
"最近", "近期", "这周", "本月", "今年",
|
||||
"实时", "动态", "热点", "发生"
|
||||
]
|
||||
|
||||
query_lower = query.lower()
|
||||
return any(kw in query_lower for kw in realtime_keywords)
|
||||
|
||||
def _web_search_flow(self, query: str, log_trace: list, emit_log, verbose: bool,
|
||||
allowed_levels: list = None) -> list:
|
||||
"""
|
||||
网络搜索流程
|
||||
|
||||
Args:
|
||||
query: 查询
|
||||
log_trace: 日志追踪列表
|
||||
emit_log: 日志发射函数
|
||||
verbose: 是否详细输出
|
||||
allowed_levels: 允许的安全级别
|
||||
|
||||
Returns:
|
||||
网络搜索结果列表
|
||||
"""
|
||||
if not self.enable_web_search or not HAS_SERPER:
|
||||
return []
|
||||
|
||||
if emit_log:
|
||||
emit_log("🌐 触发网络搜索...")
|
||||
|
||||
web_results = self._web_search(query, top_k=5)
|
||||
|
||||
if not web_results:
|
||||
if emit_log:
|
||||
emit_log("⚠️ 网络搜索未返回结果")
|
||||
return []
|
||||
|
||||
# 转换为统一上下文格式
|
||||
web_contexts = []
|
||||
for item in web_results:
|
||||
web_contexts.append({
|
||||
'doc': f"{item.get('title', '')}\n{item.get('snippet', '')}",
|
||||
'meta': {
|
||||
'source': self.SOURCE_WEB,
|
||||
'link': item.get('link', ''),
|
||||
'date': item.get('date', '')
|
||||
},
|
||||
'source_type': self.SOURCE_WEB,
|
||||
'query': query
|
||||
})
|
||||
|
||||
log_trace.append({
|
||||
'phase': 'web_search',
|
||||
'query': query,
|
||||
'results_count': len(web_contexts)
|
||||
})
|
||||
|
||||
if emit_log:
|
||||
emit_log(f"✅ 网络搜索返回 {len(web_contexts)} 条结果")
|
||||
|
||||
return web_contexts
|
||||
|
||||
def _is_kb_result_sufficient(self, query: str, docs: list) -> bool:
|
||||
"""判断知识库检索结果是否充分"""
|
||||
if not docs:
|
||||
return False
|
||||
|
||||
# 结果数量检查
|
||||
if len(docs) >= 3:
|
||||
# 至少3条结果,检查相关性
|
||||
high_rel_count = 0
|
||||
for doc in docs:
|
||||
score = doc.get('score', 0) or doc.get('distance', 1)
|
||||
# cosine 距离转相似度
|
||||
if isinstance(score, (int, float)):
|
||||
sim = 1 - score if score <= 1 else score
|
||||
if sim >= 0.6:
|
||||
high_rel_count += 1
|
||||
|
||||
if high_rel_count >= 2:
|
||||
return True
|
||||
|
||||
# 有高质量结果
|
||||
for doc in docs[:2]:
|
||||
score = doc.get('score', 0) or doc.get('distance', 1)
|
||||
if isinstance(score, (int, float)):
|
||||
sim = 1 - score if score <= 1 else score
|
||||
if sim >= 0.8:
|
||||
return True
|
||||
|
||||
return False
|
||||
@@ -189,8 +189,6 @@ class RAGCacheManager:
|
||||
# 失效旧版本缓存
|
||||
self.query_cache.invalidate_by_version(old_version)
|
||||
self.embedding_cache.invalidate_by_version(old_version)
|
||||
# Rerank cache 无 kb_version 字段,文档变更时全量清空以防过时分数
|
||||
self.rerank_cache.clear()
|
||||
|
||||
logger.info(f"知识库 {kb_name} 版本更新: {old_version} -> {new_version}")
|
||||
return new_version
|
||||
@@ -198,41 +196,87 @@ class RAGCacheManager:
|
||||
# ==================== Query Cache 方法 ====================
|
||||
|
||||
@staticmethod
|
||||
def _make_query_cache_key(query: str, kb_name: str, kb_version: int) -> str:
|
||||
def _make_query_cache_key(query: str, kb_name: str, kb_version: int, doc_hash: str = "") -> str:
|
||||
"""
|
||||
生成查询缓存 key(基于 kb_version 的粗粒度失效)
|
||||
"""
|
||||
return hashlib.md5(
|
||||
f"query:{query}:{kb_name}:{kb_version}".encode()
|
||||
).hexdigest()
|
||||
|
||||
def get_query_result(self, query: str, kb_name: str) -> Optional[Dict]:
|
||||
"""
|
||||
获取查询缓存结果
|
||||
生成查询缓存 key
|
||||
|
||||
Args:
|
||||
query: 查询文本
|
||||
kb_name: 知识库名称
|
||||
kb_version: 知识库版本号
|
||||
doc_hash: 相关文档版本哈希(细粒度失效)
|
||||
|
||||
Returns:
|
||||
缓存 key
|
||||
"""
|
||||
if doc_hash:
|
||||
# 细粒度:只失效相关文档的缓存
|
||||
return hashlib.md5(
|
||||
f"query:{query}:{kb_name}:{doc_hash}".encode()
|
||||
).hexdigest()
|
||||
else:
|
||||
# 粗粒度:整个知识库版本变化时失效
|
||||
return hashlib.md5(
|
||||
f"query:{query}:{kb_name}:{kb_version}".encode()
|
||||
).hexdigest()
|
||||
|
||||
def get_query_result(self, query: str, kb_name: str, doc_ids: List[str] = None) -> Optional[Dict]:
|
||||
"""
|
||||
获取查询缓存结果
|
||||
|
||||
始终使用粗粒度 key(基于 kb_version),确保 GET/SET key 一致。
|
||||
|
||||
Args:
|
||||
query: 查询文本
|
||||
kb_name: 知识库名称
|
||||
doc_ids: 保留参数以兼容调用方签名(当前未使用)
|
||||
"""
|
||||
kb_version = self.get_kb_version(kb_name)
|
||||
|
||||
# 使用粗粒度 key,与 SET 保持一致
|
||||
key = self._make_query_cache_key(query, kb_name, kb_version)
|
||||
return self.query_cache.get(key)
|
||||
|
||||
def set_query_result(self, query: str, kb_name: str, result: Dict) -> None:
|
||||
def set_query_result(self, query: str, kb_name: str, result: Dict, doc_ids: List[str] = None) -> None:
|
||||
"""
|
||||
设置查询缓存结果
|
||||
|
||||
始终使用粗粒度 key(基于 kb_version),确保 GET/SET key 一致。
|
||||
kb_version 在文档变更时自增,触发整个知识库的缓存失效。
|
||||
|
||||
Args:
|
||||
query: 查询文本
|
||||
kb_name: 知识库名称
|
||||
result: 缓存结果
|
||||
doc_ids: 保留参数以兼容调用方签名(当前未使用)
|
||||
"""
|
||||
kb_version = self.get_kb_version(kb_name)
|
||||
|
||||
# 使用与 GET 相同的粗粒度 key,确保缓存可命中
|
||||
key = self._make_query_cache_key(query, kb_name, kb_version)
|
||||
self.query_cache.set(key, result, kb_version=kb_version)
|
||||
|
||||
def _compute_doc_hash(self, kb_name: str, doc_ids: List[str]) -> str:
|
||||
"""
|
||||
计算文档版本哈希
|
||||
|
||||
用于细粒度缓存失效:只失效相关文档变化时的缓存
|
||||
"""
|
||||
if not doc_ids:
|
||||
return ""
|
||||
|
||||
# 从文档 ID 中提取 source(文件名)
|
||||
sources = set()
|
||||
for doc_id in doc_ids:
|
||||
# doc_id 格式通常为 "filename_text_0" 或类似
|
||||
parts = doc_id.split('_')
|
||||
if parts:
|
||||
sources.add(parts[0])
|
||||
|
||||
# 生成哈希
|
||||
sources_str = ','.join(sorted(sources))
|
||||
return hashlib.md5(f"docs:{sources_str}".encode()).hexdigest()
|
||||
|
||||
# ==================== Embedding Cache 方法 ====================
|
||||
|
||||
@staticmethod
|
||||
|
||||
332
core/engine.py
332
core/engine.py
@@ -77,10 +77,6 @@ try:
|
||||
CONTEXT_EXPANSION_MAX_CHUNKS, ENUM_QUERY_DISABLE_TOPK_SHRINK, ENUM_QUERY_MMR_LAMBDA,
|
||||
# Phase 3 扩展精细化
|
||||
EXPANSION_SCORE_THRESHOLD, MAX_EXPANDED_NEIGHBORS,
|
||||
# 章节聚类救援
|
||||
SECTION_CLUSTER_BOOST_ENABLED, CLUSTER_MIN_MEMBERS, CLUSTER_MIN_TYPES,
|
||||
CLUSTER_SEED_FLOOR, CLUSTER_MAX_BOOST_PER_SECTION, CLUSTER_MAX_SECTIONS,
|
||||
CLUSTER_SECTION_PREFIX_LEVELS,
|
||||
# 上下文与生成
|
||||
LLM_TEMPERATURE, LLM_MAX_TOKENS, RECALL_MULTIPLIER,
|
||||
# FAQ 与黑名单
|
||||
@@ -106,28 +102,20 @@ except ImportError:
|
||||
MMR_TOP_K = 30
|
||||
CONTEXT_EXPANSION_ENABLED = True
|
||||
CONTEXT_EXPANSION_BEFORE = 1
|
||||
CONTEXT_EXPANSION_AFTER = 8
|
||||
CONTEXT_EXPANSION_AFTER = 5
|
||||
CONTEXT_EXPANSION_MAX_CHUNKS = 24
|
||||
EXPANSION_SCORE_THRESHOLD = 0.3
|
||||
MAX_EXPANDED_NEIGHBORS = 8
|
||||
# 章节聚类救援默认值
|
||||
SECTION_CLUSTER_BOOST_ENABLED = True
|
||||
CLUSTER_MIN_MEMBERS = 3
|
||||
CLUSTER_MIN_TYPES = 2
|
||||
CLUSTER_SEED_FLOOR = 0.35
|
||||
CLUSTER_MAX_BOOST_PER_SECTION = 8
|
||||
CLUSTER_MAX_SECTIONS = 3
|
||||
CLUSTER_SECTION_PREFIX_LEVELS = 1
|
||||
MAX_EXPANDED_NEIGHBORS = 4
|
||||
ENUM_QUERY_DISABLE_TOPK_SHRINK = True
|
||||
ENUM_QUERY_MMR_LAMBDA = 0.85
|
||||
DYNAMIC_RRF_ENABLED = True
|
||||
EMBEDDING_DEVICE = "auto"
|
||||
RERANK_DEVICE = "auto"
|
||||
RERANK_USE_ONNX = False
|
||||
RERANK_BACKEND = "cloud"
|
||||
RERANK_BACKEND = "local"
|
||||
RERANK_CLOUD_MODEL = "xop3qwen8breranker"
|
||||
RERANK_CLOUD_API_KEY = ""
|
||||
RERANK_CLOUD_BASE_URL = "https://maas-api.cn-huabei-1.xf-yun.com/v2/rerank"
|
||||
RERANK_CLOUD_BASE_URL = "https://dashscope.aliyuncs.com/compatible-api/v1/reranks"
|
||||
RERANK_CLOUD_TIMEOUT = 15
|
||||
|
||||
|
||||
@@ -559,7 +547,8 @@ class RAGEngine:
|
||||
top_dist = result['distances'][0][0] if result.get('distances') and result['distances'][0] else 1.0
|
||||
top_score = 1.0 - top_dist
|
||||
if top_score >= CACHE_MIN_SCORE:
|
||||
cache.set_query_result(query, kb_name, result)
|
||||
doc_ids = result.get('ids', [[]])[0] if result.get('ids') else []
|
||||
cache.set_query_result(query, kb_name, result, doc_ids=doc_ids)
|
||||
result['_debug'] = _debug
|
||||
return result
|
||||
|
||||
@@ -580,7 +569,8 @@ class RAGEngine:
|
||||
top_dist = result['distances'][0][0] if result.get('distances') and result['distances'][0] else 1.0
|
||||
top_score = 1.0 - top_dist
|
||||
if top_score >= CACHE_MIN_SCORE:
|
||||
cache.set_query_result(query, kb_name, result)
|
||||
doc_ids = result.get('ids', [[]])[0] if result.get('ids') else []
|
||||
cache.set_query_result(query, kb_name, result, doc_ids=doc_ids)
|
||||
result['_debug'] = _debug
|
||||
return result
|
||||
except Exception as e:
|
||||
@@ -615,7 +605,9 @@ class RAGEngine:
|
||||
top_dist = result['distances'][0][0] if result.get('distances') and result['distances'][0] else 1.0
|
||||
top_score = 1.0 - top_dist
|
||||
if top_score >= CACHE_MIN_SCORE:
|
||||
cache.set_query_result(query, kb_name, result)
|
||||
# 传递 doc_ids 实现细粒度缓存失效
|
||||
doc_ids = result.get('ids', [[]])[0] if result.get('ids') else []
|
||||
cache.set_query_result(query, kb_name, result, doc_ids=doc_ids)
|
||||
_debug['timing']['total_ms'] = int((time.time() - _overall_start) * 1000)
|
||||
result['_debug'] = _debug
|
||||
return result
|
||||
@@ -665,7 +657,6 @@ class RAGEngine:
|
||||
|
||||
results_list = [vector_results]
|
||||
weights = [VECTOR_WEIGHT]
|
||||
bm25_results = None # 初始化,防止 NameError
|
||||
|
||||
if USE_HYBRID_SEARCH and self.bm25_index.bm25:
|
||||
bm25_results = self.bm25_index.search(query, top_k=recall_k)
|
||||
@@ -680,22 +671,6 @@ class RAGEngine:
|
||||
vector_w, bm25_w = self._get_dynamic_rrf_weights(query)
|
||||
weights = [vector_w, bm25_w]
|
||||
|
||||
# ========== 保留 BM25 原始 top-3 完整信息,用于下游分歧检测救援 ==========
|
||||
_bm25_raw_top3 = []
|
||||
if USE_HYBRID_SEARCH and bm25_results and bm25_results.get('ids') and bm25_results['ids'][0]:
|
||||
_bm25_ids = bm25_results['ids'][0][:3]
|
||||
_bm25_docs = bm25_results['documents'][0][:3]
|
||||
_bm25_metas = bm25_results['metadatas'][0][:3]
|
||||
_bm25_dists = (bm25_results.get('distances', [[]])[0] or [0]*3)[:3]
|
||||
for i in range(len(_bm25_ids)):
|
||||
_bm25_raw_top3.append({
|
||||
'id': _bm25_ids[i],
|
||||
'doc': _bm25_docs[i],
|
||||
'meta': _bm25_metas[i],
|
||||
'bm25_score': _bm25_dists[i],
|
||||
'rank': i + 1
|
||||
})
|
||||
|
||||
if len(results_list) > 1:
|
||||
fused_results = self.reciprocal_rank_fusion(results_list, weights)
|
||||
_debug['steps'].append({'name': 'rrf_fusion', 'count': len(fused_results['ids'][0]) if fused_results.get('ids') else 0, 'weights': [round(w, 2) for w in weights]})
|
||||
@@ -711,8 +686,6 @@ class RAGEngine:
|
||||
|
||||
is_enum_query = self._is_enumeration_query(query)
|
||||
fused_results['_enum_query'] = is_enum_query
|
||||
# 传递 BM25 原始 top-3 到路由层,用于分歧检测救援
|
||||
fused_results['_bm25_top3'] = _bm25_raw_top3
|
||||
|
||||
# 章节过滤(如果查询中提到了章节)
|
||||
fused_results = self._filter_by_section(fused_results, query)
|
||||
@@ -759,19 +732,11 @@ class RAGEngine:
|
||||
# 时间衰减(Time Decay)
|
||||
fused_results = self._apply_time_decay(fused_results)
|
||||
|
||||
# 提前附加 _debug,使聚类提升能写入调试步骤
|
||||
fused_results['_debug'] = _debug
|
||||
|
||||
# ========== 章节聚类提升:在扩展前将低分但聚类的切片提升至种子阈值 ==========
|
||||
if SECTION_CLUSTER_BOOST_ENABLED:
|
||||
fused_results = self._section_cluster_boost(fused_results, query)
|
||||
|
||||
# ========== 上下文扩展:补充强命中切片周围的连续文本(rerank 之后,防止被截断)==========
|
||||
# Phase 3:仅对高分种子扩展邻居
|
||||
before_exp = len(fused_results['ids'][0]) if fused_results.get('ids') else 0
|
||||
fused_results = self._expand_contiguous_chunks(fused_results, top_k=top_k,
|
||||
min_score=EXPANSION_SCORE_THRESHOLD,
|
||||
query=query)
|
||||
min_score=EXPANSION_SCORE_THRESHOLD)
|
||||
after_exp = len(fused_results['ids'][0]) if fused_results.get('ids') else 0
|
||||
_debug['steps'].append({'name': 'context_expansion', 'before': before_exp, 'after': after_exp})
|
||||
|
||||
@@ -783,15 +748,7 @@ class RAGEngine:
|
||||
and not (is_enum_query and ENUM_QUERY_DISABLE_TOPK_SHRINK)
|
||||
and fused_results.get('_score_source') != 'rrf'
|
||||
):
|
||||
# 根据分数来源计算相似度分数(越高越好)
|
||||
score_source = fused_results.get('_score_source')
|
||||
top_dist = fused_results['distances'][0][0]
|
||||
if score_source == 'rerank':
|
||||
# Rerank 后 distances 是相关性分数,越大越好,直接使用
|
||||
top_score = top_dist
|
||||
else:
|
||||
# 向量距离,越小越好,转为相似度
|
||||
top_score = 1.0 - top_dist
|
||||
top_score = 1.0 - fused_results['distances'][0][0] # 距离转相似度
|
||||
adjusted_k, should_retrieve, reason = self._adaptive_topk.adjust(top_score, top_k)
|
||||
if "high_confidence" in reason:
|
||||
# 高置信度时截断结果
|
||||
@@ -806,7 +763,9 @@ class RAGEngine:
|
||||
top_dist = fused_results['distances'][0][0] if fused_results.get('distances') and fused_results['distances'][0] else 1.0
|
||||
top_score = 1.0 - top_dist # 距离转相似度
|
||||
if top_score >= CACHE_MIN_SCORE: # 置信度阈值
|
||||
cache.set_query_result(query, kb_name, fused_results)
|
||||
# 传递 doc_ids 实现细粒度缓存失效
|
||||
doc_ids = fused_results.get('ids', [[]])[0] if fused_results.get('ids') else []
|
||||
cache.set_query_result(query, kb_name, fused_results, doc_ids=doc_ids)
|
||||
|
||||
fused_results['_debug'] = _debug
|
||||
_debug['timing']['total_ms'] = int((time.time() - _overall_start) * 1000)
|
||||
@@ -1128,7 +1087,7 @@ class RAGEngine:
|
||||
'metadatas': [f_metas],
|
||||
'distances': [f_scores]
|
||||
}
|
||||
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context', '_bm25_top3'):
|
||||
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context'):
|
||||
if key in results:
|
||||
filtered[key] = results[key]
|
||||
return filtered
|
||||
@@ -1151,7 +1110,7 @@ class RAGEngine:
|
||||
'metadatas': [f_metas],
|
||||
'distances': [f_scores]
|
||||
}
|
||||
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context', '_bm25_top3'):
|
||||
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context'):
|
||||
if key in results:
|
||||
filtered[key] = results[key]
|
||||
return filtered
|
||||
@@ -1166,7 +1125,7 @@ class RAGEngine:
|
||||
'metadatas': [results['metadatas'][0][:top_k]],
|
||||
'distances': [results['distances'][0][:top_k]]
|
||||
}
|
||||
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context', '_bm25_top3'):
|
||||
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context'):
|
||||
if key in results:
|
||||
truncated[key] = results[key]
|
||||
return truncated
|
||||
@@ -1201,139 +1160,14 @@ class RAGEngine:
|
||||
return None
|
||||
return self.collection
|
||||
|
||||
@staticmethod
|
||||
def _normalize_section_path(section_path: str, levels: int = None) -> str:
|
||||
"""归一化 section_path:取前 N 级路径,容忍 MinerU 标题检测误差。
|
||||
|
||||
例如: "第三章 吸烟场所的功能设置 > 第三条 文明吸烟..." → "第三章 吸烟场所的功能设置"
|
||||
"""
|
||||
if not section_path:
|
||||
return ''
|
||||
if levels is None:
|
||||
levels = CLUSTER_SECTION_PREFIX_LEVELS
|
||||
parts = [p.strip() for p in section_path.split('>')]
|
||||
return ' > '.join(parts[:levels])
|
||||
|
||||
def _section_cluster_boost(self, results: dict, query: str = '') -> dict:
|
||||
"""章节聚类提升:当同一 section 下多个切片(text+table)同时出现在候选集中,
|
||||
即使单个切片 CrossEncoder 分数很低,也将整组提升到种子阈值。
|
||||
|
||||
核心洞察:单个低分切片不可信,但同一 section 多个切片同时出现是强信号。
|
||||
提升后的切片可以作为 _expand_contiguous_chunks 的种子,触发邻居扩展。
|
||||
|
||||
Args:
|
||||
results: rerank 后的检索结果
|
||||
query: 用户查询(用于后续扩展)
|
||||
|
||||
Returns:
|
||||
修改后的 results(distances 被调整,meta 中标记 _cluster_boosted)
|
||||
"""
|
||||
if not results.get('ids') or not results['ids'][0]:
|
||||
return results
|
||||
|
||||
ids = results['ids'][0]
|
||||
metas = results.get('metadatas', [[]])[0]
|
||||
distances = results.get('distances', [[]])[0] if results.get('distances') else None
|
||||
|
||||
if not distances:
|
||||
return results
|
||||
|
||||
# 1. 按 (source, normalized_section) 分组
|
||||
from collections import defaultdict
|
||||
section_groups = defaultdict(list) # key → [(index, meta, dist)]
|
||||
|
||||
for i, (meta, dist) in enumerate(zip(metas, distances)):
|
||||
source = meta.get('source', '')
|
||||
section_path = meta.get('section', '') or meta.get('section_path', '')
|
||||
norm_section = self._normalize_section_path(section_path)
|
||||
if not source or not norm_section:
|
||||
continue
|
||||
key = (source, norm_section)
|
||||
section_groups[key].append((i, meta, dist))
|
||||
|
||||
# 2. 检测聚类信号并提升
|
||||
boost_target_dist = 1.0 - CLUSTER_SEED_FLOOR # score=0.35 → dist=0.65
|
||||
boosted_sections = []
|
||||
total_boosted = 0
|
||||
|
||||
# 按组成员数降序排列,优先处理最大聚类
|
||||
sorted_groups = sorted(section_groups.items(), key=lambda x: len(x[1]), reverse=True)
|
||||
|
||||
for (source, norm_section), members in sorted_groups:
|
||||
if len(boosted_sections) >= CLUSTER_MAX_SECTIONS:
|
||||
break
|
||||
|
||||
# 聚类信号检测:成员数 >= 阈值 且 类型多样性 >= 阈值
|
||||
chunk_types = set(m[1].get('chunk_type', 'text') for m in members)
|
||||
if len(members) < CLUSTER_MIN_MEMBERS or len(chunk_types) < CLUSTER_MIN_TYPES:
|
||||
continue
|
||||
|
||||
# 提升组内切片分数(仅提升低于阈值的)
|
||||
boost_count = 0
|
||||
for idx, meta, dist in members:
|
||||
if boost_count >= CLUSTER_MAX_BOOST_PER_SECTION:
|
||||
break
|
||||
# 只提升分数低于种子阈值的切片(高分切片不需要)
|
||||
if dist > boost_target_dist:
|
||||
distances[idx] = boost_target_dist
|
||||
meta['_cluster_boosted'] = True
|
||||
boost_count += 1
|
||||
total_boosted += 1
|
||||
|
||||
if boost_count > 0:
|
||||
boosted_sections.append({
|
||||
'source': source,
|
||||
'section': norm_section,
|
||||
'members': len(members),
|
||||
'types': list(chunk_types),
|
||||
'boosted': boost_count
|
||||
})
|
||||
|
||||
# 3. 写 debug 信息
|
||||
if boosted_sections:
|
||||
debug_info = results.get('_debug', {})
|
||||
if 'steps' not in debug_info:
|
||||
debug_info['steps'] = []
|
||||
debug_info['steps'].append({
|
||||
'name': 'section_cluster_boost',
|
||||
'sections': boosted_sections,
|
||||
'total_boosted': total_boosted
|
||||
})
|
||||
results['_debug'] = debug_info
|
||||
logger.info(f"[章节聚类提升] 提升 {total_boosted} 个切片,"
|
||||
f"涉及 {len(boosted_sections)} 个 section: "
|
||||
f"{[s['section'][:30] for s in boosted_sections]}")
|
||||
|
||||
return results
|
||||
|
||||
@staticmethod
|
||||
def _chunk_lexical_score(chunk_text: str, query: str) -> float:
|
||||
"""计算切片文本与查询的词法重叠度(bigram 命中率),用于辅助种子资格判定。"""
|
||||
if not chunk_text or not query:
|
||||
return 0.0
|
||||
import re
|
||||
clean_q = re.sub(r'[??!!。,,、;;::"""\'\s*#`]+', ' ', query).strip()
|
||||
if len(clean_q) < 2:
|
||||
return 0.0
|
||||
bigrams = set()
|
||||
for i in range(len(clean_q) - 1):
|
||||
w = clean_q[i:i+2].strip()
|
||||
if len(w) == 2:
|
||||
bigrams.add(w)
|
||||
if not bigrams:
|
||||
return 0.0
|
||||
matched = sum(1 for w in bigrams if w in chunk_text)
|
||||
return matched / len(bigrams)
|
||||
|
||||
def _expand_contiguous_chunks(self, results: dict, top_k: int = None,
|
||||
min_score: float = 0.0, query: str = '') -> dict:
|
||||
min_score: float = 0.0) -> dict:
|
||||
"""Add same-source same-section neighbor text chunks around strong hits.
|
||||
|
||||
Args:
|
||||
results: 检索结果
|
||||
top_k: 最大切片数
|
||||
min_score: Phase 3 最低分数阈值,仅对 Rerank 分数高于此值的种子扩展
|
||||
query: 查询文本,用于词法匹配辅助种子资格判定
|
||||
"""
|
||||
if not CONTEXT_EXPANSION_ENABLED:
|
||||
return results
|
||||
@@ -1363,7 +1197,7 @@ class RAGEngine:
|
||||
seeds = [
|
||||
(doc_id, doc, meta, dist)
|
||||
for doc_id, doc, meta, dist in items[:base_limit]
|
||||
if (meta.get('chunk_type', 'text') == 'text' or meta.get('_cluster_boosted'))
|
||||
if meta.get('chunk_type', 'text') == 'text'
|
||||
and meta.get('source')
|
||||
and self._to_int(meta.get('chunk_index')) is not None
|
||||
]
|
||||
@@ -1374,12 +1208,8 @@ class RAGEngine:
|
||||
break
|
||||
|
||||
# Phase 3:跳过分数低于阈值的种子(仅当 min_score > 0 时生效)
|
||||
# 词法匹配豁免:CrossEncoder 低分但关键词重叠度高时仍允许作为种子
|
||||
if min_score > 0 and seed_dist < min_score:
|
||||
if query and self._chunk_lexical_score(_seed_doc, query) > 0.3:
|
||||
pass # 词法匹配度高,允许作为种子
|
||||
else:
|
||||
continue
|
||||
continue
|
||||
|
||||
source = seed_meta.get('source')
|
||||
section = seed_meta.get('section', '') or seed_meta.get('section_path', '')
|
||||
@@ -1481,7 +1311,7 @@ class RAGEngine:
|
||||
'distances': [[item[3] for item in items]],
|
||||
'_expanded_context': {'added': added}
|
||||
}
|
||||
for key in ('_debug', '_score_source', '_enum_query', '_bm25_top3'):
|
||||
for key in ('_debug', '_score_source', '_enum_query'):
|
||||
if key in results:
|
||||
expanded[key] = results[key]
|
||||
return expanded
|
||||
@@ -1508,7 +1338,6 @@ class RAGEngine:
|
||||
sub_top_k = max(top_k, 5)
|
||||
|
||||
all_results = []
|
||||
_all_bm25_top3 = [] # 收集各子查询的 BM25 top3
|
||||
for sub_q in sub_queries:
|
||||
try:
|
||||
sub_result = self.search_knowledge(
|
||||
@@ -1519,9 +1348,6 @@ class RAGEngine:
|
||||
)
|
||||
if sub_result and sub_result.get('ids') and sub_result['ids'][0]:
|
||||
all_results.append(sub_result)
|
||||
# 收集子查询的 BM25 top3
|
||||
if sub_result.get('_bm25_top3'):
|
||||
_all_bm25_top3.extend(sub_result['_bm25_top3'])
|
||||
except Exception as e:
|
||||
logger.warning(f"子查询检索失败: '{sub_q}' - {e}")
|
||||
|
||||
@@ -1530,19 +1356,9 @@ class RAGEngine:
|
||||
|
||||
# 合并去重
|
||||
if len(all_results) == 1:
|
||||
merged = all_results[0]
|
||||
else:
|
||||
merged = self._merge_and_deduplicate(all_results, top_k)
|
||||
return all_results[0]
|
||||
|
||||
# 将收集的 BM25 top3 传递到合并结果中
|
||||
if _all_bm25_top3:
|
||||
_all_bm25_top3.sort(key=lambda x: x.get('bm25_score', 0), reverse=True)
|
||||
_all_bm25_top3 = _all_bm25_top3[:3]
|
||||
for rank, item in enumerate(_all_bm25_top3):
|
||||
item['rank'] = rank + 1
|
||||
merged['_bm25_top3'] = _all_bm25_top3
|
||||
|
||||
return merged
|
||||
return self._merge_and_deduplicate(all_results, top_k)
|
||||
|
||||
def _search_with_decomposition(
|
||||
self, query, decomposer, top_k=5, allowed_levels=None,
|
||||
@@ -1574,7 +1390,6 @@ class RAGEngine:
|
||||
|
||||
# 并行检索各子查询
|
||||
all_results = []
|
||||
_all_bm25_top3 = [] # 收集各子查询的 BM25 top3
|
||||
for sub_q in sub_queries:
|
||||
try:
|
||||
sub_result = self.search_knowledge(
|
||||
@@ -1585,9 +1400,6 @@ class RAGEngine:
|
||||
)
|
||||
if sub_result and sub_result.get('ids') and sub_result['ids'][0]:
|
||||
all_results.append(sub_result)
|
||||
# 收集子查询的 BM25 top3
|
||||
if sub_result.get('_bm25_top3'):
|
||||
_all_bm25_top3.extend(sub_result['_bm25_top3'])
|
||||
except Exception as e:
|
||||
logger.warning(f"子查询检索失败: '{sub_q}' - {e}")
|
||||
|
||||
@@ -1600,14 +1412,6 @@ class RAGEngine:
|
||||
else:
|
||||
merged = self._merge_and_deduplicate(all_results, top_k)
|
||||
|
||||
# 将收集的 BM25 top3 传递到合并结果中
|
||||
if _all_bm25_top3:
|
||||
_all_bm25_top3.sort(key=lambda x: x.get('bm25_score', 0), reverse=True)
|
||||
_all_bm25_top3 = _all_bm25_top3[:3]
|
||||
for rank, item in enumerate(_all_bm25_top3):
|
||||
item['rank'] = rank + 1
|
||||
merged['_bm25_top3'] = _all_bm25_top3
|
||||
|
||||
return merged
|
||||
|
||||
def _merge_and_deduplicate(self, results_list, top_k):
|
||||
@@ -1692,13 +1496,12 @@ class RAGEngine:
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
|
||||
def _query_single_collection(coll_name):
|
||||
"""查询单个向量库(向量 + BM25),返回 (coll_results, bm25_raw_items)"""
|
||||
"""查询单个向量库(向量 + BM25)"""
|
||||
coll_results = []
|
||||
bm25_raw_items = [] # 该 collection 的 BM25 原始结果
|
||||
try:
|
||||
coll = self.kb_manager.get_collection(coll_name)
|
||||
if not coll:
|
||||
return coll_results, bm25_raw_items
|
||||
return coll_results
|
||||
|
||||
query_kwargs = {
|
||||
"query_embeddings": [query_vector],
|
||||
@@ -1716,61 +1519,25 @@ class RAGEngine:
|
||||
if USE_HYBRID_SEARCH:
|
||||
try:
|
||||
bm25 = self.kb_manager.get_bm25_index(coll_name)
|
||||
if bm25 and bm25.bm25:
|
||||
if bm25.bm25:
|
||||
bm25_res = bm25.search(query, top_k=recall_k)
|
||||
# 兼容两种 BM25Index:core.bm25_index 返回 dict,knowledge.base 返回 tuple
|
||||
if isinstance(bm25_res, tuple):
|
||||
_ids, _docs, _metas, _dists = bm25_res
|
||||
bm25_res = {
|
||||
'ids': [_ids],
|
||||
'documents': [_docs],
|
||||
'metadatas': [_metas],
|
||||
'distances': [_dists]
|
||||
}
|
||||
if source_filter and bm25_res['metadatas'][0]:
|
||||
if source_filter and bm25_res['metadatas'] and bm25_res['metadatas'][0]:
|
||||
bm25_res = self._filter_results(bm25_res, lambda meta: meta.get('source') == source_filter)
|
||||
if bm25_res['metadatas'] and bm25_res['metadatas'][0]:
|
||||
for meta in bm25_res['metadatas'][0]:
|
||||
meta['_collection'] = coll_name
|
||||
coll_results.append(bm25_res)
|
||||
# 提取 BM25 原始 top-3(在此处直接捕获,避免与向量结果混淆)
|
||||
_bm25_ids = bm25_res['ids'][0][:3]
|
||||
_bm25_docs = bm25_res['documents'][0][:3]
|
||||
_bm25_metas = bm25_res['metadatas'][0][:3]
|
||||
_bm25_dists = (bm25_res.get('distances', [[]])[0] or [0]*3)[:3]
|
||||
for i in range(len(_bm25_ids)):
|
||||
# 确保 meta 包含 _collection(用于路由层注入时下游处理)
|
||||
bm25_meta = _bm25_metas[i]
|
||||
if '_collection' not in bm25_meta:
|
||||
bm25_meta = {**bm25_meta, '_collection': coll_name}
|
||||
bm25_raw_items.append({
|
||||
'id': _bm25_ids[i],
|
||||
'doc': _bm25_docs[i],
|
||||
'meta': bm25_meta,
|
||||
'bm25_score': _bm25_dists[i],
|
||||
})
|
||||
logger.debug(f"[BM25] {coll_name}: captured {len(bm25_raw_items)} raw items")
|
||||
except Exception as e:
|
||||
logger.debug(f"向量库 {coll_name} BM25检索失败: {e}")
|
||||
logger.debug(f"向量库 {coll_name} 检索失败: {e}")
|
||||
except Exception as e:
|
||||
logger.debug(f"多向量库检索失败: {e}")
|
||||
return coll_results, bm25_raw_items
|
||||
return coll_results
|
||||
|
||||
all_results = []
|
||||
_bm25_raw_top3 = []
|
||||
with ThreadPoolExecutor(max_workers=len(target_collections)) as executor:
|
||||
futures = {executor.submit(_query_single_collection, name): name for name in target_collections}
|
||||
for future in as_completed(futures):
|
||||
coll_results, bm25_raw_items = future.result()
|
||||
all_results.extend(coll_results)
|
||||
_bm25_raw_top3.extend(bm25_raw_items)
|
||||
|
||||
# 按 bm25_score 降序取全局 top-3
|
||||
if _bm25_raw_top3:
|
||||
_bm25_raw_top3.sort(key=lambda x: x['bm25_score'], reverse=True)
|
||||
_bm25_raw_top3 = _bm25_raw_top3[:3]
|
||||
for rank, item in enumerate(_bm25_raw_top3):
|
||||
item['rank'] = rank + 1
|
||||
all_results.extend(future.result())
|
||||
|
||||
# ========== FAQ 检索 ==========
|
||||
faq_results = self._search_faq_collection(query_vector, top_k=FAQ_RECALL_TOP_K)
|
||||
@@ -1818,8 +1585,6 @@ class RAGEngine:
|
||||
|
||||
is_enum_query = self._is_enumeration_query(query)
|
||||
fused_results['_enum_query'] = is_enum_query
|
||||
# 传递 BM25 原始 top-3 到路由层,用于分歧检测救援
|
||||
fused_results['_bm25_top3'] = _bm25_raw_top3
|
||||
|
||||
# 章节过滤(如果查询中提到了章节)
|
||||
fused_results = self._filter_by_section(fused_results, query)
|
||||
@@ -1862,19 +1627,11 @@ class RAGEngine:
|
||||
# 时间衰减
|
||||
fused_results = self._apply_time_decay(fused_results)
|
||||
|
||||
# 提前附加 _debug,使聚类提升能写入调试步骤
|
||||
fused_results['_debug'] = _debug
|
||||
|
||||
# ========== 章节聚类提升:在扩展前将低分但聚类的切片提升至种子阈值 ==========
|
||||
if SECTION_CLUSTER_BOOST_ENABLED:
|
||||
fused_results = self._section_cluster_boost(fused_results, query)
|
||||
|
||||
# ========== 上下文扩展:补充强命中切片周围的连续文本(rerank 之后,防止被截断)==========
|
||||
# Phase 3:仅对高分种子扩展邻居
|
||||
before_exp = len(fused_results['ids'][0]) if fused_results.get('ids') else 0
|
||||
fused_results = self._expand_contiguous_chunks(fused_results, top_k=top_k,
|
||||
min_score=EXPANSION_SCORE_THRESHOLD,
|
||||
query=query)
|
||||
min_score=EXPANSION_SCORE_THRESHOLD)
|
||||
after_exp = len(fused_results['ids'][0]) if fused_results.get('ids') else 0
|
||||
if _debug is not None:
|
||||
_debug['steps'].append({'name': 'context_expansion', 'before': before_exp, 'after': after_exp})
|
||||
@@ -1887,15 +1644,7 @@ class RAGEngine:
|
||||
and not (is_enum_query and ENUM_QUERY_DISABLE_TOPK_SHRINK)
|
||||
and fused_results.get('_score_source') != 'rrf'
|
||||
):
|
||||
# 根据分数来源计算相似度分数(越高越好)
|
||||
score_source = fused_results.get('_score_source')
|
||||
top_dist = fused_results['distances'][0][0]
|
||||
if score_source == 'rerank':
|
||||
# Rerank 后 distances 是相关性分数,越大越好,直接使用
|
||||
top_score = top_dist
|
||||
else:
|
||||
# 向量距离,越小越好,转为相似度
|
||||
top_score = 1.0 - top_dist
|
||||
top_score = 1.0 - fused_results['distances'][0][0] # 距离转相似度
|
||||
adjusted_k, should_retrieve, reason = self._adaptive_topk.adjust(top_score, top_k)
|
||||
if "high_confidence" in reason:
|
||||
# 高置信度时截断结果
|
||||
@@ -1945,7 +1694,7 @@ class RAGEngine:
|
||||
'metadatas': [filtered_metas],
|
||||
'distances': [filtered_distances]
|
||||
}
|
||||
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context', '_bm25_top3'):
|
||||
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context'):
|
||||
if key in results:
|
||||
filtered[key] = results[key]
|
||||
return filtered
|
||||
@@ -2018,7 +1767,7 @@ class RAGEngine:
|
||||
'metadatas': [filtered_metas],
|
||||
'distances': [filtered_distances]
|
||||
}
|
||||
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context', '_bm25_top3'):
|
||||
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context'):
|
||||
if key in results:
|
||||
filtered[key] = results[key]
|
||||
return filtered
|
||||
@@ -2134,7 +1883,7 @@ class RAGEngine:
|
||||
'metadatas': [[c['metadata'] for c in selected]],
|
||||
'distances': [[id_to_dist.get(doc_id, 0) for doc_id in selected_ids]]
|
||||
}
|
||||
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context', '_bm25_top3'):
|
||||
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context'):
|
||||
if key in results:
|
||||
filtered[key] = results[key]
|
||||
return filtered
|
||||
@@ -2174,7 +1923,7 @@ class RAGEngine:
|
||||
'metadatas': [[c['metadata'] for c in selected]],
|
||||
'distances': [[id_to_dist.get(c['id'], 0) for c in selected]]
|
||||
}
|
||||
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context', '_bm25_top3'):
|
||||
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context'):
|
||||
if key in results:
|
||||
filtered[key] = results[key]
|
||||
return filtered
|
||||
@@ -2283,12 +2032,9 @@ class RAGEngine:
|
||||
'_rerank_cached': cache_hit
|
||||
}
|
||||
# 保留原有标记字段
|
||||
for key in ('_debug', '_enum_query', '_expanded_context', '_bm25_top3'):
|
||||
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context'):
|
||||
if key in results:
|
||||
reranked[key] = results[key]
|
||||
# Rerank 后 distances 语义变为 CrossEncoder 分数,更新 _score_source
|
||||
# 使自适应 TopK 能正确应用(之前 _score_source='rrf' 会导致自适应 TopK 被跳过)
|
||||
reranked['_score_source'] = 'rerank'
|
||||
return reranked
|
||||
|
||||
# ---------------- 流式生成 ----------------
|
||||
|
||||
@@ -233,10 +233,10 @@ class IntentAnalyzer:
|
||||
self._exact_cache_max = 500
|
||||
|
||||
def _get_client(self):
|
||||
"""获取 LLM 客户端(百炼快速模型)"""
|
||||
"""获取 LLM 客户端"""
|
||||
if self._client is None:
|
||||
from config import get_intent_client
|
||||
self._client = get_intent_client()
|
||||
from config import get_llm_client
|
||||
self._client = get_llm_client()
|
||||
return self._client
|
||||
|
||||
def _get_cache(self):
|
||||
@@ -308,30 +308,18 @@ class IntentAnalyzer:
|
||||
return self._exact_cache[exact_key]
|
||||
|
||||
# 2. 尝试从语义缓存获取
|
||||
# 关键:语义缓存只用原始 query 做 embedding(不含历史),
|
||||
# 避免同会话中不同问题因历史上下文污染导致误命中
|
||||
cache = self._get_cache()
|
||||
if cache:
|
||||
query_emb = self._get_embedding(query)
|
||||
# 使用 query + 历史关键信息作为缓存键
|
||||
cache_key = self._build_cache_key(query, history)
|
||||
cache_emb = self._get_embedding(cache_key)
|
||||
|
||||
if query_emb is not None:
|
||||
cached = cache.get(query_emb)
|
||||
# 确保缓存条目是意图分析结果(包含式校验,避免新增缓存类型时误命中)
|
||||
if cached and cached.get("cache_type") == "intent_analysis":
|
||||
# 二次验证:检查原始 query 文本相似度
|
||||
cached_query = cached.get("_raw_query", "")
|
||||
if cached_query and self._query_text_similar(query, cached_query):
|
||||
logger.info(f"意图分析缓存命中: {cached.get('reason', '')[:50]}")
|
||||
return IntentAnalysis.from_dict(cached)
|
||||
elif cached_query:
|
||||
logger.info(
|
||||
f"意图分析缓存二次验证拒绝: "
|
||||
f"query='{query[:30]}' vs cached='{cached_query[:30]}'"
|
||||
)
|
||||
else:
|
||||
# 旧缓存无 _raw_query 字段,兼容放行
|
||||
logger.info(f"意图分析缓存命中(无验证): {cached.get('reason', '')[:50]}")
|
||||
return IntentAnalysis.from_dict(cached)
|
||||
if cache_emb is not None:
|
||||
cached = cache.get(cache_emb)
|
||||
# 确保缓存条目是意图分析结果(非 RAG 回答缓存)
|
||||
if cached and cached.get("cache_type") != "rag_answer":
|
||||
logger.info(f"意图分析缓存命中: {cached.get('reason', '')[:50]}")
|
||||
return IntentAnalysis.from_dict(cached)
|
||||
else:
|
||||
logger.debug(f"意图分析缓存未命中,缓存状态: {cache.get_stats()}")
|
||||
else:
|
||||
@@ -400,12 +388,10 @@ class IntentAnalyzer:
|
||||
)
|
||||
|
||||
# 存入语义缓存(标记类型,避免与 RAG 回答缓存混淆)
|
||||
# 使用仅含 query 的 embedding,不含历史,防止同会话误命中
|
||||
if cache and query_emb is not None:
|
||||
if cache and cache_emb is not None:
|
||||
cache_data = analysis.to_dict()
|
||||
cache_data["cache_type"] = "intent_analysis"
|
||||
cache_data["_raw_query"] = query # 供二次验证使用
|
||||
cache.set(query_emb, cache_data)
|
||||
cache.set(cache_emb, cache_data)
|
||||
|
||||
# 存入精确匹配缓存
|
||||
if len(self._exact_cache) < self._exact_cache_max:
|
||||
@@ -445,32 +431,6 @@ class IntentAnalyzer:
|
||||
|
||||
return " | ".join(parts)
|
||||
|
||||
@staticmethod
|
||||
def _query_text_similar(query: str, cached_query: str, threshold: float = 0.5) -> bool:
|
||||
"""
|
||||
判断两个 query 文本是否足够相似(字符级 Jaccard)。
|
||||
用于语义缓存命中后的二次验证,防止语义相近但实际意图不同的问题误命中。
|
||||
|
||||
Args:
|
||||
query: 当前查询
|
||||
cached_query: 缓存中的原始查询
|
||||
threshold: 相似度阈值,默认 0.5
|
||||
|
||||
Returns:
|
||||
True 表示足够相似,可以命中缓存
|
||||
"""
|
||||
# 精确匹配快速路径
|
||||
if query.strip() == cached_query.strip():
|
||||
return True
|
||||
# 字符级 Jaccard 相似度
|
||||
set_a = set(query)
|
||||
set_b = set(cached_query)
|
||||
if not set_a or not set_b:
|
||||
return False
|
||||
intersection = len(set_a & set_b)
|
||||
union = len(set_a | set_b)
|
||||
return (intersection / union) >= threshold
|
||||
|
||||
def _build_history_summary(
|
||||
self,
|
||||
history: List[dict],
|
||||
|
||||
@@ -72,35 +72,16 @@ def call_llm(
|
||||
|
||||
content = response.choices[0].message.content
|
||||
|
||||
# 推理模型兼容(mimo-v2.5 等):
|
||||
# 推理模型思考链消耗大量 token(~1000),max_tokens 不足时 content 为空,
|
||||
# 全部输出进入 reasoning_content。此处从思考链中提取有效内容。
|
||||
# 推理模型兼容:content 为空时尝试从 reasoning_content 提取
|
||||
if not content or not content.strip():
|
||||
reasoning = getattr(response.choices[0].message, 'reasoning_content', None)
|
||||
if reasoning and reasoning.strip():
|
||||
# 先去掉 <think>...</think> 标签
|
||||
cleaned = re.sub(r'', '', reasoning, flags=re.DOTALL).strip()
|
||||
if cleaned:
|
||||
logger.info("LLM: content为空,从reasoning_content提取内容")
|
||||
# 尝试提取 JSON 对象(兼容结构化响应场景)
|
||||
json_match = re.search(r'\{[\s\S]*\}', cleaned)
|
||||
if json_match:
|
||||
try:
|
||||
json.loads(json_match.group())
|
||||
return json_match.group().strip()
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
pass
|
||||
# 尝试提取 JSON 数组
|
||||
bracket_match = re.search(r'\[[\s\S]*\]', cleaned)
|
||||
if bracket_match:
|
||||
try:
|
||||
json.loads(bracket_match.group())
|
||||
return bracket_match.group().strip()
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
pass
|
||||
# 纯文本响应:直接返回清理后的内容
|
||||
return cleaned
|
||||
logger.warning("LLM 返回空 content 且 reasoning_content 也无法提取(可能需要增大 max_tokens)")
|
||||
# 从思维链中提取 JSON 块作为内容
|
||||
json_match = re.search(r'\{[\s\S]*\}', reasoning)
|
||||
if json_match:
|
||||
logger.info("LLM: content为空,从reasoning_content提取JSON")
|
||||
return json_match.group().strip()
|
||||
logger.warning("LLM 返回空 content(可能需要增大 max_tokens)")
|
||||
return None
|
||||
|
||||
return content.strip()
|
||||
@@ -114,7 +95,7 @@ def call_llm_stream(
|
||||
prompt: str,
|
||||
model: str,
|
||||
temperature: float = 0.3,
|
||||
max_tokens: int = 3000,
|
||||
max_tokens: int = 1000,
|
||||
messages: List[dict] = None,
|
||||
error_prefix: str = "[错误]",
|
||||
**kwargs
|
||||
@@ -123,14 +104,13 @@ def call_llm_stream(
|
||||
流式 LLM 调用(生成器封装)
|
||||
|
||||
自动处理流式响应,逐块 yield 文本内容。
|
||||
兼容推理模型(mimo-v2.5 等):当 content 为空时回退到 reasoning_content。
|
||||
|
||||
Args:
|
||||
client: OpenAI 客户端实例
|
||||
prompt: 用户提示
|
||||
model: 模型名称
|
||||
temperature: 温度参数
|
||||
max_tokens: 最大 token 数(推理模型需留足思考链预算)
|
||||
max_tokens: 最大 token 数
|
||||
messages: 完整消息列表
|
||||
error_prefix: 错误时的前缀
|
||||
**kwargs: 其他参数
|
||||
@@ -155,33 +135,9 @@ def call_llm_stream(
|
||||
**kwargs
|
||||
)
|
||||
|
||||
content_yielded = False
|
||||
reasoning_buffer = []
|
||||
|
||||
for chunk in stream:
|
||||
if not chunk.choices:
|
||||
continue
|
||||
delta = chunk.choices[0].delta
|
||||
|
||||
# 正常 content 输出
|
||||
if hasattr(delta, 'content') and delta.content:
|
||||
content_yielded = True
|
||||
yield delta.content
|
||||
continue
|
||||
|
||||
# 推理模型:reasoning_content(思考链)
|
||||
rc = getattr(delta, 'reasoning_content', None)
|
||||
if rc:
|
||||
reasoning_buffer.append(rc)
|
||||
|
||||
# 回退:content 为空但 reasoning_content 有内容(推理模型 token 不足时)
|
||||
if not content_yielded and reasoning_buffer:
|
||||
reasoning_text = ''.join(reasoning_buffer)
|
||||
# 去掉 <think>...</think> 标签
|
||||
cleaned = re.sub(r'', '', reasoning_text, flags=re.DOTALL).strip()
|
||||
if cleaned:
|
||||
logger.info("流式 LLM: content为空,从reasoning_content提取内容")
|
||||
yield cleaned
|
||||
if chunk.choices and chunk.choices[0].delta.content:
|
||||
yield chunk.choices[0].delta.content
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"LLM 流式调用失败: {e}")
|
||||
|
||||
81
core/mmr.py
81
core/mmr.py
@@ -108,44 +108,22 @@ def mmr_rerank(
|
||||
return selected
|
||||
|
||||
|
||||
def _tokenize_words(text: str) -> set:
|
||||
"""
|
||||
使用 jieba 分词并过滤噪声,返回有意义的词集合。
|
||||
|
||||
过滤规则:
|
||||
- 去除单字符词(如 "的", "了", "在")—— 这些是停用词,对区分文档无意义
|
||||
- 去除纯数字 / 纯标点
|
||||
- 保留 2 字及以上的实词
|
||||
"""
|
||||
import jieba
|
||||
words = set()
|
||||
for w in jieba.cut(text):
|
||||
w = w.strip()
|
||||
if len(w) >= 2 and not w.isdigit():
|
||||
words.add(w)
|
||||
return words
|
||||
|
||||
|
||||
def mmr_filter_by_content(
|
||||
candidates: List[Dict],
|
||||
top_k: int = 30,
|
||||
similarity_threshold: float = 0.85
|
||||
similarity_threshold: float = 0.9
|
||||
) -> List[Dict]:
|
||||
"""
|
||||
基于 jieba 词级 Jaccard 相似度的去重(不需要 embedding)
|
||||
|
||||
与旧版字符级 set(text) 的区别:
|
||||
- 旧版:set("安全生产管理制度") → {'安','全','生','产',...},中文文档间字符集合高度重叠
|
||||
- 新版:jieba 分词 → {"安全生产", "管理制度", ...},词级集合区分度高
|
||||
基于内容相似度的去重(简化版,不需要 embedding)
|
||||
|
||||
适用于:
|
||||
- MMR_USE_EMBEDDING=False 时的快速去重
|
||||
- 避免 CPU 编码 100+ 文档的 50 秒开销
|
||||
- 没有 embedding 的情况
|
||||
- 快速去重场景
|
||||
|
||||
Args:
|
||||
candidates: 候选文档列表
|
||||
top_k: 返回数量
|
||||
similarity_threshold: 相似度阈值,超过则视为重复(默认 0.85)
|
||||
similarity_threshold: 相似度阈值,超过则视为重复
|
||||
|
||||
Returns:
|
||||
去重后的候选文档列表
|
||||
@@ -156,42 +134,35 @@ def mmr_filter_by_content(
|
||||
if len(candidates) <= top_k:
|
||||
return candidates
|
||||
|
||||
# 预分词:对所有候选文档一次性分词,避免重复调用 jieba.cut
|
||||
word_sets = []
|
||||
for c in candidates:
|
||||
content = c.get('content', c.get('document', ''))[:500]
|
||||
word_sets.append(_tokenize_words(content))
|
||||
selected = []
|
||||
remaining = candidates.copy()
|
||||
|
||||
selected_indices = []
|
||||
|
||||
for i in range(len(candidates)):
|
||||
if len(selected_indices) >= top_k:
|
||||
break
|
||||
|
||||
current_words = word_sets[i]
|
||||
if not current_words:
|
||||
# 空内容直接保留
|
||||
selected_indices.append(i)
|
||||
continue
|
||||
while len(selected) < top_k and remaining:
|
||||
current = remaining.pop(0)
|
||||
|
||||
# 检查是否与已选内容重复
|
||||
is_duplicate = False
|
||||
for j in selected_indices:
|
||||
selected_words = word_sets[j]
|
||||
if not selected_words:
|
||||
continue
|
||||
current_content = current.get('content', current.get('document', ''))[:200]
|
||||
|
||||
intersection = len(current_words & selected_words)
|
||||
union = len(current_words | selected_words)
|
||||
similarity = intersection / union if union > 0 else 0
|
||||
for s in selected:
|
||||
s_content = s.get('content', s.get('document', ''))[:200]
|
||||
|
||||
if similarity > similarity_threshold:
|
||||
is_duplicate = True
|
||||
break
|
||||
# 简单的 Jaccard 相似度
|
||||
words1 = set(current_content)
|
||||
words2 = set(s_content)
|
||||
if words1 and words2:
|
||||
intersection = len(words1 & words2)
|
||||
union = len(words1 | words2)
|
||||
similarity = intersection / union if union > 0 else 0
|
||||
|
||||
if similarity > similarity_threshold:
|
||||
is_duplicate = True
|
||||
break
|
||||
|
||||
if not is_duplicate:
|
||||
selected_indices.append(i)
|
||||
selected.append(current)
|
||||
|
||||
return [candidates[i] for i in selected_indices]
|
||||
return selected
|
||||
|
||||
|
||||
# ==================== 测试 ====================
|
||||
|
||||
@@ -1,115 +1,30 @@
|
||||
# Agentic RAG 完整指南
|
||||
# RAG 系统完整指南
|
||||
|
||||
> **版本**: v3.2(模型/Reranker/管线更新)
|
||||
> **生产入口**: `api/chat_routes.py::rag()` → `core/engine.py`(轻量编排,当前启用)
|
||||
> **备用编排**: `core/agentic.py::AgenticRAG.process()` + 8 个 Mixin(完整决策循环,未接线)
|
||||
> **最后更新**: 2026-06-04
|
||||
> **版本**: v4.0(统一编排 + 四层缓存修复)
|
||||
> **生产入口**: `api/chat_routes.py::rag()` → `generate()` → `core/engine.py`
|
||||
> **最后更新**: 2026-06-05
|
||||
>
|
||||
> ⚠️ 项目存在两套编排,生产 `/rag` 走的不是 `AgenticRAG`——详见下方「一·五、两套编排路径」。
|
||||
> 本次更新:删除未使用的 AgenticRAG 备用编排路径(10 个文件 ~2050 行),修复 Query Cache 键不匹配与阈值问题,将语义缓存集成至生产 `/rag` 端点。
|
||||
|
||||
## 一、功能概述
|
||||
|
||||
Agentic RAG 是一个智能问答系统,基于 Mixin 模式组合 8 个功能模块,具备以下核心能力:
|
||||
本系统是一个检索增强生成(RAG)问答系统,采用**单一统一编排路径**,由 `api/chat_routes.py` 的 `generate()` 函数直接编排全流程。核心能力包括:
|
||||
|
||||
| 功能 | 说明 | Mixin 模块 |
|
||||
|------|------|-----------|
|
||||
| **意图分析** | LLM 驱动的查询改写 + 双层判断(是否需要检索) | `IntentAnalyzer`(独立模块) |
|
||||
| **查询重写** | 口语化→专业术语、实体补全、指代消解 | `QueryRewriteMixin` |
|
||||
| **混合检索** | 向量检索 + BM25 + RRF 融合 + Rerank 重排 | `SearchMixin` → `RAGEngine` |
|
||||
| **多源融合** | 知识库 + 网络搜索,智能处理冲突 | `AnswerMixin` |
|
||||
| **幻觉验证** | 基于参考信息验证答案,防止 LLM 编造 | `AnswerMixin` |
|
||||
| **引用标注** | 自动标注信息来源和引用编号 | `CitationMixin` |
|
||||
| **富媒体提取** | 图片/表格的智能提取与展示 | `RichMediaMixin` |
|
||||
| **质量评估** | 多维度质量评估(相关性/完整性/准确性/覆盖面) | `QualityMixin` |
|
||||
| **上下文压缩** | Rerank 阈值过滤 + Token 预算控制 | `ContextMixin` |
|
||||
| **元问题处理** | 文件列表、权限查询等非知识类问题 | `MetaQuestionMixin` |
|
||||
| **置信度门控** | 基于 Reranker 分数判断检索质量,低分触发补救 | `ConfidenceGate`(独立模块) |
|
||||
|
||||
---
|
||||
|
||||
## ⚠️ 一·五、两套编排路径(务必先读)
|
||||
|
||||
> **关键认知**:本项目存在**两套并存的编排(orchestration)**。生产 HTTP 接口 `/rag` 走的是**轻量编排**,而 `AgenticRAG.process()` 那套**完整决策循环目前处于备用状态、未接入任何 HTTP 路由**。
|
||||
> 阅读下方所有架构图前请先理解这一点——下面 2.1 的「整体架构图」描绘的是**备用路径(AgenticRAG.process)**,不是当前生产实际跑的流程。
|
||||
|
||||
### 路径对比
|
||||
|
||||
| 维度 | 🟢 生产路径(当前启用) | 💤 备用路径(未接线) |
|
||||
|------|----------------------|---------------------|
|
||||
| 入口 | `api/chat_routes.py` → `rag()` → `generate()` | `core/agentic.py` → `AgenticRAG.process()` |
|
||||
| 编排者 | `chat_routes` 自己的流程代码 | `AgenticRAG` 类(8 个 Mixin 组合) |
|
||||
| 意图分析 | ✅ `intent_analyzer.analyze_intent()` | ✅ `IntentAnalyzer` / `QueryRewriteMixin` |
|
||||
| 检索 | ✅ `search_hybrid()` → `engine.search_knowledge()` | ✅ `engine.search_knowledge()` |
|
||||
| 查询分解/扩展/MMR/自适应TopK | ✅ 在 `engine` 内部执行 | ✅ 同左 |
|
||||
| 答案生成 | ✅ `engine.generate_answer_stream()`(流式) | `AnswerMixin._generate_fused_answer()` |
|
||||
| 引用标注 | ✅ `chat_routes._attach_citations()`(本地版) | `CitationMixin._attach_citations()` |
|
||||
| 置信度门控 | ❌ 不调用 | `ConfidenceGate`(仅此路径用) |
|
||||
| 多维质量评估 | ❌ 不调用 | `QualityMixin._assess_quality()` |
|
||||
| 推理反思 | ❌ 不调用 | `QualityMixin._reflect_on_answer()` |
|
||||
| 循环防护 | ❌ 不调用 | `LoopGuard`(仅此路径用) |
|
||||
| 幻觉验证 | ❌ 不调用 | `AnswerMixin._verify_and_refine_answer()` |
|
||||
|
||||
### 重要结论
|
||||
|
||||
- **Agentic 的核心能力是活跃的**:意图分析+LLM改写、子查询拆分、查询扩展、自适应 TopK、MMR 去重、混合检索+Rerank——这些都在 `/rag` 中**真实运行**,只是由 `chat_routes` + `engine` 直接调用,而非通过 `AgenticRAG` 类。
|
||||
- **休眠的只是「决策循环编排类」**:`AgenticRAG.process()` 及其独有组件(置信度门控 / 质量评估 / 推理反思 / 循环防护 / 幻觉验证)未接入 `/rag`。
|
||||
- import 证据:`confidence_gate.py`、`quality_assessor.py`、`reasoning_reflector.py`、`loop_guard.py` 以及 8 个 `agentic_*` Mixin **只被 `core/agentic.py` import**;而 `AgenticRAG` 实例虽在 `api/__init__.py:90` 启动时创建,但其唯一读取入口 `_get_agentic_rag()` **零调用**。
|
||||
- **这不是死代码可删**:`AgenticRAG` 在启动时被实例化(直接删会导致启动报错),且 `_extract_rich_media` 被 `scripts/test_rag_image_recall.py` 使用。它是「**一套更重、更完整、目前未启用的 Agentic 决策闭环**」,未来可选择接入。
|
||||
|
||||
### 🔬 如何验证「系统现在到底走哪套流程」
|
||||
|
||||
**方法 1:看开发环境 SSE 调试事件(最直接)**
|
||||
|
||||
`/rag` 在 `IS_DEV=True` 时会发出一串**只有 `chat_routes` 编排才会发**的调试事件,收到它们即证明走的是生产路径:
|
||||
|
||||
```bash
|
||||
# UTF-8 payload 避免 Windows shell 编码问题
|
||||
curl -s -N -X POST http://localhost:5001/rag \
|
||||
-H "Content-Type: application/json; charset=utf-8" \
|
||||
-H "Authorization: Bearer mock-token-admin" \
|
||||
--data-binary @payload.json
|
||||
```
|
||||
|
||||
观察 SSE 事件序列,**生产路径**会依次出现这些 `type`(`AgenticRAG.process` 不发这些):
|
||||
|
||||
| SSE 事件 `type` | 来源代码 | 含义 |
|
||||
|----------------|---------|------|
|
||||
| `start` | `chat_routes.py:1232` | 请求开始处理 |
|
||||
| `intent_result` | `chat_routes.py:1228` | 意图分析结果(来自 `intent_analyzer`)[DEV] |
|
||||
| `retrieval_debug` | `chat_routes.py:1311` | 检索管线各步骤(来自 `engine.search_knowledge` 的 `_debug`)[DEV] |
|
||||
| `chunks_retrieved` | `chat_routes.py:1416` | 召回切片详情 [DEV] |
|
||||
| `sources` | `chat_routes.py:1547` | 检索到的来源列表 |
|
||||
| `images_selected` | `chat_routes.py:1574` | 图片选择详情 [DEV] |
|
||||
| `context_built` | `chat_routes.py:1622` | 最终上下文构建 [DEV] |
|
||||
| `chunk` | `chat_routes.py:1630` | 流式答案的每个 token |
|
||||
| `finish` | `chat_routes.py:1699` | 含 `timing`、`sources`、`citations` |
|
||||
| `error` | `chat_routes.py:1733` | 处理异常时的错误信息 |
|
||||
|
||||
> 标注 [DEV] 的事件仅在 `IS_DEV=True` 时发送,其余事件在生产环境也会发送。
|
||||
|
||||
**方法 2:看服务端日志**
|
||||
|
||||
- 启动时:出现一次 `Agentic RAG 引擎已初始化`(`api/__init__.py:95`,仅实例化,不代表被调用)。
|
||||
- 每次 `/rag` 请求:出现 `[意图分析] use_context=... need_retrieval=...`(`chat_routes.py:1224`)。
|
||||
- **不会**出现任何来自 `AgenticRAG.process()` 内部的日志(如查询重写 `📝 查询重写`、`🔍 知识库检索: N 条结果`)——若出现则说明走了备用路径。
|
||||
|
||||
**方法 3:埋点验证(最确定)**
|
||||
|
||||
临时在 `core/agentic.py` 的 `AgenticRAG.process()` 第一行加 `logger.warning("AgenticRAG.process CALLED")`,重启后发 `/rag` 请求——**该日志不会触发**,即证明生产不走 `AgenticRAG`。
|
||||
|
||||
**方法 4:静态确认调用链**
|
||||
|
||||
```bash
|
||||
grep -rn "_get_agentic_rag()" --include="*.py" . # 仅定义,无调用者 → AgenticRAG 实例未被请求使用
|
||||
grep -rn "\.process(" --include="*.py" api/ # /rag、/chat 均无 .process() 调用
|
||||
```
|
||||
| 功能 | 说明 | 实现位置 |
|
||||
|------|------|----------|
|
||||
| **意图分析** | LLM 驱动的双层判断(是否需要检索)+ 查询改写 | `core/intent_analyzer.py` |
|
||||
| **混合检索** | 向量检索 + BM25 + RRF 融合 + Rerank 重排 | `core/engine.py` |
|
||||
| **四层缓存** | Query + Embedding + Rerank(LRU)+ 语义缓存(FAISS) | `core/cache.py` + `core/semantic_cache.py` |
|
||||
| **流式生成** | SSE 流式答案输出,逐 token 推送 | `core/engine.py::generate_answer_stream()` |
|
||||
| **引用标注** | 自动标注信息来源和引用编号 | `api/chat_routes.py::_attach_citations()` |
|
||||
| **富媒体** | 图片/表格的智能提取与展示 | `api/chat_routes.py` |
|
||||
| **查询理解** | 查询分解、扩展、MMR 去重、自适应 TopK | `core/` 各独立模块 |
|
||||
| **安全护栏** | 敏感信息过滤、Prompt 安全守卫 | `api/response_utils.py`、`core/prompt_guard.py` |
|
||||
|
||||
---
|
||||
|
||||
## 二、系统架构
|
||||
|
||||
> ⚠️ 注意:下方 2.1「整体架构图」描绘的是**备用路径 `AgenticRAG.process()`** 的完整设计;当前生产 `/rag` 的实际流程见上方「一·五」及本节 2.3「生产 /rag 实际流程」。
|
||||
|
||||
### 2.1 整体架构图
|
||||
|
||||
```
|
||||
@@ -118,6 +33,12 @@ grep -rn "\.process(" --include="*.py" api/ # /rag、/chat 均无 .proce
|
||||
└────────────────────────────┬────────────────────────────────────────┘
|
||||
↓
|
||||
┌─────────────────────────────────────────────────────────────────────┐
|
||||
│ 语义缓存检查 (SemanticCache - FAISS) │
|
||||
│ cosine ≥ 0.92 → 命中则直接返回缓存结果 │
|
||||
│ 跳过检索 + 生成全流程(~100ms vs ~9s) │
|
||||
└────────────────────────────┬────────────────────────────────────────┘
|
||||
↓ 未命中
|
||||
┌─────────────────────────────────────────────────────────────────────┐
|
||||
│ 意图分析 (IntentAnalyzer) │
|
||||
│ ┌──────────────┐ ┌──────────────┐ ┌──────────────┐ │
|
||||
│ │ 改写查询 │ │ 双层判断 │ │ 子查询拆分 │ │
|
||||
@@ -126,26 +47,24 @@ grep -rn "\.process(" --include="*.py" api/ # /rag、/chat 均无 .proce
|
||||
└─────────┼──────────────────┼──────────────────┼─────────────────────┘
|
||||
↓ ↓ ↓
|
||||
┌──────────┐ ┌──────────────────────────────────────────┐
|
||||
│ 直接回答 │ │ AgenticRAG.process() │
|
||||
│ (LLM) │ │ 1. 元问题检查 │
|
||||
└──────────┘ │ 2. 查询重写 (QueryRewriteMixin) │
|
||||
│ 3. 知识库检索 (RAGEngine.search_knowledge)│
|
||||
│ 4. 上下文压缩 (ContextMixin) │
|
||||
│ 5. 网络搜索 (SearchMixin, 可选) │
|
||||
│ 6. (图谱检索已废弃,graph/ 目录已清空) │
|
||||
│ 7. 融合答案生成 (AnswerMixin) │
|
||||
│ 8. 幻觉验证 (AnswerMixin) │
|
||||
│ 9. 富媒体提取 (RichMediaMixin) │
|
||||
│ 10. 引用标注 (CitationMixin) │
|
||||
│ 直接回答 │ │ 统一编排流程 │
|
||||
│ (LLM) │ │ 1. 混合检索 (engine.search_knowledge) │
|
||||
└──────────┘ │ 2. 上下文提取 + 来源去重 │
|
||||
│ 3. 图片补充检索 + 打分选择 │
|
||||
│ 4. 构建上下文 │
|
||||
│ 5. 流式答案生成 (engine.generate_stream) │
|
||||
│ 6. 答案图号对齐 + 引用标注 │
|
||||
│ 7. 敏感信息过滤 │
|
||||
│ 8. 语义缓存写入 + SSE finish │
|
||||
└────────────────────┬─────────────────────┘
|
||||
↓
|
||||
┌─────────────────────────────────────────────────────────────────────┐
|
||||
│ 检索层 (RAGEngine) │
|
||||
│ ┌──────────────┐ ┌──────────────┐ ┌──────────────┐ │
|
||||
│ │ 向量检索 │ │ BM25 检索 │ │ FAQ 独立召回 │ │
|
||||
│ │ (语义匹配) │ │ (关键词匹配) │ │ (精准命中) │ │
|
||||
│ └──────┬───────┘ └──────┬───────┘ └──────┬───────┘ │
|
||||
│ └─────────────────┼─────────────────┘ │
|
||||
│ │ 查询缓存 │ │ 向量检索 │ │ BM25 检索 │ │
|
||||
│ │ (LRU 500) │ │ (语义匹配) │ │ (关键词匹配) │ │
|
||||
│ │ 命中直接返回│ └──────┬───────┘ └──────┬───────┘ │
|
||||
│ └──────────────┘ └─────────────────┘ │
|
||||
│ ↓ │
|
||||
│ ┌──────────────┐ │
|
||||
│ │ RRF 融合 │ ← 动态权重(查询类型/长度驱动)│
|
||||
@@ -167,78 +86,133 @@ grep -rn "\.process(" --include="*.py" api/ # /rag、/chat 均无 .proce
|
||||
┌─────────────────────────────────────────────────────────────────────┐
|
||||
│ 答案生成 (LLM 流式) │
|
||||
│ ┌────────────────────────────────────────────────────────────────┐ │
|
||||
│ │ 整合多源信息 + 标注来源 + 处理冲突 + 引用编号 + SSE 流式输出 │ │
|
||||
│ │ 整合多源信息 + 标注来源 + 引用编号 + SSE 流式输出 │ │
|
||||
│ └────────────────────────────────────────────────────────────────┘ │
|
||||
└─────────────────────────────────────────────────────────────────────┘
|
||||
```
|
||||
|
||||
### 2.2 Mixin 组合架构
|
||||
### 2.2 编排方式
|
||||
|
||||
```python
|
||||
class AgenticRAG(
|
||||
QueryRewriteMixin, # 查询重写:口语化→专业术语、实体补全
|
||||
SearchMixin, # 检索功能:网络搜索
|
||||
AnswerMixin, # 答案生成:融合回答、幻觉验证
|
||||
CitationMixin, # 引用处理:来源标注、引用编号
|
||||
RichMediaMixin, # 富媒体:图片/表格提取
|
||||
QualityMixin, # 质量评估:多维评估
|
||||
ContextMixin, # 上下文处理:压缩、过滤
|
||||
MetaQuestionMixin # 元问题:文件列表、权限查询
|
||||
):
|
||||
...
|
||||
```
|
||||
系统采用**函数式编排**而非类组合模式。整个 RAG 流程由 `api/chat_routes.py` 的 `generate()` 函数直接控制,按步骤调用各独立模块:
|
||||
|
||||
> 注:以上 2.1 / 2.2 是 `AgenticRAG`(备用路径)的设计。**当前生产 `/rag` 不实例化走这条链**,实际流程见下方 2.3。
|
||||
- **意图分析**:`core/intent_analyzer.py` 的 `analyze_intent()`
|
||||
- **检索管线**:`core/engine.py` 的 `search_knowledge()`
|
||||
- **流式生成**:`core/engine.py` 的 `generate_answer_stream()`
|
||||
- **引用标注**:`api/chat_routes.py` 内的 `_attach_citations()`
|
||||
- **缓存系统**:`core/cache.py` 的 `RAGCacheManager` 单例 + `core/semantic_cache.py` 的 `SemanticCache` 单例
|
||||
|
||||
### 2.3 生产 /rag 实际流程(当前启用)
|
||||
|
||||
入口 `api/chat_routes.py::rag() → generate()`,**不经过 `AgenticRAG`**:
|
||||
|
||||
```
|
||||
POST /rag (SSE 流式)
|
||||
↓
|
||||
[chat_routes.generate()] ← 轻量编排,不实例化 AgenticRAG
|
||||
│
|
||||
├─ 发 SSE: start
|
||||
│
|
||||
├─ 1. 意图分析 intent_analyzer.analyze_intent() # chat_routes:1222
|
||||
│ ├─ need_retrieval=False → 直接 LLM 回答(流式发 SSE: chunk),结束
|
||||
│ │ └─ use_context=True 时带历史上下文,use_context=False 时纯闲聊
|
||||
│ └─ 否则继续;sub_queries 传入检索
|
||||
│ └─[DEV] 发 SSE: intent_result
|
||||
│
|
||||
├─ 2. 混合检索 search_hybrid() → engine.search_knowledge() # chat_routes:1300
|
||||
│ (内部:向量+BM25+RRF+废止过滤+章节过滤
|
||||
│ +云端Rerank+MMR去重+FAQ加权+黑名单+时间衰减
|
||||
│ +上下文扩展+自适应TopK)
|
||||
│ └─[DEV] 发 SSE: retrieval_debug
|
||||
│
|
||||
├─ 3. 提取上下文/来源(按 source 去重,doc_type 驱动溯源展示) # chat_routes:1362
|
||||
│ └─[DEV] 发 SSE: chunks_retrieved
|
||||
│ └─ 发 SSE: sources
|
||||
│
|
||||
├─ 4. 图片补充检索 + 图片打分选择 (select_images)
|
||||
│ └─[DEV] 发 SSE: images_selected
|
||||
├─ 5. 构建上下文 (_order_text_contexts_for_prompt)
|
||||
│ └─[DEV] 发 SSE: context_built
|
||||
│
|
||||
├─ 6. 流式答案生成 engine.generate_answer_stream() # chat_routes:1628
|
||||
│ └─ 逐 token 发 SSE: chunk
|
||||
│
|
||||
├─ 7. 答案图号对齐过滤
|
||||
├─ 8. 引用标注 chat_routes._attach_citations()(本地版,非 CitationMixin) # chat_routes:1668
|
||||
├─ 9. 敏感信息过滤 filter_response()
|
||||
├─ 10. 发 SSE: finish(answer + sources + citations + images + timing)
|
||||
└─[异常] 发 SSE: error
|
||||
```
|
||||
|
||||
**与备用路径(AgenticRAG.process)的差异**:生产路径**没有**置信度门控、多维质量评估、推理反思、循环防护、幻觉验证这几步——它们只存在于 `AgenticRAG.process()`。
|
||||
各模块通过 `get_engine()`、`get_cache_manager()`、`get_semantic_cache()` 等工厂函数获取全局单例实例。
|
||||
|
||||
---
|
||||
|
||||
## 三、意图分析流程
|
||||
## 三、四层缓存架构
|
||||
|
||||
### 3.1 IntentAnalyzer 双层判断
|
||||
### 3.1 缓存层次概览
|
||||
|
||||
| 层次 | 缓存类型 | 存储结构 | 容量 | TTL | 作用 |
|
||||
|------|----------|----------|------|-----|------|
|
||||
| L1 | Query Cache | LRU (OrderedDict) | 500 条 | 1 小时 | 缓存完整问答结果,命中后跳过整个检索+生成 |
|
||||
| L2 | Embedding Cache | LRU (OrderedDict) | 2000 条 | 24 小时 | 缓存向量化结果,避免重复调用 embedding 模型 |
|
||||
| L3 | Rerank Cache | LRU (OrderedDict) | 1000 条 | 1 小时 | 缓存 Rerank 分数,避免重复调用 Reranker |
|
||||
| L4 | Semantic Cache | FAISS IndexFlatIP | 10000 条 | 无过期 | 语义级缓存,相似查询也能命中 |
|
||||
|
||||
### 3.2 Query Cache
|
||||
|
||||
Query Cache 是最外层的完整问答结果缓存。命中后直接返回缓存的 `answer + sources + citations`,跳过检索和生成全流程。
|
||||
|
||||
**缓存键设计**:`{query_hash}:{kb_name}:{kb_version}`
|
||||
|
||||
- 基于查询文本哈希 + 知识库名称 + 知识库版本号
|
||||
- 知识库版本变更时(如文档更新/重新索引),相关缓存自动失效
|
||||
|
||||
**已修复的问题**:
|
||||
|
||||
1. **键不匹配问题(已修复)**:此前 `set_query_result()` 在有 `doc_ids` 参数时使用 `doc_hash` 分支生成键,而 `get_query_result()` 始终使用 `kb_version` 分支——导致 GET 和 SET 的键永远不匹配,命中率始终为 0%。修复后两端统一使用 `kb_version` 分支。
|
||||
|
||||
2. **写入阈值问题(已修复)**:`CACHE_MIN_SCORE` 原值为 `0.3`,但 ChromaDB 余弦距离经 `1 - dist` 计算后得分通常在 0.03-0.06 之间,远低于阈值,导致几乎不写入缓存。修复后设为 `0.0`。
|
||||
|
||||
### 3.3 Embedding Cache
|
||||
|
||||
缓存文本向量化结果,由 `RAGEngine` 在调用 embedding 模型前后自动读写。键为查询文本哈希,避免对相同文本重复调用 embedding 模型(如 DashScope text-embedding-v3)。
|
||||
|
||||
### 3.4 Rerank Cache
|
||||
|
||||
缓存 Rerank 重排序的分数结果。键为 `query + sorted(doc_ids)` 的精确匹配。注意:由于 RRF 融合产出差异,命中率可能偏低。
|
||||
|
||||
### 3.5 语义缓存(Semantic Cache)
|
||||
|
||||
语义缓存基于 FAISS 向量索引实现语义级匹配——即使查询文字不完全相同,只要语义足够相似(cosine similarity ≥ 0.92),就能命中缓存。
|
||||
|
||||
**工作机制**:
|
||||
|
||||
```
|
||||
用户查询 → embedding 编码 → FAISS 向量检索
|
||||
→ cosine ≥ 0.92 → 命中:返回缓存的 answer + sources + citations(~100ms)
|
||||
→ cosine < 0.92 → 未命中:执行完整 RAG 流程后写入缓存
|
||||
```
|
||||
|
||||
**集成位置**:
|
||||
|
||||
- **读取**:在 `generate()` 函数的意图分析之后、混合检索之前(跳过整个检索+生成流程)
|
||||
- **写入**:在 `generate()` 函数生成完整答案后、发送 `finish` 事件之前
|
||||
- **缓存内容**:answer、sources、citations、images、tables
|
||||
|
||||
**验证结果**:语义缓存命中率约 66.7%,平均响应从 ~9.2 秒降至 ~100 毫秒(约 92 倍加速)。
|
||||
|
||||
### 3.6 缓存失效机制
|
||||
|
||||
所有基于 LRU 的缓存(L1-L3)均支持基于知识库版本号(`kb_version`)的自动失效:
|
||||
|
||||
- 每个缓存条目关联 `kb_version`
|
||||
- 知识库文档变更时 `kb_version` 递增
|
||||
- 读取时检查 `kb_version` 是否匹配,不匹配则视为过期
|
||||
|
||||
语义缓存(L4)当前无 TTL 过期机制,仅受 `max_size=10000` 容量限制。
|
||||
|
||||
### 3.7 缓存配置
|
||||
|
||||
```python
|
||||
# config.py / config.example.py
|
||||
|
||||
# 查询结果缓存
|
||||
QUERY_CACHE_ENABLED = True
|
||||
QUERY_CACHE_SIZE = 500
|
||||
QUERY_CACHE_TTL = 3600 # 秒
|
||||
|
||||
# Embedding 缓存
|
||||
EMBEDDING_CACHE_ENABLED = True
|
||||
EMBEDDING_CACHE_SIZE = 2000
|
||||
EMBEDDING_CACHE_TTL = 86400 # 24小时
|
||||
|
||||
# Rerank 缓存
|
||||
RERANK_CACHE_ENABLED = True
|
||||
RERANK_CACHE_SIZE = 1000
|
||||
RERANK_CACHE_TTL = 3600
|
||||
|
||||
# 语义缓存
|
||||
SEMANTIC_CACHE_ENABLED = True
|
||||
SEMANTIC_CACHE_THRESHOLD = 0.92 # cosine 相似度阈值
|
||||
|
||||
# 缓存写入最低置信度
|
||||
CACHE_MIN_SCORE = 0.0 # ChromaDB 余弦距离经 1-dist 后得分约 0.03-0.06,须设为 0
|
||||
```
|
||||
|
||||
### 3.8 部署注意事项
|
||||
|
||||
当前所有缓存均为**进程内内存存储**(LRU 使用 `OrderedDict`,语义缓存使用 FAISS 内存索引),有以下部署影响:
|
||||
|
||||
- **单 Worker**:Gunicorn 默认 1 个 worker,所有请求共享同一缓存实例,缓存有效
|
||||
- **多 Worker**:每个 worker 有独立缓存,不共享,缓存效率降低
|
||||
- **冷启动**:进程重启后缓存全部丢失,需重新预热
|
||||
- **`max_requests=1000`**:Gunicorn worker 定期重启会导致缓存周期性清空
|
||||
|
||||
对于生产环境多实例部署场景,已规划 Redis 外部缓存迁移方案(见 `reports/redis_migration_plan.md`)。
|
||||
|
||||
---
|
||||
|
||||
## 四、意图分析流程
|
||||
|
||||
### 4.1 IntentAnalyzer 双层判断
|
||||
|
||||
意图分析由 `core/intent_analyzer.py` 的 `IntentAnalyzer` 类完成,采用 **LLM 驱动** 的双层判断:
|
||||
|
||||
@@ -269,7 +243,7 @@ POST /rag (SSE 流式)
|
||||
- intent: factual/comparison/reasoning/instruction/other
|
||||
```
|
||||
|
||||
### 3.2 QueryClassifier 规则分类
|
||||
### 4.2 QueryClassifier 规则分类
|
||||
|
||||
`core/query_classifier.py` 提供无 LLM 调用的快速规则分类:
|
||||
|
||||
@@ -286,9 +260,9 @@ POST /rag (SSE 流式)
|
||||
|
||||
---
|
||||
|
||||
## 四、检索管线详解
|
||||
## 五、检索管线详解
|
||||
|
||||
### 4.1 完整检索流程
|
||||
### 5.1 完整检索流程
|
||||
|
||||
```
|
||||
search_knowledge(query, top_k=30)
|
||||
@@ -335,7 +309,7 @@ search_knowledge(query, top_k=30)
|
||||
└─ 15. 缓存写入 → 返回结果
|
||||
```
|
||||
|
||||
### 4.2 混合检索代码示例
|
||||
### 5.2 混合检索代码示例
|
||||
|
||||
```python
|
||||
# 向量检索(语义相似)
|
||||
@@ -348,7 +322,10 @@ bm25_results = bm25_index.search(query, top_k=recall_k)
|
||||
faq_results = faq_collection.query(query_embeddings=[query_vector], n_results=3)
|
||||
|
||||
# 图片独立召回(P0 通道)
|
||||
image_results = collection.query(query_embeddings=[query_vector], n_results=5, where={"chunk_type": {"$in": ["image", "chart", "table"]}})
|
||||
image_results = collection.query(
|
||||
query_embeddings=[query_vector], n_results=5,
|
||||
where={"chunk_type": {"$in": ["image", "chart", "table"]}}
|
||||
)
|
||||
|
||||
# RRF 融合(动态权重)
|
||||
fused = reciprocal_rank_fusion([vector_results, bm25_results], weights=[vector_w, bm25_w])
|
||||
@@ -360,7 +337,7 @@ reranked = rerank_results(query, fused, top_k=15)
|
||||
mmr_results = mmr_rerank(query_emb, reranked, top_k=30, lambda_param=0.5)
|
||||
```
|
||||
|
||||
### 4.3 RRF 融合算法
|
||||
### 5.3 RRF 融合算法
|
||||
|
||||
```
|
||||
RRF分数 = Σ (权重 / (k + 排名位置))
|
||||
@@ -376,9 +353,9 @@ RRF分数 = Σ (权重 / (k + 排名位置))
|
||||
- 查询类型驱动: FACT→BM25优先, PROCESS→向量优先
|
||||
```
|
||||
|
||||
### 4.4 Rerank 重排
|
||||
### 5.4 Rerank 重排
|
||||
|
||||
**后端**: 支持三种模式,由 `RERANK_BACKEND` 环境变量控制
|
||||
**后端**:支持三种模式,由 `RERANK_BACKEND` 环境变量控制
|
||||
|
||||
| RERANK_BACKEND | 说明 |
|
||||
|----------------|------|
|
||||
@@ -386,20 +363,20 @@ RRF分数 = Σ (权重 / (k + 排名位置))
|
||||
| `"local"` | 仅使用本地 `BAAI/bge-reranker-base`(CrossEncoder / ONNX) |
|
||||
| `"fallback"` | 优先云端,失败时自动回退本地(推荐生产环境) |
|
||||
|
||||
**云端 Reranker(推荐)**:
|
||||
**云端 Reranker(推荐)**:
|
||||
|
||||
```python
|
||||
# config.py
|
||||
RERANK_BACKEND = os.getenv("RERANK_BACKEND", "local") # local / cloud / fallback
|
||||
RERANK_CLOUD_MODEL = "qwen3-rerank" # DashScope 云端 Rerank 模型
|
||||
RERANK_BACKEND = os.getenv("RERANK_BACKEND", "local")
|
||||
RERANK_CLOUD_MODEL = "qwen3-rerank"
|
||||
RERANK_CLOUD_API_KEY = os.getenv("RERANK_CLOUD_API_KEY", DASHSCOPE_API_KEY)
|
||||
RERANK_CLOUD_BASE_URL = "https://dashscope.aliyuncs.com/compatible-api/v1/reranks"
|
||||
RERANK_CLOUD_TIMEOUT = 15 # 云端请求超时(秒)
|
||||
RERANK_CLOUD_TIMEOUT = 15
|
||||
```
|
||||
|
||||
`CloudReranker` 类(`core/engine.py`)封装 DashScope 的 `/compatible-api/v1/reranks` 接口,提供与本地 `CrossEncoder.predict()` / `ONNXReranker.predict()` 一致的调用接口。
|
||||
|
||||
**本地 Reranker(备选)**:
|
||||
**本地 Reranker(备选)**:
|
||||
|
||||
```python
|
||||
def rerank_results(self, query, results, top_k=5):
|
||||
@@ -409,69 +386,73 @@ def rerank_results(self, query, results, top_k=5):
|
||||
# 返回 top_k 个最高分结果
|
||||
```
|
||||
|
||||
**调用位置**: `core/engine.py` 的 `search_knowledge()` 和 `_search_multi_kb()` 中,RRF 融合 + 废止/章节过滤之后、MMR 去重之前执行。
|
||||
|
||||
**引擎初始化顺序**: `RAGEngine.__init__()` 中按 `RERANK_BACKEND` 决定加载策略:
|
||||
- `cloud` / `fallback`:先尝试创建 `CloudReranker`,需要 `RERANK_CLOUD_API_KEY`
|
||||
- `local` / `fallback`(云端失败时):加载本地 `BAAI/bge-reranker-base`,支持 ONNX 加速
|
||||
**调用位置**:`core/engine.py` 的 `search_knowledge()` 和 `_search_multi_kb()` 中,RRF 融合 + 废止/章节过滤之后、MMR 去重之前执行。
|
||||
|
||||
---
|
||||
|
||||
## 五、置信度门控
|
||||
## 六、生产 /rag 完整流程
|
||||
|
||||
`core/confidence_gate.py` 基于 Reranker 分数判断检索结果质量:
|
||||
入口 `api/chat_routes.py::rag() → generate()`:
|
||||
|
||||
```
|
||||
检索结果 → Reranker 计算置信度 → 阈值判断 → 决策
|
||||
│
|
||||
┌─────────────────┼─────────────────┐
|
||||
↓ ↓ ↓
|
||||
PASS (≥0.4) REWRITE (0.2~0.4) WEB_SEARCH (<0.2)
|
||||
继续生成 查询重写 网络搜索补救
|
||||
POST /rag (SSE 流式)
|
||||
↓
|
||||
[chat_routes.generate()]
|
||||
│
|
||||
├─ 发 SSE: start
|
||||
│
|
||||
├─ 1. 语义缓存检查 SemanticCache.get() # chat_routes
|
||||
│ ├─ 命中 → 直接流式发 SSE: chunk + finish,结束(~100ms)
|
||||
│ └─ 未命中 → 继续;记录 embedding 供后续写入
|
||||
│
|
||||
├─ 2. 意图分析 intent_analyzer.analyze_intent()
|
||||
│ ├─ need_retrieval=False → 直接 LLM 回答(流式发 SSE: chunk),结束
|
||||
│ │ └─ use_context=True 时带历史上下文,use_context=False 时纯闲聊
|
||||
│ └─ 否则继续;sub_queries 传入检索
|
||||
│ └─[DEV] 发 SSE: intent_result
|
||||
│
|
||||
├─ 3. 混合检索 search_hybrid() → engine.search_knowledge()
|
||||
│ (内部:查询缓存检查 → 向量+BM25+RRF+废止过滤+章节过滤
|
||||
│ +云端Rerank+MMR去重+FAQ加权+黑名单+时间衰减
|
||||
│ +上下文扩展+自适应TopK)
|
||||
│ └─[DEV] 发 SSE: retrieval_debug
|
||||
│
|
||||
├─ 4. 提取上下文/来源(按 source 去重,doc_type 驱动溯源展示)
|
||||
│ └─[DEV] 发 SSE: chunks_retrieved
|
||||
│ └─ 发 SSE: sources
|
||||
│
|
||||
├─ 5. 图片补充检索 + 图片打分选择 (select_images)
|
||||
│ └─[DEV] 发 SSE: images_selected
|
||||
├─ 6. 构建上下文 (_order_texts_for_prompt)
|
||||
│ └─[DEV] 发 SSE: context_built
|
||||
│
|
||||
├─ 7. 流式答案生成 engine.generate_answer_stream()
|
||||
│ └─ 逐 token 发 SSE: chunk
|
||||
│
|
||||
├─ 8. 答案图号对齐过滤
|
||||
├─ 9. 引用标注 _attach_citations()
|
||||
├─ 10. 敏感信息过滤 filter_response()
|
||||
├─ 11. 语义缓存写入 SemanticCache.set() # 写入缓存供后续命中
|
||||
├─ 12. 发 SSE: finish(answer + sources + citations + images + timing)
|
||||
└─[异常] 发 SSE: error
|
||||
```
|
||||
|
||||
**阈值配置**:
|
||||
- `PASS_THRESHOLD = 0.2`: 通过阈值(低于此值需要补救)
|
||||
- `GOOD_THRESHOLD = 0.4`: 良好阈值(高质量结果)
|
||||
- `EXCELLENT_THRESHOLD = 0.7`: 优秀阈值
|
||||
### SSE 事件序列
|
||||
|
||||
---
|
||||
| SSE 事件 `type` | 含义 |
|
||||
|----------------|------|
|
||||
| `start` | 请求开始处理 |
|
||||
| `intent_result` | 意图分析结果 [DEV] |
|
||||
| `retrieval_debug` | 检索管线各步骤 [DEV] |
|
||||
| `chunks_retrieved` | 召回切片详情 [DEV] |
|
||||
| `sources` | 检索到的来源列表 |
|
||||
| `images_selected` | 图片选择详情 [DEV] |
|
||||
| `context_built` | 最终上下文构建 [DEV] |
|
||||
| `chunk` | 流式答案的每个 token |
|
||||
| `finish` | 含 `timing`、`sources`、`citations`、`images` |
|
||||
| `error` | 处理异常时的错误信息 |
|
||||
|
||||
## 六、AgenticRAG 主流程
|
||||
|
||||
### 6.1 process() 方法
|
||||
|
||||
```python
|
||||
def process(self, query, verbose=True, history=None,
|
||||
allowed_levels=None, role=None, department=None,
|
||||
emit_log=None) -> dict:
|
||||
"""
|
||||
返回:
|
||||
{
|
||||
"answer": str, # 最终答案
|
||||
"sources": list, # 来源列表
|
||||
"images": list, # 图片列表
|
||||
"tables": list, # 表格列表
|
||||
"citations": list, # 引用列表
|
||||
"log_trace": list # 推理过程追踪
|
||||
}
|
||||
"""
|
||||
```
|
||||
|
||||
### 6.2 流程步骤
|
||||
|
||||
| 步骤 | 方法 | 说明 |
|
||||
|------|------|------|
|
||||
| 1 | `_is_meta_question()` | 检查元问题(文件列表、权限等) |
|
||||
| 2 | `should_rewrite()` + `_rewrite_query()` | 查询重写(口语化→专业术语、实体补全) |
|
||||
| 3 | `engine.search_knowledge()` / `engine.search_multiple()` | 知识库检索(含向量+BM25+RRF+MMR+Rerank) |
|
||||
| 4 | `_compress_contexts()` | 上下文压缩(Rerank 阈值过滤) |
|
||||
| 5 | `_web_search_flow()` | 网络搜索(可选,需 SERPER_API_KEY) |
|
||||
| 6 | ~~`_graph_search()`~~ | ~~图谱检索(已废弃,graph/ 目录已清空)~~ |
|
||||
| 7 | `_generate_fused_answer()` | 融合答案生成(多源信息+冲突处理) |
|
||||
| 8 | `_verify_and_refine_answer()` | 幻觉验证(防止 LLM 编造) |
|
||||
| 9 | `_extract_rich_media()` | 富媒体提取(图片/表格) |
|
||||
| 10 | `_attach_citations()` | 引用标注 |
|
||||
> 标注 [DEV] 的事件仅在 `IS_DEV=True` 时发送。
|
||||
|
||||
---
|
||||
|
||||
@@ -489,7 +470,7 @@ curl -X POST http://localhost:5001/rag \
|
||||
}'
|
||||
```
|
||||
|
||||
**响应格式**: SSE(Server-Sent Events)流式返回
|
||||
**响应格式**:SSE(Server-Sent Events)流式返回
|
||||
|
||||
```
|
||||
event: token
|
||||
@@ -501,7 +482,7 @@ data: {"text": "规定"}
|
||||
...
|
||||
|
||||
event: finish
|
||||
data: {"answer": "完整答案", "sources": [...], "citations": [...], "images": [...], "duration_ms": 3200}
|
||||
data: {"answer": "完整答案", "sources": [...], "citations": [...], "images": [...], "duration_ms": 200}
|
||||
```
|
||||
|
||||
### 7.2 代码调用
|
||||
@@ -520,20 +501,27 @@ for token in engine.generate_answer_stream(query, context, history=history):
|
||||
print(token, end="", flush=True)
|
||||
```
|
||||
|
||||
### 7.3 AgenticRAG 调用
|
||||
### 7.3 缓存统计查询
|
||||
|
||||
```python
|
||||
from core.agentic import AgenticRAG
|
||||
|
||||
rag = AgenticRAG(max_iterations=3, enable_web_search=True)
|
||||
result = rag.process("出差补助标准是什么?")
|
||||
|
||||
print(f"答案: {result['answer']}")
|
||||
print(f"来源: {result['sources']}")
|
||||
print(f"图片: {result['images']}")
|
||||
print(f"引用: {result['citations']}")
|
||||
```bash
|
||||
# 查看各层缓存命中率和统计
|
||||
curl http://localhost:5001/cache/stats \
|
||||
-H "Authorization: Bearer mock-token-admin"
|
||||
```
|
||||
|
||||
返回示例:
|
||||
|
||||
```json
|
||||
{
|
||||
"query_cache": {"total_entries": 50, "hits": 10, "misses": 6, "hit_rate": 0.625},
|
||||
"embedding_cache": {"total_entries": 200, "hits": 0, "misses": 0, "hit_rate": 0},
|
||||
"rerank_cache": {"total_entries": 100, "hits": 0, "misses": 0, "hit_rate": 0},
|
||||
"semantic_cache": {"total_entries": 15, "hits": 6, "misses": 3, "hit_rate": 0.667}
|
||||
}
|
||||
```
|
||||
|
||||
> 注:当 Query Cache 或 Semantic Cache 在外层拦截了重复查询时,Embedding Cache 和 Rerank Cache 的命中率为 0 是正常现象——重复查询根本不会到达这些层。
|
||||
|
||||
---
|
||||
|
||||
## 八、配置说明
|
||||
@@ -565,10 +553,10 @@ RECALL_MULTIPLIER = 3 # 候选池最小倍数
|
||||
# 重排序
|
||||
USE_RERANK = True # 启用重排序
|
||||
RERANK_BACKEND = "local" # "local"=本地模型, "cloud"=云端API, "fallback"=优先云端失败回退本地
|
||||
RERANK_CLOUD_MODEL = "qwen3-rerank" # 云端 Rerank 模型(DashScope API)
|
||||
RERANK_CLOUD_MODEL = "qwen3-rerank"
|
||||
RERANK_CANDIDATES = 20 # 送入重排序的候选数
|
||||
RERANK_TOP_K = 15 # 重排序后保留数
|
||||
RERANK_USE_ONNX = True # ONNX 加速(仅本地模式,环境变量控制,默认开启)
|
||||
RERANK_USE_ONNX = True # ONNX 加速(仅本地模式,环境变量控制)
|
||||
|
||||
# RRF 融合
|
||||
RRF_K = 60 # RRF 常数
|
||||
@@ -581,30 +569,7 @@ MMR_TOP_K = 30 # MMR 保留数
|
||||
MMR_LAMBDA = 0.5 # 相关性 vs 多样性权衡
|
||||
```
|
||||
|
||||
### 8.3 缓存配置
|
||||
|
||||
```python
|
||||
# 查询结果缓存
|
||||
QUERY_CACHE_ENABLED = True
|
||||
QUERY_CACHE_SIZE = 500
|
||||
QUERY_CACHE_TTL = 3600 # 1小时
|
||||
|
||||
# Embedding 缓存
|
||||
EMBEDDING_CACHE_ENABLED = True
|
||||
EMBEDDING_CACHE_SIZE = 2000
|
||||
EMBEDDING_CACHE_TTL = 86400 # 24小时
|
||||
|
||||
# Rerank 缓存
|
||||
RERANK_CACHE_ENABLED = True
|
||||
RERANK_CACHE_SIZE = 1000
|
||||
RERANK_CACHE_TTL = 3600 # 1小时
|
||||
|
||||
# 语义缓存
|
||||
SEMANTIC_CACHE_ENABLED = True
|
||||
SEMANTIC_CACHE_THRESHOLD = 0.92 # 相似度阈值
|
||||
```
|
||||
|
||||
### 8.4 设备配置
|
||||
### 8.3 设备配置
|
||||
|
||||
```python
|
||||
DEVICE = "auto" # auto / cuda / cpu / cuda:0
|
||||
@@ -618,17 +583,9 @@ RERANK_DEVICE = DEVICE # Rerank 模型设备
|
||||
|
||||
```
|
||||
core/ # RAG 核心引擎
|
||||
├── engine.py # RAGEngine 单例(检索主流程、Rerank、RRF)
|
||||
├── agentic.py # AgenticRAG 主类(Mixin 组合)
|
||||
├── agentic_base.py # 基础常量与条件导入
|
||||
├── agentic_query.py # QueryRewriteMixin(查询重写)
|
||||
├── agentic_search.py # SearchMixin(网络搜索)
|
||||
├── agentic_answer.py # AnswerMixin(答案生成、幻觉验证)
|
||||
├── agentic_citation.py # CitationMixin(引用标注)
|
||||
├── agentic_media.py # RichMediaMixin(富媒体提取)
|
||||
├── agentic_quality.py # QualityMixin(质量评估)
|
||||
├── agentic_context.py # ContextMixin(上下文压缩)
|
||||
├── agentic_meta.py # MetaQuestionMixin(元问题处理)
|
||||
├── engine.py # RAGEngine 单例(检索主流程、Rerank、RRF、流式生成)
|
||||
├── cache.py # 三层 LRU 缓存管理器(Query/Embedding/Rerank)
|
||||
├── semantic_cache.py # 语义缓存(FAISS IndexFlatIP)
|
||||
├── bm25_index.py # BM25Index(关键词检索)
|
||||
├── chunker.py # 文本分块器
|
||||
├── mmr.py # MMR 去重(语义向量版 + 文本 Jaccard 版)
|
||||
@@ -637,14 +594,13 @@ core/ # RAG 核心引擎
|
||||
├── query_decomposer.py # QueryDecomposer(复杂查询拆分)
|
||||
├── query_expansion.py # 查询扩展
|
||||
├── adaptive_topk.py # AdaptiveTopK(自适应 TopK)
|
||||
├── confidence_gate.py # ConfidenceGate(置信度门控)
|
||||
├── quality_assessor.py # 多维质量评估
|
||||
├── reasoning_reflector.py # 推理反思
|
||||
├── loop_guard.py # 循环防护
|
||||
├── confidence_gate.py # ConfidenceGate(置信度门控,当前未接入生产流程)
|
||||
├── quality_assessor.py # 多维质量评估(当前未接入生产流程)
|
||||
├── reasoning_reflector.py # 推理反思(当前未接入生产流程)
|
||||
├── loop_guard.py # 循环防护(当前未接入生产流程)
|
||||
├── prompt_guard.py # Prompt 安全守卫
|
||||
├── llm_budget.py # LLM 调用预算控制
|
||||
├── llm_utils.py # LLM 调用工具函数
|
||||
├── semantic_cache.py # 语义缓存
|
||||
├── cache.py # 三层缓存管理器(Query/Embedding/Rerank)
|
||||
├── status_codes.py # 状态码定义
|
||||
└── constants.py # 公共常量
|
||||
|
||||
@@ -666,7 +622,7 @@ knowledge/ # 知识库管理
|
||||
|
||||
api/ # API 路由层
|
||||
├── __init__.py # create_app() 工厂
|
||||
├── chat_routes.py # /chat, /rag(SSE), /search
|
||||
├── chat_routes.py # /chat, /rag(SSE), /search(核心编排入口)
|
||||
├── kb_routes.py # /collections
|
||||
├── document_routes.py # /documents/*
|
||||
├── sync_routes.py # /sync
|
||||
@@ -726,6 +682,10 @@ deploy/ # 部署配置
|
||||
├── gunicorn.conf.py # Gunicorn WSGI 配置
|
||||
└── wsgi.py # WSGI 入口
|
||||
|
||||
reports/ # 分析报告
|
||||
├── cache_performance_report.md # 缓存性能验证报告
|
||||
└── redis_migration_plan.md # Redis 缓存迁移方案(规划中)
|
||||
|
||||
config/ # 运行时配置
|
||||
└── banned_words.txt # 敏感词库
|
||||
```
|
||||
@@ -734,8 +694,8 @@ config/ # 运行时配置
|
||||
|
||||
## 十、与传统 RAG 对比
|
||||
|
||||
| 特性 | 传统 RAG | Agentic RAG (当前) |
|
||||
|------|---------|-------------------|
|
||||
| 特性 | 传统 RAG | 本系统 |
|
||||
|------|---------|--------|
|
||||
| 意图判断 | 无 | IntentAnalyzer LLM 双层判断 |
|
||||
| 查询改写 | 无 | 口语化→专业术语 + 实体补全 + 指代消解 |
|
||||
| 检索方式 | 单一向量检索 | 向量 + BM25 + FAQ + 图片独立召回 |
|
||||
@@ -744,14 +704,11 @@ config/ # 运行时配置
|
||||
| 重排序 | 无 | 云端 qwen3-rerank API(支持本地 BGE 回退) |
|
||||
| 问题分解 | 无 | 自动拆分对比/推理类查询 |
|
||||
| 闲聊处理 | 无 | 意图分析自动判断 |
|
||||
| 网络搜索 | 无 | 可选支持(Serper API) |
|
||||
| 知识图谱 | 无 | ~~可选支持(Neo4j)~~(已废弃,graph/ 目录已清空) |
|
||||
| 幻觉验证 | 无 | 基于参考信息的答案验证 |
|
||||
| 置信度门控 | 无 | Reranker 分数驱动,低分触发补救 |
|
||||
| 缓存 | 无 | 三层缓存 + 语义缓存 |
|
||||
| 缓存体系 | 无 | 四层缓存(Query + Embedding + Rerank + 语义缓存) |
|
||||
| 语义缓存 | 无 | FAISS 向量索引,相似查询复用(92x 加速) |
|
||||
| 自适应 TopK | 固定 top_k | 根据置信度动态调整 |
|
||||
| 上下文理解 | 无 | 多轮对话 + 历史上下文 |
|
||||
| 响应时间 | ~2秒 | ~3-8秒(取决于 Rerank + LLM) |
|
||||
| 响应时间 | ~2秒 | 首次 ~3-8秒 / 缓存命中 ~100-200毫秒 |
|
||||
|
||||
---
|
||||
|
||||
@@ -759,71 +716,71 @@ config/ # 运行时配置
|
||||
|
||||
### 11.1 Rerank 调用路径
|
||||
|
||||
Rerank 在系统中有 **两个独立调用路径**:
|
||||
Rerank 在生产流程中有 **一个调用路径**:
|
||||
|
||||
| 路径 | 位置 | 说明 |
|
||||
|------|------|------|
|
||||
| 主检索管线 | `engine.rerank_results()` | RRF 融合后、MMR 去重前执行,对候选重排取 top_k |
|
||||
| 置信度门控 | `confidence_gate._compute_scores()` | 直接调用 `reranker.predict()`,可能重复推理 |
|
||||
|
||||
### 11.2 性能瓶颈
|
||||
### 11.2 性能特征
|
||||
|
||||
| 瓶颈 | 严重程度 | 说明 |
|
||||
|------|---------|------|
|
||||
| Rerank 缓存命中率偏低 | 🟡 中 | `rerank_results()` 已正确调用缓存读写,但缓存 key 基于 `query + sorted(doc_ids)` 精确匹配,RRF 融合产出稍有不同就无法命中 |
|
||||
| 置信度门控重复推理 | 🟡 中 | 同一 query+documents 可能被 Rerank 两次(当前仅备用路径使用,暂未影响生产) |
|
||||
| ~~无性能计时~~ | ~~🟡 中~~ | 已修复:`rerank_results()` 现返回 `_rerank_time_ms` 计时字段 |
|
||||
| 查询分类器策略未生效 | 🟢 低 | `QueryClassifier` 定义的差异化 rerank 参数未传递到引擎 |
|
||||
| 项目 | 说明 |
|
||||
|------|------|
|
||||
| Rerank 缓存命中率 | 偏低——键基于 `query + sorted(doc_ids)` 精确匹配,RRF 融合产出稍有不同就无法命中 |
|
||||
| 性能计时 | `rerank_results()` 返回 `_rerank_time_ms` 计时字段 |
|
||||
| Query Cache 拦截 | 重复查询被 Query Cache 在外层拦截,不会到达 Rerank 层(正确行为) |
|
||||
|
||||
### 11.3 Rerank 配置参数
|
||||
|
||||
| 配置项 | 默认值 | 说明 |
|
||||
|--------|--------|------|
|
||||
| `USE_RERANK` | `True` | 总开关 |
|
||||
| `RERANK_BACKEND` | `"local"` | 后端选择:`local`=本地模型, `cloud`=云端API, `fallback`=优先云端失败回退本地 |
|
||||
| `RERANK_CLOUD_MODEL` | `"qwen3-rerank"` | 云端 Rerank 模型名称(DashScope API) |
|
||||
| `RERANK_BACKEND` | `"local"` | 后端选择 |
|
||||
| `RERANK_CLOUD_MODEL` | `"qwen3-rerank"` | 云端模型名称 |
|
||||
| `RERANK_CLOUD_API_KEY` | 同 `DASHSCOPE_API_KEY` | 云端 API 密钥 |
|
||||
| `RERANK_CLOUD_BASE_URL` | `https://dashscope.aliyuncs.com/compatible-api/v1/reranks` | 云端 API 地址 |
|
||||
| `RERANK_CLOUD_TIMEOUT` | `15` | 云端请求超时(秒) |
|
||||
| `RERANK_MODEL_PATH` | `models/bge-reranker-base` | 本地模型路径(仅 local/fallback 模式) |
|
||||
| `RERANK_MODEL_PATH` | `models/bge-reranker-base` | 本地模型路径 |
|
||||
| `RERANK_CANDIDATES` | `20` | 送入 Rerank 的候选数 |
|
||||
| `RERANK_TOP_K` | `15` | Rerank 后保留数 |
|
||||
| `RERANK_USE_ONNX` | `True`(环境变量默认) | ONNX 加速开关(仅本地模式) |
|
||||
| `RERANK_DEVICE` | 跟随 `DEVICE` | 设备选择(仅本地模式) |
|
||||
| `RERANK_USE_ONNX` | `True` | ONNX 加速开关 |
|
||||
| `RERANK_DEVICE` | 跟随 `DEVICE` | 设备选择 |
|
||||
| `RERANK_THRESHOLD` | `0.3` | 上下文过滤阈值 |
|
||||
| `RERANK_CACHE_ENABLED` | `True` | 缓存开关(已在 `rerank_results()` 中使用) |
|
||||
| `RERANK_CACHE_ENABLED` | `True` | 缓存开关 |
|
||||
|
||||
---
|
||||
|
||||
## 十二、最佳实践
|
||||
|
||||
### 12.1 何时使用 Agentic RAG
|
||||
### 12.1 何时使用 /rag 接口
|
||||
|
||||
✅ **推荐使用**:
|
||||
- 复杂问题需要多轮检索
|
||||
**推荐使用**:
|
||||
|
||||
- 复杂问题需要检索知识库
|
||||
- 用户表达模糊需要改写
|
||||
- 需要区分闲聊和知识问答
|
||||
- 需要多轮对话记忆
|
||||
- 需要引用来源和幻觉验证
|
||||
- 需要引用来源和证据
|
||||
|
||||
❌ **不推荐使用**:
|
||||
- 简单明确的问题(用 `/search` 接口更快)
|
||||
- 对响应时间极度敏感的场景
|
||||
**不推荐使用**(改用 `/search` 接口更快):
|
||||
|
||||
- 简单明确的问题,只需返回原始检索结果
|
||||
- 对响应时间极度敏感且不需要 LLM 生成答案的场景
|
||||
|
||||
### 12.2 性能优化
|
||||
|
||||
```python
|
||||
# 减少迭代次数
|
||||
rag = AgenticRAG(max_iterations=2)
|
||||
|
||||
# 禁用网络搜索
|
||||
rag = AgenticRAG(enable_web_search=False)
|
||||
|
||||
# ONNX 加速默认已开启;如有兼容性问题可关闭(环境变量)
|
||||
# RERANK_USE_ONNX=false
|
||||
|
||||
# 使用轻量 MMR(文本相似度代替语义向量)
|
||||
# config.py: MMR_USE_EMBEDDING = False
|
||||
|
||||
# 调整语义缓存阈值(降低阈值可提高命中率,但可能降低准确性)
|
||||
# config.py: SEMANTIC_CACHE_THRESHOLD = 0.90
|
||||
|
||||
# 调整缓存容量
|
||||
# config.py: QUERY_CACHE_SIZE = 1000 # 增大查询缓存容量
|
||||
```
|
||||
|
||||
### 12.3 调试技巧
|
||||
@@ -834,73 +791,62 @@ result = engine.search_knowledge("问题", top_k=10)
|
||||
debug = result.get('_debug', {})
|
||||
for step in debug.get('steps', []):
|
||||
print(f"步骤: {step['name']}, 详情: {step}")
|
||||
|
||||
# 查看缓存统计
|
||||
from core.cache import get_cache_manager
|
||||
cm = get_cache_manager()
|
||||
print(cm.get_stats())
|
||||
|
||||
# 查看语义缓存统计
|
||||
from core.semantic_cache import get_semantic_cache
|
||||
sc = get_semantic_cache()
|
||||
print({"hits": sc.hits, "misses": sc.misses, "total": sc.total_entries})
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 附加篇:Agentic RAG 深入优化与工作机制
|
||||
## 十三、演进记录
|
||||
|
||||
### 一、Agentic RAG 的核心架构
|
||||
### v4.0(2026-06-05)— 统一编排 + 四层缓存修复
|
||||
|
||||
Agentic RAG 构建了动态的决策闭环,核心组件包括:
|
||||
**删除未使用的备用编排路径**:移除了 `core/agentic.py` 及 8 个 Mixin 文件(共 10 个文件 ~2050 行)。这些文件实现了完整的决策循环编排(含置信度门控、质量评估、推理反思等),但从未接入任何 HTTP 路由。
|
||||
|
||||
- **意图分析器**:LLM 驱动的双层判断,替代硬编码规则
|
||||
- **查询重写器**:口语化→专业术语、实体补全、指代消解
|
||||
- **混合检索引擎**:向量 + BM25 + FAQ + 图片独立召回 + RRF 融合
|
||||
- **MMR 去重**:平衡相关性与多样性,Rerank 后进一步精炼结果
|
||||
- **Rerank 重排**:云端 qwen3-rerank API 精确排序,支持本地 BGE 回退
|
||||
- **置信度门控**:Reranker 分数驱动,低分触发补救流程
|
||||
- **幻觉验证**:基于参考信息验证答案,防止 LLM 编造
|
||||
**修复 Query Cache**:
|
||||
|
||||
### 二、分阶段优化策略
|
||||
1. 修复 GET/SET 键不匹配——此前 SET 使用 `doc_hash` 分支,GET 使用 `kb_version` 分支,两端永远不匹配,命中率始终为 0%
|
||||
2. 修复 `CACHE_MIN_SCORE = 0.3` 阈值过高——ChromaDB 余弦距离经 `1-dist` 后得分约 0.03-0.06,远低于 0.3,导致几乎不写入缓存
|
||||
|
||||
#### 1. 检索前:优化查询质量
|
||||
**集成语义缓存**:将 FAISS 语义缓存从已删除的备用路径移植到生产 `/rag` 端点,在意图分析后、混合检索前检查,命中时跳过整个检索+生成流程。验证结果:命中率 66.7%,92 倍加速。
|
||||
|
||||
- **智能查询重写**:口语化表述 → 精准检索术语
|
||||
- **复杂问题分解**:对比/推理类查询自动拆分为子查询
|
||||
- **意图分析**:LLM 双层判断,避免不必要的检索
|
||||
### v3.2 — 模型/Reranker/管线更新
|
||||
|
||||
#### 2. 检索中:提升召回精准度
|
||||
引入云端 qwen3-rerank、ONNX 加速、动态 RRF 权重等。
|
||||
|
||||
- **多路召回与融合**:向量 + BM25 + FAQ + 图片独立召回
|
||||
- **动态 RRF 权重**:查询类型/长度驱动的权重调整
|
||||
- **MMR 去重**:Rerank 后进一步精炼,平衡相关性与多样性(召回100 → Rerank取15 → MMR精炼)
|
||||
- **Rerank 重排**:云端 qwen3-rerank 精排,置信度门控过滤低质量结果
|
||||
---
|
||||
|
||||
#### 3. 检索后:质量评估与自我迭代
|
||||
## 十四、未来规划
|
||||
|
||||
- **多维质量评估**:相关性/完整性/准确性/覆盖面
|
||||
- **推理反思**:检查推理过程中未验证的假设
|
||||
- **分层补救**:低置信度 → 查询重写 → 网络搜索
|
||||
### Redis 缓存外部化
|
||||
|
||||
### 三、系统级优化
|
||||
当前四层缓存均为进程内内存存储,在多 Worker / 多实例部署时无法共享。已规划 Redis 迁移方案(详见 `reports/redis_migration_plan.md`),核心设计:
|
||||
|
||||
#### 1. 避免"循环检索"陷阱
|
||||
- `RedisCacheManager` 提供与 `RAGCacheManager` 相同的接口
|
||||
- 通过 `REDIS_CACHE_URL` 环境变量启用,向后兼容
|
||||
- 语义缓存采用混合方案:FAISS 索引保持在进程内,缓存结果存储到 Redis
|
||||
- Query Cache、Embedding Cache、Rerank Cache 全部迁移到 Redis
|
||||
|
||||
- 循环防护器(`loop_guard.py`):最多允许 N 次重写检索
|
||||
- 置信度递增检查:连续两次无提升则终止
|
||||
### 可选能力接入
|
||||
|
||||
#### 2. 平衡智能性与效率
|
||||
`core/` 目录下仍保留以下独立模块,当前未接入生产流程,可按需启用:
|
||||
|
||||
- 轻量级决策模型:意图分析使用低温度、少 token 的 LLM 调用
|
||||
- 三层缓存:Query Cache + Embedding Cache + Rerank Cache
|
||||
- 语义缓存:相似查询复用结果(threshold=0.92)
|
||||
- LLM 预算控制:`MAX_LLM_CALLS_PER_QUERY = 2`
|
||||
- `confidence_gate.py`:置信度门控,基于 Reranker 分数判断检索质量
|
||||
- `quality_assessor.py`:多维质量评估(相关性/完整性/准确性/覆盖面)
|
||||
- `reasoning_reflector.py`:推理反思,检查未验证的假设
|
||||
- `loop_guard.py`:循环防护,防止重复检索
|
||||
|
||||
#### 3. 安全与可解释性
|
||||
---
|
||||
|
||||
- 证据溯源:引用标注 + 来源编号
|
||||
- 思维链展示:`log_trace` 记录推理过程
|
||||
- 安全护栏:输入验证 + 输出过滤 + 权限控制
|
||||
|
||||
### 四、学术前沿
|
||||
|
||||
1. **RAG-Gym**:三维度系统优化(提示工程 + 执行器调优 + 评判器训练)
|
||||
2. **过程监督 vs 结果监督**:细粒度过程奖励显著提升训练效率
|
||||
3. **Re2Search**:推理反思机制,F1 score 提升 10%+
|
||||
|
||||
### 参考资料
|
||||
## 参考资料
|
||||
|
||||
1. Xiong, G., et al. (2025). RAG-Gym: Systematic Optimization of Language Agents for Retrieval-Augmented Generation. arXiv:2502.13957
|
||||
2. Zhang, W., et al. (2025). Process vs. Outcome Reward: Which is Better for Agentic RAG Reinforcement Learning. arXiv:2505.14069
|
||||
3. Agentic RAG 实战指南:从查询重写到多步重查全掌握。火山引擎 ADG 社区
|
||||
@@ -444,7 +444,7 @@ Query Rewriting: "它" → "出差补助"
|
||||
|
||||
- [后端对接规范.md](./后端对接规范.md) - API 接口规范(主要)
|
||||
- [数据库设计文档.md](./数据库设计文档.md) - 数据库结构
|
||||
- [Agentic_RAG完整指南.md](./Agentic_RAG完整指南.md) - Agentic RAG 详解
|
||||
- [RAG系统完整指南.md](./RAG系统完整指南.md) - RAG 系统架构与缓存详解
|
||||
|
||||
|
||||
---
|
||||
|
||||
@@ -382,34 +382,18 @@ class BM25Index:
|
||||
|
||||
def add_documents(self, ids: List[str], documents: List[str], metadatas: List[dict]) -> None:
|
||||
"""
|
||||
添加文档到索引(追加模式,自动去重)
|
||||
|
||||
如果 ID 已存在则更新对应文档,否则追加新文档。
|
||||
添加后自动重建 BM25 索引。
|
||||
添加文档到索引(会覆盖原有索引)
|
||||
|
||||
Args:
|
||||
ids: 文档 ID 列表
|
||||
documents: 文档内容列表
|
||||
metadatas: 文档元数据列表
|
||||
"""
|
||||
# 建立已有 ID -> 索引位置 的映射,用于去重
|
||||
existing_map = {doc_id: idx for idx, doc_id in enumerate(self.ids)}
|
||||
|
||||
for i, doc_id in enumerate(ids):
|
||||
if doc_id in existing_map:
|
||||
# 更新已有文档
|
||||
pos = existing_map[doc_id]
|
||||
self.documents[pos] = documents[i]
|
||||
self.metadatas[pos] = metadatas[i]
|
||||
else:
|
||||
# 追加新文档
|
||||
existing_map[doc_id] = len(self.ids)
|
||||
self.ids.append(doc_id)
|
||||
self.documents.append(documents[i])
|
||||
self.metadatas.append(metadatas[i])
|
||||
|
||||
if self.documents:
|
||||
tokenized = [self.tokenize(doc) for doc in self.documents]
|
||||
self.ids = ids
|
||||
self.documents = documents
|
||||
self.metadatas = metadatas
|
||||
if documents:
|
||||
tokenized = [self.tokenize(doc) for doc in documents]
|
||||
self.bm25 = BM25Okapi(tokenized)
|
||||
|
||||
def search(self, query: str, top_k: int = 10) -> Tuple[List[str], List[str], List[dict], List[float]]:
|
||||
|
||||
@@ -134,9 +134,6 @@ class CollectionMixin:
|
||||
"""
|
||||
from .base import BM25Index
|
||||
|
||||
# 从磁盘重新加载元数据,确保多 worker 进程间状态一致
|
||||
self._metadata = self._load_metadata()
|
||||
|
||||
if not kb_name or not kb_name.replace('_', '').isalnum():
|
||||
return False, "向量库名称只能包含字母、数字和下划线"
|
||||
|
||||
@@ -193,9 +190,6 @@ class CollectionMixin:
|
||||
Returns:
|
||||
更新成功返回 True,向量库不存在返回 False
|
||||
"""
|
||||
# 从磁盘重新加载元数据,确保多 worker 进程间状态一致
|
||||
self._metadata = self._load_metadata()
|
||||
|
||||
collections = self._metadata.get("collections", {})
|
||||
if kb_name not in collections:
|
||||
return False
|
||||
@@ -234,9 +228,6 @@ class CollectionMixin:
|
||||
"""
|
||||
import shutil
|
||||
|
||||
# 从磁盘重新加载元数据,确保多 worker 进程间状态一致
|
||||
self._metadata = self._load_metadata()
|
||||
|
||||
if kb_name == PUBLIC_KB_NAME:
|
||||
return False, "公开知识库不能删除"
|
||||
|
||||
@@ -331,21 +322,6 @@ class CollectionMixin:
|
||||
except Exception as e:
|
||||
logger.warning(f"清理版本记录失败: {e}")
|
||||
|
||||
# 清理不再被引用的图片和 VLM 缓存文件
|
||||
# 注意:此时 ChromaDB collection 已删除,cleanup_image_orphans 会扫描
|
||||
# 所有剩余 collection,仅该 collection 引用的图片会被识别为孤儿
|
||||
try:
|
||||
from knowledge.image_cleanup import cleanup_image_orphans
|
||||
cleanup_result = cleanup_image_orphans(self)
|
||||
if cleanup_result['deleted_images'] or cleanup_result['deleted_caches']:
|
||||
logger.info(
|
||||
f"清理孤儿文件: {cleanup_result['deleted_images']} 图片 + "
|
||||
f"{cleanup_result['deleted_caches']} VLM缓存, "
|
||||
f"释放 {cleanup_result['freed_bytes']/1024:.1f} KB"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"清理孤儿文件失败: {e}")
|
||||
|
||||
if kb_name in self._metadata.get("collections", {}):
|
||||
del self._metadata["collections"][kb_name]
|
||||
self._save_metadata()
|
||||
@@ -372,8 +348,6 @@ class CollectionMixin:
|
||||
- department: 所属部门
|
||||
- description: 描述
|
||||
"""
|
||||
# 从磁盘重新加载元数据,确保多 worker 进程间状态一致
|
||||
self._metadata = self._load_metadata()
|
||||
result = []
|
||||
|
||||
# 扫描 base_path 下的所有子目录作为向量库
|
||||
@@ -408,31 +382,16 @@ class CollectionMixin:
|
||||
|
||||
self._save_metadata()
|
||||
|
||||
stale_collections = []
|
||||
|
||||
for name, info in self._metadata.get("collections", {}).items():
|
||||
try:
|
||||
collection = self.get_collection(name)
|
||||
result.append(CollectionInfo(
|
||||
name=name,
|
||||
display_name=info.get("display_name", name),
|
||||
document_count=collection.count() if collection else 0,
|
||||
created_at=info.get("created_at", ""),
|
||||
department=info.get("department", ""),
|
||||
description=info.get("description", "")
|
||||
))
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"跳过异常向量库 '{name}': {e},可能是 ChromaDB 数据目录已丢失"
|
||||
)
|
||||
stale_collections.append(name)
|
||||
|
||||
# 清理元数据中指向已失效集合的条目
|
||||
if stale_collections:
|
||||
for name in stale_collections:
|
||||
self._metadata.get("collections", {}).pop(name, None)
|
||||
logger.info(f"清理失效向量库元数据: {name}")
|
||||
self._save_metadata()
|
||||
collection = self.get_collection(name)
|
||||
result.append(CollectionInfo(
|
||||
name=name,
|
||||
display_name=info.get("display_name", name),
|
||||
document_count=collection.count() if collection else 0,
|
||||
created_at=info.get("created_at", ""),
|
||||
department=info.get("department", ""),
|
||||
description=info.get("description", "")
|
||||
))
|
||||
|
||||
return result
|
||||
|
||||
@@ -446,6 +405,4 @@ class CollectionMixin:
|
||||
Returns:
|
||||
存在返回 True,不存在返回 False
|
||||
"""
|
||||
# 从磁盘重新加载元数据,确保多 worker 进程间状态一致
|
||||
self._metadata = self._load_metadata()
|
||||
return kb_name in self._metadata.get("collections", {})
|
||||
|
||||
@@ -72,18 +72,6 @@ class DocumentMixin:
|
||||
except Exception as e:
|
||||
logger.warning(f"清理版本记录失败: {e}")
|
||||
|
||||
# 清理不再被引用的图片和 VLM 缓存文件
|
||||
try:
|
||||
from knowledge.image_cleanup import cleanup_image_orphans
|
||||
cleanup_result = cleanup_image_orphans(self, collections=[kb_name])
|
||||
if cleanup_result['deleted_images'] or cleanup_result['deleted_caches']:
|
||||
logger.info(
|
||||
f"清理孤儿文件: {cleanup_result['deleted_images']} 图片 + "
|
||||
f"{cleanup_result['deleted_caches']} VLM缓存"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"清理孤儿文件失败: {e}")
|
||||
|
||||
logger.info(f"从 {kb_name} 删除文档: {filename}, 片段数: {deleted}")
|
||||
return deleted
|
||||
|
||||
|
||||
@@ -1,140 +0,0 @@
|
||||
"""
|
||||
图片/VLM缓存孤儿文件清理模块
|
||||
|
||||
提供可被 document.py / collection.py 调用的清理函数,
|
||||
也可被 cleanup_orphans.py 独立脚本使用。
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
IMAGES_DIR = Path(".data/images")
|
||||
VLM_CACHE_DIR = Path(".data/cache/vlm")
|
||||
|
||||
|
||||
def compute_file_hash(file_path: str) -> str:
|
||||
"""计算文件 MD5"""
|
||||
with open(file_path, 'rb') as f:
|
||||
return hashlib.md5(f.read()).hexdigest()
|
||||
|
||||
|
||||
def collect_referenced_images(manager, collections=None) -> set:
|
||||
"""
|
||||
从 ChromaDB 收集所有被引用的图片文件名。
|
||||
|
||||
Args:
|
||||
manager: KnowledgeBaseManager 实例
|
||||
collections: 限定知识库列表,None 表示全部
|
||||
|
||||
Returns:
|
||||
set of image filenames (e.g., {"185a7a75d246.png", ...})
|
||||
"""
|
||||
referenced = set()
|
||||
|
||||
if collections:
|
||||
kb_names = collections
|
||||
else:
|
||||
kb_names = [c.name if hasattr(c, 'name') else str(c)
|
||||
for c in manager.list_collections()]
|
||||
|
||||
for kb_name in kb_names:
|
||||
try:
|
||||
col = manager.get_collection(kb_name)
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
results = col.get(include=['metadatas'])
|
||||
if not results['ids']:
|
||||
continue
|
||||
|
||||
for meta in results['metadatas']:
|
||||
image_path = meta.get('image_path', '')
|
||||
if image_path:
|
||||
referenced.add(os.path.basename(image_path))
|
||||
|
||||
return referenced
|
||||
|
||||
|
||||
def cleanup_image_orphans(manager, collections=None, dry_run=False) -> dict:
|
||||
"""
|
||||
清理不再被任何 ChromaDB 切片引用的图片和 VLM 缓存文件。
|
||||
|
||||
Args:
|
||||
manager: KnowledgeBaseManager 实例
|
||||
collections: 限定知识库列表,None 表示全部
|
||||
dry_run: True 时只返回孤儿列表不实际删除
|
||||
|
||||
Returns:
|
||||
{
|
||||
'orphan_images': [(filepath, filename, size_bytes)],
|
||||
'orphan_caches': [(filepath, filename, size_bytes)],
|
||||
'deleted_images': int,
|
||||
'deleted_caches': int,
|
||||
'freed_bytes': int
|
||||
}
|
||||
"""
|
||||
result = {
|
||||
'orphan_images': [],
|
||||
'orphan_caches': [],
|
||||
'deleted_images': 0,
|
||||
'deleted_caches': 0,
|
||||
'freed_bytes': 0
|
||||
}
|
||||
|
||||
# 1. 收集引用
|
||||
referenced = collect_referenced_images(manager, collections)
|
||||
|
||||
# 2. 查找孤儿图片
|
||||
if IMAGES_DIR.exists():
|
||||
for f in IMAGES_DIR.iterdir():
|
||||
if f.is_file() and f.name not in referenced:
|
||||
result['orphan_images'].append((str(f), f.name, f.stat().st_size))
|
||||
|
||||
# 3. 查找孤儿 VLM 缓存(图片已删除则缓存也应是孤儿)
|
||||
if VLM_CACHE_DIR.exists():
|
||||
referenced_hashes = set()
|
||||
for filename in referenced:
|
||||
full_path = IMAGES_DIR / filename
|
||||
if full_path.exists():
|
||||
try:
|
||||
img_hash = compute_file_hash(str(full_path))
|
||||
referenced_hashes.add(img_hash)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
for f in VLM_CACHE_DIR.iterdir():
|
||||
if f.is_file() and f.suffix == '.txt':
|
||||
cache_hash = f.stem
|
||||
if cache_hash not in referenced_hashes:
|
||||
result['orphan_caches'].append((str(f), f.name, f.stat().st_size))
|
||||
|
||||
# 4. 删除
|
||||
if not dry_run:
|
||||
for filepath, filename, size in result['orphan_images']:
|
||||
try:
|
||||
os.remove(filepath)
|
||||
result['deleted_images'] += 1
|
||||
result['freed_bytes'] += size
|
||||
except OSError as e:
|
||||
logger.warning(f"删除孤儿图片失败: {filename} - {e}")
|
||||
|
||||
for filepath, filename, size in result['orphan_caches']:
|
||||
try:
|
||||
os.remove(filepath)
|
||||
result['deleted_caches'] += 1
|
||||
result['freed_bytes'] += size
|
||||
except OSError as e:
|
||||
logger.warning(f"删除孤儿缓存失败: {filename} - {e}")
|
||||
|
||||
if result['deleted_images'] or result['deleted_caches']:
|
||||
logger.info(
|
||||
f"清理孤儿: {result['deleted_images']} 图片 + "
|
||||
f"{result['deleted_caches']} 缓存, "
|
||||
f"释放 {result['freed_bytes']/1024:.1f} KB"
|
||||
)
|
||||
|
||||
return result
|
||||
@@ -25,20 +25,7 @@ def compute_file_hash(file_path: str) -> str:
|
||||
return hashlib.md5(file_path.encode()).hexdigest()
|
||||
|
||||
|
||||
def _get_embedding_model():
|
||||
"""从 RAGEngine 获取 embedding 模型(KnowledgeBaseManager 上没有此属性)"""
|
||||
try:
|
||||
from core.engine import get_engine
|
||||
engine = get_engine()
|
||||
if not engine._initialized:
|
||||
engine.initialize()
|
||||
return engine.embedding_model
|
||||
except Exception as e:
|
||||
logger.warning(f"获取 embedding 模型失败: {e}")
|
||||
return None
|
||||
|
||||
|
||||
async def lazy_vlm_description(chunk_id: str, image_path: str, kb_name: str, metadata: dict = None, defer_chromadb: bool = False) -> str:
|
||||
async def lazy_vlm_description(chunk_id: str, image_path: str, kb_name: str, metadata: dict = None) -> str:
|
||||
"""
|
||||
懒加载 VLM 描述
|
||||
|
||||
@@ -49,7 +36,6 @@ async def lazy_vlm_description(chunk_id: str, image_path: str, kb_name: str, met
|
||||
image_path: 图片路径(相对路径或绝对路径)
|
||||
kb_name: 知识库名称
|
||||
metadata: 图片元数据(包含 section、page、caption、上下文等)
|
||||
defer_chromadb: 为 True 时跳过 ChromaDB 更新(仅写文件缓存),避免后台线程写锁竞争
|
||||
|
||||
Returns:
|
||||
VLM 生成的图片描述
|
||||
@@ -63,45 +49,23 @@ async def lazy_vlm_description(chunk_id: str, image_path: str, kb_name: str, met
|
||||
else:
|
||||
full_image_path = image_path
|
||||
|
||||
# 1. 检查缓存(空缓存视为无效,需重新生成)
|
||||
# 1. 检查缓存
|
||||
img_hash = compute_file_hash(full_image_path)
|
||||
cache_file = VLM_CACHE_DIR / f"{img_hash}.txt"
|
||||
if cache_file.exists():
|
||||
cached = cache_file.read_text(encoding='utf-8')
|
||||
if len(cached.strip()) >= 5:
|
||||
logger.info(f"VLM 缓存命中: {image_path}")
|
||||
return cached
|
||||
else:
|
||||
logger.warning(f"VLM 缓存内容过短({len(cached.strip())}字符),删除并重新生成: {image_path}")
|
||||
try:
|
||||
cache_file.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
logger.info(f"VLM 缓存命中: {image_path}")
|
||||
return cache_file.read_text(encoding='utf-8')
|
||||
|
||||
# 2. 调用 VLM(传入元数据)
|
||||
logger.info(f"VLM 懒加载: {image_path}")
|
||||
kb_manager = get_kb_manager()
|
||||
description = kb_manager._generate_image_description(full_image_path, metadata=metadata)
|
||||
|
||||
# 3. 空描述保护:VLM 返回内容过短时不写入缓存和向量库
|
||||
if not description or len(description.strip()) < 5:
|
||||
logger.warning(f"VLM 返回描述过短({len(description.strip()) if description else 0}字符),跳过缓存和向量库更新: {image_path}")
|
||||
return description or ''
|
||||
|
||||
# 4. 写入缓存
|
||||
# 3. 写入缓存
|
||||
VLM_CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
cache_file.write_text(description, encoding='utf-8')
|
||||
|
||||
# 5. 更新向量库(metadata + embedding),需校验 chunk_id 非空
|
||||
# defer_chromadb=True 时跳过(后台线程只写缓存,避免 SQLite 写锁竞争)
|
||||
if defer_chromadb:
|
||||
logger.info(f"延迟 ChromaDB 更新(仅写缓存): {chunk_id}")
|
||||
return description
|
||||
|
||||
if not chunk_id:
|
||||
logger.warning("chunk_id 为空,跳过向量库更新")
|
||||
return description
|
||||
|
||||
# 4. 更新向量库(metadata + embedding)
|
||||
try:
|
||||
collection = kb_manager.get_collection(kb_name)
|
||||
result = collection.get(ids=[chunk_id], include=['metadatas'])
|
||||
@@ -115,7 +79,7 @@ async def lazy_vlm_description(chunk_id: str, image_path: str, kb_name: str, met
|
||||
|
||||
# 更新 embedding(使用 VLM 描述重新计算向量)
|
||||
# 这样 VLM 描述中的关键词(如"发电量")才能参与相似度检索
|
||||
embedding_model = _get_embedding_model()
|
||||
embedding_model = kb_manager.embedding_model
|
||||
if embedding_model:
|
||||
new_vector = embedding_model.encode(description).tolist()
|
||||
if isinstance(new_vector[0], list):
|
||||
@@ -127,21 +91,20 @@ async def lazy_vlm_description(chunk_id: str, image_path: str, kb_name: str, met
|
||||
embeddings=[new_vector],
|
||||
documents=[description] # 同时更新 document 字段
|
||||
)
|
||||
logger.info(f"已更新向量库(embedding+metadata): {chunk_id}")
|
||||
logger.info(f"已更新向量库 embedding: {chunk_id}")
|
||||
else:
|
||||
# 无 embedding 模型时只更新 metadata
|
||||
collection.update(
|
||||
ids=[chunk_id],
|
||||
metadatas=[new_metadata]
|
||||
)
|
||||
logger.info(f"已更新向量库(仅metadata,无embedding模型): {chunk_id}")
|
||||
except Exception as e:
|
||||
logger.warning(f"更新向量库失败: {e}")
|
||||
|
||||
return description
|
||||
|
||||
|
||||
async def lazy_table_summary(chunk_id: str, table_md: str, kb_name: str, defer_chromadb: bool = False) -> str:
|
||||
async def lazy_table_summary(chunk_id: str, table_md: str, kb_name: str) -> str:
|
||||
"""
|
||||
懒加载表格摘要
|
||||
|
||||
@@ -151,76 +114,50 @@ async def lazy_table_summary(chunk_id: str, table_md: str, kb_name: str, defer_c
|
||||
chunk_id: 切片 ID
|
||||
table_md: 表格 Markdown 内容
|
||||
kb_name: 知识库名称
|
||||
defer_chromadb: 为 True 时跳过 ChromaDB 更新(仅写文件缓存),避免后台线程写锁竞争
|
||||
|
||||
Returns:
|
||||
LLM 生成的表格摘要
|
||||
"""
|
||||
from knowledge.manager import get_kb_manager
|
||||
|
||||
# 1. 检查缓存(空缓存视为无效)
|
||||
# 1. 检查缓存
|
||||
table_hash = hashlib.md5(table_md.encode()).hexdigest()
|
||||
cache_file = LLM_CACHE_DIR / f"{table_hash}.txt"
|
||||
if cache_file.exists():
|
||||
cached = cache_file.read_text(encoding='utf-8')
|
||||
if len(cached.strip()) >= 5:
|
||||
logger.info(f"LLM 缓存命中: {chunk_id}")
|
||||
return cached
|
||||
else:
|
||||
logger.warning(f"LLM 缓存内容过短({len(cached.strip())}字符),删除并重新生成: {chunk_id}")
|
||||
try:
|
||||
cache_file.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
logger.info(f"LLM 缓存命中: {chunk_id}")
|
||||
return cache_file.read_text(encoding='utf-8')
|
||||
|
||||
# 2. 调用 LLM
|
||||
logger.info(f"LLM 懒加载: {chunk_id}")
|
||||
kb_manager = get_kb_manager()
|
||||
summary = kb_manager._generate_table_summary(table_md, None)
|
||||
|
||||
# 空摘要保护
|
||||
if not summary or len(summary.strip()) < 5:
|
||||
logger.warning(f"LLM 返回摘要过短,跳过缓存和向量库更新: {chunk_id}")
|
||||
return summary or ''
|
||||
|
||||
# 3. 写入缓存
|
||||
LLM_CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
cache_file.write_text(summary, encoding='utf-8')
|
||||
|
||||
# 4. 更新向量库,需校验 chunk_id 非空
|
||||
# defer_chromadb=True 时跳过(后台线程只写缓存,避免 SQLite 写锁竞争)
|
||||
if defer_chromadb:
|
||||
logger.info(f"延迟 ChromaDB 更新(仅写缓存): {chunk_id}")
|
||||
return summary
|
||||
|
||||
if not chunk_id:
|
||||
logger.warning("chunk_id 为空,跳过表格向量库更新")
|
||||
return summary
|
||||
# 4. 更新向量库(可选)
|
||||
try:
|
||||
collection = kb_manager.get_collection(kb_name)
|
||||
result = collection.get(ids=[chunk_id], include=['metadatas'])
|
||||
if result['metadatas']:
|
||||
# 新增摘要切片(需要 embedding 模型)
|
||||
embedding_model = _get_embedding_model()
|
||||
if embedding_model:
|
||||
vector = embedding_model.encode(summary).tolist()
|
||||
if isinstance(vector[0], list):
|
||||
vector = vector[0]
|
||||
# 新增摘要切片
|
||||
embedding_model = kb_manager.embedding_model
|
||||
vector = embedding_model.encode(summary).tolist()
|
||||
if isinstance(vector[0], list):
|
||||
vector = vector[0]
|
||||
|
||||
collection.add(
|
||||
ids=[f"{chunk_id}_summary"],
|
||||
embeddings=[vector],
|
||||
documents=[summary],
|
||||
metadatas=[{
|
||||
**result['metadatas'][0],
|
||||
'is_summary': True,
|
||||
'original_doc_id': chunk_id
|
||||
}]
|
||||
)
|
||||
logger.info(f"已新增摘要切片(embedding): {chunk_id}_summary")
|
||||
else:
|
||||
logger.info(f"跳过摘要切片(无embedding模型): {chunk_id}")
|
||||
# 更新原切片标记(不依赖 embedding 模型)
|
||||
collection.add(
|
||||
ids=[f"{chunk_id}_summary"],
|
||||
embeddings=[vector],
|
||||
documents=[summary],
|
||||
metadatas=[{
|
||||
**result['metadatas'][0],
|
||||
'is_summary': True,
|
||||
'original_doc_id': chunk_id
|
||||
}]
|
||||
)
|
||||
# 更新原切片标记
|
||||
collection.update(
|
||||
ids=[chunk_id],
|
||||
metadatas=[{**result['metadatas'][0], 'has_summary': True}]
|
||||
@@ -231,7 +168,7 @@ async def lazy_table_summary(chunk_id: str, table_md: str, kb_name: str, defer_c
|
||||
return summary
|
||||
|
||||
|
||||
async def enhance_retrieved_chunks(contexts: list, query: str, kb_name: str, defer_chromadb: bool = False):
|
||||
async def enhance_retrieved_chunks(contexts: list, query: str, kb_name: str):
|
||||
"""
|
||||
检索后增强:按需调用 LLM/VLM
|
||||
|
||||
@@ -239,24 +176,23 @@ async def enhance_retrieved_chunks(contexts: list, query: str, kb_name: str, def
|
||||
contexts: 检索上下文列表
|
||||
query: 用户查询
|
||||
kb_name: 知识库名称
|
||||
defer_chromadb: 为 True 时后台线程只写文件缓存,不更新 ChromaDB(避免写锁竞争)
|
||||
"""
|
||||
import re
|
||||
|
||||
for ctx in contexts:
|
||||
try:
|
||||
meta = ctx.get('meta', {})
|
||||
chunk_type = meta.get('chunk_type', 'text')
|
||||
image_path = meta.get('image_path', '')
|
||||
meta = ctx.get('meta', {})
|
||||
chunk_type = meta.get('chunk_type', 'text')
|
||||
image_path = meta.get('image_path', '')
|
||||
|
||||
# 图片切片:懒加载 VLM 描述
|
||||
if chunk_type in ('image', 'chart') and not meta.get('has_vlm_desc'):
|
||||
if image_path:
|
||||
# 图片切片:懒加载 VLM 描述
|
||||
if chunk_type in ('image', 'chart') and not meta.get('has_vlm_desc'):
|
||||
if image_path:
|
||||
try:
|
||||
# 从 doc 字段中提取图号(上下文可能包含"见图2.5"等)
|
||||
doc_text = ctx.get('doc', '')
|
||||
import re
|
||||
|
||||
# 提取图号(从前文/后文中)
|
||||
figure_number = ""
|
||||
# 匹配 "见图2.5"、"图2.5"、"见图 2.5" 等
|
||||
fig_match = re.search(r'[见如]?图\s*(\d+\.?\d*)', doc_text)
|
||||
if fig_match:
|
||||
figure_number = fig_match.group(1)
|
||||
@@ -274,73 +210,78 @@ async def enhance_retrieved_chunks(contexts: list, query: str, kb_name: str, def
|
||||
'page': meta.get('page'),
|
||||
'caption': meta.get('caption', ''),
|
||||
'source': meta.get('source', ''),
|
||||
'figure_number': figure_number,
|
||||
'doc_text': doc_text
|
||||
'figure_number': figure_number, # 添加提取的图号
|
||||
'doc_text': doc_text # 添加完整文档文本
|
||||
}
|
||||
vlm_desc = await lazy_vlm_description(
|
||||
meta.get('chunk_id', ''),
|
||||
meta.get('id', ''),
|
||||
image_path,
|
||||
kb_name,
|
||||
metadata=image_metadata,
|
||||
defer_chromadb=defer_chromadb
|
||||
metadata=image_metadata
|
||||
)
|
||||
if vlm_desc:
|
||||
ctx['doc'] = vlm_desc
|
||||
ctx['vlm_enhanced'] = True
|
||||
ctx['doc'] = vlm_desc
|
||||
ctx['vlm_enhanced'] = True
|
||||
except Exception as e:
|
||||
logger.warning(f"VLM 懒加载失败: {e}")
|
||||
|
||||
# 表格切片:同时处理摘要和关联图片的 VLM 描述
|
||||
elif chunk_type == 'table':
|
||||
doc_text = ctx.get('doc', '')
|
||||
# 表格切片:同时处理摘要和关联图片的 VLM 描述
|
||||
elif chunk_type == 'table':
|
||||
doc_text = ctx.get('doc', '')
|
||||
|
||||
# 1. 懒加载表格摘要(高分切片)
|
||||
if not meta.get('has_summary'):
|
||||
score = ctx.get('score', 0)
|
||||
if score > 0.7:
|
||||
# 1. 懒加载表格摘要(高分切片)
|
||||
if not meta.get('has_summary'):
|
||||
score = meta.get('score', 0)
|
||||
if score > 0.7: # 只对高相关表格生成摘要
|
||||
try:
|
||||
summary = await lazy_table_summary(
|
||||
meta.get('chunk_id', ''),
|
||||
meta.get('id', ''),
|
||||
doc_text,
|
||||
kb_name,
|
||||
defer_chromadb=defer_chromadb
|
||||
kb_name
|
||||
)
|
||||
if summary:
|
||||
ctx['summary'] = summary
|
||||
ctx['llm_enhanced'] = True
|
||||
# 摘要作为补充信息
|
||||
ctx['summary'] = summary
|
||||
ctx['llm_enhanced'] = True
|
||||
except Exception as e:
|
||||
logger.warning(f"表格摘要懒加载失败: {e}")
|
||||
|
||||
# 2. 表格有关联图片时,懒加载 VLM 描述
|
||||
if image_path and not meta.get('has_vlm_desc'):
|
||||
try:
|
||||
import re
|
||||
|
||||
# 2. 表格有关联图片时,懒加载 VLM 描述
|
||||
if image_path and not meta.get('has_vlm_desc'):
|
||||
# 提取表号(如 "表2.2"、"见表2.1")
|
||||
table_number = ""
|
||||
# 匹配 "表2.2"、"见表2.2"、"见表 2.2" 等
|
||||
table_match = re.search(r'[见如]?表\s*(\d+\.?\d*)', doc_text)
|
||||
if table_match:
|
||||
table_number = table_match.group(1)
|
||||
|
||||
# 如果 doc 中没有,尝试从 section 中提取
|
||||
section = meta.get('section') or meta.get('section_path', '')
|
||||
if not table_number and section:
|
||||
table_match = re.search(r'[见如]?表\s*(\d+\.?\d*)', section)
|
||||
if table_match:
|
||||
table_number = table_match.group(1)
|
||||
|
||||
# 构建表格图片元数据
|
||||
table_image_metadata = {
|
||||
'section': section,
|
||||
'page': meta.get('page'),
|
||||
'caption': meta.get('caption', ''),
|
||||
'source': meta.get('source', ''),
|
||||
'table_number': table_number,
|
||||
'figure_number': table_number,
|
||||
'table_number': table_number, # 表号
|
||||
'figure_number': table_number, # 兼容字段
|
||||
'doc_text': doc_text,
|
||||
'is_table': True
|
||||
'is_table': True # 标记为表格图片
|
||||
}
|
||||
vlm_desc = await lazy_vlm_description(
|
||||
meta.get('chunk_id', ''),
|
||||
meta.get('id', ''),
|
||||
image_path,
|
||||
kb_name,
|
||||
metadata=table_image_metadata,
|
||||
defer_chromadb=defer_chromadb
|
||||
metadata=table_image_metadata
|
||||
)
|
||||
if vlm_desc:
|
||||
ctx['image_description'] = vlm_desc
|
||||
ctx['vlm_enhanced'] = True
|
||||
|
||||
except Exception as e:
|
||||
chunk_id = ctx.get('meta', {}).get('chunk_id', '?')
|
||||
logger.warning(f"增强切片失败(chunk_id={chunk_id}): {e}")
|
||||
# 表格图片描述作为补充信息
|
||||
ctx['image_description'] = vlm_desc
|
||||
ctx['vlm_enhanced'] = True
|
||||
except Exception as e:
|
||||
logger.warning(f"表格图片 VLM 懒加载失败: {e}")
|
||||
|
||||
@@ -29,17 +29,11 @@
|
||||
import os
|
||||
import json
|
||||
import threading
|
||||
try:
|
||||
import fcntl
|
||||
_HAS_FCNTL = True
|
||||
except ImportError:
|
||||
_HAS_FCNTL = False # Windows 环境无 fcntl
|
||||
from typing import List, Dict, Optional, Tuple
|
||||
from pathlib import Path
|
||||
import logging
|
||||
|
||||
import chromadb
|
||||
from bs4 import BeautifulSoup
|
||||
|
||||
# 从 base.py 导入基础类和常量
|
||||
from .base import (
|
||||
@@ -140,36 +134,22 @@ class KnowledgeBaseManager(
|
||||
logger.info(f"知识库管理器初始化完成,路径: {self.base_path},发现 {len(existing_kbs)} 个向量库: {existing_kbs}")
|
||||
|
||||
def _load_metadata(self) -> dict:
|
||||
"""加载元数据(带文件锁,确保多 worker 进程间一致)"""
|
||||
"""加载元数据"""
|
||||
metadata_path = os.path.join(self.base_path, KB_METADATA_FILE)
|
||||
if os.path.exists(metadata_path):
|
||||
try:
|
||||
with open(metadata_path, 'r', encoding='utf-8') as f:
|
||||
if _HAS_FCNTL:
|
||||
fcntl.flock(f.fileno(), fcntl.LOCK_SH)
|
||||
try:
|
||||
return json.load(f)
|
||||
finally:
|
||||
if _HAS_FCNTL:
|
||||
fcntl.flock(f.fileno(), fcntl.LOCK_UN)
|
||||
return json.load(f)
|
||||
except Exception as e:
|
||||
logger.error(f"加载元数据失败: {e}")
|
||||
return {"collections": {}}
|
||||
|
||||
def _save_metadata(self):
|
||||
"""保存元数据(带文件锁,防止并发写入数据覆盖)"""
|
||||
"""保存元数据"""
|
||||
metadata_path = os.path.join(self.base_path, KB_METADATA_FILE)
|
||||
try:
|
||||
with open(metadata_path, 'w', encoding='utf-8') as f:
|
||||
if _HAS_FCNTL:
|
||||
fcntl.flock(f.fileno(), fcntl.LOCK_EX)
|
||||
try:
|
||||
json.dump(self._metadata, f, ensure_ascii=False, indent=2)
|
||||
f.flush()
|
||||
os.fsync(f.fileno())
|
||||
finally:
|
||||
if _HAS_FCNTL:
|
||||
fcntl.flock(f.fileno(), fcntl.LOCK_UN)
|
||||
json.dump(self._metadata, f, ensure_ascii=False, indent=2)
|
||||
except Exception as e:
|
||||
logger.error(f"保存元数据失败: {e}")
|
||||
|
||||
@@ -336,7 +316,6 @@ class KnowledgeBaseManager(
|
||||
"section": section_path,
|
||||
"status": "active",
|
||||
"version": "v1",
|
||||
"text_level": getattr(chunk, 'text_level', 0), # 标题级别(0=正文,1=h1,2=h2,3=h3),供检索管线层次感知
|
||||
}
|
||||
|
||||
if extra_metadata:
|
||||
@@ -349,25 +328,6 @@ class KnowledgeBaseManager(
|
||||
if hasattr(chunk, 'image_path') and chunk.image_path:
|
||||
metadata['image_path'] = chunk.image_path
|
||||
|
||||
# bbox 坐标(PDF 有,DOCX 无,需 None 保护)
|
||||
# _build_citation() 从 metadata 读取 bbox 做引用定位
|
||||
chunk_bbox = getattr(chunk, 'bbox', None)
|
||||
if chunk_bbox:
|
||||
metadata['bbox'] = json.dumps(chunk_bbox)
|
||||
|
||||
# MinerU 结构化元数据(表格类型、嵌套层级、图片子类型)
|
||||
chunk_table_type = getattr(chunk, 'table_type', '')
|
||||
if chunk_table_type:
|
||||
metadata['table_type'] = chunk_table_type
|
||||
|
||||
chunk_nest_level = getattr(chunk, 'table_nest_level', '')
|
||||
if chunk_nest_level:
|
||||
metadata['table_nest_level'] = str(chunk_nest_level)
|
||||
|
||||
chunk_sub_type = getattr(chunk, 'sub_type', '')
|
||||
if chunk_sub_type:
|
||||
metadata['sub_type'] = chunk_sub_type
|
||||
|
||||
# 生成向量
|
||||
try:
|
||||
embedding = embedding_model.encode(semantic_content).tolist()
|
||||
@@ -551,39 +511,8 @@ class KnowledgeBaseManager(
|
||||
curr_html = getattr(current, 'table_html', '') or ''
|
||||
next_html = getattr(next_chunk, 'table_html', '') or ''
|
||||
if curr_html and next_html:
|
||||
# 正确合并两个表格的 HTML:
|
||||
# 将第二个表格的 <tr> 行追加到第一个表格中
|
||||
# (而非简单拼接两个 <table>,否则 html_table_to_markdown
|
||||
# 的 soup.find('table') 只能找到第一个表格)
|
||||
try:
|
||||
soup1 = BeautifulSoup(curr_html, 'html.parser')
|
||||
soup2 = BeautifulSoup(next_html, 'html.parser')
|
||||
table1 = soup1.find('table')
|
||||
table2 = soup2.find('table')
|
||||
if table1 and table2:
|
||||
# 从第二个表格提取数据行
|
||||
next_rows = table2.find_all('tr')
|
||||
# 跳过与第一个表格表头重复的行
|
||||
# 对比第一行而非所有 th(find_all('th') 会匹配
|
||||
# 整个表格的 th,无法与单行做列表比较)
|
||||
first_row_t1 = table1.find('tr')
|
||||
if first_row_t1 and next_rows:
|
||||
row1_texts = [c.get_text(strip=True) for c in first_row_t1.find_all(['th', 'td'])]
|
||||
row2_texts = [c.get_text(strip=True) for c in next_rows[0].find_all(['th', 'td'])]
|
||||
if row1_texts and row2_texts and row1_texts == row2_texts:
|
||||
next_rows = next_rows[1:]
|
||||
logger.debug("跨页表格合并: 跳过了重复的表头行")
|
||||
for row in next_rows:
|
||||
table1.append(row)
|
||||
current.table_html = str(soup1)
|
||||
logger.debug(f"跨页表格 HTML 合并成功: 追加了 {len(next_rows)} 行")
|
||||
else:
|
||||
current.table_html = curr_html + '\n' + next_html
|
||||
except Exception as e:
|
||||
logger.warning(f"跨页表格 HTML 合并异常: {e},回退到简单拼接")
|
||||
current.table_html = curr_html + '\n' + next_html
|
||||
elif not curr_html and next_html:
|
||||
current.table_html = next_html
|
||||
# 合并两个表格的 HTML
|
||||
current.table_html = curr_html + '\n' + next_html
|
||||
|
||||
# 合并 image_path 和嵌入图片到 images
|
||||
curr_img = getattr(current, 'image_path', None)
|
||||
@@ -654,7 +583,7 @@ class KnowledgeBaseManager(
|
||||
try:
|
||||
from config import get_llm_client, DASHSCOPE_MODEL
|
||||
client = get_llm_client()
|
||||
summary = call_llm(client, prompt, DASHSCOPE_MODEL, max_tokens=2048)
|
||||
summary = call_llm(client, prompt, DASHSCOPE_MODEL, max_tokens=512)
|
||||
return summary.strip() if summary else ""
|
||||
except Exception as e:
|
||||
logger.warning(f"生成表格摘要失败: {e}")
|
||||
@@ -713,31 +642,13 @@ class KnowledgeBaseManager(
|
||||
]
|
||||
}
|
||||
],
|
||||
max_tokens=2048 # mimo-v2.5 推理模型思考链消耗 ~1000 token,需留足输出空间
|
||||
max_tokens=512
|
||||
)
|
||||
|
||||
description = response.choices[0].message.content
|
||||
|
||||
# 推理模型兼容:content 为空时从 reasoning_content 提取
|
||||
if not description or not description.strip():
|
||||
reasoning = getattr(response.choices[0].message, 'reasoning_content', None)
|
||||
if reasoning and reasoning.strip():
|
||||
import re
|
||||
# 尝试从思考链中提取有用文本(去掉 <think> 标签后的内容)
|
||||
cleaned = re.sub(r'', '', reasoning, flags=re.DOTALL).strip()
|
||||
if cleaned:
|
||||
logger.info(f"VLM content为空,从reasoning_content提取描述: {image_path}")
|
||||
description = cleaned
|
||||
else:
|
||||
description = reasoning.strip()
|
||||
|
||||
if not description:
|
||||
logger.warning(f"VLM 返回空描述: {image_path}")
|
||||
return ""
|
||||
|
||||
# 缓存结果
|
||||
import hashlib
|
||||
import re as _re
|
||||
img_hash = hashlib.md5(img_path.read_bytes()).hexdigest()
|
||||
cache_dir = Path('.data/cache/vlm')
|
||||
cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
@@ -664,17 +664,6 @@ class KnowledgeSyncService:
|
||||
except Exception as e:
|
||||
logger.warning(f"递增缓存版本号失败: {e}")
|
||||
|
||||
# 语义缓存无版本号机制,文档变更后必须清空,
|
||||
# 否则可能返回过时的 images/sources/citations(如已删除的图片 404)
|
||||
try:
|
||||
from core.semantic_cache import get_semantic_cache
|
||||
_sc = get_semantic_cache()
|
||||
if _sc:
|
||||
_sc.clear()
|
||||
logger.debug(f"已清空语义缓存(文档变更触发): {kb_name}")
|
||||
except Exception as e:
|
||||
logger.warning(f"清空语义缓存失败: {e}")
|
||||
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
|
||||
@@ -106,12 +106,11 @@ DEFAULT_HEADING_RULES: List[HeadingRule] = [
|
||||
name="chinese_chapter",
|
||||
),
|
||||
# 2. 中文条款编号 -> h2
|
||||
# 匹配:第一条、第三款 等(仅短标题,长正文段落不算标题)
|
||||
# 匹配:第一条、第三款 等
|
||||
HeadingRule(
|
||||
pattern=re.compile(r'^第[一二三四五六七八九十百千万]+[条款]'),
|
||||
level=2,
|
||||
name="chinese_article",
|
||||
max_length=30,
|
||||
),
|
||||
# 3. 数字三级标题 -> h3(必须在二级之前匹配)
|
||||
# 匹配:1.1.1 背景、2.3.4 方案 等
|
||||
@@ -122,23 +121,20 @@ DEFAULT_HEADING_RULES: List[HeadingRule] = [
|
||||
max_length=100,
|
||||
),
|
||||
# 4. 数字二级标题 -> h2(必须在一级之前匹配)
|
||||
# 匹配:1.1 背景、2.3 方案、2.1运行调度(无空格) 等
|
||||
# 使用负向前瞻排除三级标题(由 numeric_level3 处理)
|
||||
# 匹配:1.1 背景、2.3 方案 等
|
||||
HeadingRule(
|
||||
pattern=re.compile(r'^\d+\.\d+(?!\.\d)'),
|
||||
pattern=re.compile(r'^\d+\.\d+[\.、\s]'),
|
||||
level=2,
|
||||
name="numeric_level2",
|
||||
max_length=80,
|
||||
),
|
||||
# 5. 数字一级标题 -> h1
|
||||
# 匹配:1. 概述、2、背景 等
|
||||
# 排除:以 ;;。,、: 结尾的文本(这些是编号列表项/子条目,不是独立标题)
|
||||
HeadingRule(
|
||||
pattern=re.compile(r'^\d+[\.、\s]'),
|
||||
level=1,
|
||||
name="numeric_level1",
|
||||
max_length=50,
|
||||
exclude_pattern=re.compile(r'[;;。,、::]$'),
|
||||
),
|
||||
# 6. 英文章节标题 -> h1
|
||||
# 匹配:Chapter 1、Section 2、Part 3 等
|
||||
@@ -205,22 +201,6 @@ class HeadingRuleEngine:
|
||||
import copy
|
||||
self.rules = copy.deepcopy(DEFAULT_HEADING_RULES)
|
||||
|
||||
def _validate_level(self, level: int, text: str, rule_name=None):
|
||||
"""各级别标题长度防护:超长文本不应作为标题,降为正文。
|
||||
H1 > 40字, H2 > 60字, H3 > 50字 → 降为正文。
|
||||
统一覆盖 v1 常规匹配、v2 style 匹配、bold_short_text 兜底所有返回路径。"""
|
||||
text_len = len(text)
|
||||
if level == 1 and text_len > 40:
|
||||
logger.debug(f"标题识别: '{text[:30]}...' H1 但超长({text_len}字),降为正文")
|
||||
return 0, None
|
||||
if level == 2 and text_len > 60:
|
||||
logger.debug(f"标题识别: '{text[:30]}...' H2 但超长({text_len}字),降为正文")
|
||||
return 0, None
|
||||
if level == 3 and text_len > 50:
|
||||
logger.debug(f"标题识别: '{text[:30]}...' H3 但超长({text_len}字),降为正文")
|
||||
return 0, None
|
||||
return level, rule_name
|
||||
|
||||
def detect(self, text: str, style: Optional[List[str]] = None) -> Tuple[int, Optional[str]]:
|
||||
"""
|
||||
检测文本的标题级别
|
||||
@@ -247,7 +227,7 @@ class HeadingRuleEngine:
|
||||
level = rule.match(text)
|
||||
if level > 0:
|
||||
logger.debug(f"标题识别(v2 style): '{text[:30]}' -> h{level} (规则: {rule.name})")
|
||||
return self._validate_level(level, text, rule.name)
|
||||
return level, rule.name
|
||||
|
||||
# 再检查是否匹配中文章节/条款等高优先级规则
|
||||
for rule in self.rules:
|
||||
@@ -256,16 +236,9 @@ class HeadingRuleEngine:
|
||||
level = rule.match(text)
|
||||
if level > 0:
|
||||
logger.debug(f"标题识别(v2 style): '{text[:30]}' -> h{level} (规则: {rule.name})")
|
||||
return self._validate_level(level, text, rule.name)
|
||||
|
||||
# 最后兜底:加粗短文本 → h2
|
||||
# 但需先检查所有规则的 exclude_pattern,防止编号列表项被误判为标题
|
||||
# 例如 "3.完全满足品规。指..." 虽有 bold 样式,但属于列表项而非标题
|
||||
for rule in self.rules:
|
||||
if rule.enabled and rule.exclude_pattern and rule.exclude_pattern.search(text):
|
||||
logger.debug(f"标题识别(v2 style): '{text[:30]}' 被 {rule.name} 的 exclude_pattern 排除")
|
||||
return 0, None
|
||||
return level, rule.name
|
||||
|
||||
# 否则作为加粗短文本 → h2(与 bold_short_text 规则对齐,但不依赖 **...** 标记)
|
||||
logger.debug(f"标题识别(v2 style): '{text[:30]}' -> h2 (规则: bold_short_text_via_style)")
|
||||
return 2, 'bold_short_text'
|
||||
|
||||
@@ -274,7 +247,7 @@ class HeadingRuleEngine:
|
||||
level = rule.match(text)
|
||||
if level > 0:
|
||||
logger.debug(f"标题识别: '{text[:30]}' -> h{level} (规则: {rule.name})")
|
||||
return self._validate_level(level, text, rule.name)
|
||||
return level, rule.name
|
||||
|
||||
return 0, None
|
||||
|
||||
|
||||
@@ -131,21 +131,14 @@ class MinerUChunk:
|
||||
# 图片上下文(用于语义检索)
|
||||
context_before: str = "" # 图片前的文本上下文
|
||||
context_after: str = "" # 图片后的文本上下文
|
||||
# VLM 增强信息
|
||||
vlm_description: str = "" # VLM 视觉描述(图片/图表)
|
||||
chart_markdown: str = "" # VLM 提取的图表数据表(Markdown 格式)
|
||||
# MinerU 结构化元数据
|
||||
table_type: str = "" # 表格类型(cflow/table/text,来自 _v2_table_type)
|
||||
table_nest_level: str = "" # 表格嵌套层级(来自 _v2_table_nest_level)
|
||||
sub_type: str = "" # 图片/图表子类型(natural_image/table_image 等)
|
||||
|
||||
|
||||
def parse_with_mineru_online(
|
||||
file_path: str,
|
||||
api_token: str = None,
|
||||
api_url: str = None,
|
||||
model_version: str = None,
|
||||
timeout: int = None
|
||||
model_version: str = "vlm",
|
||||
timeout: int = 300
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
使用 MinerU 在线 API 解析文档
|
||||
@@ -164,8 +157,8 @@ def parse_with_mineru_online(
|
||||
file_path: 文档文件路径
|
||||
api_token: API Token(默认从 config 读取)
|
||||
api_url: API 地址
|
||||
model_version: 模型版本 (vlm / pipeline / MinerU-HTML),默认从 config 读取
|
||||
timeout: 轮询超时(秒),默认从 config 读取
|
||||
model_version: 模型版本 (vlm / pipeline / MinerU-HTML)
|
||||
timeout: 请求超时(秒)
|
||||
|
||||
Returns:
|
||||
解析结果(与 parse_with_mineru 格式相同)
|
||||
@@ -174,12 +167,10 @@ def parse_with_mineru_online(
|
||||
import time
|
||||
import zipfile
|
||||
import io
|
||||
from config import MINERU_API_TOKEN, MINERU_API_URL, MINERU_MODEL_VERSION, MINERU_ONLINE_TIMEOUT
|
||||
from config import MINERU_API_TOKEN, MINERU_API_URL
|
||||
|
||||
token = api_token or MINERU_API_TOKEN
|
||||
url = api_url or MINERU_API_URL
|
||||
model_version = model_version or MINERU_MODEL_VERSION
|
||||
timeout = timeout or MINERU_ONLINE_TIMEOUT
|
||||
|
||||
if not token:
|
||||
raise RuntimeError("MinerU 在线 API Token 未配置,请在 config.py 中设置 MINERU_API_TOKEN")
|
||||
@@ -250,15 +241,9 @@ def parse_with_mineru_online(
|
||||
result_resp.raise_for_status()
|
||||
result = result_resp.json()
|
||||
|
||||
# 检查 API 层面的错误码,快速失败而非静默等到超时
|
||||
api_code = result.get("code")
|
||||
if api_code and api_code != 0:
|
||||
api_msg = result.get("msg", "未知错误")
|
||||
raise RuntimeError(f"MinerU API 错误 (code={api_code}): {api_msg}")
|
||||
|
||||
extract_results = result.get("data", {}).get("extract_result", [])
|
||||
if not extract_results:
|
||||
logger.debug(f"等待解析结果... ({waited}s/{max_wait}s)")
|
||||
logger.debug(f"等待解析结果... ({waited}s)")
|
||||
continue
|
||||
|
||||
# 取第一个文件的结果
|
||||
@@ -291,16 +276,11 @@ def parse_with_mineru_online(
|
||||
if progress:
|
||||
extracted = progress.get("extracted_pages", 0)
|
||||
total = progress.get("total_pages", 0)
|
||||
logger.info(f"解析进度: {extracted}/{total} 页 ({waited}s/{max_wait}s, model={model_version})")
|
||||
logger.info(f"解析进度: {extracted}/{total} 页 ({waited}s)")
|
||||
else:
|
||||
logger.debug(f"状态: {state}, 等待中... ({waited}s/{max_wait}s)")
|
||||
logger.debug(f"状态: {state}, 等待中... ({waited}s)")
|
||||
|
||||
raise RuntimeError(
|
||||
f"MinerU 在线解析超时 ({max_wait}s)"
|
||||
f",当前 model_version={model_version}"
|
||||
f",可尝试: 1) 设置 MINERU_MODEL_VERSION=pipeline 加速"
|
||||
f" 2) 增大 MINERU_ONLINE_TIMEOUT"
|
||||
)
|
||||
raise RuntimeError("MinerU 在线解析超时")
|
||||
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"MinerU 在线 API 调用失败: {e}")
|
||||
@@ -315,26 +295,15 @@ def _parse_v2_content_list(v2_data: list) -> list:
|
||||
|
||||
转换规则:
|
||||
- paragraph → text(拼接 paragraph_content,提取 style 信息)
|
||||
- title → text(从 title_content 提取文本,level 信息)
|
||||
- table → table(保留 html,提取 table_caption/table_footnote 列表格式)
|
||||
- image → image(提取 image_source.path、VLM 视觉描述 content.content)
|
||||
- chart → chart(提取 image_source.path、VLM 数据表 content.content Markdown)
|
||||
- list → text(将 list_items 拼接为段落,保留 list_type)
|
||||
- equation → equation
|
||||
- table → table(保留 table_body/html)
|
||||
- title → text(提取 level 信息)
|
||||
- page_header / page_footer / page_number → 过滤掉
|
||||
|
||||
支持的额外字段(_v2_ 前缀):
|
||||
- _v2_styles: 样式列表(如 ['bold'])
|
||||
- _v2_table_type / _v2_table_nest_level / _v2_table_footnote: 表格元数据
|
||||
- _v2_vlm_description: VLM 生成的图片视觉描述
|
||||
- _v2_chart_markdown: VLM 从图表中提取的 Markdown 数据表
|
||||
- _v2_list_type: 列表类型(text_list 等)
|
||||
|
||||
Args:
|
||||
v2_data: v2 格式的嵌套列表
|
||||
|
||||
Returns:
|
||||
v1 兼容的扁平列表
|
||||
v1 兼容的扁平列表,额外包含 _v2_styles / _v2_table_type 字段
|
||||
"""
|
||||
flat_list: List[Dict] = []
|
||||
|
||||
@@ -348,8 +317,8 @@ def _parse_v2_content_list(v2_data: list) -> list:
|
||||
v2_type: str = item.get('type', '')
|
||||
content = item.get('content', {})
|
||||
|
||||
# 过滤噪音类型(页眉、页脚、页码、目录索引)
|
||||
if v2_type in ('page_header', 'page_footer', 'page_number', 'index'):
|
||||
# 过滤噪音类型(页眉、页脚、页码)
|
||||
if v2_type in ('page_header', 'page_footer', 'page_number'):
|
||||
continue
|
||||
|
||||
if v2_type == 'paragraph':
|
||||
@@ -381,68 +350,28 @@ def _parse_v2_content_list(v2_data: list) -> list:
|
||||
flat_list.append(flat_item)
|
||||
|
||||
elif v2_type == 'table':
|
||||
# table_caption 在 V2 中是列表格式: [{"type": "text", "content": "..."}]
|
||||
table_caption_list = content.get('table_caption', []) if isinstance(content, dict) else []
|
||||
table_caption = ''.join(
|
||||
p.get('content', '') for p in table_caption_list if isinstance(p, dict)
|
||||
).strip() if isinstance(table_caption_list, list) else str(table_caption_list)
|
||||
# 兜底: 如果 caption 列表为空,尝试旧的 caption 字符串字段
|
||||
if not table_caption:
|
||||
table_caption = content.get('caption', '') if isinstance(content, dict) else ''
|
||||
# table_footnote
|
||||
table_footnote_list = content.get('table_footnote', []) if isinstance(content, dict) else []
|
||||
table_footnote = ''.join(
|
||||
p.get('content', '') for p in table_footnote_list if isinstance(p, dict)
|
||||
).strip() if isinstance(table_footnote_list, list) else ''
|
||||
# image_source (V2 新格式) 兜底 img_path
|
||||
img_src = content.get('image_source', {}) if isinstance(content, dict) else {}
|
||||
img_path = img_src.get('path', '') if isinstance(img_src, dict) else ''
|
||||
if not img_path:
|
||||
img_path = content.get('img_path', '') if isinstance(content, dict) else ''
|
||||
|
||||
flat_item = {
|
||||
'type': 'table',
|
||||
'page_idx': page_idx,
|
||||
'bbox': item.get('bbox', []),
|
||||
'html': content.get('html', '') if isinstance(content, dict) else '',
|
||||
'table_body': content.get('html', '') if isinstance(content, dict) else '',
|
||||
'caption': table_caption,
|
||||
'img_path': img_path,
|
||||
'caption': content.get('caption', '') if isinstance(content, dict) else '',
|
||||
'_v2_table_type': content.get('table_type', '') if isinstance(content, dict) else '',
|
||||
'_v2_table_nest_level': content.get('table_nest_level', '') if isinstance(content, dict) else '',
|
||||
'_v2_table_footnote': table_footnote,
|
||||
}
|
||||
flat_list.append(flat_item)
|
||||
|
||||
elif v2_type == 'title':
|
||||
# v2 title 项: level 在 content.level,文本在 title_content(非 paragraph_content)
|
||||
vlm_level = content.get('level', 0) if isinstance(content, dict) else 0
|
||||
# v2 title 项有 level 信息
|
||||
level = content.get('level', 0) if isinstance(content, dict) else 0
|
||||
text_parts = []
|
||||
# PDF V2 使用 title_content,DOCX 理论上不应出现 title 类型
|
||||
title_content = content.get('title_content', []) if isinstance(content, dict) else []
|
||||
# 兜底: 如果 title_content 为空,尝试 paragraph_content
|
||||
if not title_content:
|
||||
title_content = content.get('paragraph_content', []) if isinstance(content, dict) else []
|
||||
for part in title_content:
|
||||
para_content = content.get('paragraph_content', []) if isinstance(content, dict) else []
|
||||
for part in para_content:
|
||||
if isinstance(part, dict):
|
||||
text_parts.append(part.get('content', ''))
|
||||
full_text = ''.join(text_parts).strip()
|
||||
|
||||
if full_text:
|
||||
# PDF V2: heading_rules 优先(模式匹配对编号标题可靠)
|
||||
# VLM level 仅作兜底(VLM 常给所有标题 level=1,不可靠)
|
||||
level = 0
|
||||
try:
|
||||
from parsers.heading_rules import get_heading_engine
|
||||
engine = get_heading_engine()
|
||||
detected_level, rule_name = engine.detect(full_text, style=['bold'])
|
||||
if detected_level > 0:
|
||||
level = detected_level
|
||||
except Exception:
|
||||
pass
|
||||
if level == 0 and vlm_level > 0:
|
||||
level = vlm_level
|
||||
|
||||
flat_item = {
|
||||
'type': 'text',
|
||||
'text': full_text,
|
||||
@@ -455,97 +384,16 @@ def _parse_v2_content_list(v2_data: list) -> list:
|
||||
flat_list.append(flat_item)
|
||||
|
||||
elif v2_type in ('image', 'chart'):
|
||||
# image_source (V2 新格式) 兜底 img_path
|
||||
img_src = content.get('image_source', {}) if isinstance(content, dict) else {}
|
||||
img_path = img_src.get('path', '') if isinstance(img_src, dict) else ''
|
||||
if not img_path:
|
||||
img_path = content.get('img_path', '') if isinstance(content, dict) else ''
|
||||
|
||||
# caption 在 V2 中可能是列表格式
|
||||
caption_raw = content.get('image_caption', content.get('caption', '')) if isinstance(content, dict) else ''
|
||||
if isinstance(caption_raw, list):
|
||||
caption = ''.join(p.get('content', '') for p in caption_raw if isinstance(p, dict)).strip()
|
||||
else:
|
||||
caption = str(caption_raw)
|
||||
|
||||
# VLM 视觉描述 (image 和 chart 项的 content.content 字段)
|
||||
vlm_description = ''
|
||||
if isinstance(content, dict):
|
||||
desc = content.get('content', '')
|
||||
if isinstance(desc, str) and desc and len(desc) > 10:
|
||||
# image 直接使用;chart 需排除 markdown 表格(表格走 chart_markdown)
|
||||
if v2_type == 'image':
|
||||
vlm_description = desc
|
||||
elif v2_type == 'chart' and '|' not in desc:
|
||||
vlm_description = desc
|
||||
|
||||
# chart 的 VLM 数据表 (content.content 字段,Markdown 表格)
|
||||
chart_markdown = ''
|
||||
if v2_type == 'chart' and isinstance(content, dict):
|
||||
md = content.get('content', '')
|
||||
if isinstance(md, str) and md:
|
||||
if '|' in md:
|
||||
chart_markdown = md
|
||||
# 即使没有 '|',如果有结构化数据特征也保留
|
||||
elif len(md) > 50 and any(kw in md for kw in ('数据', '合计', '总计', '年份', '单位')):
|
||||
chart_markdown = md
|
||||
# chart caption
|
||||
chart_caption_list = content.get('chart_caption', [])
|
||||
if isinstance(chart_caption_list, list) and chart_caption_list:
|
||||
caption = ''.join(
|
||||
p.get('content', '') for p in chart_caption_list if isinstance(p, dict)
|
||||
).strip() or caption
|
||||
|
||||
# sub_type (natural_image / table_image 等)
|
||||
sub_type = item.get('sub_type', '')
|
||||
|
||||
# 封面 logo 过滤:第一页无 caption 无 VLM 描述的图片通常是封面装饰
|
||||
if (v2_type == 'image' and page_idx == 0
|
||||
and not caption and not vlm_description and not img_path):
|
||||
continue
|
||||
|
||||
flat_item = {
|
||||
'type': v2_type,
|
||||
'page_idx': page_idx,
|
||||
'bbox': item.get('bbox', []),
|
||||
'img_path': img_path,
|
||||
'image_path': img_path,
|
||||
'caption': caption,
|
||||
'sub_type': sub_type,
|
||||
'_v2_vlm_description': vlm_description,
|
||||
'_v2_chart_markdown': chart_markdown,
|
||||
'img_path': content.get('img_path', '') if isinstance(content, dict) else '',
|
||||
'image_path': content.get('img_path', '') if isinstance(content, dict) else '',
|
||||
'caption': content.get('caption', '') if isinstance(content, dict) else '',
|
||||
}
|
||||
flat_list.append(flat_item)
|
||||
|
||||
elif v2_type == 'list':
|
||||
# 结构化列表:将列表项拼接为段落文本
|
||||
list_items = content.get('list_items', []) if isinstance(content, dict) else []
|
||||
item_texts = []
|
||||
for li in list_items:
|
||||
if not isinstance(li, dict):
|
||||
continue
|
||||
item_content = li.get('item_content', [])
|
||||
text = ''.join(
|
||||
p.get('content', '') for p in item_content if isinstance(p, dict)
|
||||
).strip()
|
||||
if text:
|
||||
item_texts.append(text)
|
||||
|
||||
if item_texts:
|
||||
list_type = content.get('list_type', 'text_list') if isinstance(content, dict) else 'text_list'
|
||||
full_text = '\n'.join(item_texts)
|
||||
flat_item = {
|
||||
'type': 'text',
|
||||
'text': full_text,
|
||||
'content': full_text,
|
||||
'page_idx': page_idx,
|
||||
'bbox': item.get('bbox', []),
|
||||
'text_level': 0,
|
||||
'_v2_styles': [],
|
||||
'_v2_list_type': list_type,
|
||||
}
|
||||
flat_list.append(flat_item)
|
||||
|
||||
elif v2_type == 'equation':
|
||||
flat_item = {
|
||||
'type': 'equation',
|
||||
@@ -558,100 +406,6 @@ def _parse_v2_content_list(v2_data: list) -> list:
|
||||
}
|
||||
flat_list.append(flat_item)
|
||||
|
||||
# ── TOC 残留过滤 ──
|
||||
# 目录条目可能以 list / title / paragraph 类型混入,用行尾页码模式检测
|
||||
# 模式1: 连续点号/省略号+页码(如 "1 综述..........1"、"2 三峡工程…4")
|
||||
# 模式2: 短行+空格+页码数字(如 "2.2 防洪 8"、"6.2 水位 28")
|
||||
import re
|
||||
_toc_dots = re.compile(r'(\.{2,}|…+|⋯+)\s*\d+\s*$')
|
||||
_toc_space_num = re.compile(r'\s{2,}\d{1,3}\s*$') # 2+空格+1-3位数字
|
||||
_toc_filtered = 0
|
||||
filtered_list = []
|
||||
for fi in flat_list:
|
||||
text = fi.get('text', '') or fi.get('content', '')
|
||||
# 只检查 text 类型(title/list/paragraph 产出),table/image/chart 不动
|
||||
if fi.get('type') == 'text' and text:
|
||||
lines = [l.strip() for l in text.split('\n') if l.strip()]
|
||||
if lines:
|
||||
toc_hits = 0
|
||||
for l in lines:
|
||||
if _toc_dots.search(l):
|
||||
toc_hits += 1
|
||||
elif _toc_space_num.search(l) and len(l) < 40:
|
||||
# 短行+空格+页码:典型的目录格式
|
||||
toc_hits += 1
|
||||
# TOC 判定逻辑:
|
||||
# 1) 匹配率 > 50% → 确定是目录
|
||||
# 2) 匹配率 > 25% 且平均行长 < 25字 → 目录(短行+页码是强信号)
|
||||
avg_line_len = sum(len(l) for l in lines) / len(lines) if lines else 0
|
||||
match_ratio = toc_hits / len(lines) if lines else 0
|
||||
is_toc = (match_ratio > 0.5) or (match_ratio > 0.25 and avg_line_len < 25)
|
||||
if is_toc:
|
||||
_toc_filtered += 1
|
||||
logger.debug(f"TOC 过滤命中({toc_hits}/{len(lines)}, avg={avg_line_len:.0f}): {text[:60]}...")
|
||||
continue
|
||||
elif toc_hits > 0:
|
||||
logger.debug(f"TOC 部分匹配({toc_hits}/{len(lines)}): {text[:60]}...")
|
||||
filtered_list.append(fi)
|
||||
if _toc_filtered > 0:
|
||||
logger.info(f"TOC 过滤第一轮: 移除 {_toc_filtered} 个目录块")
|
||||
flat_list = filtered_list
|
||||
|
||||
# 第二轮:清理孤立的 TOC 标题(子条目被过滤后残留的父标题)
|
||||
_toc_orphan = 0
|
||||
final_list = []
|
||||
for fi in flat_list:
|
||||
text = (fi.get('text', '') or fi.get('content', '')).strip()
|
||||
page = fi.get('page_idx', 0)
|
||||
if fi.get('type') == 'text' and text and page <= 5:
|
||||
# 孤立 "目录" 标题
|
||||
if text == '目录':
|
||||
_toc_orphan += 1
|
||||
continue
|
||||
# 短标题 + 尾部页码数字(如 "6 长江中下游河道状况 25")
|
||||
if (len(text) < 40 and not text.endswith('。') and not text.endswith(';')
|
||||
and re.search(r'\s+\d{1,3}\s*$', text)):
|
||||
_toc_orphan += 1
|
||||
continue
|
||||
final_list.append(fi)
|
||||
if _toc_orphan > 0:
|
||||
logger.info(f"TOC 过滤第二轮: 移除 {_toc_orphan} 个孤立目录标题")
|
||||
flat_list = final_list
|
||||
|
||||
# ── 单字符标题残留过滤 ──
|
||||
# VLM 可能部分识别目录标题(如 "目录" → "录"),单字符标题几乎不会是有效章节
|
||||
_single_char = 0
|
||||
_sc_list = []
|
||||
for fi in flat_list:
|
||||
text = (fi.get('text', '') or fi.get('content', '')).strip()
|
||||
if fi.get('type') == 'text' and text and len(text) <= 1 and fi.get('text_level', 0) > 0:
|
||||
_single_char += 1
|
||||
logger.debug(f"单字符标题过滤: '{text}' pg={fi.get('page_idx', 0)}")
|
||||
continue
|
||||
_sc_list.append(fi)
|
||||
if _single_char > 0:
|
||||
logger.info(f"单字符标题过滤: 移除 {_single_char} 个残留")
|
||||
flat_list = _sc_list
|
||||
|
||||
# ── 封面重复标题去重 ──
|
||||
# pg=0(封面)和 pg=1(扉页)常有完全相同的标题(如 "三峡工程公报"、"2022"),保留较后的
|
||||
_cover_seen = {} # text -> page_idx
|
||||
_cover_dup = 0
|
||||
dedup_list = []
|
||||
for fi in flat_list:
|
||||
text = (fi.get('text', '') or fi.get('content', '')).strip()
|
||||
page = fi.get('page_idx', 0)
|
||||
if fi.get('type') == 'text' and text and page <= 1 and fi.get('text_level', 0) > 0:
|
||||
if text in _cover_seen:
|
||||
_cover_dup += 1
|
||||
logger.debug(f"封面重复标题过滤: '{text}' pg={page} (首次 pg={_cover_seen[text]})")
|
||||
continue
|
||||
_cover_seen[text] = page
|
||||
dedup_list.append(fi)
|
||||
if _cover_dup > 0:
|
||||
logger.info(f"封面去重: 移除 {_cover_dup} 个重复封面标题")
|
||||
flat_list = dedup_list
|
||||
|
||||
logger.info(f"v2 格式转换: {len(v2_data)} 页 → {len(flat_list)} 项(已过滤噪音类型)")
|
||||
return flat_list
|
||||
|
||||
@@ -861,7 +615,6 @@ def _parse_mineru_online_result(result: Dict, file_path: Path) -> Dict[str, Any]
|
||||
elif item_type == "table":
|
||||
table_body = item.get("html", "") or item.get("table_body", "")
|
||||
table_caption = item.get("caption", "") or item.get("table_caption", "")
|
||||
table_footnote = item.get("_v2_table_footnote", "")
|
||||
img_path = item.get("img_path", "") or item.get("image_path", "")
|
||||
|
||||
section_path = " > ".join([s[1] for s in section_stack])
|
||||
@@ -876,13 +629,8 @@ def _parse_mineru_online_result(result: Dict, file_path: Path) -> Dict[str, Any]
|
||||
# 从原始 HTML 提取嵌入图片(md_table 经 get_text 转换后已丢失 <img> 标签)
|
||||
table_images = extract_images_from_markdown(table_body) if table_body else []
|
||||
|
||||
# 表格内容增强:caption + footnote
|
||||
table_content = table_caption or "表格"
|
||||
if table_footnote:
|
||||
table_content = f"{table_content}\n[脚注] {table_footnote}"
|
||||
|
||||
chunk = MinerUChunk(
|
||||
content=table_content,
|
||||
content=table_caption or "表格",
|
||||
chunk_type="table",
|
||||
page_start=page_idx + 1,
|
||||
page_end=page_idx + 1,
|
||||
@@ -892,9 +640,7 @@ def _parse_mineru_online_result(result: Dict, file_path: Path) -> Dict[str, Any]
|
||||
source_file=file_path.name,
|
||||
table_html=table_body,
|
||||
image_path=img_path,
|
||||
images=table_images if table_images else None,
|
||||
table_type=item.get("_v2_table_type", ""),
|
||||
table_nest_level=item.get("_v2_table_nest_level", ""),
|
||||
images=table_images if table_images else None
|
||||
)
|
||||
chunks.append(chunk)
|
||||
if table_body:
|
||||
@@ -905,29 +651,16 @@ def _parse_mineru_online_result(result: Dict, file_path: Path) -> Dict[str, Any]
|
||||
elif item_type in ("image", "chart"):
|
||||
img_path = item.get("img_path", "") or item.get("image_path", "")
|
||||
caption = item.get("caption", "")
|
||||
vlm_desc = item.get("_v2_vlm_description", "")
|
||||
chart_md = item.get("_v2_chart_markdown", "")
|
||||
sub_type = item.get("sub_type", "")
|
||||
|
||||
section_path = " > ".join([s[1] for s in section_stack])
|
||||
|
||||
markdown_parts.append(f"\n")
|
||||
# chart 的 VLM 数据表也写入 markdown 输出
|
||||
if chart_md:
|
||||
markdown_parts.append(chart_md)
|
||||
|
||||
chunk_type = "chart" if item_type == "chart" else "image"
|
||||
context_before, context_after = get_context_for_image(idx, page_idx)
|
||||
|
||||
# 图片内容增强:caption + VLM 描述
|
||||
content_text = caption or ("图表" if item_type == "chart" else "图片")
|
||||
if vlm_desc:
|
||||
content_text = f"{content_text}\n[视觉描述] {vlm_desc}"
|
||||
if chart_md:
|
||||
content_text = f"{content_text}\n[数据表]\n{chart_md}"
|
||||
|
||||
chunk = MinerUChunk(
|
||||
content=content_text,
|
||||
content=caption or ("图表" if item_type == "chart" else "图片"),
|
||||
chunk_type=chunk_type,
|
||||
page_start=page_idx + 1,
|
||||
page_end=page_idx + 1,
|
||||
@@ -937,18 +670,11 @@ def _parse_mineru_online_result(result: Dict, file_path: Path) -> Dict[str, Any]
|
||||
source_file=file_path.name,
|
||||
image_path=img_path,
|
||||
context_before=context_before,
|
||||
context_after=context_after,
|
||||
vlm_description=vlm_desc,
|
||||
chart_markdown=chart_md,
|
||||
table_html=chart_md if chart_md else None, # chart 数据表作为表格存储
|
||||
sub_type=sub_type,
|
||||
context_after=context_after
|
||||
)
|
||||
chunks.append(chunk)
|
||||
if img_path:
|
||||
images.append(img_path)
|
||||
# chart 数据表也加入 tables 列表,便于表格检索
|
||||
if chart_md:
|
||||
tables.append(chart_md)
|
||||
|
||||
elif item_type == "equation":
|
||||
# 处理公式类型
|
||||
@@ -1011,7 +737,7 @@ def parse_with_mineru(
|
||||
lang: str = "ch",
|
||||
enable_table: bool = True,
|
||||
enable_formula: bool = True,
|
||||
backend: str = None,
|
||||
backend: str = "pipeline",
|
||||
start_page: int = 0,
|
||||
end_page: int = 99999
|
||||
) -> Dict[str, Any]:
|
||||
@@ -1044,14 +770,6 @@ def parse_with_mineru(
|
||||
if not file_path.exists():
|
||||
raise FileNotFoundError(f"文件不存在: {file_path}")
|
||||
|
||||
# 从 config 读取默认 backend
|
||||
if backend is None:
|
||||
try:
|
||||
from config import MINERU_LOCAL_BACKEND
|
||||
backend = MINERU_LOCAL_BACKEND
|
||||
except ImportError:
|
||||
backend = 'pipeline'
|
||||
|
||||
# 检查文件大小
|
||||
file_size = file_path.stat().st_size
|
||||
if file_size > MAX_PDF_SIZE:
|
||||
@@ -1096,6 +814,7 @@ def parse_with_mineru(
|
||||
|
||||
cmd = [
|
||||
str(mineru_exe),
|
||||
"--",
|
||||
"-p", str(file_path),
|
||||
"-o", str(output_dir),
|
||||
"-m", "auto",
|
||||
@@ -1140,7 +859,7 @@ def parse_with_mineru(
|
||||
logger.error(f"MinerU 解析失败: {e}")
|
||||
raise
|
||||
finally:
|
||||
# 清理临时目录
|
||||
清理临时目录
|
||||
if cleanup_output and os.path.exists(output_dir):
|
||||
shutil.rmtree(output_dir, ignore_errors=True)
|
||||
|
||||
@@ -1409,7 +1128,6 @@ def _parse_mineru_output(file_path: Path, output_dir) -> Dict[str, Any]:
|
||||
elif item_type == "table":
|
||||
table_body = item.get("table_body", "")
|
||||
table_caption = item.get("table_caption", "")
|
||||
table_footnote = item.get("_v2_table_footnote", "")
|
||||
# 表格也可能有图片形式(img_path)
|
||||
img_path = item.get("img_path", "")
|
||||
|
||||
@@ -1425,13 +1143,8 @@ def _parse_mineru_output(file_path: Path, output_dir) -> Dict[str, Any]:
|
||||
# 从原始 HTML 提取嵌入图片(md_table 经 get_text 转换后已丢失 <img> 标签)
|
||||
table_images = extract_images_from_markdown(table_body) if table_body else []
|
||||
|
||||
# 表格内容增强:caption + footnote
|
||||
table_content = table_caption or "表格"
|
||||
if table_footnote:
|
||||
table_content = f"{table_content}\n[脚注] {table_footnote}"
|
||||
|
||||
chunk = MinerUChunk(
|
||||
content=table_content,
|
||||
content=table_caption or "表格",
|
||||
chunk_type="table",
|
||||
page_start=page_idx + 1,
|
||||
page_end=page_idx + 1,
|
||||
@@ -1441,9 +1154,7 @@ def _parse_mineru_output(file_path: Path, output_dir) -> Dict[str, Any]:
|
||||
source_file=file_path.name,
|
||||
table_html=table_body,
|
||||
image_path=img_path, # 表格的独立图片形式
|
||||
images=table_images if table_images else None, # 嵌入图片列表
|
||||
table_type=item.get("_v2_table_type", ""),
|
||||
table_nest_level=item.get("_v2_table_nest_level", ""),
|
||||
images=table_images if table_images else None # 嵌入图片列表
|
||||
)
|
||||
chunks.append(chunk)
|
||||
if table_body:
|
||||
@@ -1456,15 +1167,10 @@ def _parse_mineru_output(file_path: Path, output_dir) -> Dict[str, Any]:
|
||||
# 处理图片和图表类型(MinerU 将图表识别为 chart 类型)
|
||||
img_path = item.get("img_path", "")
|
||||
caption = item.get("caption", "")
|
||||
vlm_desc = item.get("_v2_vlm_description", "")
|
||||
chart_md = item.get("_v2_chart_markdown", "")
|
||||
sub_type = item.get("sub_type", "")
|
||||
|
||||
section_path = " > ".join([s[1] for s in section_stack])
|
||||
|
||||
markdown_parts.append(f"\n")
|
||||
if chart_md:
|
||||
markdown_parts.append(chart_md)
|
||||
|
||||
# 图表类型标记为 chart,便于后续区分处理
|
||||
chunk_type = "chart" if item_type == "chart" else "image"
|
||||
@@ -1472,15 +1178,8 @@ def _parse_mineru_output(file_path: Path, output_dir) -> Dict[str, Any]:
|
||||
# 获取图片上下文
|
||||
context_before, context_after = get_context_for_image(idx, page_idx)
|
||||
|
||||
# 图片内容增强:caption + VLM 描述
|
||||
content_text = caption or ("图表" if item_type == "chart" else "图片")
|
||||
if vlm_desc:
|
||||
content_text = f"{content_text}\n[视觉描述] {vlm_desc}"
|
||||
if chart_md:
|
||||
content_text = f"{content_text}\n[数据表]\n{chart_md}"
|
||||
|
||||
chunk = MinerUChunk(
|
||||
content=content_text,
|
||||
content=caption or ("图表" if item_type == "chart" else "图片"),
|
||||
chunk_type=chunk_type,
|
||||
page_start=page_idx + 1,
|
||||
page_end=page_idx + 1,
|
||||
@@ -1490,17 +1189,11 @@ def _parse_mineru_output(file_path: Path, output_dir) -> Dict[str, Any]:
|
||||
source_file=file_path.name,
|
||||
image_path=img_path,
|
||||
context_before=context_before,
|
||||
context_after=context_after,
|
||||
vlm_description=vlm_desc,
|
||||
chart_markdown=chart_md,
|
||||
table_html=chart_md if chart_md else None,
|
||||
sub_type=sub_type,
|
||||
context_after=context_after
|
||||
)
|
||||
chunks.append(chunk)
|
||||
if img_path:
|
||||
images.append(img_path)
|
||||
if chart_md:
|
||||
tables.append(chart_md)
|
||||
|
||||
elif item_type == "equation":
|
||||
# 处理公式类型
|
||||
@@ -1609,7 +1302,6 @@ def _post_process_chunks(
|
||||
# Phase 2: 合并碎片
|
||||
merged = []
|
||||
buffer = None # 当前合并缓冲
|
||||
_buffer_has_body = False # 缓冲是否已包含正文(防止标题继续合并)
|
||||
|
||||
for chunk in chunks:
|
||||
# 表格、图片和图表不参与合并,直接输出
|
||||
@@ -1617,34 +1309,13 @@ def _post_process_chunks(
|
||||
if buffer:
|
||||
merged.append(buffer)
|
||||
buffer = None
|
||||
_buffer_has_body = False
|
||||
merged.append(chunk)
|
||||
continue
|
||||
|
||||
# 标题 chunk(text_level > 0)
|
||||
# 标题 chunk(text_level > 0),开始新的合并组
|
||||
if chunk.text_level > 0:
|
||||
if buffer:
|
||||
# H1 级标题是章节边界,强制断开,不参与连续标题链合并
|
||||
if chunk.text_level == 1:
|
||||
merged.append(buffer)
|
||||
_buffer_has_body = False
|
||||
# 连续标题链合并:缓冲也是纯标题(无正文)时,合并而非刷新(仅非 H1)
|
||||
elif buffer.text_level > 0 and not _buffer_has_body:
|
||||
combined = buffer.content.rstrip() + '\n' + chunk.content
|
||||
if len(combined) <= max_merged_size:
|
||||
buffer.content = combined
|
||||
buffer.page_end = chunk.page_end
|
||||
# 取更高层级(数值更小)
|
||||
buffer.text_level = min(buffer.text_level, chunk.text_level)
|
||||
# section_path 保留第一个(更高级别)的
|
||||
continue
|
||||
# 合并后超限,刷新缓冲
|
||||
merged.append(buffer)
|
||||
_buffer_has_body = False
|
||||
else:
|
||||
# 缓冲是正文,正常刷新
|
||||
merged.append(buffer)
|
||||
_buffer_has_body = False
|
||||
merged.append(buffer)
|
||||
# 标题作为新缓冲的起点
|
||||
buffer = MinerUChunk(
|
||||
content=chunk.content,
|
||||
@@ -1657,7 +1328,6 @@ def _post_process_chunks(
|
||||
bbox=chunk.bbox,
|
||||
source_file=chunk.source_file,
|
||||
)
|
||||
_buffer_has_body = False
|
||||
continue
|
||||
|
||||
# 正文 chunk
|
||||
@@ -1677,7 +1347,6 @@ def _post_process_chunks(
|
||||
bbox=chunk.bbox,
|
||||
source_file=chunk.source_file,
|
||||
)
|
||||
_buffer_has_body = True # 正文 chunk 创建的缓冲已含正文
|
||||
else:
|
||||
# 足够长,直接输出
|
||||
merged.append(chunk)
|
||||
@@ -1688,9 +1357,6 @@ def _post_process_chunks(
|
||||
# 合并
|
||||
buffer.content = buffer.content.rstrip() + '\n' + chunk.content
|
||||
buffer.page_end = chunk.page_end
|
||||
# 正文并入标题缓冲后,标记已含正文,防止后续标题继续合并
|
||||
# 保留 text_level 不置零,使标题层级信息传递到向量库
|
||||
_buffer_has_body = True
|
||||
else:
|
||||
# 超过上限,输出缓冲,当前 chunk 开始新缓冲或直接输出
|
||||
merged.append(buffer)
|
||||
@@ -1706,10 +1372,8 @@ def _post_process_chunks(
|
||||
bbox=chunk.bbox,
|
||||
source_file=chunk.source_file,
|
||||
)
|
||||
_buffer_has_body = True # 正文 chunk 创建的缓冲已含正文
|
||||
else:
|
||||
buffer = None
|
||||
_buffer_has_body = False
|
||||
merged.append(chunk)
|
||||
|
||||
# 刷新最后的缓冲
|
||||
@@ -2024,7 +1688,7 @@ def parse_with_mineru_persistent(
|
||||
lang: str = "ch",
|
||||
enable_table: bool = True,
|
||||
enable_formula: bool = True,
|
||||
backend: str = None,
|
||||
backend: str = "pipeline",
|
||||
start_page: int = 0,
|
||||
end_page: int = 99999,
|
||||
cleanup_after_image_move: bool = True
|
||||
@@ -2063,14 +1727,6 @@ def parse_with_mineru_persistent(
|
||||
if not file_path.exists():
|
||||
raise FileNotFoundError(f"文件不存在: {file_path}")
|
||||
|
||||
# 从 config 读取默认 backend
|
||||
if backend is None:
|
||||
try:
|
||||
from config import MINERU_LOCAL_BACKEND
|
||||
backend = MINERU_LOCAL_BACKEND
|
||||
except ImportError:
|
||||
backend = 'pipeline'
|
||||
|
||||
# 检查文件大小
|
||||
file_size = file_path.stat().st_size
|
||||
if file_size > MAX_PDF_SIZE:
|
||||
|
||||
Reference in New Issue
Block a user