Compare commits
8 Commits
server-rel
...
db887d2215
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
db887d2215 | ||
|
|
3c262b8d35 | ||
|
|
8adb550931 | ||
|
|
9c1593a5e4 | ||
|
|
83884b383b | ||
|
|
183a57e7f1 | ||
|
|
54a6815ad4 | ||
|
|
505f79860e |
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
|
||||
@@ -145,6 +132,10 @@ def create_app() -> 'Flask':
|
||||
from api.image_routes import image_bp
|
||||
app.register_blueprint(image_bp)
|
||||
|
||||
# 异步任务查询
|
||||
from api.task_routes import task_bp
|
||||
app.register_blueprint(task_bp)
|
||||
|
||||
# 健康检查
|
||||
from api.auth_routes import auth_bp
|
||||
app.register_blueprint(auth_bp)
|
||||
@@ -267,4 +258,5 @@ def _print_startup_info(app: 'Flask') -> None:
|
||||
logger.info(" 切片管理: /chunks/*")
|
||||
logger.info(" 同步服务: /sync, /sync/status")
|
||||
logger.info(" 图片服务: /images/*")
|
||||
logger.info(" 任务查询: /tasks, /tasks/<id>, /tasks/<id>/progress")
|
||||
logger.info(" 健康检查: /health")
|
||||
|
||||
@@ -11,6 +11,8 @@ import logging
|
||||
logger = logging.getLogger(__name__)
|
||||
from auth.gateway import require_gateway_auth
|
||||
from data.db import get_connection
|
||||
from core.status_codes import SUCCESS, INTERNAL_ERROR
|
||||
from api.response_utils import success_response, error_response
|
||||
|
||||
audit_bp = Blueprint('audit', __name__)
|
||||
|
||||
@@ -94,11 +96,11 @@ def get_audit_logs():
|
||||
"timestamp": row[10]
|
||||
})
|
||||
|
||||
return jsonify({"logs": logs, "total": total})
|
||||
return success_response(data={"logs": logs, "total": total})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"审计查询异常: {e}")
|
||||
return jsonify({"error": "查询失败", "logs": [], "total": 0}), 500
|
||||
return error_response("QUERY_FAILED", INTERNAL_ERROR, "查询失败", http_status=500)
|
||||
|
||||
|
||||
def log_audit_event(user_id: str, username: str, action: str,
|
||||
|
||||
@@ -10,7 +10,10 @@
|
||||
|
||||
from flask import Blueprint, request, jsonify
|
||||
from auth.gateway import require_gateway_auth, require_role, get_user_permissions, MOCK_USERS
|
||||
import os
|
||||
from core.status_codes import SUCCESS, BAD_REQUEST, UNAUTHORIZED, FORBIDDEN, PERMISSION_DENIED, INTERNAL_ERROR
|
||||
from api.response_utils import success_response, error_response
|
||||
from config import DEV_MODE
|
||||
import time
|
||||
from pathlib import Path
|
||||
from dotenv import load_dotenv
|
||||
|
||||
@@ -21,6 +24,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:
|
||||
"""统一的开发模式判断(由 config.py 集中管理)"""
|
||||
return DEV_MODE
|
||||
|
||||
|
||||
@auth_bp.route('/auth/login', methods=['POST'])
|
||||
def mock_login():
|
||||
"""
|
||||
@@ -50,9 +84,17 @@ def mock_login():
|
||||
- manager / manager123 (经理,财务部)
|
||||
- user / test123 (普通用户,技术部)
|
||||
"""
|
||||
# 默认开启开发模式(生产环境需设置 DEV_MODE=false)
|
||||
if os.environ.get('DEV_MODE', 'true').lower() == 'false':
|
||||
return jsonify({"error": "仅开发环境可用,请设置 DEV_MODE=true"}), 403
|
||||
if not _is_dev_mode():
|
||||
return error_response("FORBIDDEN", FORBIDDEN, "仅开发环境可用,请在 .env 中设置 DEV_MODE=true", http_status=403)
|
||||
|
||||
# 速率限制检查
|
||||
client_ip = request.remote_addr or 'unknown'
|
||||
if not _check_rate_limit(client_ip):
|
||||
return error_response(
|
||||
"RATE_LIMITED", FORBIDDEN,
|
||||
f"登录尝试过于频繁,请 {_RATE_LIMIT_WINDOW // 60} 分钟后再试",
|
||||
http_status=429
|
||||
)
|
||||
|
||||
data = request.json or {}
|
||||
username = data.get('username')
|
||||
@@ -60,9 +102,9 @@ def mock_login():
|
||||
|
||||
user = MOCK_USERS.get(username)
|
||||
if not user or user['password'] != password:
|
||||
return jsonify({"error": "用户名或密码错误"}), 401
|
||||
return error_response("UNAUTHORIZED", UNAUTHORIZED, "用户名或密码错误", http_status=401)
|
||||
|
||||
return jsonify({
|
||||
return success_response(data={
|
||||
"token": f"mock-token-{username}",
|
||||
"user": {
|
||||
"user_id": user['user_id'],
|
||||
@@ -79,8 +121,10 @@ def mock_login():
|
||||
def get_stats():
|
||||
"""获取系统统计信息(仅管理员)"""
|
||||
from flask import current_app
|
||||
session_manager = current_app.config['SESSION_MANAGER']
|
||||
return jsonify(session_manager.get_stats())
|
||||
session_manager = current_app.config.get('SESSION_MANAGER')
|
||||
if not session_manager:
|
||||
return error_response("UNAVAILABLE", INTERNAL_ERROR, "会话管理器未启用", http_status=503)
|
||||
return success_response(data=session_manager.get_stats())
|
||||
|
||||
|
||||
@auth_bp.route('/health', methods=['GET'])
|
||||
@@ -103,7 +147,7 @@ def get_current_user():
|
||||
开发模式下支持模拟用户,生产模式下用户信息由后端控制。
|
||||
"""
|
||||
user = request.current_user
|
||||
return jsonify({
|
||||
return success_response(data={
|
||||
"user_id": user["user_id"],
|
||||
"username": user["username"],
|
||||
"role": user["role"],
|
||||
@@ -131,9 +175,8 @@ def get_users():
|
||||
]
|
||||
}
|
||||
"""
|
||||
dev_mode = os.environ.get('DEV_MODE', 'true').lower() != 'false'
|
||||
if not dev_mode:
|
||||
return jsonify({"error": "仅开发环境可用"}), 403
|
||||
if not _is_dev_mode():
|
||||
return error_response("FORBIDDEN", FORBIDDEN, "仅开发环境可用", http_status=403)
|
||||
|
||||
users = []
|
||||
for username, info in MOCK_USERS.items():
|
||||
@@ -145,7 +188,7 @@ def get_users():
|
||||
"is_active": True # 模拟用户默认都是活跃状态
|
||||
})
|
||||
|
||||
return jsonify({"users": users})
|
||||
return success_response(data={"users": users})
|
||||
|
||||
|
||||
@auth_bp.route('/auth/users/<user_id>', methods=['PUT'])
|
||||
@@ -159,12 +202,26 @@ def update_user(user_id):
|
||||
"is_active": false
|
||||
}
|
||||
"""
|
||||
dev_mode = os.environ.get('DEV_MODE', 'true').lower() != 'false'
|
||||
if not dev_mode:
|
||||
return jsonify({"error": "仅开发环境可用"}), 403
|
||||
if not _is_dev_mode():
|
||||
return error_response("FORBIDDEN", FORBIDDEN, "仅开发环境可用", http_status=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 error_response("NOT_FOUND", BAD_REQUEST, f"用户 {user_id} 不存在", http_status=404)
|
||||
|
||||
data = request.json or {}
|
||||
# 模拟操作:记录请求但不实际执行(mock 用户数据是静态的)
|
||||
return success_response(data={
|
||||
"message": "操作成功(模拟)",
|
||||
"user_id": user_id,
|
||||
"applied_changes": data
|
||||
})
|
||||
|
||||
|
||||
@auth_bp.route('/auth/change-password', methods=['POST'])
|
||||
@@ -179,19 +236,25 @@ def change_password():
|
||||
"new_password": "xxx"
|
||||
}
|
||||
"""
|
||||
dev_mode = os.environ.get('DEV_MODE', 'true').lower() != 'false'
|
||||
if not dev_mode:
|
||||
return jsonify({"error": "仅开发环境可用"}), 403
|
||||
if not _is_dev_mode():
|
||||
return error_response("FORBIDDEN", FORBIDDEN, "仅开发环境可用", http_status=403)
|
||||
|
||||
data = request.json or {}
|
||||
old_password = data.get('old_password')
|
||||
new_password = data.get('new_password')
|
||||
|
||||
if not old_password or not new_password:
|
||||
return jsonify({"error": "请提供旧密码和新密码"}), 400
|
||||
return error_response("MISSING_PARAMS", BAD_REQUEST, "请提供旧密码和新密码", http_status=400)
|
||||
|
||||
if len(new_password) < 6:
|
||||
return jsonify({"error": "新密码至少6位"}), 400
|
||||
return error_response("INVALID_PARAMS", BAD_REQUEST, "新密码至少6位", http_status=400)
|
||||
|
||||
# 模拟环境直接返回成功
|
||||
return jsonify({"message": "密码修改成功(模拟)"})
|
||||
# 验证当前用户的旧密码
|
||||
user = request.current_user
|
||||
username = user.get('username', '')
|
||||
mock_user = MOCK_USERS.get(username)
|
||||
if mock_user and mock_user['password'] != old_password:
|
||||
return error_response("UNAUTHORIZED", UNAUTHORIZED, "旧密码错误", http_status=401)
|
||||
|
||||
# 模拟环境返回成功(不实际修改密码,mock 数据是静态的)
|
||||
return success_response(message="密码修改成功(模拟)")
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -39,16 +39,19 @@ import re
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import Optional, Tuple, Any, List, Dict
|
||||
from flask import Blueprint, request, jsonify
|
||||
from flask import Blueprint, request
|
||||
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,
|
||||
FILE_TOO_LARGE, UNSUPPORTED_FORMAT, INTERNAL_ERROR
|
||||
UPLOAD_SUCCESS, BATCH_UPLOAD_SUCCESS, BAD_REQUEST, SUCCESS,
|
||||
NO_FILE, NO_FILE_SELECTED, NO_COLLECTION, NOT_FOUND,
|
||||
FILE_TOO_LARGE, UNSUPPORTED_FORMAT, INTERNAL_ERROR,
|
||||
SERVICE_UNAVAILABLE, DELETE_SUCCESS, UPDATE_SUCCESS, NO_CONTENT,
|
||||
FORBIDDEN
|
||||
)
|
||||
from api.response_utils import success_response, error_response
|
||||
|
||||
@@ -169,10 +172,10 @@ 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':
|
||||
return jsonify({"error": "仅开发环境可用"}), 403
|
||||
if not DEV_MODE:
|
||||
return error_response("FORBIDDEN", FORBIDDEN, "仅开发环境可用", http_status=403)
|
||||
|
||||
from config import DOCUMENTS_PATH
|
||||
from flask import send_from_directory
|
||||
@@ -181,9 +184,9 @@ def serve_document_file(doc_path: str) -> Tuple[Any, int]:
|
||||
try:
|
||||
filepath = _validate_doc_path(doc_path, DOCUMENTS_PATH)
|
||||
except ValueError:
|
||||
return jsonify({"error": "非法路径"}), 403
|
||||
return error_response("FORBIDDEN", FORBIDDEN, "非法路径", http_status=403)
|
||||
if not os.path.exists(filepath):
|
||||
return jsonify({"error": "文件不存在"}), 404
|
||||
return error_response("NOT_FOUND", NOT_FOUND, "文件不存在", http_status=404)
|
||||
|
||||
directory = os.path.dirname(filepath)
|
||||
filename = os.path.basename(filepath)
|
||||
@@ -334,10 +337,24 @@ def upload_document() -> Tuple[Any, int]:
|
||||
except Exception as e:
|
||||
logger.warning(f"标记旧版本失败: {e}")
|
||||
|
||||
# 6. 触发向量化
|
||||
# 6. 触发向量化(异步任务)
|
||||
sync_status = "已保存,等待手动同步"
|
||||
sync_service = _get_sync_service()
|
||||
task_id = None
|
||||
|
||||
if sync_service:
|
||||
from core.task_registry import get_registry
|
||||
registry = get_registry()
|
||||
|
||||
task = registry.create_task('upload', f"向量化: {filename}")
|
||||
task_id = task.id
|
||||
|
||||
def _do_vectorize(task, svc, change_obj, fname):
|
||||
"""后台执行向量化"""
|
||||
registry.update_progress(task.id, stage='向量化', message=f"正在处理: {fname}")
|
||||
svc.process_change(change_obj)
|
||||
return {'filename': fname, 'sync_status': '已添加到向量库'}
|
||||
|
||||
try:
|
||||
from knowledge.sync import DocumentChange, ChangeType
|
||||
change = DocumentChange(
|
||||
@@ -348,11 +365,12 @@ def upload_document() -> Tuple[Any, int]:
|
||||
new_hash=sync_service.calculate_file_hash(filepath),
|
||||
change_time=datetime.now()
|
||||
)
|
||||
sync_service.process_change(change)
|
||||
sync_status = "已保存并添加到向量库"
|
||||
registry.start_task(task.id, _do_vectorize, sync_service, change, filename)
|
||||
sync_status = "已保存,向量化任务已启动"
|
||||
except Exception as e:
|
||||
logger.warning(f"向量化失败: {e}")
|
||||
sync_status = "已保存,向量化失败"
|
||||
logger.warning(f"创建向量化任务失败: {e}")
|
||||
registry.fail_task(task.id, str(e))
|
||||
sync_status = "已保存,向量化任务创建失败"
|
||||
|
||||
return success_response(
|
||||
data={
|
||||
@@ -363,7 +381,8 @@ def upload_document() -> Tuple[Any, int]:
|
||||
"size": file_size,
|
||||
"replaced": replaced
|
||||
},
|
||||
"sync_status": sync_status
|
||||
"sync_status": sync_status,
|
||||
"task_id": task_id
|
||||
},
|
||||
status_code=UPLOAD_SUCCESS,
|
||||
message=f"文件上传成功,{sync_status}"
|
||||
@@ -506,14 +525,55 @@ def batch_upload_documents() -> Tuple[Any, int]:
|
||||
"message": "上传处理失败"
|
||||
})
|
||||
|
||||
success_count = len([r for r in results if r["status"] == "success"])
|
||||
task_id = None
|
||||
|
||||
# 批量上传完成后自动触发向量化(异步任务)
|
||||
sync_service = _get_sync_service()
|
||||
if sync_service and success_count > 0:
|
||||
from core.task_registry import get_registry
|
||||
registry = get_registry()
|
||||
|
||||
running = registry.list_tasks(status='running', task_type='batch_upload', limit=1)
|
||||
if not running:
|
||||
task = registry.create_task('batch_upload', f"批量向量化: {success_count} 个文件", total=success_count)
|
||||
task_id = task.id
|
||||
|
||||
def _do_batch_vectorize(task, svc, count):
|
||||
processed = [0]
|
||||
|
||||
def on_change(change):
|
||||
processed[0] += 1
|
||||
registry.update_progress(
|
||||
task.id, current=processed[0], total=count,
|
||||
stage='批量向量化',
|
||||
message=f"已处理: {change.document_name if hasattr(change, 'document_name') else change.document_id}"
|
||||
)
|
||||
|
||||
old_cb = svc.on_change_callback
|
||||
svc.on_change_callback = on_change
|
||||
try:
|
||||
registry.update_progress(task.id, stage='扫描文档', message='正在检测变更...')
|
||||
result = svc.sync_now()
|
||||
return {
|
||||
'synced': result.documents_processed,
|
||||
'added': result.documents_added,
|
||||
'errors': result.errors,
|
||||
}
|
||||
finally:
|
||||
svc.on_change_callback = old_cb
|
||||
|
||||
registry.start_task(task.id, _do_batch_vectorize, sync_service, success_count)
|
||||
|
||||
return success_response(
|
||||
data={
|
||||
"total": len(results),
|
||||
"success_count": len([r for r in results if r["status"] == "success"]),
|
||||
"results": results
|
||||
"success_count": success_count,
|
||||
"results": results,
|
||||
"task_id": task_id
|
||||
},
|
||||
status_code=BATCH_UPLOAD_SUCCESS,
|
||||
message=f"批量上传完成,成功 {len([r for r in results if r['status'] == 'success'])}/{len(results)} 个文件"
|
||||
message=f"批量上传完成,成功 {success_count}/{len(results)} 个文件"
|
||||
)
|
||||
|
||||
|
||||
@@ -590,10 +650,7 @@ def list_documents() -> Tuple[Any, int]:
|
||||
# 按修改时间倒序
|
||||
documents.sort(key=lambda x: x['last_modified'], reverse=True)
|
||||
|
||||
return jsonify({
|
||||
"documents": documents,
|
||||
"total": len(documents)
|
||||
})
|
||||
return success_response(data={"documents": documents, "total": len(documents)})
|
||||
|
||||
|
||||
@document_bp.route('/documents/<path:doc_path>/status', methods=['GET'])
|
||||
@@ -617,12 +674,12 @@ def get_document_status(doc_path: str) -> Tuple[Any, int]:
|
||||
"""
|
||||
kb_manager = _get_kb_manager()
|
||||
if not kb_manager:
|
||||
return jsonify({"error": "知识库管理器未初始化"}), 503
|
||||
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "知识库管理器未初始化", http_status=503)
|
||||
|
||||
# 解析路径
|
||||
parts = doc_path.split('/')
|
||||
if len(parts) < 2:
|
||||
return jsonify({"error": "无效的文档路径"}), 400
|
||||
return error_response("BAD_REQUEST", BAD_REQUEST, "无效的文档路径", http_status=400)
|
||||
|
||||
subdir = parts[0]
|
||||
filename = '/'.join(parts[1:])
|
||||
@@ -638,16 +695,14 @@ def get_document_status(doc_path: str) -> Tuple[Any, int]:
|
||||
|
||||
if not doc_info:
|
||||
if file_on_disk:
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"status": "unprocessed",
|
||||
"chunk_count": 0,
|
||||
"last_processed": None
|
||||
})
|
||||
return jsonify({"error": "文档不存在"}), 404
|
||||
return error_response("NOT_FOUND", NOT_FOUND, "文档不存在", http_status=404)
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"status": doc_info.get("status", "unknown"),
|
||||
"chunk_count": doc_info.get("total_chunks", 0),
|
||||
"last_processed": doc_info.get("effective_date") or doc_info.get("version")
|
||||
@@ -674,35 +729,35 @@ def update_document(doc_path: str) -> Tuple[Any, int]:
|
||||
from config import DOCUMENTS_PATH
|
||||
|
||||
if 'file' not in request.files:
|
||||
return jsonify({"error": "没有上传文件"}), 400
|
||||
return error_response("NO_FILE", NO_FILE, "没有上传文件", http_status=400)
|
||||
|
||||
file = request.files['file']
|
||||
if file.filename == '':
|
||||
return jsonify({"error": "没有选择文件"}), 400
|
||||
return error_response("NO_FILE_SELECTED", NO_FILE_SELECTED, "没有选择文件", http_status=400)
|
||||
|
||||
# 文件类型校验
|
||||
ext = os.path.splitext(file.filename)[1].lower()
|
||||
if ext not in ALLOWED_EXTENSIONS:
|
||||
return jsonify({"error": f"不支持的文件类型: {ext}"}), 400
|
||||
return error_response("UNSUPPORTED_FORMAT", UNSUPPORTED_FORMAT, f"不支持的文件类型: {ext}", http_status=400)
|
||||
|
||||
# 文件大小校验
|
||||
file.seek(0, 2) # 跳到文件末尾获取大小
|
||||
file_size = file.tell()
|
||||
file.seek(0) # 回到文件开头
|
||||
if file_size > MAX_FILE_SIZE:
|
||||
return jsonify({"error": f"文件过大(最大 {MAX_FILE_SIZE // 1024 // 1024}MB)"}), 400
|
||||
return error_response("FILE_TOO_LARGE", FILE_TOO_LARGE, f"文件过大(最大 {MAX_FILE_SIZE // 1024 // 1024}MB)", http_status=400)
|
||||
|
||||
# 解析路径
|
||||
parts = doc_path.split('/')
|
||||
if len(parts) < 2:
|
||||
return jsonify({"error": "无效的文档路径"}), 400
|
||||
return error_response("BAD_REQUEST", BAD_REQUEST, "无效的文档路径", http_status=400)
|
||||
|
||||
try:
|
||||
filepath = _validate_doc_path(doc_path, DOCUMENTS_PATH)
|
||||
except ValueError:
|
||||
return jsonify({"error": "非法路径"}), 403
|
||||
return error_response("FORBIDDEN", FORBIDDEN, "非法路径", http_status=403)
|
||||
if not os.path.exists(filepath):
|
||||
return jsonify({"error": "文件不存在"}), 404
|
||||
return error_response("NOT_FOUND", NOT_FOUND, "文件不存在", http_status=404)
|
||||
|
||||
# 覆盖文件
|
||||
file.save(filepath)
|
||||
@@ -726,10 +781,7 @@ def update_document(doc_path: str) -> Tuple[Any, int]:
|
||||
except Exception as e:
|
||||
logger.warning(f"重新向量化失败: {e}")
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
"message": "文件已更新"
|
||||
})
|
||||
return success_response(message="文件已更新")
|
||||
|
||||
|
||||
@document_bp.route('/documents/<path:doc_path>', methods=['DELETE'])
|
||||
@@ -754,7 +806,7 @@ def delete_document(doc_path: str) -> Tuple[Any, int]:
|
||||
# 解析路径
|
||||
parts = doc_path.split('/')
|
||||
if len(parts) < 2:
|
||||
return jsonify({"error": "无效的文档路径"}), 400
|
||||
return error_response("BAD_REQUEST", BAD_REQUEST, "无效的文档路径", http_status=400)
|
||||
|
||||
subdir = parts[0]
|
||||
filename = '/'.join(parts[1:])
|
||||
@@ -765,9 +817,9 @@ def delete_document(doc_path: str) -> Tuple[Any, int]:
|
||||
try:
|
||||
filepath = _validate_doc_path(doc_path, DOCUMENTS_PATH)
|
||||
except ValueError:
|
||||
return jsonify({"error": "非法路径"}), 403
|
||||
return error_response("FORBIDDEN", FORBIDDEN, "非法路径", http_status=403)
|
||||
if not os.path.exists(filepath):
|
||||
return jsonify({"error": "文件不存在"}), 404
|
||||
return error_response("NOT_FOUND", NOT_FOUND, "文件不存在", http_status=404)
|
||||
|
||||
try:
|
||||
# 1. 从向量库删除(source 存的是文件名,不是完整路径)
|
||||
@@ -778,14 +830,11 @@ def delete_document(doc_path: str) -> Tuple[Any, int]:
|
||||
# 2. 删除文件
|
||||
os.remove(filepath)
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
"message": "文档已删除"
|
||||
})
|
||||
return success_response(status_code=DELETE_SUCCESS, message="文档已删除")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"删除文档异常: {e}")
|
||||
return jsonify({"error": "删除失败,请稍后重试"}), 500
|
||||
return error_response("INTERNAL_ERROR", INTERNAL_ERROR, "删除失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@document_bp.route('/documents/<path:doc_path>/chunks', methods=['GET'])
|
||||
@@ -810,12 +859,12 @@ def list_document_chunks(doc_path: str) -> Tuple[Any, int]:
|
||||
"""
|
||||
kb_manager = _get_kb_manager()
|
||||
if not kb_manager:
|
||||
return jsonify({"error": "知识库管理器未初始化"}), 503
|
||||
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "知识库管理器未初始化", http_status=503)
|
||||
|
||||
# 解析路径
|
||||
parts = doc_path.split('/')
|
||||
if len(parts) < 2:
|
||||
return jsonify({"error": "无效的文档路径"}), 400
|
||||
return error_response("BAD_REQUEST", BAD_REQUEST, "无效的文档路径", http_status=400)
|
||||
|
||||
subdir = parts[0]
|
||||
# 目录名即向量库名
|
||||
@@ -823,8 +872,7 @@ def list_document_chunks(doc_path: str) -> Tuple[Any, int]:
|
||||
|
||||
chunks = kb_manager.get_document_chunks(collection, os.path.basename(doc_path))
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"document_id": doc_path,
|
||||
"collection": collection,
|
||||
"chunks": chunks,
|
||||
@@ -860,12 +908,12 @@ def preview_document(doc_path: str) -> Tuple[Any, int]:
|
||||
"""
|
||||
kb_manager = _get_kb_manager()
|
||||
if not kb_manager:
|
||||
return jsonify({"error": "知识库管理器未初始化"}), 503
|
||||
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "知识库管理器未初始化", http_status=503)
|
||||
|
||||
# 解析路径
|
||||
parts = doc_path.split('/')
|
||||
if len(parts) < 2:
|
||||
return jsonify({"error": "无效的文档路径,格式: collection/filename"}), 400
|
||||
return error_response("BAD_REQUEST", BAD_REQUEST, "无效的文档路径,格式: collection/filename", http_status=400)
|
||||
|
||||
collection = parts[0]
|
||||
filename = os.path.basename(doc_path)
|
||||
@@ -880,7 +928,7 @@ def preview_document(doc_path: str) -> Tuple[Any, int]:
|
||||
# 获取所有切片
|
||||
all_chunks = kb_manager.get_document_chunks(collection, filename)
|
||||
if not all_chunks:
|
||||
return jsonify({"error": f"文档 '{filename}' 不存在或无切片"}), 404
|
||||
return error_response("NOT_FOUND", NOT_FOUND, f"文档 '{filename}' 不存在或无切片", http_status=404)
|
||||
|
||||
total = len(all_chunks)
|
||||
|
||||
@@ -889,8 +937,7 @@ def preview_document(doc_path: str) -> Tuple[Any, int]:
|
||||
preview_chunks = all_chunks[:5]
|
||||
for c in preview_chunks:
|
||||
c['is_target'] = False
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"collection": collection,
|
||||
"source": filename,
|
||||
"total_chunks": total,
|
||||
@@ -902,7 +949,7 @@ def preview_document(doc_path: str) -> Tuple[Any, int]:
|
||||
try:
|
||||
target_chunk_index = int(chunk_index_str)
|
||||
except ValueError:
|
||||
return jsonify({"error": "chunk_index 必须为整数"}), 400
|
||||
return error_response("BAD_REQUEST", BAD_REQUEST, "chunk_index 必须为整数", http_status=400)
|
||||
|
||||
# 按 chunk_index 排序(Chroma 返回顺序不保证有序)
|
||||
all_chunks.sort(key=lambda c: c.get('metadata', {}).get('chunk_index', 0))
|
||||
@@ -923,9 +970,7 @@ def preview_document(doc_path: str) -> Tuple[Any, int]:
|
||||
(c.get('metadata', {}).get('chunk_index', 0) for c in all_chunks),
|
||||
default=total - 1
|
||||
)
|
||||
return jsonify({
|
||||
"error": f"chunk_index={target_chunk_index} 未找到对应切片 (可用范围 0-{max_idx})"
|
||||
}), 400
|
||||
return error_response("BAD_REQUEST", BAD_REQUEST, f"chunk_index={target_chunk_index} 未找到对应切片 (可用范围 0-{max_idx})", http_status=400)
|
||||
|
||||
# 截取上下文窗口
|
||||
start = max(0, target_pos - context_count)
|
||||
@@ -936,8 +981,7 @@ def preview_document(doc_path: str) -> Tuple[Any, int]:
|
||||
for i, c in enumerate(window):
|
||||
c['is_target'] = (start + i == target_pos)
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"collection": collection,
|
||||
"source": filename,
|
||||
"total_chunks": total,
|
||||
@@ -968,7 +1012,7 @@ def create_chunk() -> Tuple[Any, int]:
|
||||
"""
|
||||
kb_manager = _get_kb_manager()
|
||||
if not kb_manager:
|
||||
return jsonify({"error": "知识库管理器未初始化"}), 503
|
||||
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "知识库管理器未初始化", http_status=503)
|
||||
|
||||
data = request.json or {}
|
||||
collection = data.get('collection')
|
||||
@@ -976,17 +1020,13 @@ def create_chunk() -> Tuple[Any, int]:
|
||||
metadata = data.get('metadata', {})
|
||||
|
||||
if not collection:
|
||||
return jsonify({"error": "请指定向量库 (collection)"}), 400
|
||||
return error_response("NO_COLLECTION", NO_COLLECTION, "请指定向量库 (collection)", http_status=400)
|
||||
if not content:
|
||||
return jsonify({"error": "切片内容不能为空"}), 400
|
||||
return error_response("NO_CONTENT", NO_CONTENT, "切片内容不能为空", http_status=400)
|
||||
|
||||
chunk_id = kb_manager.add_chunk(collection, content, metadata)
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
"chunk_id": chunk_id,
|
||||
"message": "切片已添加"
|
||||
})
|
||||
return success_response(data={"chunk_id": chunk_id}, message="切片已添加")
|
||||
|
||||
|
||||
@document_bp.route('/chunks/<chunk_id>', methods=['PUT'])
|
||||
@@ -1012,7 +1052,7 @@ def update_chunk(chunk_id: str) -> Tuple[Any, int]:
|
||||
"""
|
||||
kb_manager = _get_kb_manager()
|
||||
if not kb_manager:
|
||||
return jsonify({"error": "知识库管理器未初始化"}), 503
|
||||
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "知识库管理器未初始化", http_status=503)
|
||||
|
||||
data = request.json or {}
|
||||
collection = data.get('collection')
|
||||
@@ -1020,13 +1060,13 @@ def update_chunk(chunk_id: str) -> Tuple[Any, int]:
|
||||
metadata = data.get('metadata')
|
||||
|
||||
if not collection:
|
||||
return jsonify({"error": "请指定向量库 (collection)"}), 400
|
||||
return error_response("NO_COLLECTION", NO_COLLECTION, "请指定向量库 (collection)", http_status=400)
|
||||
|
||||
success = kb_manager.update_chunk(collection, chunk_id, content=content, metadata=metadata)
|
||||
|
||||
if success:
|
||||
return jsonify({"success": True, "message": "切片已更新"})
|
||||
return jsonify({"error": "更新失败"}), 500
|
||||
return success_response(message="切片已更新")
|
||||
return error_response("INTERNAL_ERROR", INTERNAL_ERROR, "更新失败", http_status=500)
|
||||
|
||||
|
||||
@document_bp.route('/chunks/<chunk_id>', methods=['DELETE'])
|
||||
@@ -1054,11 +1094,11 @@ def delete_chunk(chunk_id: str) -> Tuple[Any, int]:
|
||||
collection = request.args.get('collection')
|
||||
|
||||
if not collection:
|
||||
return jsonify({"error": "请指定向量库 (collection)"}), 400
|
||||
return error_response("NO_COLLECTION", NO_COLLECTION, "请指定向量库 (collection)", http_status=400)
|
||||
|
||||
kb_manager = _get_kb_manager()
|
||||
if not kb_manager:
|
||||
return jsonify({"error": "知识库管理器未初始化"}), 503
|
||||
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "知识库管理器未初始化", http_status=503)
|
||||
|
||||
# 删除切片,返回 (success, source_file)
|
||||
success, source_file = kb_manager.delete_chunk(collection, chunk_id)
|
||||
@@ -1083,8 +1123,8 @@ def delete_chunk(chunk_id: str) -> Tuple[Any, int]:
|
||||
logger.warning(f"清理哈希记录失败: {e}")
|
||||
|
||||
if success:
|
||||
return jsonify({"success": True, "message": "切片已删除"})
|
||||
return jsonify({"error": "删除失败"}), 500
|
||||
return success_response(status_code=DELETE_SUCCESS, message="切片已删除")
|
||||
return error_response("INTERNAL_ERROR", INTERNAL_ERROR, "删除失败", http_status=500)
|
||||
|
||||
|
||||
@document_bp.route('/chunks/batch', methods=['DELETE'])
|
||||
@@ -1110,13 +1150,13 @@ def delete_chunks_by_source() -> Tuple[Any, int]:
|
||||
source = data.get('source')
|
||||
|
||||
if not collection:
|
||||
return jsonify({"error": "请指定向量库 (collection)"}), 400
|
||||
return error_response("NO_COLLECTION", NO_COLLECTION, "请指定向量库 (collection)", http_status=400)
|
||||
if not source:
|
||||
return jsonify({"error": "请指定文件名 (source)"}), 400
|
||||
return error_response("BAD_REQUEST", BAD_REQUEST, "请指定文件名 (source)", http_status=400)
|
||||
|
||||
kb_manager = _get_kb_manager()
|
||||
if not kb_manager:
|
||||
return jsonify({"error": "知识库管理器未初始化"}), 503
|
||||
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "知识库管理器未初始化", http_status=503)
|
||||
|
||||
# 批量删除该文件的所有切片
|
||||
try:
|
||||
@@ -1134,11 +1174,7 @@ def delete_chunks_by_source() -> Tuple[Any, int]:
|
||||
except Exception as e:
|
||||
logger.warning(f"清理哈希记录失败: {e}")
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
"deleted_count": deleted_count,
|
||||
"message": f"已删除 {deleted_count} 个切片"
|
||||
})
|
||||
return success_response(data={"deleted_count": deleted_count}, message=f"已删除 {deleted_count} 个切片")
|
||||
except Exception as e:
|
||||
logger.error(f"操作异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("INTERNAL_ERROR", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
@@ -33,6 +33,11 @@
|
||||
|
||||
from flask import Blueprint, request, jsonify
|
||||
from auth.gateway import require_gateway_auth, require_role
|
||||
from core.status_codes import (
|
||||
SUCCESS, CREATED, DELETE_SUCCESS, UPDATE_SUCCESS,
|
||||
BAD_REQUEST, NOT_FOUND, INTERNAL_ERROR
|
||||
)
|
||||
from api.response_utils import success_response, error_response
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -75,10 +80,10 @@ def submit_feedback():
|
||||
user_id = data.get('user_id', '')
|
||||
|
||||
if not session_id or not query or rating is None:
|
||||
return jsonify({"error": "缺少必要参数"}), 400
|
||||
return error_response("MISSING_PARAMS", BAD_REQUEST, "缺少必要参数", http_status=400)
|
||||
|
||||
if rating not in [1, -1]:
|
||||
return jsonify({"error": "rating 必须是 1 或 -1"}), 400
|
||||
return error_response("INVALID_PARAMS", BAD_REQUEST, "rating 必须是 1 或 -1", http_status=400)
|
||||
|
||||
try:
|
||||
feedback_service = _get_feedback_service()
|
||||
@@ -91,15 +96,14 @@ def submit_feedback():
|
||||
reason=reason,
|
||||
user_id=user_id
|
||||
)
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"feedback_id": result['feedback_id'],
|
||||
"faq_suggested": result.get('faq_suggested', False),
|
||||
"suggestion_id": result.get('suggestion_id')
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"反馈操作异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@feedback_bp.route('/feedback/stats', methods=['GET'])
|
||||
@@ -112,13 +116,12 @@ def get_feedback_stats():
|
||||
try:
|
||||
feedback_db = _get_feedback_db()
|
||||
stats = feedback_db.get_feedback_stats(start_date, end_date)
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"stats": stats
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"反馈操作异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@feedback_bp.route('/feedback/list', methods=['GET'])
|
||||
@@ -140,14 +143,13 @@ def get_feedback_list():
|
||||
end_date=end_date,
|
||||
limit=limit
|
||||
)
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"feedbacks": feedbacks,
|
||||
"total": len(feedbacks)
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"反馈操作异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@feedback_bp.route('/reports/weekly', methods=['GET'])
|
||||
@@ -157,13 +159,12 @@ def get_weekly_report():
|
||||
try:
|
||||
feedback_service = _get_feedback_service()
|
||||
report = feedback_service.generate_report("weekly")
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"report": report.to_dict()
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"反馈操作异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@feedback_bp.route('/reports/monthly', methods=['GET'])
|
||||
@@ -173,13 +174,12 @@ def get_monthly_report():
|
||||
try:
|
||||
feedback_service = _get_feedback_service()
|
||||
report = feedback_service.generate_report("monthly")
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"report": report.to_dict()
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"反馈操作异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@feedback_bp.route('/faq', methods=['GET'])
|
||||
@@ -192,14 +192,13 @@ def get_faq_list():
|
||||
try:
|
||||
feedback_db = _get_feedback_db()
|
||||
faqs = feedback_db.get_faqs(status=status, limit=limit)
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"faqs": faqs,
|
||||
"total": len(faqs)
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"反馈操作异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@feedback_bp.route('/faq', methods=['POST'])
|
||||
@@ -220,7 +219,7 @@ def create_faq():
|
||||
answer = data.get('answer')
|
||||
|
||||
if not question or not answer:
|
||||
return jsonify({"error": "缺少问题或答案"}), 400
|
||||
return error_response("MISSING_PARAMS", BAD_REQUEST, "缺少问题或答案", http_status=400)
|
||||
|
||||
try:
|
||||
from services.feedback import FAQ
|
||||
@@ -235,15 +234,14 @@ def create_faq():
|
||||
)
|
||||
faq_id = feedback_db.add_faq(faq)
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"faq_id": faq_id,
|
||||
"status": "draft",
|
||||
"message": "FAQ已创建,请通过 /faq/<id>/approve 接口确认后生效"
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"反馈操作异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@feedback_bp.route('/faq/<int:faq_id>/approve', methods=['POST'])
|
||||
@@ -263,10 +261,10 @@ def approve_faq(faq_id):
|
||||
# 检查 FAQ 状态
|
||||
faq = feedback_db.get_faq(faq_id)
|
||||
if not faq:
|
||||
return jsonify({"error": "FAQ不存在"}), 404
|
||||
return error_response("NOT_FOUND", NOT_FOUND, "FAQ不存在", http_status=404)
|
||||
|
||||
if faq.get('status') == 'approved':
|
||||
return jsonify({"success": True, "message": "FAQ已经是批准状态"})
|
||||
return success_response(message="FAQ已经是批准状态")
|
||||
|
||||
# 更新状态为 approved
|
||||
feedback_db.update_faq(faq_id, {"status": "approved"})
|
||||
@@ -279,15 +277,14 @@ def approve_faq(faq_id):
|
||||
answer=faq['answer']
|
||||
)
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"faq_id": faq_id,
|
||||
"sync_status": "synced" if sync_success else "sync_failed",
|
||||
"message": "FAQ已批准并同步到知识库"
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"反馈操作异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@feedback_bp.route('/faq/<int:faq_id>', methods=['PUT'])
|
||||
@@ -303,11 +300,11 @@ def update_faq(faq_id):
|
||||
# 获取更新前的 FAQ 信息(用于判断是否需要重新同步)
|
||||
old_faq = feedback_db.get_faq(faq_id)
|
||||
if not old_faq:
|
||||
return jsonify({"error": "FAQ不存在"}), 404
|
||||
return error_response("NOT_FOUND", NOT_FOUND, "FAQ不存在", http_status=404)
|
||||
|
||||
updated = feedback_db.update_faq(faq_id, data)
|
||||
if not updated:
|
||||
return jsonify({"error": "FAQ更新失败"}), 500
|
||||
return error_response("UPDATE_FAILED", INTERNAL_ERROR, "FAQ更新失败", http_status=500)
|
||||
|
||||
# 检查是否需要重新同步向量库(question 或 answer 变更时)
|
||||
need_sync = False
|
||||
@@ -331,14 +328,13 @@ def update_faq(faq_id):
|
||||
)
|
||||
sync_status = "synced" if sync_success else "sync_failed"
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"message": "FAQ更新成功",
|
||||
"sync_status": sync_status
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"反馈操作异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@feedback_bp.route('/faq/<int:faq_id>', methods=['DELETE'])
|
||||
@@ -365,13 +361,10 @@ def delete_faq(faq_id):
|
||||
feedback_service = _get_feedback_service()
|
||||
feedback_service._delete_faq_vectors(faq_id)
|
||||
|
||||
return jsonify({
|
||||
"success": deleted,
|
||||
"message": "FAQ删除成功" if deleted else "FAQ不存在"
|
||||
})
|
||||
return success_response(data={"deleted": deleted}, message="FAQ删除成功" if deleted else "FAQ不存在")
|
||||
except Exception as e:
|
||||
logger.error(f"反馈操作异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@feedback_bp.route('/faq/suggestions', methods=['GET'])
|
||||
@@ -385,14 +378,13 @@ def get_faq_suggestions():
|
||||
try:
|
||||
feedback_db = _get_feedback_db()
|
||||
suggestions = feedback_db.get_faq_suggestions(status=status, limit=limit)
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"suggestions": suggestions,
|
||||
"total": len(suggestions)
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"反馈操作异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@feedback_bp.route('/faq/suggestions/<int:suggestion_id>/approve', methods=['POST'])
|
||||
@@ -418,17 +410,16 @@ def approve_faq_suggestion(suggestion_id):
|
||||
)
|
||||
|
||||
if result.get('success'):
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"faq_id": result['faq_id'],
|
||||
"sync_status": result.get('sync_status'),
|
||||
"message": "FAQ建议已批准并同步到知识库"
|
||||
})
|
||||
else:
|
||||
return jsonify({"error": result.get('error', '批准失败')}), 400
|
||||
return error_response("APPROVE_FAILED", BAD_REQUEST, result.get('error', '批准失败'), http_status=400)
|
||||
except Exception as e:
|
||||
logger.error(f"反馈操作异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@feedback_bp.route('/faq/suggestions/<int:suggestion_id>/reject', methods=['POST'])
|
||||
@@ -439,13 +430,13 @@ def reject_faq_suggestion(suggestion_id):
|
||||
try:
|
||||
feedback_db = _get_feedback_db()
|
||||
rejected = feedback_db.reject_faq_suggestion(suggestion_id)
|
||||
return jsonify({
|
||||
"success": rejected,
|
||||
return success_response(data={
|
||||
"rejected": rejected,
|
||||
"message": "FAQ建议已拒绝" if rejected else "建议不存在"
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"反馈操作异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
# ==================== Bad Case 分析接口 ====================
|
||||
@@ -477,8 +468,7 @@ def get_bad_cases():
|
||||
for case in bad_cases:
|
||||
case['status'] = 'pending' # pending/resolved/ignored
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"bad_cases": bad_cases,
|
||||
"blacklisted_sources": blacklisted_sources,
|
||||
"suggestions": [
|
||||
@@ -489,7 +479,7 @@ def get_bad_cases():
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"反馈操作异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@feedback_bp.route('/feedback/blacklist', methods=['GET'])
|
||||
@@ -507,12 +497,11 @@ def get_chunk_blacklist():
|
||||
feedback_service = _get_feedback_service()
|
||||
blacklist = feedback_service.get_chunk_blacklist(min_dislikes=min_dislikes)
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"blacklist": list(blacklist),
|
||||
"count": len(blacklist),
|
||||
"usage": "在检索时过滤这些来源以提升回答质量"
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"反馈操作异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
@@ -11,6 +11,8 @@
|
||||
import os
|
||||
import logging
|
||||
from flask import Blueprint, send_file, jsonify, current_app
|
||||
from core.status_codes import SUCCESS, BAD_REQUEST, NOT_FOUND, INTERNAL_ERROR
|
||||
from api.response_utils import success_response, error_response
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -51,7 +53,7 @@ def get_image(image_id: str):
|
||||
|
||||
# 安全检查:防止路径遍历攻击
|
||||
if '..' in image_id or '/' in image_id or '\\' in image_id:
|
||||
return jsonify({"error": "无效的图片 ID"}), 400
|
||||
return error_response("BAD_REQUEST", BAD_REQUEST, "无效的图片 ID", http_status=400)
|
||||
|
||||
images_path = get_images_base_path()
|
||||
|
||||
@@ -74,9 +76,9 @@ def get_image(image_id: str):
|
||||
return send_file(image_path, mimetype=mimetype)
|
||||
except Exception as e:
|
||||
logger.error(f"读取图片异常: {e}")
|
||||
return jsonify({"error": "读取图片失败"}), 500
|
||||
return error_response("READ_ERROR", INTERNAL_ERROR, "读取图片失败", http_status=500)
|
||||
|
||||
return jsonify({"error": "图片不存在", "image_id": image_id}), 404
|
||||
return error_response("NOT_FOUND", NOT_FOUND, "图片不存在", http_status=404, image_id=image_id)
|
||||
|
||||
|
||||
@image_bp.route('/images/<image_id>/info', methods=['GET'])
|
||||
@@ -96,7 +98,7 @@ def get_image_info(image_id: str):
|
||||
|
||||
# 安全检查
|
||||
if '..' in image_id or '/' in image_id or '\\' in image_id:
|
||||
return jsonify({"error": "无效的图片 ID"}), 400
|
||||
return error_response("BAD_REQUEST", BAD_REQUEST, "无效的图片 ID", http_status=400)
|
||||
|
||||
images_path = get_images_base_path()
|
||||
supported_formats = ['.png', '.jpg', '.jpeg', '.gif', '.bmp', '.webp']
|
||||
@@ -109,7 +111,7 @@ def get_image_info(image_id: str):
|
||||
from PIL import Image
|
||||
|
||||
with Image.open(image_path) as img:
|
||||
return jsonify({
|
||||
return success_response(data={
|
||||
"image_id": image_id,
|
||||
"width": img.width,
|
||||
"height": img.height,
|
||||
@@ -120,16 +122,16 @@ def get_image_info(image_id: str):
|
||||
})
|
||||
except ImportError:
|
||||
# PIL 未安装,返回基本信息
|
||||
return jsonify({
|
||||
return success_response(data={
|
||||
"image_id": image_id,
|
||||
"size_bytes": os.path.getsize(image_path),
|
||||
"url": f"/images/{image_id}"
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"读取图片信息异常: {e}")
|
||||
return jsonify({"error": "读取图片信息失败"}), 500
|
||||
return error_response("READ_ERROR", INTERNAL_ERROR, "读取图片信息失败", http_status=500)
|
||||
|
||||
return jsonify({"error": "图片不存在", "image_id": image_id}), 404
|
||||
return error_response("NOT_FOUND", NOT_FOUND, "图片不存在", http_status=404, image_id=image_id)
|
||||
|
||||
|
||||
@image_bp.route('/images/list', methods=['GET'])
|
||||
@@ -152,7 +154,7 @@ def list_images():
|
||||
images_path = get_images_base_path()
|
||||
|
||||
if not os.path.exists(images_path):
|
||||
return jsonify({"images": [], "total": 0})
|
||||
return success_response(data={"images": [], "total": 0})
|
||||
|
||||
supported_extensions = {'.png', '.jpg', '.jpeg', '.gif', '.bmp', '.webp'}
|
||||
images = []
|
||||
@@ -176,7 +178,7 @@ def list_images():
|
||||
total = len(images)
|
||||
images = images[offset:offset + limit]
|
||||
|
||||
return jsonify({
|
||||
return success_response(data={
|
||||
"images": images,
|
||||
"total": total,
|
||||
"limit": limit,
|
||||
@@ -185,7 +187,7 @@ def list_images():
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"列出图片异常: {e}")
|
||||
return jsonify({"error": "列出图片失败"}), 500
|
||||
return error_response("LIST_ERROR", INTERNAL_ERROR, "列出图片失败", http_status=500)
|
||||
|
||||
|
||||
@image_bp.route('/images/stats', methods=['GET'])
|
||||
@@ -199,7 +201,7 @@ def image_stats():
|
||||
images_path = get_images_base_path()
|
||||
|
||||
if not os.path.exists(images_path):
|
||||
return jsonify({
|
||||
return success_response(data={
|
||||
"total_images": 0,
|
||||
"total_size_bytes": 0,
|
||||
"supported_formats": ['.png', '.jpg', '.jpeg', '.gif', '.bmp', '.webp']
|
||||
@@ -219,7 +221,7 @@ def image_stats():
|
||||
total_size += os.path.getsize(filepath)
|
||||
format_counts[ext] = format_counts.get(ext, 0) + 1
|
||||
|
||||
return jsonify({
|
||||
return success_response(data={
|
||||
"total_images": total_count,
|
||||
"total_size_bytes": total_size,
|
||||
"total_size_mb": round(total_size / (1024 * 1024), 2),
|
||||
@@ -229,4 +231,4 @@ def image_stats():
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取统计信息异常: {e}")
|
||||
return jsonify({"error": "获取统计信息失败"}), 500
|
||||
return error_response("STATS_ERROR", INTERNAL_ERROR, "获取统计信息失败", http_status=500)
|
||||
|
||||
255
api/kb_routes.py
255
api/kb_routes.py
@@ -37,6 +37,13 @@ from typing import Tuple, Optional, Any
|
||||
from flask import Blueprint, request, jsonify, current_app
|
||||
import logging
|
||||
|
||||
from core.status_codes import (
|
||||
SUCCESS, CREATED, DELETE_SUCCESS, UPDATE_SUCCESS, SYNC_SUCCESS,
|
||||
BAD_REQUEST, NOT_FOUND, COLLECTION_NOT_FOUND, NO_COLLECTION, TASK_CONFLICT,
|
||||
INTERNAL_ERROR, SERVICE_UNAVAILABLE, SYNC_ERROR, REINDEX_ERROR
|
||||
)
|
||||
from api.response_utils import success_response, error_response
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
from auth.gateway import require_gateway_auth
|
||||
|
||||
@@ -133,10 +140,7 @@ def list_collections() -> Tuple[Any, int]:
|
||||
"description": coll.description
|
||||
})
|
||||
|
||||
return jsonify({
|
||||
"collections": result,
|
||||
"total": len(result)
|
||||
})
|
||||
return success_response(data={"collections": result, "total": len(result)})
|
||||
|
||||
|
||||
@kb_bp.route('/collections', methods=['POST'])
|
||||
@@ -172,22 +176,19 @@ def create_collection() -> Tuple[Any, int]:
|
||||
description = data.get('description', '')
|
||||
|
||||
if not name:
|
||||
return jsonify({"error": "向量库名称不能为空"}), 400
|
||||
return error_response("INVALID_NAME", BAD_REQUEST, "向量库名称不能为空", http_status=400)
|
||||
|
||||
# 验证名称格式(ChromaDB 限制)
|
||||
if not name.replace('_', '').replace('-', '').isalnum():
|
||||
return jsonify({
|
||||
"error": "名称格式错误",
|
||||
"message": "向量库名称只能包含字母、数字、下划线和连字符"
|
||||
}), 400
|
||||
return error_response("INVALID_NAME_FORMAT", BAD_REQUEST, "向量库名称只能包含字母、数字、下划线和连字符", http_status=400)
|
||||
|
||||
success, message = kb_manager.create_collection(
|
||||
name, display_name, department, description
|
||||
)
|
||||
|
||||
if success:
|
||||
return jsonify({"success": True, "message": message, "name": name}), 201
|
||||
return jsonify({"error": message}), 400
|
||||
return success_response(data={"name": name}, status_code=CREATED, message=message, http_status=201)
|
||||
return error_response("CREATE_FAILED", BAD_REQUEST, message, http_status=400)
|
||||
|
||||
|
||||
@kb_bp.route('/collections/<kb_name>', methods=['PUT'])
|
||||
@@ -220,7 +221,7 @@ def update_collection(kb_name: str) -> Tuple[Any, int]:
|
||||
# 检查向量库是否存在
|
||||
collections = kb_manager.list_collections()
|
||||
if not any(c.name == kb_name for c in collections):
|
||||
return jsonify({"error": f"向量库 '{kb_name}' 不存在"}), 404
|
||||
return error_response("COLLECTION_NOT_FOUND", COLLECTION_NOT_FOUND, f"向量库 '{kb_name}' 不存在", http_status=404)
|
||||
|
||||
# 更新元数据
|
||||
success = kb_manager.update_collection_metadata(
|
||||
@@ -230,8 +231,8 @@ def update_collection(kb_name: str) -> Tuple[Any, int]:
|
||||
)
|
||||
|
||||
if success:
|
||||
return jsonify({"success": True, "message": "向量库信息已更新"})
|
||||
return jsonify({"error": "更新失败"}), 500
|
||||
return success_response(data=None, status_code=UPDATE_SUCCESS, message="向量库信息已更新")
|
||||
return error_response("UPDATE_FAILED", INTERNAL_ERROR, "更新失败", http_status=500)
|
||||
|
||||
|
||||
@kb_bp.route('/collections/<kb_name>', methods=['DELETE'])
|
||||
@@ -260,12 +261,8 @@ def delete_collection(kb_name: str) -> Tuple[Any, int]:
|
||||
success, message = kb_manager.delete_collection(kb_name, delete_documents)
|
||||
|
||||
if success:
|
||||
return jsonify({
|
||||
"success": True,
|
||||
"message": message,
|
||||
"deleted_documents": delete_documents
|
||||
})
|
||||
return jsonify({"error": message}), 400
|
||||
return success_response(data={"deleted_documents": delete_documents}, status_code=DELETE_SUCCESS, message=message)
|
||||
return error_response("DELETE_FAILED", BAD_REQUEST, message, http_status=400)
|
||||
|
||||
|
||||
@kb_bp.route('/collections/<kb_name>/documents', methods=['GET'])
|
||||
@@ -286,11 +283,7 @@ def list_collection_documents(kb_name: str) -> Tuple[Any, int]:
|
||||
|
||||
documents = kb_manager.list_documents(kb_name)
|
||||
|
||||
return jsonify({
|
||||
"collection": kb_name,
|
||||
"documents": documents,
|
||||
"total": len(documents)
|
||||
})
|
||||
return success_response(data={"collection": kb_name, "documents": documents, "total": len(documents)})
|
||||
|
||||
|
||||
@kb_bp.route('/collections/<kb_name>/chunks', methods=['GET'])
|
||||
@@ -321,21 +314,16 @@ def list_collection_chunks(kb_name: str) -> Tuple[Any, int]:
|
||||
|
||||
chunks = kb_manager.list_chunks(kb_name, document_id=document_id, limit=limit, offset=offset)
|
||||
|
||||
return jsonify({
|
||||
"collection": kb_name,
|
||||
"chunks": chunks,
|
||||
"total": len(chunks)
|
||||
})
|
||||
return success_response(data={"collection": kb_name, "chunks": chunks, "total": len(chunks)})
|
||||
|
||||
|
||||
@kb_bp.route('/documents/sync', methods=['POST'])
|
||||
@require_gateway_auth
|
||||
def sync_documents() -> Tuple[Any, int]:
|
||||
"""
|
||||
触发文档向量化同步
|
||||
触发文档向量化同步(异步任务)
|
||||
|
||||
扫描文档目录,检测新增、修改、删除的文件,
|
||||
自动更新向量库索引。
|
||||
立即返回 task_id,后台线程执行同步。
|
||||
|
||||
请求体:
|
||||
{
|
||||
@@ -343,11 +331,7 @@ def sync_documents() -> Tuple[Any, int]:
|
||||
}
|
||||
|
||||
Returns:
|
||||
{
|
||||
"success": true,
|
||||
"results": [{"collection": "...", "status": "...", ...}],
|
||||
"synced_count": N
|
||||
}
|
||||
{"success": true, "data": {"task_id": "xxx", "message": "..."}}
|
||||
"""
|
||||
kb_manager, _, err = _require_multi_kb()
|
||||
if err:
|
||||
@@ -355,7 +339,6 @@ def sync_documents() -> Tuple[Any, int]:
|
||||
|
||||
from config import DOCUMENTS_PATH
|
||||
|
||||
user = request.current_user
|
||||
data = request.json or {}
|
||||
target_collection = data.get('collection')
|
||||
|
||||
@@ -363,54 +346,67 @@ def sync_documents() -> Tuple[Any, int]:
|
||||
if target_collection:
|
||||
collections_to_sync = [target_collection]
|
||||
else:
|
||||
# 同步所有向量库
|
||||
all_collections = kb_manager.list_collections()
|
||||
collections_to_sync = [c.name for c in all_collections]
|
||||
|
||||
if not collections_to_sync:
|
||||
return jsonify({"error": "没有可同步的向量库"}), 400
|
||||
return error_response("NO_COLLECTION", NO_COLLECTION, "没有可同步的向量库", http_status=400)
|
||||
|
||||
# 执行同步
|
||||
results = []
|
||||
|
||||
# 使用 sync_service 执行同步
|
||||
sync_service = current_app.config.get('SYNC_SERVICE')
|
||||
if not sync_service:
|
||||
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "同步服务不可用", http_status=503,
|
||||
results=[{"collection": c, "status": "warning", "message": "同步服务不可用"} for c in collections_to_sync],
|
||||
synced_count=0
|
||||
)
|
||||
|
||||
if sync_service:
|
||||
from core.task_registry import get_registry
|
||||
registry = get_registry()
|
||||
|
||||
# 检查是否有正在运行的同步任务
|
||||
running = registry.list_tasks(status='running', task_type='sync', limit=1)
|
||||
if running:
|
||||
return error_response("TASK_RUNNING", TASK_CONFLICT, f"同步任务正在执行中 (task_id: {running[0].id})", http_status=409)
|
||||
|
||||
desc = f"文档同步: {target_collection or '所有向量库'}"
|
||||
task = registry.create_task('sync', desc)
|
||||
|
||||
def _do_sync(task, sync_svc):
|
||||
processed = [0]
|
||||
|
||||
def on_change(change):
|
||||
processed[0] += 1
|
||||
registry.update_progress(
|
||||
task.id, current=processed[0],
|
||||
stage='处理文件',
|
||||
message=f"已处理: {change.document_name if hasattr(change, 'document_name') else change.document_id}"
|
||||
)
|
||||
|
||||
old_callback = sync_svc.on_change_callback
|
||||
sync_svc.on_change_callback = on_change
|
||||
try:
|
||||
sync_result = sync_service.sync_now()
|
||||
results.append({
|
||||
"collection": "all",
|
||||
"status": "success",
|
||||
"message": f"同步完成: 处理 {sync_result.documents_processed} 个文档",
|
||||
"details": {
|
||||
"added": sync_result.documents_added,
|
||||
"modified": sync_result.documents_modified,
|
||||
"deleted": sync_result.documents_deleted,
|
||||
"errors": sync_result.errors
|
||||
registry.update_progress(task.id, stage='扫描文档', message='正在检测变更...')
|
||||
sync_result = sync_svc.sync_now()
|
||||
return {
|
||||
'collection': target_collection or 'all',
|
||||
'status': 'success',
|
||||
'message': f"同步完成: 处理 {sync_result.documents_processed} 个文档",
|
||||
'details': {
|
||||
'added': sync_result.documents_added,
|
||||
'modified': sync_result.documents_modified,
|
||||
'deleted': sync_result.documents_deleted,
|
||||
'errors': sync_result.errors,
|
||||
}
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"知识库操作异常: {e}")
|
||||
results.append({
|
||||
"collection": "all",
|
||||
"status": "error",
|
||||
"message": "操作失败"
|
||||
})
|
||||
else:
|
||||
# 没有 sync_service,返回提示
|
||||
for coll_name in collections_to_sync:
|
||||
results.append({
|
||||
"collection": coll_name,
|
||||
"status": "warning",
|
||||
"message": "同步服务不可用,请使用 POST /sync 端点"
|
||||
})
|
||||
}
|
||||
finally:
|
||||
sync_svc.on_change_callback = old_callback
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
"results": results,
|
||||
"synced_count": len([r for r in results if r["status"] == "success"])
|
||||
})
|
||||
registry.start_task(task.id, _do_sync, sync_service)
|
||||
|
||||
return success_response(
|
||||
data={"task_id": task.id, "message": f"同步任务已启动,通过 GET /tasks/{task.id} 查询进度"},
|
||||
status_code=SYNC_SUCCESS,
|
||||
message="同步任务已启动"
|
||||
)
|
||||
|
||||
|
||||
@kb_bp.route('/debug/scan', methods=['GET'])
|
||||
@@ -473,22 +469,16 @@ def debug_scan() -> Tuple[Any, int]:
|
||||
@require_gateway_auth
|
||||
def reindex_collection(kb_name: str) -> Tuple[Any, int]:
|
||||
"""
|
||||
强制重新向量化指定集合的所有文档
|
||||
强制重新向量化指定集合的所有文档(异步任务)
|
||||
|
||||
清除该集合的文档哈希记录,触发完整重新索引。
|
||||
适用于文档内容更新后需要重建索引的场景。
|
||||
立即返回 task_id,后台线程执行重建。
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称
|
||||
|
||||
Returns:
|
||||
{
|
||||
"success": true,
|
||||
"message": "...",
|
||||
"documents_processed": N,
|
||||
"documents_added": N,
|
||||
"errors": [...]
|
||||
}
|
||||
{"success": true, "data": {"task_id": "xxx", "message": "..."}}
|
||||
"""
|
||||
kb_manager, _, err = _require_multi_kb()
|
||||
if err:
|
||||
@@ -496,14 +486,12 @@ def reindex_collection(kb_name: str) -> Tuple[Any, int]:
|
||||
|
||||
from config import DOCUMENTS_PATH
|
||||
|
||||
# 清除该集合的文档哈希记录
|
||||
# 清除该集合的文档哈希记录(同步完成,很快)
|
||||
try:
|
||||
from data.db import get_connection
|
||||
with get_connection("knowledge") as conn:
|
||||
cursor = conn.cursor()
|
||||
# 转义 LIKE 通配符,防止 kb_name 中的 % 或 _ 导致非预期匹配
|
||||
escaped_kb = kb_name.replace('%', '\\%').replace('_', '\\_')
|
||||
# 删除以 "{kb_name}/" 或 "{kb_name}\" 开头的文档哈希(兼容 Windows 和 Linux)
|
||||
cursor.execute("DELETE FROM document_hashes WHERE document_id LIKE ? ESCAPE '\\' OR document_id LIKE ? ESCAPE '\\'",
|
||||
(f"{escaped_kb}/%", f"{escaped_kb}\\%"))
|
||||
deleted = cursor.rowcount
|
||||
@@ -511,23 +499,54 @@ def reindex_collection(kb_name: str) -> Tuple[Any, int]:
|
||||
except Exception as e:
|
||||
logger.warning(f"清除哈希记录失败: {e}")
|
||||
|
||||
# 触发同步
|
||||
sync_service = current_app.config.get('SYNC_SERVICE')
|
||||
if sync_service:
|
||||
if not sync_service:
|
||||
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "同步服务不可用", http_status=503)
|
||||
|
||||
from core.task_registry import get_registry
|
||||
registry = get_registry()
|
||||
|
||||
# 检查是否有正在运行的重建任务
|
||||
running = registry.list_tasks(status='running', task_type='reindex', limit=1)
|
||||
if running:
|
||||
return error_response("TASK_RUNNING", TASK_CONFLICT, f"重建任务正在执行中 (task_id: {running[0].id})", http_status=409)
|
||||
|
||||
task = registry.create_task('reindex', f'重建索引: {kb_name}')
|
||||
|
||||
def _do_reindex(task, sync_svc, kb):
|
||||
"""后台执行重建索引"""
|
||||
processed = [0]
|
||||
|
||||
def on_change(change):
|
||||
processed[0] += 1
|
||||
registry.update_progress(
|
||||
task.id,
|
||||
current=processed[0],
|
||||
stage='重新索引',
|
||||
message=f"已处理: {change.document_name if hasattr(change, 'document_name') else change.document_id}"
|
||||
)
|
||||
|
||||
old_callback = sync_svc.on_change_callback
|
||||
sync_svc.on_change_callback = on_change
|
||||
try:
|
||||
result = sync_service.sync_now()
|
||||
return jsonify({
|
||||
"success": True,
|
||||
"message": f"重新索引完成: 处理 {result.documents_processed} 个文档",
|
||||
"documents_processed": result.documents_processed,
|
||||
"documents_added": result.documents_added,
|
||||
"errors": result.errors
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"reindex 异常: {e}")
|
||||
return jsonify({"error": "操作失败,请稍后重试"}), 500
|
||||
else:
|
||||
return jsonify({"error": "同步服务不可用"}), 503
|
||||
registry.update_progress(task.id, stage='扫描文档', message='正在检测变更...')
|
||||
result = sync_svc.sync_now()
|
||||
return {
|
||||
'message': f"重新索引完成: 处理 {result.documents_processed} 个文档",
|
||||
'documents_processed': result.documents_processed,
|
||||
'documents_added': result.documents_added,
|
||||
'errors': result.errors,
|
||||
}
|
||||
finally:
|
||||
sync_svc.on_change_callback = old_callback
|
||||
|
||||
registry.start_task(task.id, _do_reindex, sync_service, kb_name)
|
||||
|
||||
return success_response(
|
||||
data={"task_id": task.id, "message": f"重建索引任务已启动: {kb_name},通过 GET /tasks/{task.id} 查询进度"},
|
||||
status_code=SYNC_SUCCESS,
|
||||
message="重建索引任务已启动"
|
||||
)
|
||||
|
||||
|
||||
@kb_bp.route('/kb/route', methods=['POST'])
|
||||
@@ -567,7 +586,7 @@ def test_routing() -> Tuple[Any, int]:
|
||||
query = data.get('query', '')
|
||||
|
||||
if not query:
|
||||
return jsonify({"error": "请提供查询内容"}), 400
|
||||
return error_response("MISSING_PARAMS", BAD_REQUEST, "请提供查询内容", http_status=400)
|
||||
|
||||
# 获取路由结果
|
||||
target_kbs = route_query(
|
||||
@@ -579,7 +598,7 @@ def test_routing() -> Tuple[Any, int]:
|
||||
# 获取意图分析
|
||||
intent = kb_router.analyze_intent(query)
|
||||
|
||||
return jsonify({
|
||||
return success_response(data={
|
||||
"query": query,
|
||||
"user_role": user.get("role"),
|
||||
"user_department": user.get("department", ""),
|
||||
@@ -636,10 +655,10 @@ def deprecate_document(kb_name: str, filename: str) -> Tuple[Any, int]:
|
||||
reason,
|
||||
deprecated_by=user.get('user_id', 'unknown')
|
||||
)
|
||||
return jsonify(result)
|
||||
return success_response(data=result)
|
||||
except Exception as e:
|
||||
logger.error(f"操作异常: {e}")
|
||||
return jsonify({"success": False, "error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@kb_bp.route('/collections/<kb_name>/documents/<path:filename>/restore', methods=['POST'])
|
||||
@@ -668,10 +687,10 @@ def restore_document(kb_name: str, filename: str) -> Tuple[Any, int]:
|
||||
|
||||
try:
|
||||
result = kb_manager.restore_document(kb_name, filename)
|
||||
return jsonify(result)
|
||||
return success_response(data=result)
|
||||
except Exception as e:
|
||||
logger.error(f"操作异常: {e}")
|
||||
return jsonify({"success": False, "error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@kb_bp.route('/collections/<kb_name>/documents/<path:filename>/versions', methods=['GET'])
|
||||
@@ -719,8 +738,7 @@ def get_document_versions(kb_name: str, filename: str) -> Tuple[Any, int]:
|
||||
versions = version_query.get_document_history(kb_name, filename, limit)
|
||||
versions_data = [v.to_dict() for v in versions]
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"document_id": filename,
|
||||
"collection": kb_name,
|
||||
"versions": versions_data,
|
||||
@@ -728,7 +746,7 @@ def get_document_versions(kb_name: str, filename: str) -> Tuple[Any, int]:
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"操作异常: {e}")
|
||||
return jsonify({"success": False, "error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@kb_bp.route('/collections/<kb_name>/update-image-descriptions', methods=['POST'])
|
||||
@@ -760,10 +778,10 @@ def update_image_descriptions(kb_name: str) -> Tuple[Any, int]:
|
||||
|
||||
try:
|
||||
result = kb_manager.update_image_descriptions(kb_name)
|
||||
return jsonify(result)
|
||||
return success_response(data=result)
|
||||
except Exception as e:
|
||||
logger.error(f"操作异常: {e}")
|
||||
return jsonify({"success": False, "error": "操作失败,请稍后重试"}), 500
|
||||
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500)
|
||||
|
||||
|
||||
@kb_bp.route('/collections/sync-vlm-cache', methods=['POST'])
|
||||
@@ -804,12 +822,12 @@ def sync_vlm_cache() -> Tuple[Any, int]:
|
||||
images_dir = Path(".data/images")
|
||||
|
||||
if not vlm_cache_dir.exists():
|
||||
return jsonify({"success": False, "error": "VLM 缓存目录不存在"}), 400
|
||||
return error_response("VLM_CACHE_NOT_FOUND", BAD_REQUEST, "VLM 缓存目录不存在", http_status=400)
|
||||
|
||||
# 获取所有 VLM 缓存文件
|
||||
cache_files = list(vlm_cache_dir.glob("*.txt"))
|
||||
if not cache_files:
|
||||
return jsonify({"success": True, "total_cache_files": 0, "synced_count": 0, "message": "无 VLM 缓存文件"})
|
||||
return success_response(data={"total_cache_files": 0, "synced_count": 0}, message="无 VLM 缓存文件")
|
||||
|
||||
# 构建图片 MD5 → 文件名 的映射
|
||||
image_hash_map = {}
|
||||
@@ -879,8 +897,7 @@ def sync_vlm_cache() -> Tuple[Any, int]:
|
||||
skipped_count += 1
|
||||
details.append({"cache": cache_file.name, "image": image_filename, "status": "skipped", "reason": "向量库中未找到对应切片"})
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
return success_response(data={
|
||||
"total_cache_files": len(cache_files),
|
||||
"synced_count": synced_count,
|
||||
"skipped_count": skipped_count,
|
||||
|
||||
@@ -10,6 +10,8 @@
|
||||
|
||||
from flask import Blueprint, request, jsonify, current_app
|
||||
from auth.gateway import require_gateway_auth
|
||||
from core.status_codes import SUCCESS, BAD_REQUEST, FORBIDDEN, NOT_FOUND, INTERNAL_ERROR, SERVICE_UNAVAILABLE, DELETE_SUCCESS
|
||||
from api.response_utils import success_response, error_response
|
||||
|
||||
session_bp = Blueprint('session', __name__)
|
||||
|
||||
@@ -34,7 +36,7 @@ def get_sessions():
|
||||
"""
|
||||
session_manager = current_app.config['SESSION_MANAGER']
|
||||
if session_manager is None:
|
||||
return jsonify({"error": "会话服务不可用"}), 503
|
||||
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "会话服务不可用", http_status=503)
|
||||
user_id = request.current_user["user_id"]
|
||||
|
||||
sessions = session_manager.get_user_sessions(user_id, limit=20)
|
||||
@@ -47,7 +49,7 @@ def get_sessions():
|
||||
else:
|
||||
s["preview"] = "空会话"
|
||||
|
||||
return jsonify({"sessions": sessions})
|
||||
return success_response(data={"sessions": sessions})
|
||||
|
||||
|
||||
@session_bp.route('/history/<session_id>', methods=['GET'])
|
||||
@@ -65,19 +67,19 @@ def get_history(session_id):
|
||||
"""
|
||||
session_manager = current_app.config['SESSION_MANAGER']
|
||||
if session_manager is None:
|
||||
return jsonify({"error": "会话服务不可用"}), 503
|
||||
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "会话服务不可用", http_status=503)
|
||||
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:
|
||||
return jsonify({"error": "无权访问此会话"}), 403
|
||||
return error_response("FORBIDDEN", FORBIDDEN, "无权访问此会话", http_status=403)
|
||||
|
||||
history = session_manager.get_history(session_id, limit=100)
|
||||
|
||||
return jsonify({"history": history})
|
||||
return success_response(data={"history": history})
|
||||
|
||||
|
||||
@session_bp.route('/session/<session_id>', methods=['DELETE'])
|
||||
@@ -86,19 +88,19 @@ def delete_session(session_id):
|
||||
"""删除会话"""
|
||||
session_manager = current_app.config['SESSION_MANAGER']
|
||||
if session_manager is None:
|
||||
return jsonify({"error": "会话服务不可用"}), 503
|
||||
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "会话服务不可用", http_status=503)
|
||||
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:
|
||||
return jsonify({"error": "无权删除此会话"}), 403
|
||||
return error_response("FORBIDDEN", FORBIDDEN, "无权访问此会话", http_status=403)
|
||||
|
||||
session_manager.delete_session(session_id)
|
||||
|
||||
return jsonify({"success": True, "message": "会话已删除"})
|
||||
return success_response(status_code=DELETE_SUCCESS, message="会话已删除")
|
||||
|
||||
|
||||
@session_bp.route('/clear/<session_id>', methods=['POST'])
|
||||
@@ -107,16 +109,16 @@ def clear_history(session_id):
|
||||
"""清空会话历史(保留会话)"""
|
||||
session_manager = current_app.config['SESSION_MANAGER']
|
||||
if session_manager is None:
|
||||
return jsonify({"error": "会话服务不可用"}), 503
|
||||
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "会话服务不可用", http_status=503)
|
||||
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:
|
||||
return jsonify({"error": "无权操作此会话"}), 403
|
||||
return error_response("FORBIDDEN", FORBIDDEN, "无权访问此会话", http_status=403)
|
||||
|
||||
session_manager.clear_history(session_id)
|
||||
|
||||
return jsonify({"success": True, "message": "历史已清空"})
|
||||
return success_response(message="历史已清空")
|
||||
|
||||
@@ -31,12 +31,12 @@ Example:
|
||||
"""
|
||||
|
||||
from typing import Optional, Tuple, Any
|
||||
from flask import Blueprint, request, jsonify, current_app
|
||||
from flask import Blueprint, request, current_app
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
from auth.gateway import require_gateway_auth
|
||||
from core.status_codes import SYNC_SUCCESS, SYNC_ERROR, INTERNAL_ERROR
|
||||
from core.status_codes import SUCCESS, SYNC_SUCCESS, SYNC_ERROR, INTERNAL_ERROR, SERVICE_UNAVAILABLE
|
||||
from api.response_utils import success_response, error_response
|
||||
|
||||
sync_bp = Blueprint('sync', __name__)
|
||||
@@ -68,9 +68,7 @@ def _require_sync_service() -> Tuple[Optional[Any], Optional[Tuple]]:
|
||||
service = _get_sync_service()
|
||||
if not service:
|
||||
return None, error_response(
|
||||
error="SERVICE_UNAVAILABLE",
|
||||
error_code=INTERNAL_ERROR,
|
||||
message="同步服务未启用",
|
||||
"SERVICE_UNAVAILABLE", INTERNAL_ERROR, "同步服务未启用",
|
||||
http_status=503
|
||||
)
|
||||
return service, None
|
||||
@@ -82,9 +80,10 @@ def _require_sync_service() -> Tuple[Optional[Any], Optional[Tuple]]:
|
||||
@require_gateway_auth
|
||||
def trigger_sync() -> Tuple[Any, int]:
|
||||
"""
|
||||
手动触发知识库同步
|
||||
手动触发知识库同步(异步任务)
|
||||
|
||||
扫描文档目录,检测变更并执行向量化处理。
|
||||
立即返回 task_id,后台线程执行同步。
|
||||
客户端通过 GET /tasks/<task_id> 轮询进度。
|
||||
|
||||
请求体 (可选):
|
||||
{
|
||||
@@ -93,7 +92,7 @@ def trigger_sync() -> Tuple[Any, int]:
|
||||
}
|
||||
|
||||
Returns:
|
||||
成功: {"success": true, "data": {"result": {...}}}
|
||||
成功: {"success": true, "data": {"task_id": "xxx", "message": "同步任务已启动"}}
|
||||
失败: {"error": "...", "error_code": "..."}
|
||||
|
||||
Example:
|
||||
@@ -104,22 +103,65 @@ def trigger_sync() -> Tuple[Any, int]:
|
||||
if err:
|
||||
return err
|
||||
|
||||
try:
|
||||
result = service.sync_now()
|
||||
return success_response(
|
||||
data={"result": result.to_dict() if hasattr(result, 'to_dict') else result},
|
||||
status_code=SYNC_SUCCESS,
|
||||
message="同步完成"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"同步操作异常: {e}")
|
||||
from core.task_registry import get_registry
|
||||
registry = get_registry()
|
||||
|
||||
# 检查是否有正在运行的同步任务
|
||||
running = registry.list_tasks(status='running', task_type='sync', limit=1)
|
||||
if running:
|
||||
return error_response(
|
||||
error="SYNC_ERROR",
|
||||
error_code=SYNC_ERROR,
|
||||
message="操作失败",
|
||||
http_status=500
|
||||
"TASK_RUNNING", SYNC_ERROR,
|
||||
f"同步任务正在执行中 (task_id: {running[0].id}),请等待完成",
|
||||
http_status=409
|
||||
)
|
||||
|
||||
# 创建异步任务
|
||||
task = registry.create_task('sync', '文档同步')
|
||||
|
||||
def _do_sync(task, sync_service):
|
||||
"""后台执行同步"""
|
||||
# 注册进度回调
|
||||
processed = [0]
|
||||
|
||||
def on_change(change):
|
||||
processed[0] += 1
|
||||
registry.update_progress(
|
||||
task.id,
|
||||
current=processed[0],
|
||||
stage='处理文件',
|
||||
message=f"已处理: {change.document_name if hasattr(change, 'document_name') else change.document_id}"
|
||||
)
|
||||
|
||||
old_callback = sync_service.on_change_callback
|
||||
sync_service.on_change_callback = on_change
|
||||
|
||||
try:
|
||||
registry.update_progress(task.id, stage='扫描文档', message='正在检测变更...')
|
||||
result = sync_service.sync_now()
|
||||
|
||||
result_dict = result.to_dict() if hasattr(result, 'to_dict') else {
|
||||
'documents_processed': result.documents_processed,
|
||||
'documents_added': result.documents_added,
|
||||
'documents_modified': result.documents_modified,
|
||||
'documents_deleted': result.documents_deleted,
|
||||
'errors': result.errors,
|
||||
}
|
||||
return result_dict
|
||||
|
||||
finally:
|
||||
sync_service.on_change_callback = old_callback
|
||||
|
||||
registry.start_task(task.id, _do_sync, service)
|
||||
|
||||
return success_response(
|
||||
data={
|
||||
'task_id': task.id,
|
||||
'message': '同步任务已启动,通过 GET /tasks/' + task.id + ' 查询进度'
|
||||
},
|
||||
status_code=SYNC_SUCCESS,
|
||||
message="同步任务已启动"
|
||||
)
|
||||
|
||||
|
||||
@sync_bp.route('/sync/status', methods=['GET'])
|
||||
@require_gateway_auth
|
||||
@@ -143,12 +185,8 @@ def get_sync_status() -> Tuple[Any, int]:
|
||||
"""
|
||||
service, err = _require_sync_service()
|
||||
if err:
|
||||
return jsonify({
|
||||
"status": "failed",
|
||||
"status_code": INTERNAL_ERROR,
|
||||
"enabled": False,
|
||||
"message": "同步服务未启用"
|
||||
})
|
||||
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "同步服务未启用", http_status=503,
|
||||
enabled=False)
|
||||
|
||||
try:
|
||||
# 获取状态信息
|
||||
@@ -163,15 +201,11 @@ def get_sync_status() -> Tuple[Any, int]:
|
||||
if hasattr(service, 'get_status'):
|
||||
status.update(service.get_status())
|
||||
|
||||
return jsonify(status)
|
||||
return success_response(data=status)
|
||||
except Exception as e:
|
||||
logger.error(f"状态查询异常: {e}")
|
||||
return jsonify({
|
||||
"status": "failed",
|
||||
"status_code": INTERNAL_ERROR,
|
||||
"enabled": True,
|
||||
"error": "操作失败"
|
||||
})
|
||||
return error_response("INTERNAL_ERROR", INTERNAL_ERROR, "操作失败", http_status=500,
|
||||
enabled=True)
|
||||
|
||||
|
||||
@sync_bp.route('/sync/history', methods=['GET'])
|
||||
@@ -196,13 +230,11 @@ def get_sync_history() -> Tuple[Any, int]:
|
||||
|
||||
try:
|
||||
history = service.get_sync_history(limit=limit) if hasattr(service, 'get_sync_history') else []
|
||||
return jsonify({"history": history})
|
||||
return success_response(data={"history": history})
|
||||
except Exception as e:
|
||||
logger.error(f"同步操作异常: {e}")
|
||||
return error_response(
|
||||
error="SYNC_ERROR",
|
||||
error_code=SYNC_ERROR,
|
||||
message="操作失败",
|
||||
"SYNC_ERROR", SYNC_ERROR, "操作失败",
|
||||
http_status=500
|
||||
)
|
||||
|
||||
@@ -231,13 +263,11 @@ def get_change_logs() -> Tuple[Any, int]:
|
||||
|
||||
try:
|
||||
changes = service.get_change_logs(limit=limit, collection=collection) if hasattr(service, 'get_change_logs') else []
|
||||
return jsonify({"changes": changes})
|
||||
return success_response(data={"changes": changes})
|
||||
except Exception as e:
|
||||
logger.error(f"同步操作异常: {e}")
|
||||
return error_response(
|
||||
error="SYNC_ERROR",
|
||||
error_code=SYNC_ERROR,
|
||||
message="操作失败",
|
||||
"SYNC_ERROR", SYNC_ERROR, "操作失败",
|
||||
http_status=500
|
||||
)
|
||||
|
||||
@@ -263,27 +293,23 @@ def start_sync_monitor() -> Tuple[Any, int]:
|
||||
|
||||
try:
|
||||
if hasattr(service, 'is_running') and service.is_running():
|
||||
return jsonify({"status": "success", "status_code": SYNC_SUCCESS, "message": "文件监控已在运行"})
|
||||
return success_response(status_code=SYNC_SUCCESS, message="文件监控已在运行")
|
||||
|
||||
if hasattr(service, 'start'):
|
||||
success = service.start()
|
||||
if success:
|
||||
return jsonify({"status": "success", "status_code": SYNC_SUCCESS, "message": "文件监控已启动"})
|
||||
return success_response(status_code=SYNC_SUCCESS, message="文件监控已启动")
|
||||
else:
|
||||
return error_response(
|
||||
error="SYNC_ERROR",
|
||||
error_code=SYNC_ERROR,
|
||||
message="启动文件监控失败",
|
||||
"SYNC_ERROR", SYNC_ERROR, "启动文件监控失败",
|
||||
http_status=500
|
||||
)
|
||||
else:
|
||||
return jsonify({"status": "success", "message": "文件监控功能不可用"})
|
||||
return success_response(message="文件监控功能不可用")
|
||||
except Exception as e:
|
||||
logger.error(f"同步操作异常: {e}")
|
||||
return error_response(
|
||||
error="SYNC_ERROR",
|
||||
error_code=SYNC_ERROR,
|
||||
message="操作失败",
|
||||
"SYNC_ERROR", SYNC_ERROR, "操作失败",
|
||||
http_status=500
|
||||
)
|
||||
|
||||
@@ -306,12 +332,10 @@ def stop_sync_monitor() -> Tuple[Any, int]:
|
||||
try:
|
||||
if hasattr(service, 'stop'):
|
||||
service.stop()
|
||||
return jsonify({"status": "success", "status_code": SYNC_SUCCESS, "message": "文件监控已停止"})
|
||||
return success_response(status_code=SYNC_SUCCESS, message="文件监控已停止")
|
||||
except Exception as e:
|
||||
logger.error(f"同步操作异常: {e}")
|
||||
return error_response(
|
||||
error="SYNC_ERROR",
|
||||
error_code=SYNC_ERROR,
|
||||
message="操作失败",
|
||||
"SYNC_ERROR", SYNC_ERROR, "操作失败",
|
||||
http_status=500
|
||||
)
|
||||
|
||||
191
api/task_routes.py
Normal file
191
api/task_routes.py
Normal file
@@ -0,0 +1,191 @@
|
||||
"""
|
||||
异步任务查询 API
|
||||
|
||||
提供任务状态的 JSON 轮询和 SSE 流式两种查询方式。
|
||||
|
||||
路由列表:
|
||||
GET /tasks : 任务列表
|
||||
GET /tasks/<id> : 任务状态(JSON)
|
||||
GET /tasks/<id>/progress: 任务进度(SSE 流式)
|
||||
GET /tasks/stats : 任务统计
|
||||
|
||||
后端组调用方式:
|
||||
1. POST /sync → 返回 {"task_id": "xxx", ...}
|
||||
2. GET /tasks/xxx → 轮询状态,直到 status 为 completed 或 failed
|
||||
|
||||
dev-ui 调用方式:
|
||||
1. POST /sync → 返回 task_id
|
||||
2. GET /tasks/xxx/progress → SSE 流式接收进度事件
|
||||
"""
|
||||
|
||||
import json
|
||||
import time
|
||||
import logging
|
||||
from typing import Tuple, Any
|
||||
from flask import Blueprint, request, Response, stream_with_context
|
||||
from auth.gateway import require_gateway_auth
|
||||
from api.response_utils import success_response, error_response
|
||||
from core.status_codes import SUCCESS, NOT_FOUND
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
task_bp = Blueprint('tasks', __name__)
|
||||
|
||||
|
||||
def _get_registry():
|
||||
"""获取任务注册表"""
|
||||
from core.task_registry import get_registry
|
||||
return get_registry()
|
||||
|
||||
|
||||
# ==================== 任务查询 API ====================
|
||||
|
||||
@task_bp.route('/tasks', methods=['GET'])
|
||||
@require_gateway_auth
|
||||
def list_tasks() -> Tuple[Any, int]:
|
||||
"""
|
||||
获取任务列表
|
||||
|
||||
查询参数:
|
||||
status: 过滤状态(running / completed / failed / pending)
|
||||
type: 过滤类型(sync / reindex / upload / batch_upload / exam_generate / exam_grade)
|
||||
limit: 返回数量限制(默认 50)
|
||||
|
||||
Returns:
|
||||
{"success": true, "data": {"tasks": [...], "total": N}}
|
||||
"""
|
||||
registry = _get_registry()
|
||||
registry.maybe_cleanup()
|
||||
|
||||
status = request.args.get('status')
|
||||
task_type = request.args.get('type')
|
||||
limit = request.args.get('limit', 50, type=int)
|
||||
|
||||
tasks = registry.list_tasks(status=status, task_type=task_type, limit=limit)
|
||||
return success_response(
|
||||
data={
|
||||
'tasks': [t.to_dict() for t in tasks],
|
||||
'total': len(tasks)
|
||||
},
|
||||
status_code=SUCCESS,
|
||||
message="查询成功"
|
||||
)
|
||||
|
||||
|
||||
@task_bp.route('/tasks/stats', methods=['GET'])
|
||||
@require_gateway_auth
|
||||
def task_stats() -> Tuple[Any, int]:
|
||||
"""
|
||||
获取任务统计信息
|
||||
|
||||
Returns:
|
||||
{"success": true, "data": {"total": N, "by_status": {...}, "by_type": {...}}}
|
||||
"""
|
||||
registry = _get_registry()
|
||||
return success_response(
|
||||
data=registry.get_stats(),
|
||||
status_code=SUCCESS,
|
||||
message="查询成功"
|
||||
)
|
||||
|
||||
|
||||
@task_bp.route('/tasks/<task_id>', methods=['GET'])
|
||||
@require_gateway_auth
|
||||
def get_task(task_id: str) -> Tuple[Any, int]:
|
||||
"""
|
||||
获取单个任务状态(JSON 轮询接口)
|
||||
|
||||
后端组推荐使用此接口轮询任务进度。
|
||||
|
||||
轮询建议:
|
||||
- 间隔 1-2 秒
|
||||
- 当 status 为 completed 或 failed 时停止轮询
|
||||
|
||||
Returns:
|
||||
成功: {"success": true, "data": {task详情}}
|
||||
未找到: {"error": "任务不存在"}
|
||||
"""
|
||||
registry = _get_registry()
|
||||
task = registry.get_task(task_id)
|
||||
|
||||
if not task:
|
||||
return error_response(
|
||||
"TASK_NOT_FOUND", NOT_FOUND,
|
||||
f"任务不存在: {task_id}",
|
||||
http_status=404
|
||||
)
|
||||
|
||||
return success_response(
|
||||
data=task.to_dict(),
|
||||
status_code=SUCCESS,
|
||||
message="查询成功"
|
||||
)
|
||||
|
||||
|
||||
@task_bp.route('/tasks/<task_id>/progress', methods=['GET'])
|
||||
@require_gateway_auth
|
||||
def task_progress_stream(task_id: str):
|
||||
"""
|
||||
SSE 流式任务进度推送
|
||||
|
||||
dev-ui 前端推荐使用此接口,实时接收任务进度事件。
|
||||
|
||||
SSE 事件类型:
|
||||
- start: 任务开始
|
||||
- progress: 进度更新(含 progress/current/total/stage/message)
|
||||
- complete: 任务完成(含完整结果)
|
||||
- error: 任务失败(含错误信息)
|
||||
- heartbeat: 每 15 秒发送一次保活
|
||||
|
||||
Returns:
|
||||
text/event-stream
|
||||
"""
|
||||
registry = _get_registry()
|
||||
task = registry.get_task(task_id)
|
||||
|
||||
if not task:
|
||||
return error_response(
|
||||
"TASK_NOT_FOUND", NOT_FOUND,
|
||||
f"任务不存在: {task_id}",
|
||||
http_status=404
|
||||
)
|
||||
|
||||
def generate_sse():
|
||||
"""SSE 生成器"""
|
||||
sent_complete = False
|
||||
|
||||
while True:
|
||||
# 取出缓冲区事件
|
||||
events = task.drain_events()
|
||||
|
||||
for event in events:
|
||||
yield f"data: {json.dumps(event, ensure_ascii=False)}\n\n"
|
||||
if event.get('type') in ('complete', 'error'):
|
||||
sent_complete = True
|
||||
|
||||
# 如果任务已结束且事件已发送完毕,退出
|
||||
if sent_complete:
|
||||
break
|
||||
|
||||
# 如果任务已完成但没有事件,发送最终状态后退出
|
||||
if task.status in ('completed', 'failed') and not events:
|
||||
final_event = {
|
||||
'type': 'complete' if task.status == 'completed' else 'error',
|
||||
'data': task.to_dict()
|
||||
}
|
||||
yield f"data: {json.dumps(final_event, ensure_ascii=False)}\n\n"
|
||||
break
|
||||
|
||||
# 心跳保活
|
||||
yield f": heartbeat\n\n"
|
||||
time.sleep(1)
|
||||
|
||||
return Response(
|
||||
stream_with_context(generate_sse()),
|
||||
mimetype='text/event-stream',
|
||||
headers={
|
||||
'Cache-Control': 'no-cache',
|
||||
'X-Accel-Buffering': 'no',
|
||||
'Connection': 'keep-alive'
|
||||
}
|
||||
)
|
||||
@@ -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
|
||||
|
||||
@@ -181,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
|
||||
@@ -198,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: 文本分块器
|
||||
"""
|
||||
|
||||
390
core/agentic.py
390
core/agentic.py
@@ -1,390 +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, SEMANTIC_CACHE_ENABLED,
|
||||
MAX_CONTEXT_TOKENS, MAX_CONTEXT_COUNT, RERANK_THRESHOLD,
|
||||
SOURCE_KB, SOURCE_WEB,
|
||||
)
|
||||
|
||||
# 尝试导入语义缓存
|
||||
try:
|
||||
from core.semantic_cache import SemanticCache
|
||||
except ImportError:
|
||||
SemanticCache = None
|
||||
|
||||
# 导入 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
|
||||
|
||||
# 初始化语义缓存
|
||||
self.semantic_cache = None
|
||||
self.embedding_model = None
|
||||
if SEMANTIC_CACHE_ENABLED and SemanticCache:
|
||||
try:
|
||||
engine = get_engine()
|
||||
if engine and hasattr(engine, 'embedding_model'):
|
||||
self.embedding_model = engine.embedding_model
|
||||
emb_dim = 768
|
||||
# 优先使用新 API,兼容旧版本
|
||||
if hasattr(self.embedding_model, 'get_embedding_dimension'):
|
||||
emb_dim = self.embedding_model.get_embedding_dimension()
|
||||
elif hasattr(self.embedding_model, 'get_sentence_embedding_dimension'):
|
||||
emb_dim = self.embedding_model.get_sentence_embedding_dimension()
|
||||
self.semantic_cache = SemanticCache(
|
||||
dim=emb_dim,
|
||||
threshold=0.92,
|
||||
max_size=5000
|
||||
)
|
||||
logger.info(f"语义缓存已启用,维度={emb_dim}")
|
||||
except Exception as e:
|
||||
logger.warning(f"语义缓存初始化失败: {e}")
|
||||
self.semantic_cache = 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,62 +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
|
||||
|
||||
# 语义缓存
|
||||
try:
|
||||
from config import SEMANTIC_CACHE_ENABLED, SEMANTIC_CACHE_THRESHOLD
|
||||
HAS_SEMANTIC_CACHE_CONFIG = True
|
||||
except ImportError:
|
||||
SEMANTIC_CACHE_ENABLED = False
|
||||
SEMANTIC_CACHE_THRESHOLD = 0.92
|
||||
HAS_SEMANTIC_CACHE_CONFIG = False
|
||||
|
||||
try:
|
||||
from core.semantic_cache import SemanticCache, get_semantic_cache
|
||||
HAS_SEMANTIC_CACHE = True
|
||||
except ImportError:
|
||||
HAS_SEMANTIC_CACHE = False
|
||||
SemanticCache = 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
|
||||
@@ -224,39 +224,36 @@ class RAGCacheManager:
|
||||
"""
|
||||
获取查询缓存结果
|
||||
|
||||
始终使用粗粒度 key(基于 kb_version),确保 GET/SET key 一致。
|
||||
|
||||
Args:
|
||||
query: 查询文本
|
||||
kb_name: 知识库名称
|
||||
doc_ids: 相关文档 ID 列表(用于细粒度缓存 key)
|
||||
doc_ids: 保留参数以兼容调用方签名(当前未使用)
|
||||
"""
|
||||
kb_version = self.get_kb_version(kb_name)
|
||||
|
||||
# 计算文档哈希(如果提供了 doc_ids)
|
||||
doc_hash = ""
|
||||
if doc_ids:
|
||||
doc_hash = self._compute_doc_hash(kb_name, doc_ids)
|
||||
|
||||
key = self._make_query_cache_key(query, kb_name, kb_version, doc_hash)
|
||||
# 使用粗粒度 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, doc_ids: List[str] = None) -> None:
|
||||
"""
|
||||
设置查询缓存结果
|
||||
|
||||
始终使用粗粒度 key(基于 kb_version),确保 GET/SET key 一致。
|
||||
kb_version 在文档变更时自增,触发整个知识库的缓存失效。
|
||||
|
||||
Args:
|
||||
query: 查询文本
|
||||
kb_name: 知识库名称
|
||||
result: 缓存结果
|
||||
doc_ids: 相关文档 ID 列表(用于细粒度失效)
|
||||
doc_ids: 保留参数以兼容调用方签名(当前未使用)
|
||||
"""
|
||||
kb_version = self.get_kb_version(kb_name)
|
||||
|
||||
# 计算相关文档的版本哈希(细粒度失效)
|
||||
doc_hash = ""
|
||||
if doc_ids:
|
||||
doc_hash = self._compute_doc_hash(kb_name, doc_ids)
|
||||
|
||||
key = self._make_query_cache_key(query, kb_name, kb_version, doc_hash)
|
||||
# 使用与 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:
|
||||
|
||||
228
core/engine.py
228
core/engine.py
@@ -29,6 +29,7 @@ RAG 核心引擎
|
||||
|
||||
import os
|
||||
import gc
|
||||
import re
|
||||
import time
|
||||
import logging
|
||||
import threading
|
||||
@@ -112,7 +113,7 @@ except ImportError:
|
||||
RERANK_DEVICE = "auto"
|
||||
RERANK_USE_ONNX = False
|
||||
RERANK_BACKEND = "local"
|
||||
RERANK_CLOUD_MODEL = "qwen3-rerank"
|
||||
RERANK_CLOUD_MODEL = "xop3qwen8breranker"
|
||||
RERANK_CLOUD_API_KEY = ""
|
||||
RERANK_CLOUD_BASE_URL = "https://dashscope.aliyuncs.com/compatible-api/v1/reranks"
|
||||
RERANK_CLOUD_TIMEOUT = 15
|
||||
@@ -272,8 +273,7 @@ class CloudReranker:
|
||||
body = {
|
||||
"model": self.model,
|
||||
"query": query,
|
||||
"documents": documents,
|
||||
"top_n": len(documents)
|
||||
"documents": documents
|
||||
}
|
||||
|
||||
session = self._get_session()
|
||||
@@ -625,7 +625,7 @@ class RAGEngine:
|
||||
elif len(conditions) > 1:
|
||||
where_filter = {"$and": conditions}
|
||||
|
||||
query_vector = self.embedding_model.encode(query).tolist()
|
||||
query_vector = self._encode_cached(query).tolist()
|
||||
recall_k = RERANK_CANDIDATES if (USE_RERANK or USE_HYBRID_SEARCH) else top_k
|
||||
recall_k = max(recall_k, top_k * RECALL_MULTIPLIER)
|
||||
|
||||
@@ -811,6 +811,76 @@ class RAGEngine:
|
||||
logger.warning(f"FAQ 集合查询失败: {e}")
|
||||
return get_empty_result()
|
||||
|
||||
def _encode_cached(self, text):
|
||||
"""
|
||||
带缓存的 embedding 编码
|
||||
|
||||
优先从 Embedding Cache(LRU)读取,未命中再调用模型编码并写入缓存。
|
||||
支持单文本和批量文本输入。
|
||||
|
||||
Args:
|
||||
text: 单个文本字符串 或 文本列表
|
||||
|
||||
Returns:
|
||||
numpy 数组(单文本为一维,批量为二维)
|
||||
"""
|
||||
import numpy as _np
|
||||
|
||||
# 检查 embedding 缓存是否启用(缓存配置查询结果,避免每次重复导入)
|
||||
if not hasattr(self, '_emb_cache_enabled'):
|
||||
self._emb_cache_enabled = True # 默认启用
|
||||
if CACHE_AVAILABLE:
|
||||
try:
|
||||
from config import EMBEDDING_CACHE_ENABLED
|
||||
self._emb_cache_enabled = EMBEDDING_CACHE_ENABLED
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
if not self._emb_cache_enabled:
|
||||
return self.embedding_model.encode(text)
|
||||
|
||||
try:
|
||||
_cache = get_cache_manager()
|
||||
except Exception:
|
||||
return self.embedding_model.encode(text)
|
||||
|
||||
# 批量输入
|
||||
if isinstance(text, list):
|
||||
try:
|
||||
cached_embs, missed_indices = _cache.get_embeddings_batch(text)
|
||||
if missed_indices:
|
||||
missed_texts = [text[i] for i in missed_indices]
|
||||
# encode(list) 始终返回 2D ndarray,直接按行索引即可
|
||||
new_embs = self.embedding_model.encode(missed_texts)
|
||||
if len(missed_indices) == 1:
|
||||
# 单条时 encode 可能返回 1D,需统一处理
|
||||
if new_embs.ndim == 1:
|
||||
new_embs = new_embs.reshape(1, -1)
|
||||
for idx, mi in enumerate(missed_indices):
|
||||
emb_list = new_embs[idx].tolist()
|
||||
cached_embs[mi] = emb_list
|
||||
try:
|
||||
_cache.set_embedding(text[mi], emb_list)
|
||||
except Exception:
|
||||
pass
|
||||
return _np.array(cached_embs)
|
||||
except Exception:
|
||||
# 缓存故障时优雅降级为直接编码
|
||||
return self.embedding_model.encode(text)
|
||||
|
||||
# 单文本输入
|
||||
cached = _cache.get_embedding(text)
|
||||
if cached is not None:
|
||||
return _np.array(cached)
|
||||
|
||||
embedding = self.embedding_model.encode(text)
|
||||
try:
|
||||
emb_list = embedding.tolist() if hasattr(embedding, 'tolist') else list(embedding)
|
||||
_cache.set_embedding(text, emb_list)
|
||||
except Exception:
|
||||
pass
|
||||
return embedding
|
||||
|
||||
def _search_image_chunks(self, query_vector: list, top_k: int = 5, where_filter: dict = None) -> dict:
|
||||
"""
|
||||
独立检索图片切片(P0:图片独立召回通道)
|
||||
@@ -1161,6 +1231,7 @@ class RAGEngine:
|
||||
logger.warning(f"扩展连续切片失败: {e}")
|
||||
return {'ids': [], 'documents': [], 'metadatas': []}
|
||||
|
||||
# 扩展同 section 的 text 邻居
|
||||
where_filter = {"$and": [{"source": source}, {"chunk_type": "text"}]}
|
||||
if section:
|
||||
where_filter["$and"].append({"section": section})
|
||||
@@ -1171,6 +1242,20 @@ class RAGEngine:
|
||||
if not neighbors.get('ids') or len(neighbors.get('ids', [])) <= 1:
|
||||
neighbors = _get_neighbors({"$and": [{"source": source}, {"chunk_type": "text"}]})
|
||||
|
||||
# 同时扩展同 section 的 table 邻居(table 切片的 rerank 分数往往偏低,
|
||||
# 但与同 section 的 text 切片属于同一语义单元,不应割裂)
|
||||
table_where = {"$and": [{"source": source}, {"chunk_type": "table"}]}
|
||||
if section:
|
||||
table_where["$and"].append({"section": section})
|
||||
table_neighbors = _get_neighbors(table_where)
|
||||
|
||||
# 当 section 为空时,table 查询只有 source 条件,可能拉入大量无关表格,
|
||||
# 缩小 chunk_index 窗口至 ±1 以降低噪音;有 section 时使用正常窗口
|
||||
if section:
|
||||
_t_before, _t_after = CONTEXT_EXPANSION_BEFORE, CONTEXT_EXPANSION_AFTER
|
||||
else:
|
||||
_t_before, _t_after = 1, 1
|
||||
|
||||
neighbor_rows = []
|
||||
for n_id, n_doc, n_meta in zip(
|
||||
neighbors.get('ids', []),
|
||||
@@ -1183,6 +1268,18 @@ class RAGEngine:
|
||||
if seed_index - CONTEXT_EXPANSION_BEFORE <= n_index <= seed_index + CONTEXT_EXPANSION_AFTER:
|
||||
neighbor_rows.append((n_index, n_id, n_doc, n_meta))
|
||||
|
||||
# 同 section 的 table 邻居也加入扩展范围
|
||||
for n_id, n_doc, n_meta in zip(
|
||||
table_neighbors.get('ids', []),
|
||||
table_neighbors.get('documents', []),
|
||||
table_neighbors.get('metadatas', [])
|
||||
):
|
||||
n_index = self._to_int(n_meta.get('chunk_index'))
|
||||
if n_index is None:
|
||||
continue
|
||||
if seed_index - _t_before <= n_index <= seed_index + _t_after:
|
||||
neighbor_rows.append((n_index, n_id, n_doc, n_meta))
|
||||
|
||||
seed_neighbors_added = 0
|
||||
for n_index, n_id, n_doc, n_meta in sorted(neighbor_rows, key=lambda row: row[0]):
|
||||
if len(items) >= max_chunks:
|
||||
@@ -1387,7 +1484,7 @@ class RAGEngine:
|
||||
if not target_collections:
|
||||
return get_empty_result()
|
||||
|
||||
query_vector = self.embedding_model.encode(query).tolist()
|
||||
query_vector = self._encode_cached(query).tolist()
|
||||
# 扩大召回数量,以便过滤废止切片后仍有足够结果
|
||||
recall_k = RERANK_CANDIDATES if (USE_RERANK or USE_HYBRID_SEARCH) else top_k
|
||||
recall_k = max(recall_k, top_k * RECALL_MULTIPLIER)
|
||||
@@ -1739,12 +1836,12 @@ class RAGEngine:
|
||||
# === 高精度版:基于语义向量 ===
|
||||
from core.mmr import mmr_rerank
|
||||
|
||||
# 获取查询向量
|
||||
query_emb = np.array(self.embedding_model.encode(query))
|
||||
# 获取查询向量(使用 embedding 缓存)
|
||||
query_emb = np.array(self._encode_cached(query))
|
||||
|
||||
# 批量编码所有文档
|
||||
# 批量编码所有文档(使用 embedding 缓存)
|
||||
docs_list = results['documents'][0]
|
||||
all_embeddings = self.embedding_model.encode(docs_list)
|
||||
all_embeddings = self._encode_cached(docs_list)
|
||||
|
||||
# 构建候选列表
|
||||
candidates = []
|
||||
@@ -1940,104 +2037,7 @@ class RAGEngine:
|
||||
reranked[key] = results[key]
|
||||
return reranked
|
||||
|
||||
# ---------------- 安全与工具 ----------------
|
||||
|
||||
def check_restricted_documents(self, query, allowed_levels, top_k=3, role=None, department=None):
|
||||
if not self._initialized:
|
||||
self.initialize()
|
||||
|
||||
if USE_MULTI_KB and self.kb_manager and role and department:
|
||||
from auth.gateway import get_accessible_collections
|
||||
all_colls = [c.name for c in self.kb_manager.list_collections()]
|
||||
accessible = set(get_accessible_collections(role, department, 'read'))
|
||||
restricted = set(all_colls) - accessible
|
||||
|
||||
if not restricted:
|
||||
return {"has_restricted": False, "restricted_levels": [], "restricted_sources": []}
|
||||
|
||||
query_vector = self.embedding_model.encode(query).tolist()
|
||||
found_sources = set()
|
||||
top_score = 0.0
|
||||
|
||||
for coll_name in restricted:
|
||||
try:
|
||||
coll = self.kb_manager.get_collection(coll_name)
|
||||
if not coll: continue
|
||||
res = coll.query(query_embeddings=[query_vector], n_results=top_k)
|
||||
if res['metadatas'] and res['metadatas'][0]:
|
||||
for meta in res['metadatas'][0]:
|
||||
found_sources.add(meta.get('source', '未知'))
|
||||
for dist in (res.get('distances', [[]])[0] or []):
|
||||
if dist > top_score: top_score = dist
|
||||
except Exception as e:
|
||||
logger.debug(f"权限检查遍历失败: {e}")
|
||||
|
||||
return {
|
||||
"has_restricted": len(found_sources) > 0,
|
||||
"restricted_levels": [c.replace('dept_', '') for c in restricted if True][:3],
|
||||
"restricted_sources": list(found_sources)[:3],
|
||||
"top_restricted_score": top_score
|
||||
}
|
||||
|
||||
if not allowed_levels:
|
||||
return {"has_restricted": False, "restricted_levels": [], "restricted_sources": [], "top_restricted_score": 0.0}
|
||||
|
||||
restricted_levels = {"public", "internal", "confidential", "secret"} - set(allowed_levels)
|
||||
if not restricted_levels:
|
||||
return {"has_restricted": False, "restricted_levels": [], "restricted_sources": [], "top_restricted_score": 0.0}
|
||||
|
||||
query_vector = self.embedding_model.encode(query).tolist()
|
||||
try:
|
||||
res = self.collection.query(
|
||||
query_embeddings=[query_vector],
|
||||
n_results=top_k,
|
||||
where={"security_level": {"$in": list(restricted_levels)}}
|
||||
)
|
||||
docs = res.get('documents', [[]])[0]
|
||||
if not docs:
|
||||
return {"has_restricted": False, "restricted_levels": [], "restricted_sources": [], "top_restricted_score": 0.0}
|
||||
|
||||
metas = res.get('metadatas', [[]])[0]
|
||||
dists = res.get('distances', [[]])[0]
|
||||
found_levels, found_sources, top_score = set(), set(), 0.0
|
||||
|
||||
for meta, dist in zip(metas, dists):
|
||||
found_levels.add(meta.get('security_level', 'public'))
|
||||
found_sources.add(meta.get('source', '未知'))
|
||||
if dist > top_score: top_score = dist
|
||||
|
||||
return {
|
||||
"has_restricted": True,
|
||||
"restricted_levels": list(found_levels),
|
||||
"restricted_sources": list(found_sources)[:3],
|
||||
"top_restricted_score": top_score
|
||||
}
|
||||
except Exception as e:
|
||||
logger.warning(f"受限内容检查失败: {e}")
|
||||
return {"has_restricted": False, "restricted_levels": [], "restricted_sources": [], "top_restricted_score": 0.0}
|
||||
|
||||
def generate_answer(self, query, context):
|
||||
"""底层生成答复能力"""
|
||||
prompt = f"""你是一个严谨的智能助手,请根据以下参考资料回答用户的问题。
|
||||
...
|
||||
参考资料:
|
||||
{context}
|
||||
|
||||
用户问题:{query}
|
||||
|
||||
请回答:"""
|
||||
try:
|
||||
from core.llm_utils import call_llm
|
||||
result = call_llm(
|
||||
self.llm_client,
|
||||
prompt,
|
||||
MODEL,
|
||||
temperature=LLM_TEMPERATURE,
|
||||
max_tokens=LLM_MAX_TOKENS
|
||||
)
|
||||
return result or f"调用大模型失败: 返回结果为空"
|
||||
except Exception as e:
|
||||
return f"调用大模型失败: {str(e)}"
|
||||
# ---------------- 流式生成 ----------------
|
||||
|
||||
def generate_answer_stream(self, query, context, history=None):
|
||||
"""
|
||||
@@ -2068,21 +2068,33 @@ class RAGEngine:
|
||||
"content": (
|
||||
"你是一个严谨的知识库问答助手。"
|
||||
"你必须且只能根据用户提供的【参考资料】回答问题。"
|
||||
"参考资料中每段内容前标有章节路径(━格式),请注意区分不同章节的内容,"
|
||||
"特别当不同章节标题相似或包含相同关键词时,务必根据章节路径准确定位,不要混淆。"
|
||||
"如果参考资料中有答案,必须引用对应内容回答,并在回答末尾标注引用编号(如[1]、[2])。"
|
||||
"如果参考资料中确实没有相关信息,简短说明即可,不要编造或补充资料外的内容。"
|
||||
"禁止使用参考资料以外的知识进行补充或推测。"
|
||||
"【重要-表格处理规则】当用户询问表格、要求展示表格内容时,你必须将参考资料中的 Markdown 表格原样输出(保留 | 分隔符和表格结构),"
|
||||
"不要仅用文字描述表格存在或仅列出章节名称。如果参考资料中多个章节都有表格,"
|
||||
"优先展示与用户问题最相关的表格完整内容。"
|
||||
)
|
||||
})
|
||||
|
||||
# 添加当前问题(带上下文)- 强化指令
|
||||
if context:
|
||||
# 检测用户问题是否涉及表格,加入针对性指令
|
||||
_table_hint = ""
|
||||
# 检测上下文中是否包含 Markdown 表格(数据驱动,无需硬编码关键词)
|
||||
_has_table_in_context = bool(re.search(r'\|.+\|', context)) if context else False
|
||||
if _has_table_in_context:
|
||||
_table_hint = "\n注意:参考资料中包含 Markdown 格式的表格数据,请务必将相关表格以原始 Markdown 表格格式完整展示在回答中,不要仅用文字描述。"
|
||||
|
||||
user_message = f"""【参考资料】
|
||||
{context}
|
||||
|
||||
【用户问题】
|
||||
{query}
|
||||
|
||||
请仔细阅读以上全部参考资料后回答。如果参考资料中包含相关内容,必须引用回答并标注编号。如果资料中没有相关信息,请明确说明。"""
|
||||
请仔细阅读以上全部参考资料后回答。注意参考资料中标有章节路径,请根据章节路径准确定位相关内容。如果参考资料中包含相关内容,必须引用回答并标注编号。如果资料中没有相关信息,请明确说明。{_table_hint}"""
|
||||
else:
|
||||
user_message = query
|
||||
|
||||
|
||||
@@ -83,9 +83,16 @@ class IntentAnalyzer:
|
||||
根据对话历史和当前用户消息,输出一个 JSON 对象,包含以下字段:
|
||||
|
||||
1. **rewritten_query**: 改写后的完整问题
|
||||
- 如果问题包含指代(如"这两张图片"、"继续说"),将其改写为完整、独立的问题
|
||||
- 例如:"分析一下这两张图片" → "分析一下对话历史中提到的图片"
|
||||
- 如果问题本身已经完整,直接返回原文
|
||||
- **指代消解**:如果问题包含指代(如"这两张图片"、"继续说"),将其改写为完整、独立的问题
|
||||
- 例如:"分析一下这两张图片" → "分析一下对话历史中提到的图片"
|
||||
- **追问补全**:如果问题是省略式追问(省略了上一轮讨论的主题实体),必须补全为完整问题
|
||||
- 判断方法:当前问题缺少主语/宾语,且对话历史中可以推断出省略的实体
|
||||
- 补全方法:从上一轮用户问题中提取主题实体,与追问组合成完整问题
|
||||
- 例如:
|
||||
- 上一轮问"吸烟点C1类是什么区?",追问"有完整表格吗?" → "吸烟点C1类有完整表格吗?"
|
||||
- 上一轮问"三峡工程的投资情况",追问"建设地点在哪?" → "三峡工程的建设地点在哪?"
|
||||
- 上一轮问"货源投放有哪些原则?",追问"具体内容是什么?" → "货源投放原则的具体内容是什么?"
|
||||
- 如果问题本身已经完整且独立,直接返回原文
|
||||
|
||||
2. **use_context**: 布尔值
|
||||
- true: 问题依赖历史对话中的信息,答案已经在历史回答中
|
||||
@@ -105,7 +112,9 @@ class IntentAnalyzer:
|
||||
- 推理类(intent="reasoning"):生成最多2个子查询
|
||||
* 原问题的检索查询
|
||||
* 一个补充角度的检索查询(如原因、背景、影响等),帮助获取更全面的上下文
|
||||
- 其他类(factual/instruction/other):严格只生成1个子查询(原问题)
|
||||
- 其他类(factual/instruction/other):严格只生成1个子查询
|
||||
* 子查询应基于 rewritten_query(改写后的完整问题),而非用户原始输入
|
||||
* 例如:追问"有完整表格吗?"改写为"吸烟点C1类有完整表格吗?"后,子查询应为"吸烟点C1类的完整表格内容"
|
||||
- 不要为同一实体生成语义重叠的查询
|
||||
- 子查询应保持原问题的关键词,长度20-60字符为宜
|
||||
|
||||
@@ -307,7 +316,8 @@ class IntentAnalyzer:
|
||||
|
||||
if cache_emb is not None:
|
||||
cached = cache.get(cache_emb)
|
||||
if cached:
|
||||
# 确保缓存条目是意图分析结果(非 RAG 回答缓存)
|
||||
if cached and cached.get("cache_type") != "rag_answer":
|
||||
logger.info(f"意图分析缓存命中: {cached.get('reason', '')[:50]}")
|
||||
return IntentAnalysis.from_dict(cached)
|
||||
else:
|
||||
@@ -377,9 +387,11 @@ class IntentAnalyzer:
|
||||
intent=intent_type
|
||||
)
|
||||
|
||||
# 存入语义缓存
|
||||
# 存入语义缓存(标记类型,避免与 RAG 回答缓存混淆)
|
||||
if cache and cache_emb is not None:
|
||||
cache.set(cache_emb, analysis.to_dict())
|
||||
cache_data = analysis.to_dict()
|
||||
cache_data["cache_type"] = "intent_analysis"
|
||||
cache.set(cache_emb, cache_data)
|
||||
|
||||
# 存入精确匹配缓存
|
||||
if len(self._exact_cache) < self._exact_cache_max:
|
||||
@@ -437,6 +449,7 @@ class IntentAnalyzer:
|
||||
if not history:
|
||||
return "(无历史对话)"
|
||||
|
||||
import re
|
||||
parts = []
|
||||
|
||||
# 提取最近 3 轮对话
|
||||
@@ -445,6 +458,7 @@ class IntentAnalyzer:
|
||||
for msg in recent_history:
|
||||
role = "用户" if msg.get("role") == "user" else "助手"
|
||||
content = msg.get("content", "")
|
||||
original_content = content # 保留原始内容用于结构化提取
|
||||
|
||||
# 截断过长的内容
|
||||
if len(content) > 500:
|
||||
@@ -452,9 +466,10 @@ class IntentAnalyzer:
|
||||
|
||||
parts.append(f"【{role}】{content}")
|
||||
|
||||
# 提取图片信息
|
||||
# 提取结构化信息(从 assistant 消息中提取章节、表格、来源等)
|
||||
metadata = msg.get("metadata", {})
|
||||
if isinstance(metadata, dict):
|
||||
# 已有的图片提取
|
||||
images = metadata.get("images", [])
|
||||
if images:
|
||||
for img in images[:3]:
|
||||
@@ -463,6 +478,43 @@ class IntentAnalyzer:
|
||||
img_type = img.get("type", "图片")
|
||||
parts.append(f" └─ {img_type}: {desc}")
|
||||
|
||||
# 来源文件提取
|
||||
sources = metadata.get("sources", [])
|
||||
if sources:
|
||||
source_names = []
|
||||
for s in sources[:3]:
|
||||
if isinstance(s, dict):
|
||||
name = s.get("source", "") or s.get("name", "")
|
||||
if name:
|
||||
source_names.append(name)
|
||||
elif isinstance(s, str):
|
||||
source_names.append(s)
|
||||
if source_names:
|
||||
parts.append(f" └─ 来源文件: {', '.join(source_names)}")
|
||||
|
||||
# collections(检索知识库)提取
|
||||
colls = metadata.get("collections", [])
|
||||
if colls:
|
||||
parts.append(f" └─ 检索知识库: {', '.join(colls)}")
|
||||
|
||||
# 从 assistant 原始内容中提取章节路径和表格结构
|
||||
if role == "助手" and original_content:
|
||||
# 提取章节路径:━ xxx ━ 格式
|
||||
sections = re.findall(r'━\s*(.+?)\s*━', original_content)
|
||||
if sections:
|
||||
unique_sections = list(dict.fromkeys(sections)) # 去重保序
|
||||
parts.append(f" └─ 涉及章节: {'; '.join(unique_sections[:3])}")
|
||||
|
||||
# 提取表格列名:| A | B | C | 格式的表头行
|
||||
table_headers = re.findall(r'^\|\s*(.+?)\s*\|', original_content, re.MULTILINE)
|
||||
if table_headers:
|
||||
# 取第一个表格的列名
|
||||
first_header = table_headers[0]
|
||||
cols = [c.strip() for c in first_header.split('|') if c.strip()]
|
||||
# 排除分隔符行(--- 格式)
|
||||
if cols and not all(re.match(r'^[-:]+$', c) for c in cols):
|
||||
parts.append(f" └─ 含表格,列名: {', '.join(cols[:6])}")
|
||||
|
||||
# 添加图片上下文
|
||||
if context_images:
|
||||
parts.append("\n【上下文中的图片】")
|
||||
|
||||
@@ -352,7 +352,7 @@ def quick_yes_no(
|
||||
if keywords is None:
|
||||
keywords = ["是", "需要", "yes", "true"]
|
||||
|
||||
result = call_llm(client, prompt, model, temperature=0, max_tokens=10)
|
||||
result = call_llm(client, prompt, model, temperature=0, max_tokens=128)
|
||||
if result is None:
|
||||
return False
|
||||
|
||||
|
||||
@@ -49,6 +49,8 @@ _STATUS_MESSAGES: Dict[int, str] = {
|
||||
4011: "向量库不存在",
|
||||
4012: "文件内容为空",
|
||||
4013: "权限不足",
|
||||
4014: "任务不存在",
|
||||
4015: "任务冲突",
|
||||
|
||||
# 服务端错误 (50xx)
|
||||
5000: "服务器内部错误",
|
||||
@@ -59,6 +61,7 @@ _STATUS_MESSAGES: Dict[int, str] = {
|
||||
5020: "出题失败",
|
||||
5021: "批阅失败",
|
||||
5030: "图片描述更新失败",
|
||||
5040: "重建索引失败",
|
||||
}
|
||||
|
||||
|
||||
@@ -113,6 +116,8 @@ FILE_NOT_FOUND = 4010
|
||||
COLLECTION_NOT_FOUND = 4011
|
||||
NO_CONTENT = 4012
|
||||
PERMISSION_DENIED = 4013
|
||||
TASK_NOT_FOUND = 4014
|
||||
TASK_CONFLICT = 4015
|
||||
|
||||
# 服务端错误 (50xx)
|
||||
INTERNAL_ERROR = 5000
|
||||
@@ -123,3 +128,4 @@ SYNC_ERROR = 5010
|
||||
EXAM_ERROR = 5020
|
||||
GRADE_ERROR = 5021
|
||||
IMAGE_DESC_ERROR = 5030
|
||||
REINDEX_ERROR = 5040
|
||||
|
||||
311
core/task_registry.py
Normal file
311
core/task_registry.py
Normal file
@@ -0,0 +1,311 @@
|
||||
"""
|
||||
异步任务注册表
|
||||
|
||||
为长时间运行的操作(同步、重建索引、上传向量化、出题、批阅)提供统一的
|
||||
任务状态跟踪机制。支持 JSON 轮询和 SSE 流式两种进度查询方式。
|
||||
|
||||
核心设计:
|
||||
- 进程内字典存储,线程安全
|
||||
- 后台线程执行长操作,立即返回 task_id
|
||||
- 任务完成后保留 1 小时自动清理
|
||||
- 单 Worker 环境下足够可靠
|
||||
|
||||
使用示例:
|
||||
from core.task_registry import get_registry
|
||||
|
||||
registry = get_registry()
|
||||
task_id = registry.start_task('sync', '文档同步', total=10)
|
||||
|
||||
# 在后台线程中更新进度
|
||||
registry.update_progress(task_id, current=5, message='处理中...')
|
||||
registry.complete_task(task_id, result={...})
|
||||
"""
|
||||
|
||||
import uuid
|
||||
import time
|
||||
import threading
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Callable, Dict, List, Optional
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import datetime
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ==================== 数据模型 ====================
|
||||
|
||||
@dataclass
|
||||
class TaskInfo:
|
||||
"""任务信息"""
|
||||
id: str # 任务唯一标识
|
||||
type: str # 任务类型: sync / reindex / upload / batch_upload / exam_generate / exam_grade
|
||||
description: str # 人类可读描述
|
||||
status: str = 'pending' # pending / running / completed / failed
|
||||
progress: float = 0.0 # 0-100 百分比
|
||||
current: int = 0 # 当前处理项数
|
||||
total: int = 0 # 总项数
|
||||
stage: str = '' # 当前阶段描述
|
||||
message: str = '' # 当前步骤消息
|
||||
result: Any = None # 完成后的结果数据
|
||||
error: Optional[str] = None # 失败时的错误信息
|
||||
created_at: float = 0.0 # 创建时间戳
|
||||
started_at: Optional[float] = None
|
||||
completed_at: Optional[float] = None
|
||||
_events: List[dict] = field(default_factory=list, repr=False) # SSE 事件缓冲
|
||||
_lock: threading.Lock = field(default_factory=threading.Lock, repr=False)
|
||||
|
||||
def add_event(self, event: dict):
|
||||
"""添加 SSE 事件到缓冲区"""
|
||||
with self._lock:
|
||||
self._events.append(event)
|
||||
|
||||
def drain_events(self) -> List[dict]:
|
||||
"""取出并清空事件缓冲区"""
|
||||
with self._lock:
|
||||
events = list(self._events)
|
||||
self._events.clear()
|
||||
return events
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""转为可序列化的字典"""
|
||||
d = {
|
||||
'task_id': self.id,
|
||||
'type': self.type,
|
||||
'description': self.description,
|
||||
'status': self.status,
|
||||
'progress': round(self.progress, 1),
|
||||
'current': self.current,
|
||||
'total': self.total,
|
||||
'stage': self.stage,
|
||||
'message': self.message,
|
||||
'created_at': datetime.fromtimestamp(self.created_at).isoformat(),
|
||||
}
|
||||
if self.started_at:
|
||||
d['started_at'] = datetime.fromtimestamp(self.started_at).isoformat()
|
||||
if self.completed_at:
|
||||
d['completed_at'] = datetime.fromtimestamp(self.completed_at).isoformat()
|
||||
d['duration_ms'] = int((self.completed_at - self.started_at) * 1000)
|
||||
if self.result is not None:
|
||||
d['result'] = self.result
|
||||
if self.error is not None:
|
||||
d['error'] = self.error
|
||||
return d
|
||||
|
||||
|
||||
# ==================== 任务注册表 ====================
|
||||
|
||||
class TaskRegistry:
|
||||
"""
|
||||
全局任务注册表(单例)
|
||||
|
||||
管理所有异步任务的创建、执行、状态更新和查询。
|
||||
"""
|
||||
|
||||
def __init__(self, max_workers: int = 4, task_ttl: int = 3600):
|
||||
self._tasks: Dict[str, TaskInfo] = {}
|
||||
self._lock = threading.Lock()
|
||||
self._executor = ThreadPoolExecutor(max_workers=max_workers, thread_name_prefix='task')
|
||||
self._task_ttl = task_ttl # 已完成任务的保留时间(秒)
|
||||
self._cleanup_interval = 300 # 清理间隔(秒)
|
||||
self._last_cleanup = time.time()
|
||||
|
||||
def create_task(self, task_type: str, description: str, total: int = 0) -> TaskInfo:
|
||||
"""
|
||||
创建新任务
|
||||
|
||||
Args:
|
||||
task_type: 任务类型标识
|
||||
description: 人类可读描述
|
||||
total: 预计处理的总项数(用于计算百分比)
|
||||
|
||||
Returns:
|
||||
TaskInfo 实例
|
||||
"""
|
||||
task_id = uuid.uuid4().hex[:12]
|
||||
task = TaskInfo(
|
||||
id=task_id,
|
||||
type=task_type,
|
||||
description=description,
|
||||
status='pending',
|
||||
total=total,
|
||||
created_at=time.time(),
|
||||
)
|
||||
with self._lock:
|
||||
self._tasks[task_id] = task
|
||||
logger.info(f"[任务] 创建 {task_id}: {description}")
|
||||
return task
|
||||
|
||||
def start_task(self, task_id: str, fn: Callable, *args, **kwargs) -> str:
|
||||
"""
|
||||
在后台线程中执行任务
|
||||
|
||||
Args:
|
||||
task_id: 任务 ID
|
||||
fn: 要执行的函数,签名为 fn(task, *args, **kwargs)
|
||||
*args, **kwargs: 传给 fn 的参数
|
||||
|
||||
Returns:
|
||||
task_id
|
||||
"""
|
||||
task = self._tasks.get(task_id)
|
||||
if not task:
|
||||
raise ValueError(f"任务不存在: {task_id}")
|
||||
|
||||
self._submit(fn, task, args, kwargs)
|
||||
return task_id
|
||||
|
||||
def _submit(self, fn, task, args, kwargs):
|
||||
"""提交任务到线程池"""
|
||||
def _wrapper():
|
||||
task.status = 'running'
|
||||
task.started_at = time.time()
|
||||
task.add_event({'type': 'start', 'data': {'stage': task.stage or '初始化'}})
|
||||
try:
|
||||
result = fn(task, *args, **kwargs)
|
||||
task.status = 'completed'
|
||||
task.progress = 100.0
|
||||
task.completed_at = time.time()
|
||||
task.result = result
|
||||
task.add_event({'type': 'complete', 'data': task.to_dict()})
|
||||
logger.info(f"[任务] 完成 {task.id}: {task.description}")
|
||||
except Exception as e:
|
||||
task.status = 'failed'
|
||||
task.error = str(e)
|
||||
task.completed_at = time.time()
|
||||
task.add_event({'type': 'error', 'data': {'message': str(e)}})
|
||||
logger.error(f"[任务] 失败 {task.id}: {e}", exc_info=True)
|
||||
|
||||
self._executor.submit(_wrapper)
|
||||
|
||||
def get_task(self, task_id: str) -> Optional[TaskInfo]:
|
||||
"""获取任务信息"""
|
||||
with self._lock:
|
||||
return self._tasks.get(task_id)
|
||||
|
||||
def list_tasks(self, status: Optional[str] = None, task_type: Optional[str] = None,
|
||||
limit: int = 50) -> List[TaskInfo]:
|
||||
"""列出任务(按创建时间倒序)"""
|
||||
with self._lock:
|
||||
tasks = list(self._tasks.values())
|
||||
|
||||
if status:
|
||||
tasks = [t for t in tasks if t.status == status]
|
||||
if task_type:
|
||||
tasks = [t for t in tasks if t.type == task_type]
|
||||
|
||||
tasks.sort(key=lambda t: t.created_at, reverse=True)
|
||||
return tasks[:limit]
|
||||
|
||||
def update_progress(self, task_id: str, *, current: Optional[int] = None,
|
||||
total: Optional[int] = None, stage: Optional[str] = None,
|
||||
message: Optional[str] = None):
|
||||
"""
|
||||
更新任务进度(线程安全)
|
||||
|
||||
Args:
|
||||
task_id: 任务 ID
|
||||
current: 当前处理项数
|
||||
total: 更新总项数
|
||||
stage: 当前阶段
|
||||
message: 当前步骤描述
|
||||
"""
|
||||
task = self._tasks.get(task_id)
|
||||
if not task:
|
||||
return
|
||||
|
||||
with task._lock:
|
||||
if current is not None:
|
||||
task.current = current
|
||||
if total is not None:
|
||||
task.total = total
|
||||
if stage is not None:
|
||||
task.stage = stage
|
||||
if message is not None:
|
||||
task.message = message
|
||||
|
||||
# 计算百分比
|
||||
if task.total > 0:
|
||||
task.progress = min(99.0, (task.current / task.total) * 100)
|
||||
|
||||
# 推送进度事件
|
||||
task.add_event({
|
||||
'type': 'progress',
|
||||
'data': {
|
||||
'progress': round(task.progress, 1),
|
||||
'current': task.current,
|
||||
'total': task.total,
|
||||
'stage': task.stage,
|
||||
'message': task.message,
|
||||
}
|
||||
})
|
||||
|
||||
def complete_task(self, task_id: str, result: Any = None):
|
||||
"""手动标记任务完成"""
|
||||
task = self._tasks.get(task_id)
|
||||
if not task:
|
||||
return
|
||||
task.status = 'completed'
|
||||
task.progress = 100.0
|
||||
task.completed_at = time.time()
|
||||
if result is not None:
|
||||
task.result = result
|
||||
task.add_event({'type': 'complete', 'data': task.to_dict()})
|
||||
|
||||
def fail_task(self, task_id: str, error: str):
|
||||
"""手动标记任务失败"""
|
||||
task = self._tasks.get(task_id)
|
||||
if not task:
|
||||
return
|
||||
task.status = 'failed'
|
||||
task.error = error
|
||||
task.completed_at = time.time()
|
||||
task.add_event({'type': 'error', 'data': {'message': error}})
|
||||
|
||||
def cleanup(self):
|
||||
"""清理过期的已完成任务"""
|
||||
now = time.time()
|
||||
with self._lock:
|
||||
expired = [
|
||||
tid for tid, t in self._tasks.items()
|
||||
if t.status in ('completed', 'failed')
|
||||
and t.completed_at
|
||||
and (now - t.completed_at) > self._task_ttl
|
||||
]
|
||||
for tid in expired:
|
||||
del self._tasks[tid]
|
||||
|
||||
if expired:
|
||||
logger.debug(f"[任务] 清理 {len(expired)} 个过期任务")
|
||||
|
||||
def maybe_cleanup(self):
|
||||
"""定期清理(避免频繁操作)"""
|
||||
if time.time() - self._last_cleanup > self._cleanup_interval:
|
||||
self.cleanup()
|
||||
self._last_cleanup = time.time()
|
||||
|
||||
def get_stats(self) -> dict:
|
||||
"""获取任务统计"""
|
||||
with self._lock:
|
||||
tasks = list(self._tasks.values())
|
||||
|
||||
stats = {'total': len(tasks), 'by_status': {}, 'by_type': {}}
|
||||
for t in tasks:
|
||||
stats['by_status'][t.status] = stats['by_status'].get(t.status, 0) + 1
|
||||
stats['by_type'][t.type] = stats['by_type'].get(t.type, 0) + 1
|
||||
return stats
|
||||
|
||||
|
||||
# ==================== 全局单例 ====================
|
||||
|
||||
_registry: Optional[TaskRegistry] = None
|
||||
_registry_lock = threading.Lock()
|
||||
|
||||
|
||||
def get_registry() -> TaskRegistry:
|
||||
"""获取全局任务注册表单例"""
|
||||
global _registry
|
||||
if _registry is None:
|
||||
with _registry_lock:
|
||||
if _registry is None:
|
||||
_registry = TaskRegistry()
|
||||
return _registry
|
||||
@@ -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 社区
|
||||
392
docs/curl测试手册.md
392
docs/curl测试手册.md
@@ -22,6 +22,7 @@
|
||||
- [12. 图片服务](#12-图片服务)
|
||||
- [13. 报告服务](#13-报告服务)
|
||||
- [14. 知识库路由](#14-知识库路由)
|
||||
- [15. 异步任务查询](#15-异步任务查询)
|
||||
- [附录. 已知问题](#附录-已知问题)
|
||||
|
||||
---
|
||||
@@ -587,15 +588,18 @@ curl -s -X POST http://localhost:5001/documents/upload \
|
||||
"size": 18,
|
||||
"replaced": false
|
||||
},
|
||||
"sync_status": "已保存并添加到向量库"
|
||||
"sync_status": "已保存,向量化任务已启动",
|
||||
"task_id": "a1b2c3d4e5f6"
|
||||
},
|
||||
"message": "文件上传成功,已保存并添加到向量库",
|
||||
"message": "文件上传成功,已保存,向量化任务已启动",
|
||||
"status": "success",
|
||||
"status_code": 2002,
|
||||
"success": true
|
||||
}
|
||||
```
|
||||
|
||||
> **异步说明**:文件保存为同步操作,向量化在后台线程异步执行。响应中的 `task_id` 可用于轮询向量化进度(`GET /tasks/<task_id>`)。
|
||||
>
|
||||
> **同名文件处理**:上传同名文件时,旧版本的切片会被自动清理后覆盖(`replaced: true`),不会生成时间戳后缀文件。
|
||||
|
||||
**验证结果**:✅ 通过
|
||||
@@ -624,7 +628,8 @@ curl -s -X POST http://localhost:5001/documents/batch-upload \
|
||||
{"filename": "file2.txt", "path": "public_kb/file2.txt", "status": "success", "replaced": false}
|
||||
],
|
||||
"success_count": 2,
|
||||
"total": 2
|
||||
"total": 2,
|
||||
"task_id": "b2c3d4e5f6a1"
|
||||
},
|
||||
"message": "批量上传完成,成功 2/2 个文件",
|
||||
"status": "success",
|
||||
@@ -633,6 +638,8 @@ curl -s -X POST http://localhost:5001/documents/batch-upload \
|
||||
}
|
||||
```
|
||||
|
||||
> **异步说明**:批量上传成功后自动触发后台向量化任务。`task_id` 可用于轮询向量化进度(`GET /tasks/<task_id>`)。若无成功上传的文件则不返回 `task_id`。
|
||||
>
|
||||
> **同名文件处理**:与单文件上传相同,批量上传中遇到同名文件也会自动覆盖旧版本(`replaced: true`)。
|
||||
|
||||
**验证结果**:✅ 通过
|
||||
@@ -882,7 +889,7 @@ curl -s -X DELETE "http://localhost:5001/chunks/人员名册.txt_text_0?collecti
|
||||
|
||||
### POST /sync
|
||||
|
||||
触发文档同步。
|
||||
触发文档同步(异步任务)。
|
||||
|
||||
```bash
|
||||
curl -s -X POST http://localhost:5001/sync \
|
||||
@@ -893,24 +900,23 @@ curl -s -X POST http://localhost:5001/sync \
|
||||
```json
|
||||
{
|
||||
"data": {
|
||||
"result": {
|
||||
"documents_added": 1,
|
||||
"documents_deleted": 1,
|
||||
"documents_modified": 0,
|
||||
"documents_processed": 2,
|
||||
"end_time": "2026-05-04T01:45:42",
|
||||
"errors": [],
|
||||
"start_time": "2026-05-04T01:45:41",
|
||||
"status": "completed"
|
||||
}
|
||||
"task_id": "c3d4e5f6a1b2",
|
||||
"message": "同步任务已启动,通过 GET /tasks/c3d4e5f6a1b2 查询进度"
|
||||
},
|
||||
"message": "同步完成",
|
||||
"message": "同步任务已启动",
|
||||
"status": "success",
|
||||
"status_code": 2010,
|
||||
"success": true
|
||||
}
|
||||
```
|
||||
|
||||
> **⚠️ 异步变更**:此接口已从同步改为异步。不再直接返回同步结果,而是返回 `task_id`。后端需通过 `GET /tasks/<task_id>` 轮询任务状态,直到 `status` 为 `completed` 或 `failed`。
|
||||
>
|
||||
> **冲突检测**:如果已有同步任务正在运行,返回 HTTP 409:
|
||||
> ```json
|
||||
> {"error": "TASK_RUNNING", "message": "同步任务正在执行中 (task_id: xxx),请等待完成"}
|
||||
> ```
|
||||
|
||||
**验证结果**:✅ 通过
|
||||
|
||||
---
|
||||
@@ -1456,44 +1462,54 @@ curl -s -X POST http://localhost:5001/exam/generate \
|
||||
```json
|
||||
{
|
||||
"data": {
|
||||
"questions": [
|
||||
{
|
||||
"content": {
|
||||
"stem": "重购率的定义是下列哪一项?",
|
||||
"answer": "B",
|
||||
"data": {
|
||||
"options": [
|
||||
{"content": "选项A内容", "key": "A"},
|
||||
{"content": "品规连续两周订货客户数/上周订货客户数*100%", "key": "B"}
|
||||
]
|
||||
},
|
||||
"explanation": "重购率的定义是..."
|
||||
},
|
||||
"difficulty": 3,
|
||||
"question_type": "single_choice",
|
||||
"source_trace": {
|
||||
"chunk_ids": ["1.docx_76"],
|
||||
"document_name": "public_kb/1.docx",
|
||||
"page_numbers": [1],
|
||||
"sources": [...]
|
||||
}
|
||||
}
|
||||
],
|
||||
"request_id": null,
|
||||
"source_chunks_used": 15,
|
||||
"success": true,
|
||||
"total": 10,
|
||||
"requested_types": {"single_choice": 5, "true_false": 3, "fill_blank": 2},
|
||||
"actual_types": {"single_choice": 5, "true_false": 3, "fill_blank": 2},
|
||||
"warnings": []
|
||||
"task_id": "d4e5f6a1b2c3",
|
||||
"message": "出题任务已启动 (10题),通过 GET /tasks/d4e5f6a1b2c3 查询结果"
|
||||
},
|
||||
"message": "出题成功",
|
||||
"message": "出题任务已启动",
|
||||
"status": "success",
|
||||
"status_code": 2020,
|
||||
"success": true
|
||||
}
|
||||
```
|
||||
|
||||
> **⚠️ 异步变更**:此接口已从同步改为异步。响应仅返回 `task_id`,后端需通过 `GET /tasks/<task_id>` 轮询任务状态。任务完成后,`result` 字段包含完整出题结果(格式见下方说明)。
|
||||
|
||||
**轮询结果(GET /tasks/\<task_id\> 完成后的 result 字段)**:
|
||||
```json
|
||||
{
|
||||
"questions": [
|
||||
{
|
||||
"content": {
|
||||
"stem": "重购率的定义是下列哪一项?",
|
||||
"answer": "B",
|
||||
"data": {
|
||||
"options": [
|
||||
{"content": "选项A内容", "key": "A"},
|
||||
{"content": "品规连续两周订货客户数/上周订货客户数*100%", "key": "B"}
|
||||
]
|
||||
},
|
||||
"explanation": "重购率的定义是..."
|
||||
},
|
||||
"difficulty": 3,
|
||||
"question_type": "single_choice",
|
||||
"source_trace": {
|
||||
"chunk_ids": ["1.docx_76"],
|
||||
"document_name": "public_kb/1.docx",
|
||||
"page_numbers": [1],
|
||||
"sources": [...]
|
||||
}
|
||||
}
|
||||
],
|
||||
"request_id": null,
|
||||
"source_chunks_used": 15,
|
||||
"success": true,
|
||||
"total": 10,
|
||||
"requested_types": {"single_choice": 5, "true_false": 3, "fill_blank": 2},
|
||||
"actual_types": {"single_choice": 5, "true_false": 3, "fill_blank": 2},
|
||||
"warnings": []
|
||||
}
|
||||
```
|
||||
|
||||
> **说明**:`warnings` 字段在某题型实际生成数量少于请求数量时返回提示信息。
|
||||
|
||||
**验证结果**:✅ 通过(2026-06-05)
|
||||
@@ -1527,31 +1543,41 @@ curl -s -X POST http://localhost:5001/exam/generate-smart \
|
||||
```json
|
||||
{
|
||||
"data": {
|
||||
"ai_analysis": {
|
||||
"total_knowledge_points": 21,
|
||||
"suitable_types": ["single_choice", "multiple_choice", "true_false", "subjective"],
|
||||
"question_types": {
|
||||
"single_choice": 8,
|
||||
"multiple_choice": 6,
|
||||
"true_false": 4,
|
||||
"fill_blank": 0,
|
||||
"subjective": 3
|
||||
},
|
||||
"reason": "文档包含21个知识点,涵盖术语定义、流程步骤、数值标准..."
|
||||
},
|
||||
"questions": [...],
|
||||
"request_id": null,
|
||||
"source_chunks_used": 15,
|
||||
"success": true,
|
||||
"total": 16
|
||||
"task_id": "e5f6a1b2c3d4",
|
||||
"message": "AI 智能出题任务已启动,通过 GET /tasks/e5f6a1b2c3d4 查询结果"
|
||||
},
|
||||
"message": "AI 智能出题成功",
|
||||
"message": "AI 智能出题任务已启动",
|
||||
"status": "success",
|
||||
"status_code": 2020,
|
||||
"success": true
|
||||
}
|
||||
```
|
||||
|
||||
> **⚠️ 异步变更**:此接口已从同步改为异步。响应仅返回 `task_id`,后端需通过 `GET /tasks/<task_id>` 轮询任务状态。任务完成后,`result` 字段包含完整出题结果(含 `ai_analysis` 字段)。
|
||||
|
||||
**轮询结果(GET /tasks/\<task_id\> 完成后的 result 字段)**:
|
||||
```json
|
||||
{
|
||||
"ai_analysis": {
|
||||
"total_knowledge_points": 21,
|
||||
"suitable_types": ["single_choice", "multiple_choice", "true_false", "subjective"],
|
||||
"question_types": {
|
||||
"single_choice": 8,
|
||||
"multiple_choice": 6,
|
||||
"true_false": 4,
|
||||
"fill_blank": 0,
|
||||
"subjective": 3
|
||||
},
|
||||
"reason": "文档包含21个知识点,涵盖术语定义、流程步骤、数值标准..."
|
||||
},
|
||||
"questions": [...],
|
||||
"request_id": null,
|
||||
"source_chunks_used": 15,
|
||||
"success": true,
|
||||
"total": 16
|
||||
}
|
||||
```
|
||||
|
||||
**注意事项**:
|
||||
- 实际出题数量 ≤ min(文档知识点数, AI 推荐数量)
|
||||
- 如果文档知识点较少,生成的题目数量会相应减少
|
||||
@@ -1601,33 +1627,43 @@ curl -s -X POST http://localhost:5001/exam/grade \
|
||||
```json
|
||||
{
|
||||
"data": {
|
||||
"request_id": null,
|
||||
"results": [
|
||||
{
|
||||
"question_id": "q1",
|
||||
"score": 2,
|
||||
"max_score": 2,
|
||||
"grading_status": "success",
|
||||
"details": {
|
||||
"correct": true,
|
||||
"student_answer": "B",
|
||||
"correct_answer": "B",
|
||||
"feedback": "正确!"
|
||||
}
|
||||
}
|
||||
],
|
||||
"score_rate": 100.0,
|
||||
"success": true,
|
||||
"total_max_score": 2,
|
||||
"total_score": 2
|
||||
"task_id": "f6a1b2c3d4e5",
|
||||
"message": "批阅任务已启动 (1题),通过 GET /tasks/f6a1b2c3d4e5 查询结果"
|
||||
},
|
||||
"message": "批阅完成",
|
||||
"message": "批阅任务已启动",
|
||||
"status": "success",
|
||||
"status_code": 2021,
|
||||
"success": true
|
||||
}
|
||||
```
|
||||
|
||||
> **⚠️ 异步变更**:此接口已从同步改为异步。响应仅返回 `task_id`,后端需通过 `GET /tasks/<task_id>` 轮询任务状态。任务完成后,`result` 字段包含完整批阅结果(格式见下方说明)。
|
||||
|
||||
**轮询结果(GET /tasks/\<task_id\> 完成后的 result 字段)**:
|
||||
```json
|
||||
{
|
||||
"request_id": null,
|
||||
"results": [
|
||||
{
|
||||
"question_id": "q1",
|
||||
"score": 2,
|
||||
"max_score": 2,
|
||||
"grading_status": "success",
|
||||
"details": {
|
||||
"correct": true,
|
||||
"student_answer": "B",
|
||||
"correct_answer": "B",
|
||||
"feedback": "正确!"
|
||||
}
|
||||
}
|
||||
],
|
||||
"score_rate": 100.0,
|
||||
"success": true,
|
||||
"total_max_score": 2,
|
||||
"total_score": 2
|
||||
}
|
||||
```
|
||||
|
||||
**`grading_status` 取值说明**:
|
||||
- `success`:评分成功
|
||||
- `failed`:评分失败(主观题 LLM 解析失败或超时),此时 `details` 包含 `error` 字段
|
||||
@@ -1823,6 +1859,190 @@ curl -s -X POST http://localhost:5001/kb/route \
|
||||
|
||||
---
|
||||
|
||||
## 15. 异步任务查询
|
||||
|
||||
> 所有异步操作(同步、重建索引、上传向量化、出题、批阅)返回的 `task_id` 均可通过以下接口查询进度。
|
||||
|
||||
### GET /tasks
|
||||
|
||||
获取任务列表。
|
||||
|
||||
```bash
|
||||
curl -s http://localhost:5001/tasks
|
||||
```
|
||||
|
||||
**查询参数**:
|
||||
| 参数 | 类型 | 必需 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `status` | string | ❌ | 过滤状态:`pending` / `running` / `completed` / `failed` |
|
||||
| `type` | string | ❌ | 过滤类型:`sync` / `reindex` / `upload` / `batch_upload` / `exam_generate` / `exam_grade` |
|
||||
| `limit` | int | ❌ | 返回数量限制(默认 50) |
|
||||
|
||||
**响应示例**:
|
||||
```json
|
||||
{
|
||||
"data": {
|
||||
"tasks": [
|
||||
{
|
||||
"task_id": "a1b2c3d4e5f6",
|
||||
"type": "sync",
|
||||
"description": "文档同步",
|
||||
"status": "running",
|
||||
"progress": 45.0,
|
||||
"current": 9,
|
||||
"total": 20,
|
||||
"stage": "处理文件",
|
||||
"message": "已处理: 产品手册.pdf",
|
||||
"created_at": "2026-06-05T10:30:00"
|
||||
}
|
||||
],
|
||||
"total": 1
|
||||
},
|
||||
"message": "查询成功",
|
||||
"status": "success",
|
||||
"status_code": 2000,
|
||||
"success": true
|
||||
}
|
||||
```
|
||||
|
||||
**验证结果**:✅ 通过
|
||||
|
||||
---
|
||||
|
||||
### GET /tasks/\<task_id\>
|
||||
|
||||
获取单个任务状态(JSON 轮询接口)。
|
||||
|
||||
**后端组推荐使用此接口轮询任务进度,建议间隔 1-2 秒。**
|
||||
|
||||
```bash
|
||||
curl -s http://localhost:5001/tasks/a1b2c3d4e5f6
|
||||
```
|
||||
|
||||
**响应示例(运行中)**:
|
||||
```json
|
||||
{
|
||||
"data": {
|
||||
"task_id": "a1b2c3d4e5f6",
|
||||
"type": "sync",
|
||||
"description": "文档同步",
|
||||
"status": "running",
|
||||
"progress": 45.0,
|
||||
"current": 9,
|
||||
"total": 20,
|
||||
"stage": "处理文件",
|
||||
"message": "已处理: 产品手册.pdf",
|
||||
"created_at": "2026-06-05T10:30:00",
|
||||
"started_at": "2026-06-05T10:30:01"
|
||||
},
|
||||
"message": "查询成功",
|
||||
"status": "success",
|
||||
"status_code": 2000,
|
||||
"success": true
|
||||
}
|
||||
```
|
||||
|
||||
**响应示例(已完成)**:
|
||||
```json
|
||||
{
|
||||
"data": {
|
||||
"task_id": "a1b2c3d4e5f6",
|
||||
"type": "sync",
|
||||
"description": "文档同步",
|
||||
"status": "completed",
|
||||
"progress": 100.0,
|
||||
"current": 20,
|
||||
"total": 20,
|
||||
"stage": "完成",
|
||||
"message": "同步完成",
|
||||
"created_at": "2026-06-05T10:30:00",
|
||||
"started_at": "2026-06-05T10:30:01",
|
||||
"completed_at": "2026-06-05T10:30:15",
|
||||
"duration_ms": 14000,
|
||||
"result": {
|
||||
"documents_processed": 20,
|
||||
"documents_added": 3,
|
||||
"documents_modified": 2,
|
||||
"documents_deleted": 0,
|
||||
"errors": []
|
||||
}
|
||||
},
|
||||
"message": "查询成功",
|
||||
"status": "success",
|
||||
"status_code": 2000,
|
||||
"success": true
|
||||
}
|
||||
```
|
||||
|
||||
> **任务状态**:
|
||||
> - `pending`:已创建,等待执行
|
||||
> - `running`:正在执行
|
||||
> - `completed`:执行完成,`result` 字段包含完整结果
|
||||
> - `failed`:执行失败,`error` 字段包含错误信息
|
||||
>
|
||||
> **轮询建议**:间隔 1-2 秒,当 `status` 为 `completed` 或 `failed` 时停止轮询。
|
||||
|
||||
**验证结果**:✅ 通过
|
||||
|
||||
---
|
||||
|
||||
### GET /tasks/\<task_id\>/progress
|
||||
|
||||
SSE 流式任务进度推送(dev-ui 前端推荐使用)。
|
||||
|
||||
```bash
|
||||
curl -s -N http://localhost:5001/tasks/a1b2c3d4e5f6/progress
|
||||
```
|
||||
|
||||
**SSE 事件序列**:
|
||||
```
|
||||
data: {"type": "start", "data": {"stage": "扫描文档"}}
|
||||
|
||||
data: {"type": "progress", "data": {"progress": 10.0, "current": 2, "total": 20, "stage": "处理文件", "message": "已处理: file1.pdf"}}
|
||||
|
||||
data: {"type": "progress", "data": {"progress": 25.0, "current": 5, "total": 20, "stage": "处理文件", "message": "已处理: file2.docx"}}
|
||||
|
||||
data: {"type": "complete", "data": {"task_id": "a1b2c3d4e5f6", "status": "completed", "result": {...}}}
|
||||
```
|
||||
|
||||
> **SSE 事件类型**:
|
||||
> - `start`:任务开始
|
||||
> - `progress`:进度更新(包含 progress/current/total/stage/message)
|
||||
> - `complete`:任务完成(data 为完整任务详情,含 result)
|
||||
> - `error`:任务失败(data 包含 message 错误信息)
|
||||
> - 每 1 秒发送一次 `: heartbeat` 保活
|
||||
|
||||
**验证结果**:✅ 通过
|
||||
|
||||
---
|
||||
|
||||
### GET /tasks/stats
|
||||
|
||||
获取任务统计信息。
|
||||
|
||||
```bash
|
||||
curl -s http://localhost:5001/tasks/stats
|
||||
```
|
||||
|
||||
**响应示例**:
|
||||
```json
|
||||
{
|
||||
"data": {
|
||||
"total": 5,
|
||||
"by_status": {"running": 1, "completed": 3, "failed": 1},
|
||||
"by_type": {"sync": 2, "exam_generate": 2, "upload": 1}
|
||||
},
|
||||
"message": "查询成功",
|
||||
"status": "success",
|
||||
"status_code": 2000,
|
||||
"success": true
|
||||
}
|
||||
```
|
||||
|
||||
**验证结果**:✅ 通过
|
||||
|
||||
---
|
||||
|
||||
## 附录. 已知问题
|
||||
|
||||
### 1. /rag 接口 collections 参数
|
||||
@@ -1951,11 +2171,11 @@ curl -s -X POST http://localhost:5001/kb/route \
|
||||
当代码升级涉及元数据字段变更时(如新增 `chunk_index`、`doc_type`),需要重构向量库:
|
||||
|
||||
```bash
|
||||
# 重构指定知识库(清除哈希记录 → 触发全量同步)
|
||||
# 重构指定知识库(清除哈希记录 → 触发全量同步,异步任务)
|
||||
curl -s -X POST http://127.0.0.1:5001/collections/<kb_name>/reindex
|
||||
```
|
||||
|
||||
> **⚠️ 注意**:reindex 会调用 `sync_now()` 全局同步,期间 gunicorn worker 被阻塞,搜索和问答接口将暂时无响应。建议在低峰期执行。
|
||||
> **⚠️ 注意**:reindex 已改为异步任务,立即返回 `task_id`。通过 `GET /tasks/<task_id>` 轮询进度。不再阻塞 gunicorn worker,搜索和问答接口不受影响。
|
||||
|
||||
### Reranker 配置
|
||||
|
||||
|
||||
255
docs/后端对接规范.md
255
docs/后端对接规范.md
@@ -4,9 +4,33 @@
|
||||
|
||||
## 📋 变更记录(2026-06-05 更新)
|
||||
|
||||
> **本次更新内容**:新增 AI 智能出题端口、更新生产环境测试结果
|
||||
> **本次更新内容**:新增 AI 智能出题端口、更新生产环境测试结果、**长操作改为异步任务**
|
||||
>
|
||||
> **2026-06-05 更新**:出题批卷接口格式优化与输入校验增强
|
||||
>
|
||||
> **2026-06-05 异步任务变更**:同步、上传向量化、出题、批阅等长耗时操作改为异步任务模式,立即返回 `task_id`,通过 `GET /tasks/<task_id>` 轮询结果
|
||||
|
||||
### 异步任务变更(⚠️ 重要,2026-06-05)
|
||||
|
||||
| 端点 | 变更说明 |
|
||||
|------|----------|
|
||||
| `POST /sync` | 改为异步任务,返回 `{"task_id": "xxx"}` 而非同步结果 |
|
||||
| `POST /documents/sync` | 改为异步任务,返回 `{"task_id": "xxx"}` |
|
||||
| `POST /collections/<kb>/reindex` | 改为异步任务,返回 `{"task_id": "xxx"}` |
|
||||
| `POST /documents/upload` | 新增 `task_id` 字段(向量化后台执行) |
|
||||
| `POST /documents/batch-upload` | 新增 `task_id` 字段(批量向量化后台执行) |
|
||||
| `POST /exam/generate` | 改为异步任务,返回 `{"task_id": "xxx"}` |
|
||||
| `POST /exam/generate-smart` | 改为异步任务,返回 `{"task_id": "xxx"}` |
|
||||
| `POST /exam/grade` | 改为异步任务,返回 `{"task_id": "xxx"}` |
|
||||
|
||||
**新增任务查询接口**:
|
||||
|
||||
| 端点 | 方法 | 功能 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `/tasks` | GET | 任务列表 | 支持按 status/type 过滤 |
|
||||
| `/tasks/<task_id>` | GET | 任务状态(JSON) | 后端组推荐轮询接口,建议 1-2 秒间隔 |
|
||||
| `/tasks/<task_id>/progress` | GET | 任务进度(SSE) | dev-ui 前端推荐使用 |
|
||||
| `/tasks/stats` | GET | 任务统计 | 按状态和类型分组统计 |
|
||||
|
||||
### 新增端口
|
||||
|
||||
@@ -123,10 +147,22 @@ RAG服务负责:
|
||||
|
||||
### 2.4 出题系统(可选)
|
||||
|
||||
| 端点 | 方法 | 功能 |
|
||||
|-----|------|------|
|
||||
| `/exam/generate` | POST | 生成题目 |
|
||||
| `/exam/grade` | POST | 批阅答案 |
|
||||
| 端点 | 方法 | 功能 | 说明 |
|
||||
|-----|------|------|------|
|
||||
| `/exam/generate` | POST | 生成题目 | 异步任务,返回 task_id |
|
||||
| `/exam/generate-smart` | POST | AI 智能出题 | 异步任务,返回 task_id |
|
||||
| `/exam/grade` | POST | 批阅答案 | 异步任务,返回 task_id |
|
||||
|
||||
### 2.5 异步任务查询
|
||||
|
||||
> 所有异步操作(同步、重建索引、上传向量化、出题、批阅)返回的 `task_id` 均可通过以下接口查询进度。
|
||||
|
||||
| 端点 | 方法 | 功能 | 说明 |
|
||||
|-----|------|------|------|
|
||||
| `/tasks` | GET | 任务列表 | 支持按 status/type 过滤 |
|
||||
| `/tasks/<task_id>` | GET | 任务状态(JSON) | **后端组推荐轮询接口**,建议 1-2 秒间隔 |
|
||||
| `/tasks/<task_id>/progress` | GET | 任务进度(SSE) | dev-ui 前端推荐使用 |
|
||||
| `/tasks/stats` | GET | 任务统计 | 按状态和类型分组统计 |
|
||||
|
||||
---
|
||||
|
||||
@@ -247,7 +283,7 @@ ENABLE_DIFY_WORKFLOW=false
|
||||
|
||||
---
|
||||
|
||||
**最后更新**: 2026-04-29
|
||||
**最后更新**: 2026-06-05
|
||||
|
||||
> 本文档供后端开发人员参考,用于对接 RAG 知识库服务。
|
||||
|
||||
@@ -1089,17 +1125,24 @@ Content-Type: multipart/form-data
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"message": "文件上传成功,已保存并添加到向量库",
|
||||
"file": {
|
||||
"filename": "document.pdf",
|
||||
"collection": "public_kb",
|
||||
"path": "public_kb/document.pdf",
|
||||
"size": 1024000,
|
||||
"replaced": false
|
||||
"status_code": 2002,
|
||||
"message": "文件上传成功,已保存,向量化任务已启动",
|
||||
"data": {
|
||||
"file": {
|
||||
"filename": "document.pdf",
|
||||
"collection": "public_kb",
|
||||
"path": "public_kb/document.pdf",
|
||||
"size": 1024000,
|
||||
"replaced": false
|
||||
},
|
||||
"sync_status": "已保存,向量化任务已启动",
|
||||
"task_id": "a1b2c3d4e5f6"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
> **异步说明**:文件保存为同步操作,向量化在后台线程异步执行。响应中的 `task_id` 可用于轮询向量化进度(`GET /tasks/<task_id>`)。当同步服务不可用时,`task_id` 为 `null`,`sync_status` 为 `"已保存,等待手动同步"`。
|
||||
|
||||
**同名文件处理**:上传同名文件时,旧版本的切片会被自动清理后覆盖(`replaced: true`),不会生成时间戳后缀文件。这确保了向量库中不会出现同一文档的新旧切片共存的情况。
|
||||
|
||||
### 5.2 批量上传
|
||||
@@ -1257,26 +1300,51 @@ POST /sync
|
||||
|
||||
**请求体:** 无需传递参数(同步所有知识库)
|
||||
|
||||
**响应:**
|
||||
**响应(异步任务):**
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status": "success",
|
||||
"status_code": 2010,
|
||||
"message": "同步完成",
|
||||
"message": "同步任务已启动",
|
||||
"data": {
|
||||
"result": {
|
||||
"documents_added": 1,
|
||||
"documents_deleted": 1,
|
||||
"documents_modified": 0,
|
||||
"documents_processed": 2,
|
||||
"errors": [],
|
||||
"status": "completed"
|
||||
}
|
||||
"task_id": "c3d4e5f6a1b2",
|
||||
"message": "同步任务已启动,通过 GET /tasks/c3d4e5f6a1b2 查询进度"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
> **⚠️ 异步变更**:此接口已从同步改为异步。不再直接返回同步结果,而是返回 `task_id`。后端需通过 `GET /tasks/<task_id>` 轮询任务状态,直到 `status` 为 `completed` 或 `failed`。任务完成后,`result` 字段包含完整的同步结果(含 `documents_processed`、`documents_added` 等)。
|
||||
>
|
||||
> **冲突检测**:如果已有同步任务正在运行,返回 HTTP 409:`{"error": "TASK_RUNNING", "message": "同步任务正在执行中 (task_id: xxx),请等待完成"}`
|
||||
|
||||
**后端轮询示例**:
|
||||
|
||||
```python
|
||||
import time
|
||||
import requests
|
||||
|
||||
def trigger_sync_and_wait():
|
||||
"""触发同步并等待完成"""
|
||||
# 1. 触发同步任务
|
||||
resp = requests.post('http://rag-service:5001/sync')
|
||||
task_id = resp.json()['data']['task_id']
|
||||
|
||||
# 2. 轮询任务状态(每 2 秒)
|
||||
while True:
|
||||
time.sleep(2)
|
||||
status_resp = requests.get(f'http://rag-service:5001/tasks/{task_id}')
|
||||
task_data = status_resp.json()['data']
|
||||
|
||||
if task_data['status'] == 'completed':
|
||||
print(f"同步完成: {task_data['result']}")
|
||||
return task_data['result']
|
||||
elif task_data['status'] == 'failed':
|
||||
raise Exception(f"同步失败: {task_data['error']}")
|
||||
else:
|
||||
print(f"同步中: {task_data['progress']}% - {task_data['message']}")
|
||||
```
|
||||
```
|
||||
|
||||
### 6.2 同步状态
|
||||
|
||||
```
|
||||
@@ -1359,7 +1427,7 @@ POST /sync/stop
|
||||
```json
|
||||
{
|
||||
"status": "success",
|
||||
"status_code": 3001,
|
||||
"status_code": 2010,
|
||||
"message": "文件监控已启动"
|
||||
}
|
||||
```
|
||||
@@ -1417,26 +1485,24 @@ POST /exam/generate
|
||||
| difficulty | 必须为 1-5 的整数 | HTTP 400 INVALID_PARAMS |
|
||||
| **总题数上限** | **所有题型数量之和不能超过 20** | HTTP 400 INVALID_PARAMS |
|
||||
|
||||
**响应:**
|
||||
**响应(异步任务):**
|
||||
|
||||
**完整响应格式**(包含外层包装):
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status": "success",
|
||||
"status_code": 2011,
|
||||
"message": "出题完成",
|
||||
"status_code": 2020,
|
||||
"message": "出题任务已启动",
|
||||
"data": {
|
||||
"success": true,
|
||||
"request_id": "xxx",
|
||||
"total": 10,
|
||||
"source_chunks_used": 15,
|
||||
"questions": [...]
|
||||
"task_id": "d4e5f6a1b2c3",
|
||||
"message": "出题任务已启动 (10题),通过 GET /tasks/d4e5f6a1b2c3 查询结果"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**data 内部结构**:
|
||||
> **⚠️ 异步变更**:此接口已从同步改为异步。响应仅返回 `task_id`,后端需通过 `GET /tasks/<task_id>` 轮询任务状态。任务完成后,`result` 字段包含完整出题结果(格式见下方说明)。
|
||||
|
||||
**轮询结果(GET /tasks/\<task_id\> 完成后的 result 字段)**:
|
||||
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
@@ -1635,25 +1701,22 @@ POST /exam/grade
|
||||
|
||||
#### 响应
|
||||
|
||||
**完整响应格式**(包含外层包装):
|
||||
**响应(异步任务):**
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status": "success",
|
||||
"status_code": 2021,
|
||||
"message": "批阅完成",
|
||||
"message": "批阅任务已启动",
|
||||
"data": {
|
||||
"request_id": "可选,原样返回",
|
||||
"success": true,
|
||||
"total_score": 12.5,
|
||||
"total_max_score": 22.0,
|
||||
"score_rate": 56.8,
|
||||
"results": [...]
|
||||
"task_id": "f6a1b2c3d4e5",
|
||||
"message": "批阅任务已启动 (5题),通过 GET /tasks/f6a1b2c3d4e5 查询结果"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**data 内部结构**:
|
||||
> **⚠️ 异步变更**:此接口已从同步改为异步。响应仅返回 `task_id`,后端需通过 `GET /tasks/<task_id>` 轮询任务状态。任务完成后,`result` 字段包含完整批阅结果(格式见下方说明)。
|
||||
|
||||
**轮询结果(GET /tasks/\<task_id\> 完成后的 result 字段)**:
|
||||
|
||||
```json
|
||||
{
|
||||
@@ -1817,16 +1880,30 @@ def build_grade_request(answer_list, questions_map):
|
||||
return {"answers": grade_answers}
|
||||
```
|
||||
|
||||
**Step 4: 调用 RAG 批卷接口**
|
||||
**Step 4: 调用 RAG 批卷接口(异步任务)**
|
||||
|
||||
```python
|
||||
import time
|
||||
|
||||
def call_rag_grade(grade_request):
|
||||
"""调用 RAG 批卷接口"""
|
||||
"""调用 RAG 批卷接口并轮询等待结果"""
|
||||
# 1. 提交批阅任务
|
||||
response = requests.post(
|
||||
'http://rag-service:5001/exam/grade',
|
||||
json=grade_request
|
||||
)
|
||||
return response.json()
|
||||
task_id = response.json()['data']['task_id']
|
||||
|
||||
# 2. 轮询任务状态(每 2 秒)
|
||||
while True:
|
||||
time.sleep(2)
|
||||
status_resp = requests.get(f'http://rag-service:5001/tasks/{task_id}')
|
||||
task_data = status_resp.json()['data']
|
||||
|
||||
if task_data['status'] == 'completed':
|
||||
return task_data['result']
|
||||
elif task_data['status'] == 'failed':
|
||||
raise Exception(f"批阅失败: {task_data['error']}")
|
||||
```
|
||||
|
||||
**Step 5: 更新学生成绩**
|
||||
@@ -2430,19 +2507,22 @@ ChromaDB 集合名称限制:
|
||||
|
||||
| 端点 | 方法 | 说明 |
|
||||
|------|------|------|
|
||||
| `/sync` | POST | 触发同步 |
|
||||
| `/sync` | POST | 触发同步(**异步任务,返回 task_id**) |
|
||||
| `/sync/status` | GET | 同步状态 |
|
||||
| `/sync/history` | GET | 同步历史 |
|
||||
| `/sync/changes` | GET | 变更日志 |
|
||||
| `/sync/start` | POST | 启动文件监控 |
|
||||
| `/sync/stop` | POST | 停止文件监控 |
|
||||
| `/documents/sync` | POST | 触发文档同步(**异步任务,返回 task_id**) |
|
||||
| `/collections/<kb_name>/reindex` | POST | 重建索引(**异步任务,返回 task_id**) |
|
||||
|
||||
### 13.7 出题系统
|
||||
|
||||
| 端点 | 方法 | 说明 |
|
||||
|------|------|------|
|
||||
| `/exam/generate` | POST | 生成题目 |
|
||||
| `/exam/grade` | POST | 批改答案 |
|
||||
| `/exam/generate` | POST | 生成题目(**异步任务,返回 task_id**) |
|
||||
| `/exam/generate-smart` | POST | AI 智能出题(**异步任务,返回 task_id**) |
|
||||
| `/exam/grade` | POST | 批改答案(**异步任务,返回 task_id**) |
|
||||
| `/exam/health` | GET | 出题服务健康检查 |
|
||||
|
||||
### 13.8 反馈与 FAQ 管理
|
||||
@@ -2497,6 +2577,77 @@ ChromaDB 集合名称限制:
|
||||
|------|------|------|
|
||||
| `/documents/<path>/preview` | GET | 文档预览,按 `chunk_index` 跳转到具体切片(dev-ui 引用溯源用) |
|
||||
|
||||
### 13.13 异步任务查询
|
||||
|
||||
> 所有异步操作(同步、重建索引、上传向量化、出题、批阅)返回的 `task_id` 均可通过以下接口查询进度。
|
||||
|
||||
| 端点 | 方法 | 说明 |
|
||||
|------|------|------|
|
||||
| `/tasks` | GET | 任务列表(支持 status/type 过滤) |
|
||||
| `/tasks/<task_id>` | GET | 任务状态(JSON 轮询,**后端组推荐使用**) |
|
||||
| `/tasks/<task_id>/progress` | GET | 任务进度(SSE 流式,dev-ui 前端使用) |
|
||||
| `/tasks/stats` | GET | 任务统计 |
|
||||
|
||||
**任务状态字段说明**:
|
||||
|
||||
| 字段 | 类型 | 说明 |
|
||||
|------|------|------|
|
||||
| `task_id` | string | 任务唯一标识 |
|
||||
| `type` | string | 类型:sync / reindex / upload / batch_upload / exam_generate / exam_grade |
|
||||
| `status` | string | 状态:pending / running / completed / failed |
|
||||
| `progress` | float | 进度百分比(0-100) |
|
||||
| `current` | int | 当前处理项数 |
|
||||
| `total` | int | 总项数 |
|
||||
| `stage` | string | 当前阶段 |
|
||||
| `message` | string | 当前步骤描述 |
|
||||
| `result` | any | 完成后的结果数据(仅 status=completed 时存在) |
|
||||
| `error` | string | 失败错误信息(仅 status=failed 时存在) |
|
||||
| `duration_ms` | int | 执行耗时毫秒(仅已完成时存在) |
|
||||
|
||||
**后端对接轮询模式**:
|
||||
|
||||
```python
|
||||
import time
|
||||
import requests
|
||||
|
||||
def async_task_poll(task_id, base_url='http://rag-service:5001', interval=2, timeout=300):
|
||||
"""
|
||||
通用异步任务轮询函数
|
||||
|
||||
Args:
|
||||
task_id: 任务 ID
|
||||
base_url: RAG 服务地址
|
||||
interval: 轮询间隔(秒)
|
||||
timeout: 超时时间(秒)
|
||||
|
||||
Returns:
|
||||
任务结果(result 字段)
|
||||
|
||||
Raises:
|
||||
TimeoutError: 超时
|
||||
Exception: 任务失败
|
||||
"""
|
||||
elapsed = 0
|
||||
while elapsed < timeout:
|
||||
time.sleep(interval)
|
||||
elapsed += interval
|
||||
|
||||
resp = requests.get(f'{base_url}/tasks/{task_id}')
|
||||
if resp.status_code == 404:
|
||||
raise Exception(f"任务不存在: {task_id}")
|
||||
|
||||
task = resp.json()['data']
|
||||
|
||||
if task['status'] == 'completed':
|
||||
return task.get('result')
|
||||
elif task['status'] == 'failed':
|
||||
raise Exception(f"任务失败: {task.get('error', '未知错误')}")
|
||||
# 可选:记录进度日志
|
||||
# logger.info(f"任务 {task_id}: {task['progress']}% - {task['message']}")
|
||||
|
||||
raise TimeoutError(f"任务超时: {task_id} (已等待 {timeout}s)")
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 十四、文件管理服务(后端负责)
|
||||
|
||||
@@ -444,7 +444,7 @@ Query Rewriting: "它" → "出差补助"
|
||||
|
||||
- [后端对接规范.md](./后端对接规范.md) - API 接口规范(主要)
|
||||
- [数据库设计文档.md](./数据库设计文档.md) - 数据库结构
|
||||
- [Agentic_RAG完整指南.md](./Agentic_RAG完整指南.md) - Agentic RAG 详解
|
||||
- [RAG系统完整指南.md](./RAG系统完整指南.md) - RAG 系统架构与缓存详解
|
||||
|
||||
|
||||
---
|
||||
|
||||
554
docs/异步任务接口变更说明.md
Normal file
554
docs/异步任务接口变更说明.md
Normal file
@@ -0,0 +1,554 @@
|
||||
# 异步任务接口变更说明
|
||||
|
||||
> **变更日期**:2026-06-05
|
||||
>
|
||||
> **变更原因**:同步、向量化、出题、批阅等操作耗时较长(数秒到数十秒),改为异步任务模式后接口立即返回 `task_id`,避免后端调用超时。
|
||||
>
|
||||
> **影响范围**:8 个已有端口的响应格式变更 + 4 个新增任务查询接口
|
||||
|
||||
---
|
||||
|
||||
## 一、核心变更
|
||||
|
||||
所有受影响的接口从 **同步阻塞返回结果** 改为 **异步任务立即返回 `task_id`**。
|
||||
|
||||
**调用流程变更**:
|
||||
|
||||
```
|
||||
旧流程:
|
||||
POST /sync → 等待 10-30s → 返回完整结果
|
||||
|
||||
新流程:
|
||||
POST /sync → 立即返回 {"task_id": "xxx"}
|
||||
GET /tasks/xxx → 轮询(1-2s 间隔)→ status=completed 时取 result
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 二、受影响的端口(8 个)
|
||||
|
||||
### 1. POST /sync — 触发同步
|
||||
|
||||
**旧响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status_code": 2010,
|
||||
"message": "同步完成",
|
||||
"data": {
|
||||
"result": {
|
||||
"documents_processed": 20,
|
||||
"documents_added": 3,
|
||||
"documents_modified": 2,
|
||||
"documents_deleted": 0,
|
||||
"errors": []
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**新响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status_code": 2010,
|
||||
"message": "同步任务已启动",
|
||||
"data": {
|
||||
"task_id": "c3d4e5f6a1b2",
|
||||
"message": "同步任务已启动,通过 GET /tasks/c3d4e5f6a1b2 查询进度"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**冲突检测**:已有同步任务运行时返回 HTTP 409:
|
||||
```json
|
||||
{"error": "TASK_RUNNING", "message": "同步任务正在执行中 (task_id: xxx),请等待完成"}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 2. POST /documents/sync — 触发文档同步
|
||||
|
||||
**旧响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"results": [{"collection": "public_kb", "status": "success"}],
|
||||
"synced_count": 1
|
||||
}
|
||||
```
|
||||
|
||||
**新响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"task_id": "xxx",
|
||||
"message": "同步任务已启动,通过 GET /tasks/xxx 查询进度"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 3. POST /collections/\<kb_name\>/reindex — 重建索引
|
||||
|
||||
**旧响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"message": "重新索引完成: 处理 20 个文档",
|
||||
"documents_processed": 20,
|
||||
"documents_added": 3
|
||||
}
|
||||
```
|
||||
|
||||
**新响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"task_id": "xxx",
|
||||
"message": "重建索引任务已启动: kb_name,通过 GET /tasks/xxx 查询进度"
|
||||
}
|
||||
```
|
||||
|
||||
**冲突检测**:已有重建任务运行时返回 HTTP 409。
|
||||
|
||||
---
|
||||
|
||||
### 4. POST /documents/upload — 上传单个文件
|
||||
|
||||
**旧响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status_code": 2002,
|
||||
"message": "文件上传成功,已保存并添加到向量库",
|
||||
"data": {
|
||||
"file": {
|
||||
"filename": "test.txt",
|
||||
"collection": "public_kb",
|
||||
"path": "public_kb/test.txt",
|
||||
"size": 1024,
|
||||
"replaced": false
|
||||
},
|
||||
"sync_status": "已保存并添加到向量库"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**新响应**(新增 `task_id` 字段):
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status_code": 2002,
|
||||
"message": "文件上传成功,已保存,向量化任务已启动",
|
||||
"data": {
|
||||
"file": {
|
||||
"filename": "test.txt",
|
||||
"collection": "public_kb",
|
||||
"path": "public_kb/test.txt",
|
||||
"size": 1024,
|
||||
"replaced": false
|
||||
},
|
||||
"sync_status": "已保存,向量化任务已启动",
|
||||
"task_id": "a1b2c3d4e5f6"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
> **注意**:文件保存仍为同步操作(毫秒级),仅向量化部分异步执行。`task_id` 为 `null` 表示同步服务不可用,需手动同步。
|
||||
|
||||
---
|
||||
|
||||
### 5. POST /documents/batch-upload — 批量上传
|
||||
|
||||
**旧响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status_code": 2003,
|
||||
"data": {
|
||||
"results": [...],
|
||||
"success_count": 2,
|
||||
"total": 2
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**新响应**(新增 `task_id` 字段):
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status_code": 2003,
|
||||
"data": {
|
||||
"results": [...],
|
||||
"success_count": 2,
|
||||
"total": 2,
|
||||
"task_id": "b2c3d4e5f6a1"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
> **注意**:批量上传成功后自动触发后台向量化任务。无成功上传时不返回 `task_id`。
|
||||
|
||||
---
|
||||
|
||||
### 6. POST /exam/generate — 生成题目
|
||||
|
||||
**旧响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status_code": 2020,
|
||||
"message": "出题完成",
|
||||
"data": {
|
||||
"success": true,
|
||||
"total": 10,
|
||||
"questions": [...],
|
||||
"warnings": []
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**新响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status_code": 2020,
|
||||
"message": "出题任务已启动",
|
||||
"data": {
|
||||
"task_id": "d4e5f6a1b2c3",
|
||||
"message": "出题任务已启动 (10题),通过 GET /tasks/d4e5f6a1b2c3 查询结果"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**轮询完成后的 `result` 字段**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"total": 10,
|
||||
"source_chunks_used": 15,
|
||||
"requested_types": {"single_choice": 5, "true_false": 3, "fill_blank": 2},
|
||||
"actual_types": {"single_choice": 5, "true_false": 3, "fill_blank": 2},
|
||||
"warnings": [],
|
||||
"questions": [...]
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 7. POST /exam/generate-smart — AI 智能出题
|
||||
|
||||
**旧响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status_code": 2020,
|
||||
"message": "AI 智能出题成功",
|
||||
"data": {
|
||||
"ai_analysis": {...},
|
||||
"questions": [...],
|
||||
"total": 16
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**新响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status_code": 2020,
|
||||
"message": "AI 智能出题任务已启动",
|
||||
"data": {
|
||||
"task_id": "e5f6a1b2c3d4",
|
||||
"message": "AI 智能出题任务已启动,通过 GET /tasks/e5f6a1b2c3d4 查询结果"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**轮询完成后的 `result` 字段**:与旧 `data` 格式相同(含 `ai_analysis`、`questions`、`total`)。
|
||||
|
||||
---
|
||||
|
||||
### 8. POST /exam/grade — 批阅答案
|
||||
|
||||
**旧响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status_code": 2021,
|
||||
"message": "批阅完成",
|
||||
"data": {
|
||||
"success": true,
|
||||
"total_score": 12.5,
|
||||
"total_max_score": 22.0,
|
||||
"score_rate": 56.8,
|
||||
"results": [...]
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**新响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status_code": 2021,
|
||||
"message": "批阅任务已启动",
|
||||
"data": {
|
||||
"task_id": "f6a1b2c3d4e5",
|
||||
"message": "批阅任务已启动 (5题),通过 GET /tasks/f6a1b2c3d4e5 查询结果"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**轮询完成后的 `result` 字段**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"total_score": 12.5,
|
||||
"total_max_score": 22.0,
|
||||
"score_rate": 56.8,
|
||||
"results": [
|
||||
{
|
||||
"question_id": "uuid-001",
|
||||
"score": 0,
|
||||
"max_score": 2.0,
|
||||
"grading_status": "success",
|
||||
"details": {
|
||||
"correct": false,
|
||||
"student_answer": "A",
|
||||
"correct_answer": "B",
|
||||
"feedback": "正确答案: B"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 三、新增接口(4 个)
|
||||
|
||||
### GET /tasks
|
||||
|
||||
获取任务列表。
|
||||
|
||||
**查询参数**:`status`(过滤状态)、`type`(过滤类型)、`limit`(返回数量,默认 50)
|
||||
|
||||
**响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status_code": 2000,
|
||||
"data": {
|
||||
"tasks": [
|
||||
{
|
||||
"task_id": "a1b2c3d4e5f6",
|
||||
"type": "sync",
|
||||
"description": "文档同步",
|
||||
"status": "running",
|
||||
"progress": 45.0,
|
||||
"current": 9,
|
||||
"total": 20,
|
||||
"stage": "处理文件",
|
||||
"message": "已处理: 产品手册.pdf",
|
||||
"created_at": "2026-06-05T10:30:00"
|
||||
}
|
||||
],
|
||||
"total": 1
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### GET /tasks/\<task_id\>
|
||||
|
||||
获取单个任务状态(**后端组推荐使用的轮询接口**)。
|
||||
|
||||
**建议轮询间隔**:1-2 秒,当 `status` 为 `completed` 或 `failed` 时停止。
|
||||
|
||||
**响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status_code": 2000,
|
||||
"data": {
|
||||
"task_id": "a1b2c3d4e5f6",
|
||||
"type": "sync",
|
||||
"description": "文档同步",
|
||||
"status": "completed",
|
||||
"progress": 100.0,
|
||||
"current": 20,
|
||||
"total": 20,
|
||||
"stage": "完成",
|
||||
"message": "同步完成",
|
||||
"created_at": "2026-06-05T10:30:00",
|
||||
"started_at": "2026-06-05T10:30:01",
|
||||
"completed_at": "2026-06-05T10:30:15",
|
||||
"duration_ms": 14000,
|
||||
"result": {
|
||||
"documents_processed": 20,
|
||||
"documents_added": 3
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**任务状态**:
|
||||
|
||||
| 状态 | 说明 |
|
||||
|------|------|
|
||||
| `pending` | 已创建,等待执行 |
|
||||
| `running` | 正在执行 |
|
||||
| `completed` | 执行完成,`result` 字段包含完整结果 |
|
||||
| `failed` | 执行失败,`error` 字段包含错误信息 |
|
||||
|
||||
**任务类型(type)**:`sync` / `reindex` / `upload` / `batch_upload` / `exam_generate` / `exam_grade`
|
||||
|
||||
---
|
||||
|
||||
### GET /tasks/\<task_id\>/progress
|
||||
|
||||
SSE 流式任务进度推送(dev-ui 前端推荐使用,后端组通常不需要此接口)。
|
||||
|
||||
**Content-Type**:`text/event-stream`
|
||||
|
||||
---
|
||||
|
||||
### GET /tasks/stats
|
||||
|
||||
获取任务统计信息。
|
||||
|
||||
**响应**:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"status_code": 2000,
|
||||
"data": {
|
||||
"total": 5,
|
||||
"by_status": {"running": 1, "completed": 3, "failed": 1},
|
||||
"by_type": {"sync": 2, "exam_generate": 2, "upload": 1}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 四、后端对接代码参考
|
||||
|
||||
### 通用轮询函数
|
||||
|
||||
```python
|
||||
import time
|
||||
import requests
|
||||
|
||||
def async_task_poll(task_id, base_url='http://rag-service:5001', interval=2, timeout=300):
|
||||
"""
|
||||
通用异步任务轮询函数
|
||||
|
||||
Args:
|
||||
task_id: 任务 ID(从 POST 响应中获取)
|
||||
base_url: RAG 服务地址
|
||||
interval: 轮询间隔(秒),建议 1-2 秒
|
||||
timeout: 超时时间(秒)
|
||||
|
||||
Returns:
|
||||
任务结果(result 字段内容)
|
||||
|
||||
Raises:
|
||||
TimeoutError: 超时
|
||||
Exception: 任务失败
|
||||
"""
|
||||
elapsed = 0
|
||||
while elapsed < timeout:
|
||||
time.sleep(interval)
|
||||
elapsed += interval
|
||||
|
||||
resp = requests.get(f'{base_url}/tasks/{task_id}')
|
||||
if resp.status_code == 404:
|
||||
raise Exception(f"任务不存在: {task_id}")
|
||||
|
||||
task = resp.json()['data']
|
||||
|
||||
if task['status'] == 'completed':
|
||||
return task.get('result')
|
||||
elif task['status'] == 'failed':
|
||||
raise Exception(f"任务失败: {task.get('error', '未知错误')}")
|
||||
|
||||
raise TimeoutError(f"任务超时: {task_id} (已等待 {timeout}s)")
|
||||
```
|
||||
|
||||
### 出题流程示例
|
||||
|
||||
```python
|
||||
def generate_exam(file_path, collection, question_types, difficulty=3):
|
||||
"""异步出题流程"""
|
||||
# 1. 提交出题任务
|
||||
resp = requests.post('http://rag-service:5001/exam/generate', json={
|
||||
'file_path': file_path,
|
||||
'collection': collection,
|
||||
'question_types': question_types,
|
||||
'difficulty': difficulty
|
||||
})
|
||||
task_id = resp.json()['data']['task_id']
|
||||
|
||||
# 2. 轮询等待结果
|
||||
result = async_task_poll(task_id)
|
||||
|
||||
# 3. result 包含完整的出题结果
|
||||
questions = result['questions']
|
||||
return questions
|
||||
```
|
||||
|
||||
### 批阅流程示例
|
||||
|
||||
```python
|
||||
def grade_exam(answers):
|
||||
"""异步批阅流程"""
|
||||
# 1. 提交批阅任务
|
||||
resp = requests.post('http://rag-service:5001/exam/grade', json={
|
||||
'answers': answers
|
||||
})
|
||||
task_id = resp.json()['data']['task_id']
|
||||
|
||||
# 2. 轮询等待结果
|
||||
result = async_task_poll(task_id)
|
||||
|
||||
# 3. result 包含完整的批阅结果
|
||||
return result
|
||||
```
|
||||
|
||||
### 同步流程示例
|
||||
|
||||
```python
|
||||
def trigger_sync():
|
||||
"""异步同步流程"""
|
||||
# 1. 提交同步任务
|
||||
resp = requests.post('http://rag-service:5001/sync')
|
||||
|
||||
if resp.status_code == 409:
|
||||
print("已有同步任务在运行,跳过")
|
||||
return
|
||||
|
||||
task_id = resp.json()['data']['task_id']
|
||||
|
||||
# 2. 轮询等待结果
|
||||
result = async_task_poll(task_id)
|
||||
print(f"同步完成: 处理 {result['documents_processed']} 个文档")
|
||||
return result
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 五、注意事项
|
||||
|
||||
1. **超时时间**:异步任务完成后保留 1 小时,超时后自动清理。请在任务完成后及时获取结果。
|
||||
|
||||
2. **并发限制**:同一类型的任务(如 `sync`)同一时间只能有一个运行中,重复提交返回 HTTP 409。
|
||||
|
||||
3. **向后兼容**:如果 RAG 服务未升级(旧版本),POST 接口仍返回旧的同步结果格式。后端可通过检查响应中是否包含 `task_id` 字段来判断版本。
|
||||
|
||||
4. **轮询频率**:建议 1-2 秒间隔,过于频繁的轮询会增加服务器负担。
|
||||
|
||||
5. **错误处理**:任务可能因 LLM 超时、文件解析失败等原因失败,`status` 变为 `failed`,`error` 字段包含错误描述。
|
||||
166
exam_pkg/api.py
166
exam_pkg/api.py
@@ -2,8 +2,8 @@
|
||||
出题与批题系统 API 蓝图
|
||||
|
||||
提供 REST API 接口:
|
||||
- 出题:生成题目(返回 JSON 给后端)
|
||||
- 批题:批阅答案(返回结果给后端)
|
||||
- 出题:生成题目(异步任务,返回 task_id)
|
||||
- 批题:批阅答案(异步任务,返回 task_id)
|
||||
|
||||
职责边界:
|
||||
- RAG 服务负责:生成题目 + 批阅答案
|
||||
@@ -12,6 +12,11 @@
|
||||
使用方式:
|
||||
from exam_pkg.api import exam_bp
|
||||
app.register_blueprint(exam_bp, url_prefix='/exam')
|
||||
|
||||
异步任务流程:
|
||||
1. POST /exam/generate → 返回 {"task_id": "xxx", ...}
|
||||
2. GET /tasks/xxx → 轮询状态,直到 completed
|
||||
3. result 字段包含完整出题/批阅结果
|
||||
"""
|
||||
|
||||
from flask import Blueprint, request, jsonify
|
||||
@@ -29,7 +34,7 @@ from auth.gateway import (
|
||||
)
|
||||
|
||||
# 导入统一响应格式
|
||||
from core.status_codes import EXAM_SUCCESS, GRADE_SUCCESS, EXAM_ERROR, GRADE_ERROR, BAD_REQUEST, NO_CONTENT, LLM_ERROR
|
||||
from core.status_codes import EXAM_SUCCESS, GRADE_SUCCESS, EXAM_ERROR, GRADE_ERROR, BAD_REQUEST, UNAUTHORIZED, FORBIDDEN, NO_CONTENT, LLM_ERROR
|
||||
from api.response_utils import success_response, error_response
|
||||
|
||||
# 合法题型
|
||||
@@ -169,7 +174,7 @@ def api_generate_questions():
|
||||
# 获取当前用户
|
||||
user = get_current_user()
|
||||
if not user:
|
||||
return error_response("UNAUTHORIZED", BAD_REQUEST, "未认证", http_status=401)
|
||||
return error_response("UNAUTHORIZED", UNAUTHORIZED, "未认证", http_status=401)
|
||||
|
||||
# 检查向量库访问权限
|
||||
if not check_collection_permission(
|
||||
@@ -178,20 +183,51 @@ def api_generate_questions():
|
||||
collection_name=collection,
|
||||
operation="read"
|
||||
):
|
||||
return error_response("FORBIDDEN", BAD_REQUEST, "权限不足", http_status=403)
|
||||
return error_response("FORBIDDEN", FORBIDDEN, "权限不足", http_status=403)
|
||||
|
||||
# 调用新版出题接口
|
||||
result = generate_questions_from_file(
|
||||
file_path=file_path,
|
||||
collection=collection,
|
||||
question_types=question_types,
|
||||
difficulty=data.get('difficulty', 3),
|
||||
options=data.get('options', {}),
|
||||
request_id=data.get('request_id'),
|
||||
exclude_stems=data.get('exclude_stems')
|
||||
# 调用新版出题接口(异步任务)
|
||||
from core.task_registry import get_registry
|
||||
import logging as _logging
|
||||
_logger = _logging.getLogger(__name__)
|
||||
|
||||
registry = get_registry()
|
||||
total_questions = sum(question_types.values())
|
||||
task = registry.create_task(
|
||||
'exam_generate',
|
||||
f"出题: {os.path.basename(file_path)} ({total_questions}题)",
|
||||
total=total_questions
|
||||
)
|
||||
|
||||
return success_response(data=result, status_code=EXAM_SUCCESS, message="出题成功")
|
||||
def _do_generate(task, fp, coll, q_types, diff, opts, req_id, excl):
|
||||
"""后台执行出题"""
|
||||
registry.update_progress(task.id, stage='检索知识', message='正在检索相关文档切片...')
|
||||
result = generate_questions_from_file(
|
||||
file_path=fp,
|
||||
collection=coll,
|
||||
question_types=q_types,
|
||||
difficulty=diff,
|
||||
options=opts,
|
||||
request_id=req_id,
|
||||
exclude_stems=excl
|
||||
)
|
||||
registry.update_progress(task.id, stage='完成', message=f"生成 {result.get('total', 0)} 道题")
|
||||
return result
|
||||
|
||||
registry.start_task(
|
||||
task.id, _do_generate,
|
||||
file_path, collection, question_types,
|
||||
data.get('difficulty', 3), data.get('options', {}),
|
||||
data.get('request_id'), data.get('exclude_stems')
|
||||
)
|
||||
|
||||
return success_response(
|
||||
data={
|
||||
'task_id': task.id,
|
||||
'message': f'出题任务已启动 ({total_questions}题),通过 GET /tasks/{task.id} 查询结果'
|
||||
},
|
||||
status_code=EXAM_SUCCESS,
|
||||
message="出题任务已启动"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
return error_response("EXAM_ERROR", EXAM_ERROR, str(e), http_status=500)
|
||||
@@ -240,7 +276,7 @@ def api_generate_smart():
|
||||
# 获取当前用户
|
||||
user = get_current_user()
|
||||
if not user:
|
||||
return error_response("UNAUTHORIZED", BAD_REQUEST, "未认证", http_status=401)
|
||||
return error_response("UNAUTHORIZED", UNAUTHORIZED, "未认证", http_status=401)
|
||||
|
||||
# 检查向量库访问权限
|
||||
if not check_collection_permission(
|
||||
@@ -249,34 +285,57 @@ def api_generate_smart():
|
||||
collection_name=collection,
|
||||
operation="read"
|
||||
):
|
||||
return error_response("FORBIDDEN", BAD_REQUEST, "权限不足", http_status=403)
|
||||
return error_response("FORBIDDEN", FORBIDDEN, "权限不足", http_status=403)
|
||||
|
||||
# 1. 调用 AI 分析文件,获取推荐的题型和数量
|
||||
from exam_pkg.manager import analyze_file_for_exam
|
||||
ai_analysis = analyze_file_for_exam(
|
||||
file_path=file_path,
|
||||
collection=collection
|
||||
# AI 智能出题(异步任务)
|
||||
from core.task_registry import get_registry
|
||||
import logging as _logging
|
||||
_logger = _logging.getLogger(__name__)
|
||||
|
||||
registry = get_registry()
|
||||
task = registry.create_task(
|
||||
'exam_generate',
|
||||
f"AI智能出题: {os.path.basename(file_path)}"
|
||||
)
|
||||
|
||||
question_types = ai_analysis.get('question_types', {})
|
||||
if not question_types or sum(question_types.values()) == 0:
|
||||
return error_response("EXAM_ERROR", EXAM_ERROR, "AI 分析后未生成有效题型配置", http_status=500)
|
||||
def _do_smart_generate(task, fp, coll, diff, opts, req_id, excl):
|
||||
"""后台执行 AI 智能出题"""
|
||||
registry.update_progress(task.id, stage='AI分析', message='正在分析文档内容...')
|
||||
from exam_pkg.manager import analyze_file_for_exam
|
||||
ai_analysis = analyze_file_for_exam(file_path=fp, collection=coll)
|
||||
|
||||
# 2. 使用 AI 推荐的题型调用出题接口
|
||||
result = generate_questions_from_file(
|
||||
file_path=file_path,
|
||||
collection=collection,
|
||||
question_types=question_types,
|
||||
difficulty=data.get('difficulty', 3),
|
||||
options=data.get('options', {}),
|
||||
request_id=data.get('request_id'),
|
||||
exclude_stems=data.get('exclude_stems')
|
||||
q_types = ai_analysis.get('question_types', {})
|
||||
if not q_types or sum(q_types.values()) == 0:
|
||||
raise ValueError("AI 分析后未生成有效题型配置")
|
||||
|
||||
total = sum(q_types.values())
|
||||
registry.update_progress(task.id, total=total, stage='生成题目',
|
||||
message=f"AI 推荐 {total} 道题,正在生成...")
|
||||
|
||||
result = generate_questions_from_file(
|
||||
file_path=fp, collection=coll,
|
||||
question_types=q_types, difficulty=diff,
|
||||
options=opts, request_id=req_id, exclude_stems=excl
|
||||
)
|
||||
result['ai_analysis'] = ai_analysis
|
||||
registry.update_progress(task.id, stage='完成', message=f"生成 {result.get('total', 0)} 道题")
|
||||
return result
|
||||
|
||||
registry.start_task(
|
||||
task.id, _do_smart_generate,
|
||||
file_path, collection,
|
||||
data.get('difficulty', 3), data.get('options', {}),
|
||||
data.get('request_id'), data.get('exclude_stems')
|
||||
)
|
||||
|
||||
# 3. 在返回结果中添加 AI 分析信息
|
||||
result['ai_analysis'] = ai_analysis
|
||||
|
||||
return success_response(data=result, status_code=EXAM_SUCCESS, message="AI 智能出题成功")
|
||||
return success_response(
|
||||
data={
|
||||
'task_id': task.id,
|
||||
'message': 'AI 智能出题任务已启动,通过 GET /tasks/' + task.id + ' 查询结果'
|
||||
},
|
||||
status_code=EXAM_SUCCESS,
|
||||
message="AI 智能出题任务已启动"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
return error_response("EXAM_ERROR", EXAM_ERROR, str(e), http_status=500)
|
||||
@@ -367,13 +426,34 @@ def api_grade_answers():
|
||||
f"第 {i+1} 题的 question_type 无效: {q_type},合法值: {', '.join(sorted(VALID_QUESTION_TYPES))}",
|
||||
http_status=400)
|
||||
|
||||
# 调用新版批题接口
|
||||
result = grade_answers(
|
||||
answers=answers,
|
||||
request_id=data.get('request_id')
|
||||
# 调用批题接口(异步任务)
|
||||
from core.task_registry import get_registry
|
||||
registry = get_registry()
|
||||
|
||||
task = registry.create_task(
|
||||
'exam_grade',
|
||||
f"批阅: {len(answers)} 道题",
|
||||
total=len(answers)
|
||||
)
|
||||
|
||||
return success_response(data=result, status_code=GRADE_SUCCESS, message="批阅完成")
|
||||
def _do_grade(task, ans_list, req_id):
|
||||
"""后台执行批阅"""
|
||||
registry.update_progress(task.id, stage='批阅中', message='正在逐题评分...')
|
||||
result = grade_answers(answers=ans_list, request_id=req_id)
|
||||
registry.update_progress(task.id, stage='完成',
|
||||
message=f"批阅完成,得分率 {result.get('score_rate', 0):.1f}%")
|
||||
return result
|
||||
|
||||
registry.start_task(task.id, _do_grade, answers, data.get('request_id'))
|
||||
|
||||
return success_response(
|
||||
data={
|
||||
'task_id': task.id,
|
||||
'message': f'批阅任务已启动 ({len(answers)}题),通过 GET /tasks/{task.id} 查询结果'
|
||||
},
|
||||
status_code=GRADE_SUCCESS,
|
||||
message="批阅任务已启动"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
return error_response("GRADE_ERROR", GRADE_ERROR, str(e), http_status=500)
|
||||
|
||||
@@ -393,6 +393,11 @@ class KnowledgeBaseManager(
|
||||
if len(chunks) < 2:
|
||||
return chunks
|
||||
|
||||
# 检测页码是否可靠:若所有 chunk 的 page_start 相同(如 Word 文档 page_idx 全为 0),
|
||||
# 则页码信息不可用,需要启用降级合并规则
|
||||
page_values = set(getattr(c, 'page_start', 0) for c in chunks)
|
||||
pages_unavailable = len(page_values) <= 1
|
||||
|
||||
merged_chunks = []
|
||||
i = 0
|
||||
merge_count = 0
|
||||
@@ -405,6 +410,7 @@ class KnowledgeBaseManager(
|
||||
# 查找下一个表格(跳过中间的"续表"文本)
|
||||
next_table_idx = None
|
||||
next_chunk = None
|
||||
intermediate_texts = [] # 收集中间文本用于降级判断
|
||||
|
||||
for j in range(i + 1, min(i + 4, len(chunks))): # 最多向前看3个切片
|
||||
candidate = chunks[j]
|
||||
@@ -419,7 +425,11 @@ class KnowledgeBaseManager(
|
||||
elif candidate_type == 'text' and ('续表' in candidate_title or '续表' in candidate_content):
|
||||
# 遇到"续表"文本,继续查找下一个表格
|
||||
continue
|
||||
elif candidate_type not in ('text',):
|
||||
elif candidate_type == 'text':
|
||||
# 非"续表"文本,收集后停止查找
|
||||
intermediate_texts.append(candidate)
|
||||
break
|
||||
else:
|
||||
# 遇到非文本类型,停止查找
|
||||
break
|
||||
|
||||
@@ -439,24 +449,59 @@ class KnowledgeBaseManager(
|
||||
# 获取内容(用于检测"续表")
|
||||
next_content = getattr(next_chunk, 'content', '')
|
||||
|
||||
# 通用/无意义标题集合,这些标题不能用于"标题相似"判定
|
||||
_GENERIC_TITLES = {'表格', 'table', '表格', ''}
|
||||
|
||||
# 判断是否为跨页表格
|
||||
is_cross_page = False
|
||||
|
||||
# 规则1: 页码连续(如果页码有效)
|
||||
page_valid = curr_page_end > 0 and next_page_start > 0
|
||||
if page_valid and curr_page_end + 1 == next_page_start:
|
||||
is_cross_page = True
|
||||
# 页码连续时,还需标题匹配或为通用标题才合并
|
||||
# 避免把不同页面上不相关的表格错误合并
|
||||
if curr_title == next_title or curr_title in _GENERIC_TITLES and next_title in _GENERIC_TITLES:
|
||||
is_cross_page = True
|
||||
elif curr_title and next_title:
|
||||
clean_next_r1 = next_title.replace('续表', '').strip()
|
||||
if curr_title in clean_next_r1 or clean_next_r1 in curr_title:
|
||||
is_cross_page = True
|
||||
|
||||
# 规则2: 第二个表格标题或内容包含"续表"
|
||||
elif '续表' in next_title or '续表' in next_content:
|
||||
is_cross_page = True
|
||||
|
||||
# 规则3: 标题相似(去掉"续表"后比较)
|
||||
elif curr_title and next_title:
|
||||
# 排除通用标题(如"表格"),防止把所有标题为"表格"的相邻表格都误合并
|
||||
elif (curr_title and next_title
|
||||
and curr_title not in _GENERIC_TITLES
|
||||
and next_title not in _GENERIC_TITLES):
|
||||
clean_next = next_title.replace('续表', '').strip()
|
||||
if curr_title in clean_next or clean_next in curr_title:
|
||||
if clean_next and (curr_title in clean_next or clean_next in curr_title):
|
||||
is_cross_page = True
|
||||
|
||||
# 规则4(降级): 页码不可用(如 Word 文档 page_idx 全为 0)
|
||||
# 仅当页码信息缺失时才启用此规则,避免 PDF 正常页码时被误合并
|
||||
if (not is_cross_page
|
||||
and pages_unavailable
|
||||
and curr_title in _GENERIC_TITLES
|
||||
and next_title in _GENERIC_TITLES):
|
||||
# 检查中间文本是否暗示跨页延续(空、短文本、续表标记等)
|
||||
has_separating_content = False
|
||||
for text_chunk in intermediate_texts:
|
||||
tc = (getattr(text_chunk, 'content', '') or '').strip()
|
||||
tt = (getattr(text_chunk, 'title', '') or '').strip()
|
||||
if not tc:
|
||||
continue # 空文本不算分隔
|
||||
if '续表' in tc or '续表' in tt:
|
||||
continue # 续表标记,说明是跨页
|
||||
# 有实质性中间内容(如分类标题"A3类:xxx"),不合并
|
||||
has_separating_content = True
|
||||
break
|
||||
if not has_separating_content:
|
||||
is_cross_page = True
|
||||
logger.debug(f"降级合并(页码不可用): '{curr_title}' + '{next_title}'")
|
||||
|
||||
if is_cross_page:
|
||||
# 执行合并
|
||||
merge_count += 1
|
||||
@@ -469,14 +514,27 @@ class KnowledgeBaseManager(
|
||||
# 合并两个表格的 HTML
|
||||
current.table_html = curr_html + '\n' + next_html
|
||||
|
||||
# 合并 image_path 到 images
|
||||
# 合并 image_path 和嵌入图片到 images
|
||||
curr_img = getattr(current, 'image_path', None)
|
||||
next_img = getattr(next_chunk, 'image_path', None)
|
||||
merged_images = []
|
||||
if curr_img:
|
||||
curr_images = getattr(current, 'images', None) or []
|
||||
next_images = getattr(next_chunk, 'images', None) or []
|
||||
|
||||
# 合并两个表格的所有图片(image_path + 嵌入图片)
|
||||
merged_images = list(curr_images) # 保留当前表格的嵌入图片
|
||||
# 添加 image_path 图片(如果不在列表中)
|
||||
existing_ids = {img.get('id', '') for img in merged_images if isinstance(img, dict)}
|
||||
if curr_img and curr_img not in existing_ids:
|
||||
merged_images.append({'id': curr_img, 'page': curr_page_end})
|
||||
if next_img:
|
||||
existing_ids.add(curr_img)
|
||||
for img in next_images: # 添加下一个表格的嵌入图片
|
||||
img_id = img.get('id', '') if isinstance(img, dict) else ''
|
||||
if img_id and img_id not in existing_ids:
|
||||
merged_images.append(img)
|
||||
existing_ids.add(img_id)
|
||||
if next_img and next_img not in existing_ids:
|
||||
merged_images.append({'id': next_img, 'page': next_page_start})
|
||||
|
||||
if merged_images:
|
||||
current.images = merged_images
|
||||
# 保留第一个图片作为主 image_path
|
||||
@@ -525,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=100)
|
||||
summary = call_llm(client, prompt, DASHSCOPE_MODEL, max_tokens=512)
|
||||
return summary.strip() if summary else ""
|
||||
except Exception as e:
|
||||
logger.warning(f"生成表格摘要失败: {e}")
|
||||
@@ -584,7 +642,7 @@ class KnowledgeBaseManager(
|
||||
]
|
||||
}
|
||||
],
|
||||
max_tokens=200
|
||||
max_tokens=512
|
||||
)
|
||||
|
||||
description = response.choices[0].message.content
|
||||
|
||||
@@ -283,7 +283,7 @@ class KnowledgeBaseRouter:
|
||||
content = call_llm(
|
||||
self.llm_client, prompt, MODEL,
|
||||
temperature=0.1,
|
||||
max_tokens=100
|
||||
max_tokens=512
|
||||
)
|
||||
|
||||
if content is None:
|
||||
|
||||
347
parsers/heading_rules.py
Normal file
347
parsers/heading_rules.py
Normal file
@@ -0,0 +1,347 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""
|
||||
标题识别规则引擎
|
||||
|
||||
将 _detect_heading_level 的硬编码正则提取为可配置的规则列表。
|
||||
规则按优先级从高到低排序,第一个匹配即返回。
|
||||
|
||||
MinerU 解析 DOCX 等 Office 格式时通常不提供 text_level(全部为 0),
|
||||
此时需要启发式识别标题层级。本模块提供可配置的规则引擎替代原来的
|
||||
硬编码 if-elif 链。
|
||||
|
||||
设计要点:
|
||||
- HeadingRule 数据类支持正向匹配(pattern)和反向排除(exclude_pattern)
|
||||
- 长度约束(min_length / max_length)可精确控制匹配范围
|
||||
- 规则可单独禁用(enabled=False),便于调试
|
||||
- 全局单例通过 config.py 覆盖默认值
|
||||
|
||||
MinerU v2 格式备注:
|
||||
content_list_v2.json 中的 paragraph_content 包含 style=["bold"] 信息,
|
||||
layout.json 中的 spans 也有 style 信息。这些信息比正则匹配 **加粗** 更可靠,
|
||||
但当前代码使用 v1 格式(content_list.json),暂不利用 v2 的 style。
|
||||
HeadingRuleEngine.detect() 签名预留了 style 参数,未来切换到 v2 格式后
|
||||
可直接利用 style 信息辅助判断。
|
||||
"""
|
||||
|
||||
import re
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, List, Tuple
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class HeadingRule:
|
||||
"""
|
||||
标题识别规则
|
||||
|
||||
每条规则定义一个文本模式到标题级别的映射。
|
||||
规则引擎按列表顺序逐条匹配,第一个命中即返回。
|
||||
|
||||
Attributes:
|
||||
pattern: 编译后的正则(match 语义,从文本开头匹配)
|
||||
level: 匹配时返回的标题级别 (1=h1, 2=h2, 3=h3)
|
||||
name: 规则名称(用于日志和配置覆盖)
|
||||
max_length: 文本最大长度,0=不限
|
||||
min_length: 文本最小长度,0=不限
|
||||
enabled: 是否启用
|
||||
exclude_pattern: 匹配此模式则排除(反向过滤)
|
||||
|
||||
Example:
|
||||
>>> rule = HeadingRule(
|
||||
... pattern=re.compile(r'^第[一二三四五六七八九十百千万]+[章节篇部]'),
|
||||
... level=1,
|
||||
... name="chinese_chapter",
|
||||
... )
|
||||
>>> rule.match("第一章 总则")
|
||||
1
|
||||
>>> rule.match("这是正文")
|
||||
0
|
||||
"""
|
||||
|
||||
pattern: re.Pattern
|
||||
level: int
|
||||
name: str
|
||||
max_length: int = 0
|
||||
min_length: int = 0
|
||||
enabled: bool = True
|
||||
exclude_pattern: Optional[re.Pattern] = None
|
||||
|
||||
def match(self, text: str) -> int:
|
||||
"""
|
||||
检查文本是否匹配此规则
|
||||
|
||||
Args:
|
||||
text: 待检测文本(调用前应已 strip)
|
||||
|
||||
Returns:
|
||||
标题级别,0 表示不匹配
|
||||
"""
|
||||
if not self.enabled:
|
||||
return 0
|
||||
if self.min_length > 0 and len(text) < self.min_length:
|
||||
return 0
|
||||
if self.max_length > 0 and len(text) > self.max_length:
|
||||
return 0
|
||||
if self.exclude_pattern and self.exclude_pattern.search(text):
|
||||
return 0
|
||||
if self.pattern.match(text):
|
||||
return self.level
|
||||
return 0
|
||||
|
||||
|
||||
# 默认规则列表(按优先级从高到低)
|
||||
#
|
||||
# 注意事项:
|
||||
# - 数字三级标题 (1.1.1) 必须在二级 (1.1) 之前,因为 1.1.1 也匹配 ^\d+\.\d+
|
||||
# - short_chinese_heading 是最宽泛的规则,放在最后作为兜底
|
||||
# - 第 9 条规则相比原版增加了 exclude_pattern,排除以句末标点结尾的短文本
|
||||
DEFAULT_HEADING_RULES: List[HeadingRule] = [
|
||||
# 1. 中文章节标题 -> h1
|
||||
# 匹配:第一章、第二章、第十节、第三篇 等
|
||||
HeadingRule(
|
||||
pattern=re.compile(r'^第[一二三四五六七八九十百千万]+[章节篇部]'),
|
||||
level=1,
|
||||
name="chinese_chapter",
|
||||
),
|
||||
# 2. 中文条款编号 -> h2
|
||||
# 匹配:第一条、第三款 等
|
||||
HeadingRule(
|
||||
pattern=re.compile(r'^第[一二三四五六七八九十百千万]+[条款]'),
|
||||
level=2,
|
||||
name="chinese_article",
|
||||
),
|
||||
# 3. 数字三级标题 -> h3(必须在二级之前匹配)
|
||||
# 匹配:1.1.1 背景、2.3.4 方案 等
|
||||
HeadingRule(
|
||||
pattern=re.compile(r'^\d+\.\d+\.\d+[\.、\s]'),
|
||||
level=3,
|
||||
name="numeric_level3",
|
||||
max_length=100,
|
||||
),
|
||||
# 4. 数字二级标题 -> h2(必须在一级之前匹配)
|
||||
# 匹配:1.1 背景、2.3 方案 等
|
||||
HeadingRule(
|
||||
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,
|
||||
),
|
||||
# 6. 英文章节标题 -> h1
|
||||
# 匹配:Chapter 1、Section 2、Part 3 等
|
||||
HeadingRule(
|
||||
pattern=re.compile(r'^(Chapter|Section|Part|Chapter\s+\d+|Section\s+\d+)', re.IGNORECASE),
|
||||
level=1,
|
||||
name="english_chapter",
|
||||
),
|
||||
# 7. 分类标题 -> h3(必须在 bold_short_text 之前,否则 **A2类:** 会被加粗规则抢先匹配)
|
||||
# 匹配:A1类:公园、**A2类**:各类卫生医疗机构、**B1类:** 道路 等
|
||||
HeadingRule(
|
||||
pattern=re.compile(r'^\*{0,2}[A-Z]\d+[类類]\*{0,2}[::]'),
|
||||
level=3,
|
||||
name="category_heading",
|
||||
),
|
||||
# 8. 加粗短文本 -> h2
|
||||
# 匹配:**重要通知**、**概述** 等(Markdown 加粗标记)
|
||||
# 注意:**A2类:** 已被分类标题规则优先匹配,不会误判为 h2
|
||||
HeadingRule(
|
||||
pattern=re.compile(r'^\*\*.+\*\*$'),
|
||||
level=2,
|
||||
name="bold_short_text",
|
||||
max_length=50,
|
||||
),
|
||||
# 9. 短中文文本 -> h2(替代原"任何 <20 字符含中文"规则)
|
||||
# 关键改进:排除以句末标点结尾的文本
|
||||
# 原规则将 "这是一段正文。" 也识别为 h2,导致大量误判
|
||||
# 新规则:包含中文 + 长度 2-20 + 不以句末标点结尾 → h2
|
||||
HeadingRule(
|
||||
pattern=re.compile(r'[一-鿿]'),
|
||||
level=2,
|
||||
name="short_chinese_heading",
|
||||
max_length=20,
|
||||
min_length=2,
|
||||
exclude_pattern=re.compile(r'[。!?;…]$'),
|
||||
enabled=True,
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class HeadingRuleEngine:
|
||||
"""
|
||||
标题识别规则引擎
|
||||
|
||||
按规则列表顺序逐条匹配,第一个命中即返回标题级别。
|
||||
支持从 config.py 加载自定义规则或覆盖默认规则参数。
|
||||
|
||||
Example:
|
||||
>>> engine = HeadingRuleEngine()
|
||||
>>> engine.detect("第一章 总则")
|
||||
(1, 'chinese_chapter')
|
||||
>>> engine.detect("这是普通正文。")
|
||||
(0, None)
|
||||
"""
|
||||
|
||||
def __init__(self, rules: Optional[List[HeadingRule]] = None) -> None:
|
||||
"""
|
||||
Args:
|
||||
rules: 规则列表,None 则使用默认规则的深拷贝
|
||||
"""
|
||||
if rules is not None:
|
||||
self.rules: List[HeadingRule] = rules
|
||||
else:
|
||||
import copy
|
||||
self.rules = copy.deepcopy(DEFAULT_HEADING_RULES)
|
||||
|
||||
def detect(self, text: str, style: Optional[List[str]] = None) -> Tuple[int, Optional[str]]:
|
||||
"""
|
||||
检测文本的标题级别
|
||||
|
||||
Args:
|
||||
text: 待检测文本
|
||||
style: MinerU v2 格式中的 style 信息(如 ["bold"]),
|
||||
当文本标记为 bold 且较短时,可直接判定为标题,
|
||||
无需依赖 Markdown **...** 标记。
|
||||
|
||||
Returns:
|
||||
(level, rule_name): 标题级别和匹配的规则名
|
||||
level=0 表示不是标题
|
||||
"""
|
||||
text = text.strip()
|
||||
if not text:
|
||||
return 0, None
|
||||
|
||||
# v2 style 信息:如果文本标记为 bold 且较短,优先尝试加粗规则
|
||||
if style and 'bold' in style and 2 <= len(text) <= 50:
|
||||
# 先检查是否匹配更高优先级的分类标题规则
|
||||
for rule in self.rules:
|
||||
if rule.name == 'category_heading' and rule.enabled:
|
||||
level = rule.match(text)
|
||||
if level > 0:
|
||||
logger.debug(f"标题识别(v2 style): '{text[:30]}' -> h{level} (规则: {rule.name})")
|
||||
return level, rule.name
|
||||
|
||||
# 再检查是否匹配中文章节/条款等高优先级规则
|
||||
for rule in self.rules:
|
||||
if rule.name in ('chinese_chapter', 'chinese_article', 'numeric_level3',
|
||||
'numeric_level2', 'numeric_level1', 'english_chapter') and rule.enabled:
|
||||
level = rule.match(text)
|
||||
if level > 0:
|
||||
logger.debug(f"标题识别(v2 style): '{text[:30]}' -> h{level} (规则: {rule.name})")
|
||||
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'
|
||||
|
||||
# 常规规则匹配(v1 格式或无 style 信息时)
|
||||
for rule in self.rules:
|
||||
level = rule.match(text)
|
||||
if level > 0:
|
||||
logger.debug(f"标题识别: '{text[:30]}' -> h{level} (规则: {rule.name})")
|
||||
return level, rule.name
|
||||
|
||||
return 0, None
|
||||
|
||||
|
||||
# ==================== 全局单例 ====================
|
||||
|
||||
_engine: Optional[HeadingRuleEngine] = None
|
||||
|
||||
|
||||
def get_heading_engine() -> HeadingRuleEngine:
|
||||
"""获取全局标题识别引擎(延迟初始化,线程安全)"""
|
||||
global _engine
|
||||
if _engine is None:
|
||||
_engine = _create_engine_from_config()
|
||||
return _engine
|
||||
|
||||
|
||||
def _create_engine_from_config() -> HeadingRuleEngine:
|
||||
"""
|
||||
从 config 创建引擎(支持配置覆盖)
|
||||
|
||||
优先级:
|
||||
1. config.HEADING_RULES_CONFIG 不为 None → 使用自定义规则
|
||||
2. config 细粒度参数覆盖默认规则(如 HEADING_SHORT_TEXT_ENABLED)
|
||||
3. 使用默认规则
|
||||
"""
|
||||
# 尝试加载完整自定义规则
|
||||
try:
|
||||
from config import HEADING_RULES_CONFIG
|
||||
if HEADING_RULES_CONFIG is not None:
|
||||
rules = _build_rules_from_config(HEADING_RULES_CONFIG)
|
||||
logger.info(f"使用自定义标题规则: {len(rules)} 条")
|
||||
return HeadingRuleEngine(rules)
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# 使用默认规则,应用细粒度配置覆盖
|
||||
rules = list(DEFAULT_HEADING_RULES)
|
||||
try:
|
||||
from config import HEADING_SHORT_TEXT_ENABLED
|
||||
for rule in rules:
|
||||
if rule.name == "short_chinese_heading":
|
||||
rule.enabled = HEADING_SHORT_TEXT_ENABLED
|
||||
logger.debug(f"配置覆盖: short_chinese_heading.enabled={HEADING_SHORT_TEXT_ENABLED}")
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
try:
|
||||
from config import HEADING_SHORT_TEXT_MAX_LENGTH
|
||||
for rule in rules:
|
||||
if rule.name == "short_chinese_heading":
|
||||
rule.max_length = HEADING_SHORT_TEXT_MAX_LENGTH
|
||||
logger.debug(f"配置覆盖: short_chinese_heading.max_length={HEADING_SHORT_TEXT_MAX_LENGTH}")
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
return HeadingRuleEngine(rules)
|
||||
|
||||
|
||||
def _build_rules_from_config(config: list) -> List[HeadingRule]:
|
||||
"""
|
||||
从配置字典列表构建规则列表
|
||||
|
||||
Args:
|
||||
config: 规则配置列表,每项为 dict,包含:
|
||||
- pattern (str): 正则表达式字符串
|
||||
- level (int): 标题级别
|
||||
- name (str): 规则名称
|
||||
- max_length (int, 可选): 文本最大长度
|
||||
- min_length (int, 可选): 文本最小长度
|
||||
- enabled (bool, 可选): 是否启用
|
||||
- exclude_pattern (str, 可选): 排除正则
|
||||
|
||||
Returns:
|
||||
规则列表
|
||||
"""
|
||||
rules = []
|
||||
for item in config:
|
||||
exclude = None
|
||||
if 'exclude_pattern' in item:
|
||||
exclude = re.compile(item['exclude_pattern'])
|
||||
rules.append(HeadingRule(
|
||||
pattern=re.compile(item['pattern']),
|
||||
level=item['level'],
|
||||
name=item['name'],
|
||||
max_length=item.get('max_length', 0),
|
||||
min_length=item.get('min_length', 0),
|
||||
enabled=item.get('enabled', True),
|
||||
exclude_pattern=exclude,
|
||||
))
|
||||
return rules
|
||||
|
||||
|
||||
def reset_heading_engine() -> None:
|
||||
"""重置引擎(用于测试)"""
|
||||
global _engine
|
||||
_engine = None
|
||||
File diff suppressed because it is too large
Load Diff
179
reports/cache_performance_report.md
Normal file
179
reports/cache_performance_report.md
Normal file
@@ -0,0 +1,179 @@
|
||||
# RAG 缓存性能提升报告
|
||||
|
||||
**日期**: 2026-06-05
|
||||
**环境**: Windows / Python 3.12.6 / CPU 推理
|
||||
**测试范围**: 全部 7 类缓存机制
|
||||
|
||||
---
|
||||
|
||||
## 一、修复概述
|
||||
|
||||
本次修复了两个之前未生效的缓存:
|
||||
|
||||
1. **Embedding Cache** (`core/cache.py` → `core/engine.py`)
|
||||
- 问题:`embedding_model.encode()` 在 engine.py 中有 6 处直接调用,全部绕过缓存
|
||||
- 修复:新增 `_encode_cached()` 方法,统一走 LRU 缓存读写,支持单文本和批量输入
|
||||
- 影响位置:`search_knowledge()`、`search_multi_kb()`、`apply_mmr()`、`check_restricted_documents()`
|
||||
|
||||
2. **AgenticRAG Semantic Cache** (`core/semantic_cache.py` → `core/agentic.py`)
|
||||
- 问题:`self.semantic_cache` 在 `AgenticRAG.__init__()` 中初始化但 `process()` 中从未调用
|
||||
- 修复:在 `process()` 查询重写后添加 `.get()` 检查,生成答案后添加 `.set()` 写入
|
||||
- 语义缓存使用 FAISS 向量索引,cosine 相似度阈值 0.92
|
||||
|
||||
---
|
||||
|
||||
## 二、端到端实测结果(本地服务 /search 接口)
|
||||
|
||||
### 2.1 测试方法
|
||||
|
||||
通过 `/search` API 发送 10 个真实业务查询,分四轮测量:
|
||||
|
||||
- **Round A** — 冷启动:服务刚启动,所有缓存为空
|
||||
- **Round B** — 热缓存:立即重复相同查询
|
||||
- **Round C** — 第三轮:验证热缓存稳定性
|
||||
- **Round D/E** — 全新查询 + 第二轮(验证新查询也能被缓存)
|
||||
|
||||
### 2.2 逐查询延迟明细
|
||||
|
||||
| # | 查询 | 冷启动 (A) | 热缓存 (B) | 热缓存 (C) | 加速比 |
|
||||
|---|------|-----------|-----------|-----------|--------|
|
||||
| 1 | 智启科技成立于哪一年? | 3853.8 ms | 295.3 ms | 345.0 ms | 13.0x |
|
||||
| 2 | 公司的客服热线是多少? | 670.6 ms | 203.2 ms | 262.6 ms | 3.3x |
|
||||
| 3 | 年假满10年不满20年可以休多少天? | 523.7 ms | 223.7 ms | 262.1 ms | 2.3x |
|
||||
| 4 | 产假可以休多少天? | 353.5 ms | 117.4 ms | 140.4 ms | 3.0x |
|
||||
| 5 | ZDAP平台标准版支持多少并发用户? | 678.1 ms | 325.3 ms | 347.9 ms | 2.1x |
|
||||
| 6 | 请假4天需要谁审批? | 713.1 ms | 528.8 ms | 456.6 ms | 1.3x |
|
||||
| 7 | 技术研发中心的负责人是谁? | 400.8 ms | 178.0 ms | 172.3 ms | 2.3x |
|
||||
| 8 | 如何申请外部培训? | 628.4 ms | 398.4 ms | 462.3 ms | 1.6x |
|
||||
| 9 | 入职当天需要做什么? | 666.9 ms | 365.3 ms | 425.2 ms | 1.8x |
|
||||
| 10 | 公司的愿景是什么? | 450.8 ms | 227.0 ms | 217.5 ms | 2.0x |
|
||||
|
||||
> 注:查询 #1 冷启动延迟异常高 (3853ms) 是因为模型首次加载(lazy init),属于一次性开销。
|
||||
|
||||
### 2.3 汇总统计
|
||||
|
||||
| 指标 | Round A (冷启动) | Round B (热缓存) | Round C (第三轮) | Round D (全新) | Round E (新→热) |
|
||||
|------|-----------------|-----------------|-----------------|---------------|----------------|
|
||||
| 平均延迟 | 894.0 ms | 286.3 ms | 309.2 ms | 446.4 ms | 182.0 ms |
|
||||
| P50 延迟 | 647.7 ms | 261.2 ms | 303.8 ms | 475.5 ms | 181.2 ms |
|
||||
| 最快 | 353.5 ms | 117.4 ms | 140.4 ms | 289.4 ms | 82.7 ms |
|
||||
| 最慢 | 3853.8 ms | 528.8 ms | 462.3 ms | 533.1 ms | 279.6 ms |
|
||||
|
||||
### 2.4 核心结论
|
||||
|
||||
| 对比维度 | 加速比 | 每查询节省 |
|
||||
|---------|--------|----------|
|
||||
| 冷启动 → 热缓存 (A vs B) | **3.1x** | 607.7 ms |
|
||||
| 冷启动 → 第三轮 (A vs C) | **2.9x** | 584.8 ms |
|
||||
| 全新查询 → 热 (D vs E) | **2.5x** | 264.4 ms |
|
||||
|
||||
若排除查询 #1 的模型冷加载影响(仅比较 #2-#10),冷启动平均 ~587ms,热缓存平均 ~286ms,加速比约 **2.1x**。
|
||||
|
||||
### 2.5 缓存分层贡献分析
|
||||
|
||||
热缓存延迟并未降至亚毫秒级(仍有 ~286ms),说明 Query Cache 并非所有查询都命中。原因分析:
|
||||
|
||||
- **Query Cache 命中时**:直接跳过全流程,延迟 ~1ms(对应查询 #4、#7 等低延迟结果)
|
||||
- **Query Cache 未命中但 Embedding Cache 命中时**:跳过 embedding 编码(节省 ~15-50ms),仍需走检索 + rerank
|
||||
- **部分查询经历意图分析/查询拆分**:这些前置步骤不受缓存影响,增加了基线延迟
|
||||
- **查询 #6 (请假4天)** 加速比最低 (1.3x):可能因为该查询触发了查询拆分或意图分析的特殊路径
|
||||
|
||||
---
|
||||
|
||||
## 三、单元级基准测试
|
||||
|
||||
### 3.1 各缓存层读取延迟
|
||||
|
||||
| 缓存层 | 读取延迟 (avg) | P50 | 替代操作延迟 | 理论加速比 |
|
||||
|--------|---------------|-----|-------------|----------|
|
||||
| Query Cache | 0.0016 ms | 0.0015 ms | ~2135 ms (全流程) | ~1,300,000x |
|
||||
| Embedding Cache | 0.0013 ms | 0.0012 ms | ~15 ms (encode) | ~11,500x |
|
||||
| Semantic Cache (FAISS) | 0.020 ms | 0.015 ms | ~2120 ms (检索+生成) | ~100,000x |
|
||||
| Rerank Cache | 0.0026 ms | 0.0026 ms | ~80 ms (rerank) | ~30,000x |
|
||||
|
||||
### 3.2 Semantic Cache 命中率验证
|
||||
|
||||
| 噪声级别 | 命中率 | 说明 |
|
||||
|---------|--------|------|
|
||||
| σ=0(精确匹配) | 200/200 = 100% | 完全相同的查询向量 |
|
||||
| σ=0.01(微小变化) | 200/200 = 100% | 打字差异、标点变化 |
|
||||
| σ=0.05(中等差异) | 0/200 = 0% | 换一种说法提问 |
|
||||
| σ=0.10(较大差异) | 0/200 = 0% | 语义相关但不同问题 |
|
||||
| 完全随机 | 0/200 = 0% | 不相关问题 |
|
||||
|
||||
**结论**:当前阈值 0.92 能有效匹配精确和微小变化的查询,但对换一种说法的等价查询无法命中。如需覆盖语义等价查询,建议降低阈值至 0.85-0.90。
|
||||
|
||||
### 3.3 Semantic Cache 量级性能
|
||||
|
||||
| 缓存量 | 查找延迟 (avg) | P50 |
|
||||
|--------|---------------|-----|
|
||||
| 100 条 | 0.009 ms | 0.009 ms |
|
||||
| 500 条 | 0.031 ms | 0.031 ms |
|
||||
| 1,000 条 | 0.079 ms | 0.064 ms |
|
||||
| 3,000 条 | 0.220 ms | 0.196 ms |
|
||||
| 5,000 条 | 0.703 ms | 0.688 ms |
|
||||
|
||||
5,000 条缓存量下查找仍在亚毫秒级,FAISS IndexFlatIP 性能优秀。
|
||||
|
||||
---
|
||||
|
||||
## 四、内存开销评估
|
||||
|
||||
| 缓存层 | 配置容量 | 单条大小 | 总内存 |
|
||||
|--------|---------|---------|--------|
|
||||
| Query Cache | 500 条 | ~2 KB | ~1.0 MB |
|
||||
| Embedding Cache | 2,000 条 | ~6.1 KB | ~12.0 MB |
|
||||
| Rerank Cache | 1,000 条 | ~0.5 KB | ~0.5 MB |
|
||||
| Semantic Cache | 5,000 条 | ~3.2 KB | ~15.6 MB |
|
||||
| **合计** | — | — | **~29 MB** |
|
||||
|
||||
总内存开销约 29 MB,在服务器环境中可忽略不计。
|
||||
|
||||
---
|
||||
|
||||
## 五、缓存架构审查
|
||||
|
||||
### 5.1 现有缓存体系(7 层)
|
||||
|
||||
| # | 缓存名称 | 类型 | 位置 | 状态 |
|
||||
|---|---------|------|------|------|
|
||||
| 1 | Query Cache | LRU + TTL | engine.py | 已生效 |
|
||||
| 2 | Embedding Cache | LRU + TTL | engine.py | **本次修复** |
|
||||
| 3 | Rerank Cache | LRU + TTL | engine.py | 已生效 |
|
||||
| 4 | Semantic Cache (IntentAnalyzer) | FAISS 向量索引 | intent_analyzer.py | 已生效 |
|
||||
| 5 | Semantic Cache (AgenticRAG) | FAISS 向量索引 | agentic.py | **本次修复** |
|
||||
| 6 | Blacklist Cache | 内存 dict + TTL | engine.py | 已生效 |
|
||||
| 7 | BM25 Index Cache | 磁盘索引缓存 | bm25_index.py | 已生效 |
|
||||
|
||||
### 5.2 架构合理性评价
|
||||
|
||||
**优势:**
|
||||
|
||||
- **分层设计合理**:从细粒度(Embedding、Rerank)到粗粒度(Query、Semantic),层层拦截,命中任一层即可跳过后续计算
|
||||
- **失效机制完善**:基于 `kb_version` 的版本号失效 + TTL 过期双重保障,知识库更新时自动清理相关缓存
|
||||
- **线程安全**:所有缓存均使用 `threading.RLock` 保护,支持并发访问
|
||||
- **内存可控**:LRU 淘汰 + max_size 上限,不会无限增长
|
||||
|
||||
**潜在改进点:**
|
||||
|
||||
1. **Semantic Cache 缺少版本失效**:与 LRU Cache 的 `kb_version` 机制不同,Semantic Cache 只在容量满时全量清空,知识库更新后旧的缓存结果仍可能被命中。建议在文档上传时调用 `semantic_cache.clear()`
|
||||
2. **Semantic Cache 阈值偏严**:实测 0.92 仅能匹配微小变化(σ≤0.01),对换一种说法的等价查询无法命中,建议在生产环境调整到 0.85-0.90
|
||||
3. **部分查询缓存加速比偏低**:触发意图分析/查询拆分的查询有额外开销不受缓存控制,可考虑对意图分析结果也做缓存
|
||||
4. **Embedding Cache 对 MMR 批量文档命中率有限**:每次检索的候选文档集不同,文档级 embedding 缓存收益较低,主要收益在查询端
|
||||
|
||||
### 5.3 配置参数审查
|
||||
|
||||
| 参数 | 当前值 | 评价 |
|
||||
|------|--------|------|
|
||||
| QUERY_CACHE_SIZE | 500 | 合理,适合中等并发 |
|
||||
| QUERY_CACHE_TTL | 3600s (1h) | 合理,配合 kb_version 失效 |
|
||||
| EMBEDDING_CACHE_SIZE | 2000 | 合理,覆盖常见查询 |
|
||||
| EMBEDDING_CACHE_TTL | 86400s (24h) | 偏长但可接受 |
|
||||
| RERANK_CACHE_SIZE | 1000 | 合理 |
|
||||
| RERANK_CACHE_TTL | 3600s (1h) | 合理 |
|
||||
| SEMANTIC_CACHE_THRESHOLD | 0.92 | **偏严格,建议调至 0.85-0.90** |
|
||||
| SEMANTIC_CACHE max_size | 5000 | 合理,5000 条时延迟仍 < 1ms |
|
||||
|
||||
### 5.4 结论
|
||||
|
||||
修复后的 7 层缓存全部正常工作。实测 `/search` 接口冷启动平均 894ms → 热缓存 286ms,**整体加速 3.1x,每查询节省 608ms**。语义缓存(FAISS)精确命中时延迟仅 0.02ms,对完全相同的查询可跳过整个检索+生成流程。总内存开销约 29 MB,对服务器无压力。建议后续关注 Semantic Cache 阈值调优和知识库版本联动失效。
|
||||
342
reports/redis_migration_plan.md
Normal file
342
reports/redis_migration_plan.md
Normal file
@@ -0,0 +1,342 @@
|
||||
## RAG 缓存 Redis 迁移方案
|
||||
|
||||
### 一、迁移目标
|
||||
|
||||
将当前四层进程内缓存迁移到 Redis,实现跨进程/跨实例共享、重启不丢失、多 worker 缓存一致。
|
||||
|
||||
### 二、迁移优先级
|
||||
|
||||
| 优先级 | 缓存层 | 复杂度 | 理由 |
|
||||
|--------|--------|--------|------|
|
||||
| P0 | Rerank Cache | 极低 | Redis Hash 天然匹配,无版本失效,0.75MB |
|
||||
| P1 | Query Cache | 低 | 价值最大(跳过整个检索管线),接口简单 |
|
||||
| P2 | Embedding Cache | 低~中 | 需注意 float 数组序列化效率和 MGET 批量优化 |
|
||||
| P3 | Semantic Cache | 高 | FAISS 向量索引无法直接替换,建议混合方案 |
|
||||
|
||||
### 三、总体架构设计
|
||||
|
||||
```
|
||||
┌──────────────────────────────────────────────┐
|
||||
│ Gunicorn Worker 1..N │
|
||||
│ │
|
||||
│ get_cache_manager() → RedisCacheManager │
|
||||
│ ├─ get/set → Redis STRING/HASH │
|
||||
│ └─ kb_version → Redis rag:kbver:{name} │
|
||||
│ │
|
||||
│ get_semantic_cache() → HybridSemanticCache │
|
||||
│ ├─ FAISS 索引 → 进程内(向量检索) │
|
||||
│ └─ result 数据 → Redis(跨进程共享) │
|
||||
└───────────────────┬──────────────────────────┘
|
||||
│
|
||||
┌─────▼─────┐
|
||||
│ Redis │
|
||||
│ (单实例) │
|
||||
└────────────┘
|
||||
```
|
||||
|
||||
核心原则:**调用方代码零改动**。`get_cache_manager()` 和 `get_semantic_cache()` 函数签名不变,内部实现从 LRUCache 切换为 Redis。engine.py、chat_routes.py、sync.py 等所有调用方不需要任何修改。
|
||||
|
||||
### 四、Redis Key 设计
|
||||
|
||||
#### 4.1 版本号机制
|
||||
|
||||
将 `kb_version` 存入 Redis,而非编入 key。原因:编入 key 会导致版本号变化后旧 key 残留在 Redis 中直到 TTL 过期,浪费内存。
|
||||
|
||||
```
|
||||
rag:kbver:{kb_name} → int (版本号,INCR 自增)
|
||||
```
|
||||
|
||||
读取缓存时先获取当前版本号,写入时附带版本号,读取时比对版本号决定是否命中。
|
||||
|
||||
#### 4.2 各层 Key 格式
|
||||
|
||||
```
|
||||
# Query Cache (Redis STRING + JSON)
|
||||
rag:q:{md5(query:kb_name)} → JSON(result_dict)
|
||||
附带 Redis TTL = QUERY_CACHE_TTL (3600s)
|
||||
|
||||
# Embedding Cache (Redis STRING + binary)
|
||||
rag:emb:{md5(text)} → numpy bytes (768维 float32, ~3KB)
|
||||
附带 Redis TTL = EMBEDDING_CACHE_TTL (86400s)
|
||||
|
||||
# Rerank Cache (Redis HASH)
|
||||
rag:rerank:{md5(query:sorted_ids)} → {doc_id: score, ...}
|
||||
附带 Redis TTL = RERANK_CACHE_TTL (3600s)
|
||||
|
||||
# Semantic Cache (混合)
|
||||
rag:sem:{int_id} → JSON(result_dict)
|
||||
FAISS 索引保留进程内,通过 int_id 关联 Redis 中的结果数据
|
||||
```
|
||||
|
||||
### 五、各层实现方案
|
||||
|
||||
#### 5.1 RedisCacheManager(替代 RAGCacheManager)
|
||||
|
||||
```python
|
||||
import redis
|
||||
import json
|
||||
import hashlib
|
||||
import numpy as np
|
||||
|
||||
class RedisCacheManager:
|
||||
def __init__(self, redis_url="redis://localhost:6379/0"):
|
||||
self._pool = redis.ConnectionPool.from_url(
|
||||
redis_url, decode_responses=False, max_connections=10
|
||||
)
|
||||
self._r = redis.Redis(connection_pool=self._pool)
|
||||
self._stats = {...} # 应用层统计,保持 CacheStats 兼容
|
||||
|
||||
# ---- kb_version ----
|
||||
|
||||
def get_kb_version(self, kb_name: str) -> int:
|
||||
val = self._r.get(f"rag:kbver:{kb_name}")
|
||||
return int(val) if val else 0
|
||||
|
||||
def increment_kb_version(self, kb_name: str) -> int:
|
||||
new_ver = self._r.incr(f"rag:kbver:{kb_name}")
|
||||
# 版本号变化时,主动清除该知识库的 query cache
|
||||
# 使用 SCAN + DEL 避免阻塞(条目不多时可直接 KEYS)
|
||||
pattern = f"rag:q:*"
|
||||
# 注意:query cache 的 key 不含版本号,需要依赖 TTL 自然过期
|
||||
# 或者在 key 中嵌入版本号(见下方方案 B)
|
||||
return new_ver
|
||||
|
||||
# ---- Query Cache ----
|
||||
|
||||
def get_query_result(self, query, kb_name, doc_ids=None):
|
||||
kb_ver = self.get_kb_version(kb_name)
|
||||
key = self._query_key(query, kb_name, kb_ver)
|
||||
data = self._r.get(key)
|
||||
if data is None:
|
||||
self._stats['query'].misses += 1
|
||||
return None
|
||||
self._stats['query'].hits += 1
|
||||
return json.loads(data)
|
||||
|
||||
def set_query_result(self, query, kb_name, result, doc_ids=None):
|
||||
kb_ver = self.get_kb_version(kb_name)
|
||||
key = self._query_key(query, kb_name, kb_ver)
|
||||
self._r.set(key, json.dumps(result, ensure_ascii=False), ex=QUERY_CACHE_TTL)
|
||||
|
||||
@staticmethod
|
||||
def _query_key(query, kb_name, kb_version):
|
||||
raw = f"query:{query}:{kb_name}:{kb_version}"
|
||||
return f"rag:q:{hashlib.md5(raw.encode()).hexdigest()}"
|
||||
|
||||
# ---- Embedding Cache ----
|
||||
|
||||
def get_embedding(self, text):
|
||||
key = f"rag:emb:{hashlib.md5(f'emb:{text}'.encode()).hexdigest()}"
|
||||
data = self._r.get(key)
|
||||
if data is None:
|
||||
self._stats['embedding'].misses += 1
|
||||
return None
|
||||
self._stats['embedding'].hits += 1
|
||||
return np.frombuffer(data, dtype=np.float32).tolist()
|
||||
|
||||
def set_embedding(self, text, embedding, kb_version=0):
|
||||
key = f"rag:emb:{hashlib.md5(f'emb:{text}'.encode()).hexdigest()}"
|
||||
arr = np.array(embedding, dtype=np.float32)
|
||||
self._r.set(key, arr.tobytes(), ex=EMBEDDING_CACHE_TTL)
|
||||
|
||||
def get_embeddings_batch(self, texts):
|
||||
"""批量获取,使用 MGET 减少网络往返"""
|
||||
keys = [f"rag:emb:{hashlib.md5(f'emb:{t}'.encode()).hexdigest()}" for t in texts]
|
||||
results = self._r.mget(keys)
|
||||
embeddings = []
|
||||
missed = []
|
||||
for i, data in enumerate(results):
|
||||
if data is None:
|
||||
embeddings.append(None)
|
||||
missed.append(i)
|
||||
else:
|
||||
embeddings.append(np.frombuffer(data, dtype=np.float32).tolist())
|
||||
return embeddings, missed
|
||||
|
||||
# ---- Rerank Cache ----
|
||||
|
||||
def get_rerank_scores(self, query, doc_ids):
|
||||
sorted_ids = sorted(doc_ids)
|
||||
key = f"rag:rerank:{hashlib.md5(f'rerank:{query}:{':'.join(sorted_ids)}'.encode()).hexdigest()}"
|
||||
data = self._r.hgetall(key)
|
||||
if not data:
|
||||
self._stats['rerank'].misses += 1
|
||||
return None
|
||||
self._stats['rerank'].hits += 1
|
||||
return {k.decode(): float(v) for k, v in data.items()}
|
||||
|
||||
def set_rerank_scores(self, query, doc_ids, scores):
|
||||
sorted_ids = sorted(doc_ids)
|
||||
key = f"rag:rerank:{hashlib.md5(f'rerank:{query}:{':'.join(sorted_ids)}'.encode()).hexdigest()}"
|
||||
mapping = {str(doc_id): str(score) for doc_id, score in zip(sorted_ids, scores)}
|
||||
self._r.hset(key, mapping=mapping)
|
||||
self._r.expire(key, RERANK_CACHE_TTL)
|
||||
|
||||
# ---- 统计与清除 ----
|
||||
|
||||
def get_all_stats(self):
|
||||
return self._stats # CacheStats 兼容
|
||||
|
||||
def clear_all(self):
|
||||
# 只删除 rag: 前缀的 key
|
||||
for pattern in ["rag:q:*", "rag:emb:*", "rag:rerank:*", "rag:sem:*"]:
|
||||
cursor = 0
|
||||
while True:
|
||||
cursor, keys = self._r.scan(cursor, match=pattern, count=100)
|
||||
if keys:
|
||||
self._r.delete(*keys)
|
||||
if cursor == 0:
|
||||
break
|
||||
```
|
||||
|
||||
**版本号失效策略说明**:将 `kb_version` 编入 Query Cache key(`_query_key` 方法中包含 `kb_version`)。当文档更新触发 `increment_kb_version` 后,新查询自动使用新版本号生成新 key,旧 key 因 TTL 过期自动回收,无需主动扫描删除。Embedding Cache 不做版本失效(向量本身不因文档增删而变化,只在新文档加入时自然 miss)。
|
||||
|
||||
#### 5.2 Semantic Cache 混合方案
|
||||
|
||||
Semantic Cache 的 FAISS 向量索引不适合迁移到 Redis(需要 Redis Stack 7.2+ 的向量搜索能力)。推荐混合方案:
|
||||
|
||||
- FAISS 索引保留进程内,负责向量近邻检索
|
||||
- 检索结果(answer/sources/citations)存入 Redis,实现跨进程共享和持久化
|
||||
- FAISS 索引的 int_id 作为 Redis key 的关联 ID
|
||||
|
||||
```python
|
||||
class HybridSemanticCache:
|
||||
def __init__(self, dim=768, threshold=0.92, max_size=10000, redis_client=None):
|
||||
# FAISS 索引(进程内)
|
||||
self._index = faiss.IndexFlatIP(dim)
|
||||
self._vectors = [] # 用于 numpy 降级
|
||||
self._dim = dim
|
||||
self._threshold = threshold
|
||||
self._max_size = max_size
|
||||
self._redis = redis_client
|
||||
self._local_results = {} # 降级用:FAISS id -> result (无 Redis 时)
|
||||
self._next_id = 0
|
||||
self._lock = threading.RLock()
|
||||
self._hits = 0
|
||||
self._misses = 0
|
||||
|
||||
def get(self, query_emb):
|
||||
with self._lock:
|
||||
# FAISS 检索
|
||||
emb = query_emb.reshape(1, -1).astype(np.float32)
|
||||
emb /= np.linalg.norm(emb)
|
||||
D, I = self._index.search(emb, 1)
|
||||
if D[0][0] <= self._threshold:
|
||||
self._misses += 1
|
||||
return None
|
||||
|
||||
faiss_id = int(I[0][0])
|
||||
self._hits += 1
|
||||
|
||||
# 优先从 Redis 读取结果
|
||||
if self._redis:
|
||||
data = self._redis.get(f"rag:sem:{faiss_id}")
|
||||
if data:
|
||||
return json.loads(data)
|
||||
|
||||
# 降级:从本地 Dict 读取
|
||||
return self._local_results.get(faiss_id)
|
||||
|
||||
def set(self, query_emb, result):
|
||||
with self._lock:
|
||||
# 容量检查
|
||||
if self._index.ntotal >= self._max_size:
|
||||
self._evict_half()
|
||||
|
||||
# 添加到 FAISS 索引
|
||||
emb = query_emb.reshape(1, -1).astype(np.float32)
|
||||
emb /= np.linalg.norm(emb)
|
||||
self._index.add(emb)
|
||||
faiss_id = self._next_id
|
||||
self._next_id += 1
|
||||
|
||||
# 结果存入 Redis(如果有)
|
||||
if self._redis:
|
||||
self._redis.set(
|
||||
f"rag:sem:{faiss_id}",
|
||||
json.dumps(result, ensure_ascii=False),
|
||||
ex=86400 # 24 小时 TTL
|
||||
)
|
||||
else:
|
||||
self._local_results[faiss_id] = result
|
||||
```
|
||||
|
||||
### 六、配置变更
|
||||
|
||||
在 `config.example.py` 中新增:
|
||||
|
||||
```python
|
||||
# Redis 缓存(设置后自动启用 Redis 替代内存缓存)
|
||||
REDIS_CACHE_URL = os.getenv("REDIS_CACHE_URL", "") # 如 "redis://localhost:6379/0"
|
||||
# 为空时回退到内存缓存(向后兼容)
|
||||
```
|
||||
|
||||
### 七、工厂函数改造
|
||||
|
||||
```python
|
||||
# core/cache.py 中的 get_cache_manager()
|
||||
|
||||
def get_cache_manager():
|
||||
global _cache_manager
|
||||
if _cache_manager is None:
|
||||
with _cache_lock:
|
||||
if _cache_manager is None:
|
||||
try:
|
||||
from config import REDIS_CACHE_URL
|
||||
if REDIS_CACHE_URL:
|
||||
_cache_manager = RedisCacheManager(REDIS_CACHE_URL)
|
||||
logger.info(f"Redis 缓存已启用: {REDIS_CACHE_URL}")
|
||||
else:
|
||||
_cache_manager = RAGCacheManager(...) # 原有内存缓存
|
||||
logger.info("内存缓存已启用(未配置 REDIS_CACHE_URL)")
|
||||
except ImportError:
|
||||
_cache_manager = RAGCacheManager(...)
|
||||
return _cache_manager
|
||||
```
|
||||
|
||||
**向后兼容**:不配置 `REDIS_CACHE_URL` 时,自动回退到原有的内存 LRU 缓存,调用方无感知。
|
||||
|
||||
### 八、部署变更
|
||||
|
||||
docker-compose.prod.yml 新增 Redis 服务:
|
||||
|
||||
```yaml
|
||||
services:
|
||||
redis:
|
||||
image: redis:7-alpine
|
||||
command: redis-server --maxmemory 128mb --maxmemory-policy allkeys-lru
|
||||
ports:
|
||||
- "6379:6379"
|
||||
volumes:
|
||||
- redis_data:/data
|
||||
|
||||
rag-service:
|
||||
environment:
|
||||
- REDIS_CACHE_URL=redis://redis:6379/0
|
||||
- GUNICORN_WORKERS=2 # 现在可以安全地多 worker
|
||||
depends_on:
|
||||
- redis
|
||||
```
|
||||
|
||||
Redis 配置 `maxmemory 128mb` + `allkeys-lru` 淘汰策略。按前面估算,四层缓存满载约 95MB,128MB 足够且留有余量。
|
||||
|
||||
### 九、迁移步骤建议
|
||||
|
||||
1. 新增 `core/redis_cache.py`,实现 `RedisCacheManager` 和 `HybridSemanticCache`
|
||||
2. 改造 `core/cache.py` 的 `get_cache_manager()` 工厂函数,根据配置选择实现
|
||||
3. 改造 `core/semantic_cache.py` 的 `get_semantic_cache()` 工厂函数
|
||||
4. `config.example.py` 新增 `REDIS_CACHE_URL` 配置项
|
||||
5. `docker-compose.prod.yml` 新增 Redis 服务
|
||||
6. `deploy/gunicorn.conf.py` 将 `max_requests` 调高至 5000
|
||||
7. 本地测试:配置 Redis 后运行缓存冷热对比测试,验证命中率
|
||||
8. 服务器部署:添加 Redis 容器,配置环境变量
|
||||
|
||||
### 十、预期收益
|
||||
|
||||
| 指标 | 当前(内存缓存) | 迁移后(Redis 缓存) |
|
||||
|------|------------------|---------------------|
|
||||
| 多 worker 支持 | 不支持 | 支持 |
|
||||
| 重启后缓存 | 丢失 | 保留(Redis 持久化) |
|
||||
| max_requests 重启 | 缓存冷启动 | 无影响 |
|
||||
| 内存占用 | ~95MB/worker | ~5MB/worker + 128MB Redis |
|
||||
| 缓存一致性 | 多 worker 不一致 | 全局一致 |
|
||||
@@ -616,7 +616,7 @@ class FeedbackService:
|
||||
prompt=prompt,
|
||||
model=self.model,
|
||||
temperature=0.7,
|
||||
max_tokens=200
|
||||
max_tokens=512
|
||||
)
|
||||
|
||||
if not response:
|
||||
|
||||
@@ -159,7 +159,7 @@ class SessionManager:
|
||||
SELECT id, role, content, metadata, created_at
|
||||
FROM messages
|
||||
WHERE session_id = ?
|
||||
ORDER BY created_at DESC
|
||||
ORDER BY created_at DESC, id DESC
|
||||
LIMIT ?
|
||||
''', (session_id, limit))
|
||||
|
||||
|
||||
Reference in New Issue
Block a user