7 Commits

Author SHA1 Message Date
lacerate551
279e2bf47c chore: 配置集中化、LLM 参数调整与 gitignore 更新
- config.example: 新增 MINERU_PREFER_V2、标题规则引擎、表单二次校正等配置项
- document_routes: DEV_MODE 判断统一收归 config.py
- llm_utils: quick_yes_no max_tokens 10→128,避免截断过短回答
- knowledge/router: 路由 LLM 调用 max_tokens 100→512
- feedback: 反馈分析 LLM 调用 max_tokens 200→512
- .gitignore: 新增 scripts/ 和 plans/ 目录忽略规则

🤖 Generated with [Qoder][https://qoder.com]
2026-06-08 16:14:17 +08:00
lacerate551
acb84b804d fix(auth): 登录速率限制与用户会话隔离增强
- auth_routes: 新增内存级登录速率限制(5min/10次),防止暴力破解
- auth_routes: 统一 DEV_MODE 判断收归 config.py 集中管理
- auth_routes: update_user 增加用户存在性校验
- gateway: require_role 装饰器恢复实际角色校验(原为占位实现)
- session_routes: get_user_sessions 统一传入 limit=20 参数
- session: get_history SQL 增加 id DESC 排序,修复同秒消息顺序错乱

🤖 Generated with [Qoder][https://qoder.com]
2026-06-08 16:12:56 +08:00
lacerate551
0b2ef8c161 feat(rag): 数据驱动的检索优化与表格/图片处理增强
- 移除硬编码关键词列表,改为数据驱动的图片意图检测
- 新增 _section_similarity() 层级章节相似度函数,替代正则匹配
- 新增 _rescue_table_chunks() 表格救援机制,基于 _in_budget_context 标记
- 修复 _build_context_with_budget 表格组超预算不 break 的问题
- 新增 _filter_images_by_answer() 后置图片过滤
- 表格切片 rerank 分数保护策略(同 section 表格不被误删)
- 跨页表格合并逻辑改进(通用标题检测 + 页码不可用降级策略)
- 意图分析增强追问补全 + 缓存类型隔离
- 语义缓存 key 加入 collections 防止跨上下文误命中
- 使用 retrieval_query 替代原始 message 进行检索
- engine.py: CloudReranker 模型更新 + 移除 top_n + table 邻居扩展
- LLM prompt 增加表格处理规则和章节路径提示

🤖 Generated with [Qoder][https://qoder.com]
2026-06-08 16:08:32 +08:00
lacerate551
8e3e9832ff feat(parser): 新增标题识别规则引擎
- 新增 parsers/heading_rules.py,支持可配置的标题层级检测规则
- 支持短中文文本标题识别的独立开关控制

🤖 Generated with [Qoder][https://qoder.com]
2026-06-08 16:06:01 +08:00
lacerate551
43261e9aff refactor(parser): 重构 MinerU 文档解析器
- 重写 MinerU 解析模块,改进文档解析流程和结构化输出
- 优化表格识别、标题层级提取和多格式文档处理
- 支持 v2 格式偏好(MINERU_PREFER_V2)

🤖 Generated with [Qoder][https://qoder.com]
2026-06-08 16:05:54 +08:00
lacerate551
a340eaaeee docs: 重写 RAG 系统指南并重命名(移除 AgenticRAG 内容)
- 全面重写文档,反映统一编排路径和四层缓存架构
- 添加 Query Cache 修复说明和语义缓存集成文档
- 重命名 Agentic_RAG完整指南.md → RAG系统完整指南.md
- 同步更新开发与系统模块说明.md 中的引用链接
2026-06-08 16:05:43 +08:00
lacerate551
8af8d38c01 perf(cache): 修复缓存 key 不匹配 Bug + embedding 缓存 + 清理 AgenticRAG 死代码
- cache.py: get/set 统一使用粗粒度 key(基于 kb_version),修复缓存永远不命中的问题
- engine.py: 新增 _encode_cached() 方法,embedding 编码走 LRU 缓存
- 删除 10 个 core/agentic_*.py 死文件(~1950 行)
- 清理 api/__init__.py 和 chat_routes.py 中的 AgenticRAG 残留引用
2026-06-08 16:05:28 +08:00
31 changed files with 3743 additions and 4651 deletions

5
.gitignore vendored
View File

@@ -120,8 +120,9 @@ test_*.json
rag_response.json rag_response.json
nul nul
# 临时调试脚本(下划线开头 # 调试脚本和临时计划(仅本地使用
scripts/_*.py scripts/
plans/
# Qoder 工具目录 # Qoder 工具目录
.qoder/ .qoder/

View File

@@ -132,10 +132,6 @@ def create_app() -> 'Flask':
from api.image_routes import image_bp from api.image_routes import image_bp
app.register_blueprint(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 from api.auth_routes import auth_bp
app.register_blueprint(auth_bp) app.register_blueprint(auth_bp)
@@ -258,5 +254,4 @@ def _print_startup_info(app: 'Flask') -> None:
logger.info(" 切片管理: /chunks/*") logger.info(" 切片管理: /chunks/*")
logger.info(" 同步服务: /sync, /sync/status") logger.info(" 同步服务: /sync, /sync/status")
logger.info(" 图片服务: /images/*") logger.info(" 图片服务: /images/*")
logger.info(" 任务查询: /tasks, /tasks/<id>, /tasks/<id>/progress")
logger.info(" 健康检查: /health") logger.info(" 健康检查: /health")

View File

@@ -11,8 +11,6 @@ import logging
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
from auth.gateway import require_gateway_auth from auth.gateway import require_gateway_auth
from data.db import get_connection 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__) audit_bp = Blueprint('audit', __name__)
@@ -96,11 +94,11 @@ def get_audit_logs():
"timestamp": row[10] "timestamp": row[10]
}) })
return success_response(data={"logs": logs, "total": total}) return jsonify({"logs": logs, "total": total})
except Exception as e: except Exception as e:
logger.error(f"审计查询异常: {e}") logger.error(f"审计查询异常: {e}")
return error_response("QUERY_FAILED", INTERNAL_ERROR, "查询失败", http_status=500) return jsonify({"error": "查询失败", "logs": [], "total": 0}), 500
def log_audit_event(user_id: str, username: str, action: str, def log_audit_event(user_id: str, username: str, action: str,

View File

@@ -10,9 +10,8 @@
from flask import Blueprint, request, jsonify from flask import Blueprint, request, jsonify
from auth.gateway import require_gateway_auth, require_role, get_user_permissions, MOCK_USERS from auth.gateway import require_gateway_auth, require_role, get_user_permissions, MOCK_USERS
from core.status_codes import SUCCESS, BAD_REQUEST, UNAUTHORIZED, FORBIDDEN, PERMISSION_DENIED, INTERNAL_ERROR
from api.response_utils import success_response, error_response
import os import os
import time
from pathlib import Path from pathlib import Path
from dotenv import load_dotenv from dotenv import load_dotenv
@@ -23,6 +22,37 @@ load_dotenv(env_path)
auth_bp = Blueprint('auth', __name__) auth_bp = Blueprint('auth', __name__)
# ==================== 登录速率限制 ====================
# 简单的内存级速率限制:每个 IP 在时间窗口内最多允许 N 次登录尝试
_login_attempts = {} # {ip: [(timestamp, ...), ...]}
_RATE_LIMIT_WINDOW = 300 # 5 分钟窗口
_RATE_LIMIT_MAX = 10 # 窗口内最多 10 次尝试
def _check_rate_limit(client_ip: str) -> bool:
"""检查是否超出登录速率限制,返回 True 表示允许"""
now = time.time()
if client_ip not in _login_attempts:
_login_attempts[client_ip] = []
# 清理过期记录
_login_attempts[client_ip] = [
t for t in _login_attempts[client_ip]
if now - t < _RATE_LIMIT_WINDOW
]
if len(_login_attempts[client_ip]) >= _RATE_LIMIT_MAX:
return False
_login_attempts[client_ip].append(now)
return True
def _is_dev_mode() -> bool:
"""统一的开发模式判断"""
return os.environ.get('DEV_MODE', 'true').lower() != 'false'
@auth_bp.route('/auth/login', methods=['POST']) @auth_bp.route('/auth/login', methods=['POST'])
def mock_login(): def mock_login():
""" """
@@ -52,9 +82,13 @@ def mock_login():
- manager / manager123 (经理,财务部) - manager / manager123 (经理,财务部)
- user / test123 (普通用户,技术部) - user / test123 (普通用户,技术部)
""" """
# 默认开启开发模式(生产环境需设置 DEV_MODE=false if not _is_dev_mode():
if os.environ.get('DEV_MODE', 'true').lower() == 'false': return jsonify({"error": "仅开发环境可用,请设置 DEV_MODE=true"}), 403
return error_response("FORBIDDEN", FORBIDDEN, "仅开发环境可用,请设置 DEV_MODE=true", http_status=403)
# 速率限制检查
client_ip = request.remote_addr or 'unknown'
if not _check_rate_limit(client_ip):
return jsonify({"error": f"登录尝试过于频繁,请 {_RATE_LIMIT_WINDOW // 60} 分钟后再试"}), 429
data = request.json or {} data = request.json or {}
username = data.get('username') username = data.get('username')
@@ -62,9 +96,9 @@ def mock_login():
user = MOCK_USERS.get(username) user = MOCK_USERS.get(username)
if not user or user['password'] != password: if not user or user['password'] != password:
return error_response("UNAUTHORIZED", UNAUTHORIZED, "用户名或密码错误", http_status=401) return jsonify({"error": "用户名或密码错误"}), 401
return success_response(data={ return jsonify({
"token": f"mock-token-{username}", "token": f"mock-token-{username}",
"user": { "user": {
"user_id": user['user_id'], "user_id": user['user_id'],
@@ -81,8 +115,10 @@ def mock_login():
def get_stats(): def get_stats():
"""获取系统统计信息(仅管理员)""" """获取系统统计信息(仅管理员)"""
from flask import current_app from flask import current_app
session_manager = current_app.config['SESSION_MANAGER'] session_manager = current_app.config.get('SESSION_MANAGER')
return success_response(data=session_manager.get_stats()) if not session_manager:
return jsonify({"error": "会话管理器未启用"}), 503
return jsonify(session_manager.get_stats())
@auth_bp.route('/health', methods=['GET']) @auth_bp.route('/health', methods=['GET'])
@@ -105,7 +141,7 @@ def get_current_user():
开发模式下支持模拟用户,生产模式下用户信息由后端控制。 开发模式下支持模拟用户,生产模式下用户信息由后端控制。
""" """
user = request.current_user user = request.current_user
return success_response(data={ return jsonify({
"user_id": user["user_id"], "user_id": user["user_id"],
"username": user["username"], "username": user["username"],
"role": user["role"], "role": user["role"],
@@ -133,9 +169,8 @@ def get_users():
] ]
} }
""" """
dev_mode = os.environ.get('DEV_MODE', 'true').lower() != 'false' if not _is_dev_mode():
if not dev_mode: return jsonify({"error": "仅开发环境可用"}), 403
return error_response("FORBIDDEN", FORBIDDEN, "仅开发环境可用", http_status=403)
users = [] users = []
for username, info in MOCK_USERS.items(): for username, info in MOCK_USERS.items():
@@ -147,7 +182,7 @@ def get_users():
"is_active": True # 模拟用户默认都是活跃状态 "is_active": True # 模拟用户默认都是活跃状态
}) })
return success_response(data={"users": users}) return jsonify({"users": users})
@auth_bp.route('/auth/users/<user_id>', methods=['PUT']) @auth_bp.route('/auth/users/<user_id>', methods=['PUT'])
@@ -161,12 +196,26 @@ def update_user(user_id):
"is_active": false "is_active": false
} }
""" """
dev_mode = os.environ.get('DEV_MODE', 'true').lower() != 'false' if not _is_dev_mode():
if not dev_mode: return jsonify({"error": "仅开发环境可用"}), 403
return error_response("FORBIDDEN", FORBIDDEN, "仅开发环境可用", http_status=403)
# 模拟用户不支持真正的状态切换,直接返回成功 # 验证目标用户是否存在
return success_response(data={"message": "操作成功(模拟)", "user_id": user_id}) target_user = None
for username, info in MOCK_USERS.items():
if info['user_id'] == user_id:
target_user = info
break
if not target_user:
return jsonify({"error": f"用户 {user_id} 不存在"}), 404
data = request.json or {}
# 模拟操作记录请求但不实际执行mock 用户数据是静态的)
return jsonify({
"message": "操作成功(模拟)",
"user_id": user_id,
"applied_changes": data
})
@auth_bp.route('/auth/change-password', methods=['POST']) @auth_bp.route('/auth/change-password', methods=['POST'])
@@ -181,19 +230,25 @@ def change_password():
"new_password": "xxx" "new_password": "xxx"
} }
""" """
dev_mode = os.environ.get('DEV_MODE', 'true').lower() != 'false' if not _is_dev_mode():
if not dev_mode: return jsonify({"error": "仅开发环境可用"}), 403
return error_response("FORBIDDEN", FORBIDDEN, "仅开发环境可用", http_status=403)
data = request.json or {} data = request.json or {}
old_password = data.get('old_password') old_password = data.get('old_password')
new_password = data.get('new_password') new_password = data.get('new_password')
if not old_password or not new_password: if not old_password or not new_password:
return error_response("MISSING_PARAMS", BAD_REQUEST, "请提供旧密码和新密码", http_status=400) return jsonify({"error": "请提供旧密码和新密码"}), 400
if len(new_password) < 6: if len(new_password) < 6:
return error_response("INVALID_PARAMS", BAD_REQUEST, "新密码至少6位", http_status=400) return jsonify({"error": "新密码至少6位"}), 400
# 模拟环境直接返回成功 # 验证当前用户的旧密码
return success_response(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 jsonify({"error": "旧密码错误"}), 401
# 模拟环境返回成功不实际修改密码mock 数据是静态的)
return jsonify({"message": "密码修改成功(模拟)"})

File diff suppressed because it is too large Load Diff

View File

@@ -39,18 +39,17 @@ import re
import uuid import uuid
from datetime import datetime from datetime import datetime
from typing import Optional, Tuple, Any, List, Dict from typing import Optional, Tuple, Any, List, Dict
from flask import Blueprint, request from flask import Blueprint, request, jsonify
import logging import logging
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
from werkzeug.utils import secure_filename from werkzeug.utils import secure_filename
from auth.gateway import require_gateway_auth from auth.gateway import require_gateway_auth
from config import DEV_MODE
from core.status_codes import ( from core.status_codes import (
UPLOAD_SUCCESS, BATCH_UPLOAD_SUCCESS, BAD_REQUEST, SUCCESS, UPLOAD_SUCCESS, BATCH_UPLOAD_SUCCESS, BAD_REQUEST,
NO_FILE, NO_FILE_SELECTED, NO_COLLECTION, NOT_FOUND, NO_FILE, NO_FILE_SELECTED, NO_COLLECTION,
FILE_TOO_LARGE, UNSUPPORTED_FORMAT, INTERNAL_ERROR, 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 from api.response_utils import success_response, error_response
@@ -171,10 +170,10 @@ def serve_document_file(doc_path: str) -> Tuple[Any, int]:
文件内容或错误响应 文件内容或错误响应
Note: Note:
仅在 DEV_MODE=true 时可用 仅在 DEV_MODE=true 时可用(需在 .env 中显式设置)
""" """
if os.environ.get('DEV_MODE', 'true').lower() == 'false': if not DEV_MODE:
return error_response("FORBIDDEN", FORBIDDEN, "仅开发环境可用", http_status=403) return jsonify({"error": "仅开发环境可用"}), 403
from config import DOCUMENTS_PATH from config import DOCUMENTS_PATH
from flask import send_from_directory from flask import send_from_directory
@@ -183,9 +182,9 @@ def serve_document_file(doc_path: str) -> Tuple[Any, int]:
try: try:
filepath = _validate_doc_path(doc_path, DOCUMENTS_PATH) filepath = _validate_doc_path(doc_path, DOCUMENTS_PATH)
except ValueError: except ValueError:
return error_response("FORBIDDEN", FORBIDDEN, "非法路径", http_status=403) return jsonify({"error": "非法路径"}), 403
if not os.path.exists(filepath): if not os.path.exists(filepath):
return error_response("NOT_FOUND", NOT_FOUND, "文件不存在", http_status=404) return jsonify({"error": "文件不存在"}), 404
directory = os.path.dirname(filepath) directory = os.path.dirname(filepath)
filename = os.path.basename(filepath) filename = os.path.basename(filepath)
@@ -336,24 +335,10 @@ def upload_document() -> Tuple[Any, int]:
except Exception as e: except Exception as e:
logger.warning(f"标记旧版本失败: {e}") logger.warning(f"标记旧版本失败: {e}")
# 6. 触发向量化(异步任务) # 6. 触发向量化
sync_status = "已保存,等待手动同步" sync_status = "已保存,等待手动同步"
sync_service = _get_sync_service() sync_service = _get_sync_service()
task_id = None
if sync_service: 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: try:
from knowledge.sync import DocumentChange, ChangeType from knowledge.sync import DocumentChange, ChangeType
change = DocumentChange( change = DocumentChange(
@@ -364,12 +349,11 @@ def upload_document() -> Tuple[Any, int]:
new_hash=sync_service.calculate_file_hash(filepath), new_hash=sync_service.calculate_file_hash(filepath),
change_time=datetime.now() change_time=datetime.now()
) )
registry.start_task(task.id, _do_vectorize, sync_service, change, filename) sync_service.process_change(change)
sync_status = "已保存,向量化任务已启动" sync_status = "已保存并添加到向量库"
except Exception as e: except Exception as e:
logger.warning(f"创建向量化任务失败: {e}") logger.warning(f"向量化失败: {e}")
registry.fail_task(task.id, str(e)) sync_status = "已保存,向量化失败"
sync_status = "已保存,向量化任务创建失败"
return success_response( return success_response(
data={ data={
@@ -380,8 +364,7 @@ def upload_document() -> Tuple[Any, int]:
"size": file_size, "size": file_size,
"replaced": replaced "replaced": replaced
}, },
"sync_status": sync_status, "sync_status": sync_status
"task_id": task_id
}, },
status_code=UPLOAD_SUCCESS, status_code=UPLOAD_SUCCESS,
message=f"文件上传成功,{sync_status}" message=f"文件上传成功,{sync_status}"
@@ -524,55 +507,14 @@ def batch_upload_documents() -> Tuple[Any, int]:
"message": "上传处理失败" "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( return success_response(
data={ data={
"total": len(results), "total": len(results),
"success_count": success_count, "success_count": len([r for r in results if r["status"] == "success"]),
"results": results, "results": results
"task_id": task_id
}, },
status_code=BATCH_UPLOAD_SUCCESS, status_code=BATCH_UPLOAD_SUCCESS,
message=f"批量上传完成,成功 {success_count}/{len(results)} 个文件" message=f"批量上传完成,成功 {len([r for r in results if r['status'] == 'success'])}/{len(results)} 个文件"
) )
@@ -649,7 +591,10 @@ def list_documents() -> Tuple[Any, int]:
# 按修改时间倒序 # 按修改时间倒序
documents.sort(key=lambda x: x['last_modified'], reverse=True) documents.sort(key=lambda x: x['last_modified'], reverse=True)
return success_response(data={"documents": documents, "total": len(documents)}) return jsonify({
"documents": documents,
"total": len(documents)
})
@document_bp.route('/documents/<path:doc_path>/status', methods=['GET']) @document_bp.route('/documents/<path:doc_path>/status', methods=['GET'])
@@ -673,12 +618,12 @@ def get_document_status(doc_path: str) -> Tuple[Any, int]:
""" """
kb_manager = _get_kb_manager() kb_manager = _get_kb_manager()
if not kb_manager: if not kb_manager:
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "知识库管理器未初始化", http_status=503) return jsonify({"error": "知识库管理器未初始化"}), 503
# 解析路径 # 解析路径
parts = doc_path.split('/') parts = doc_path.split('/')
if len(parts) < 2: if len(parts) < 2:
return error_response("BAD_REQUEST", BAD_REQUEST, "无效的文档路径", http_status=400) return jsonify({"error": "无效的文档路径"}), 400
subdir = parts[0] subdir = parts[0]
filename = '/'.join(parts[1:]) filename = '/'.join(parts[1:])
@@ -694,14 +639,16 @@ def get_document_status(doc_path: str) -> Tuple[Any, int]:
if not doc_info: if not doc_info:
if file_on_disk: if file_on_disk:
return success_response(data={ return jsonify({
"success": True,
"status": "unprocessed", "status": "unprocessed",
"chunk_count": 0, "chunk_count": 0,
"last_processed": None "last_processed": None
}) })
return error_response("NOT_FOUND", NOT_FOUND, "文档不存在", http_status=404) return jsonify({"error": "文档不存在"}), 404
return success_response(data={ return jsonify({
"success": True,
"status": doc_info.get("status", "unknown"), "status": doc_info.get("status", "unknown"),
"chunk_count": doc_info.get("total_chunks", 0), "chunk_count": doc_info.get("total_chunks", 0),
"last_processed": doc_info.get("effective_date") or doc_info.get("version") "last_processed": doc_info.get("effective_date") or doc_info.get("version")
@@ -728,35 +675,35 @@ def update_document(doc_path: str) -> Tuple[Any, int]:
from config import DOCUMENTS_PATH from config import DOCUMENTS_PATH
if 'file' not in request.files: if 'file' not in request.files:
return error_response("NO_FILE", NO_FILE, "没有上传文件", http_status=400) return jsonify({"error": "没有上传文件"}), 400
file = request.files['file'] file = request.files['file']
if file.filename == '': if file.filename == '':
return error_response("NO_FILE_SELECTED", NO_FILE_SELECTED, "没有选择文件", http_status=400) return jsonify({"error": "没有选择文件"}), 400
# 文件类型校验 # 文件类型校验
ext = os.path.splitext(file.filename)[1].lower() ext = os.path.splitext(file.filename)[1].lower()
if ext not in ALLOWED_EXTENSIONS: if ext not in ALLOWED_EXTENSIONS:
return error_response("UNSUPPORTED_FORMAT", UNSUPPORTED_FORMAT, f"不支持的文件类型: {ext}", http_status=400) return jsonify({"error": f"不支持的文件类型: {ext}"}), 400
# 文件大小校验 # 文件大小校验
file.seek(0, 2) # 跳到文件末尾获取大小 file.seek(0, 2) # 跳到文件末尾获取大小
file_size = file.tell() file_size = file.tell()
file.seek(0) # 回到文件开头 file.seek(0) # 回到文件开头
if file_size > MAX_FILE_SIZE: if file_size > MAX_FILE_SIZE:
return error_response("FILE_TOO_LARGE", FILE_TOO_LARGE, f"文件过大(最大 {MAX_FILE_SIZE // 1024 // 1024}MB", http_status=400) return jsonify({"error": f"文件过大(最大 {MAX_FILE_SIZE // 1024 // 1024}MB"}), 400
# 解析路径 # 解析路径
parts = doc_path.split('/') parts = doc_path.split('/')
if len(parts) < 2: if len(parts) < 2:
return error_response("BAD_REQUEST", BAD_REQUEST, "无效的文档路径", http_status=400) return jsonify({"error": "无效的文档路径"}), 400
try: try:
filepath = _validate_doc_path(doc_path, DOCUMENTS_PATH) filepath = _validate_doc_path(doc_path, DOCUMENTS_PATH)
except ValueError: except ValueError:
return error_response("FORBIDDEN", FORBIDDEN, "非法路径", http_status=403) return jsonify({"error": "非法路径"}), 403
if not os.path.exists(filepath): if not os.path.exists(filepath):
return error_response("NOT_FOUND", NOT_FOUND, "文件不存在", http_status=404) return jsonify({"error": "文件不存在"}), 404
# 覆盖文件 # 覆盖文件
file.save(filepath) file.save(filepath)
@@ -780,7 +727,10 @@ def update_document(doc_path: str) -> Tuple[Any, int]:
except Exception as e: except Exception as e:
logger.warning(f"重新向量化失败: {e}") logger.warning(f"重新向量化失败: {e}")
return success_response(message="文件已更新") return jsonify({
"success": True,
"message": "文件已更新"
})
@document_bp.route('/documents/<path:doc_path>', methods=['DELETE']) @document_bp.route('/documents/<path:doc_path>', methods=['DELETE'])
@@ -805,7 +755,7 @@ def delete_document(doc_path: str) -> Tuple[Any, int]:
# 解析路径 # 解析路径
parts = doc_path.split('/') parts = doc_path.split('/')
if len(parts) < 2: if len(parts) < 2:
return error_response("BAD_REQUEST", BAD_REQUEST, "无效的文档路径", http_status=400) return jsonify({"error": "无效的文档路径"}), 400
subdir = parts[0] subdir = parts[0]
filename = '/'.join(parts[1:]) filename = '/'.join(parts[1:])
@@ -816,9 +766,9 @@ def delete_document(doc_path: str) -> Tuple[Any, int]:
try: try:
filepath = _validate_doc_path(doc_path, DOCUMENTS_PATH) filepath = _validate_doc_path(doc_path, DOCUMENTS_PATH)
except ValueError: except ValueError:
return error_response("FORBIDDEN", FORBIDDEN, "非法路径", http_status=403) return jsonify({"error": "非法路径"}), 403
if not os.path.exists(filepath): if not os.path.exists(filepath):
return error_response("NOT_FOUND", NOT_FOUND, "文件不存在", http_status=404) return jsonify({"error": "文件不存在"}), 404
try: try:
# 1. 从向量库删除source 存的是文件名,不是完整路径) # 1. 从向量库删除source 存的是文件名,不是完整路径)
@@ -829,11 +779,14 @@ def delete_document(doc_path: str) -> Tuple[Any, int]:
# 2. 删除文件 # 2. 删除文件
os.remove(filepath) os.remove(filepath)
return success_response(status_code=DELETE_SUCCESS, message="文档已删除") return jsonify({
"success": True,
"message": "文档已删除"
})
except Exception as e: except Exception as e:
logger.error(f"删除文档异常: {e}") logger.error(f"删除文档异常: {e}")
return error_response("INTERNAL_ERROR", INTERNAL_ERROR, "删除失败,请稍后重试", http_status=500) return jsonify({"error": "删除失败,请稍后重试"}), 500
@document_bp.route('/documents/<path:doc_path>/chunks', methods=['GET']) @document_bp.route('/documents/<path:doc_path>/chunks', methods=['GET'])
@@ -858,12 +811,12 @@ def list_document_chunks(doc_path: str) -> Tuple[Any, int]:
""" """
kb_manager = _get_kb_manager() kb_manager = _get_kb_manager()
if not kb_manager: if not kb_manager:
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "知识库管理器未初始化", http_status=503) return jsonify({"error": "知识库管理器未初始化"}), 503
# 解析路径 # 解析路径
parts = doc_path.split('/') parts = doc_path.split('/')
if len(parts) < 2: if len(parts) < 2:
return error_response("BAD_REQUEST", BAD_REQUEST, "无效的文档路径", http_status=400) return jsonify({"error": "无效的文档路径"}), 400
subdir = parts[0] subdir = parts[0]
# 目录名即向量库名 # 目录名即向量库名
@@ -871,7 +824,8 @@ def list_document_chunks(doc_path: str) -> Tuple[Any, int]:
chunks = kb_manager.get_document_chunks(collection, os.path.basename(doc_path)) chunks = kb_manager.get_document_chunks(collection, os.path.basename(doc_path))
return success_response(data={ return jsonify({
"success": True,
"document_id": doc_path, "document_id": doc_path,
"collection": collection, "collection": collection,
"chunks": chunks, "chunks": chunks,
@@ -907,12 +861,12 @@ def preview_document(doc_path: str) -> Tuple[Any, int]:
""" """
kb_manager = _get_kb_manager() kb_manager = _get_kb_manager()
if not kb_manager: if not kb_manager:
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "知识库管理器未初始化", http_status=503) return jsonify({"error": "知识库管理器未初始化"}), 503
# 解析路径 # 解析路径
parts = doc_path.split('/') parts = doc_path.split('/')
if len(parts) < 2: if len(parts) < 2:
return error_response("BAD_REQUEST", BAD_REQUEST, "无效的文档路径,格式: collection/filename", http_status=400) return jsonify({"error": "无效的文档路径,格式: collection/filename"}), 400
collection = parts[0] collection = parts[0]
filename = os.path.basename(doc_path) filename = os.path.basename(doc_path)
@@ -927,7 +881,7 @@ def preview_document(doc_path: str) -> Tuple[Any, int]:
# 获取所有切片 # 获取所有切片
all_chunks = kb_manager.get_document_chunks(collection, filename) all_chunks = kb_manager.get_document_chunks(collection, filename)
if not all_chunks: if not all_chunks:
return error_response("NOT_FOUND", NOT_FOUND, f"文档 '{filename}' 不存在或无切片", http_status=404) return jsonify({"error": f"文档 '{filename}' 不存在或无切片"}), 404
total = len(all_chunks) total = len(all_chunks)
@@ -936,7 +890,8 @@ def preview_document(doc_path: str) -> Tuple[Any, int]:
preview_chunks = all_chunks[:5] preview_chunks = all_chunks[:5]
for c in preview_chunks: for c in preview_chunks:
c['is_target'] = False c['is_target'] = False
return success_response(data={ return jsonify({
"success": True,
"collection": collection, "collection": collection,
"source": filename, "source": filename,
"total_chunks": total, "total_chunks": total,
@@ -948,7 +903,7 @@ def preview_document(doc_path: str) -> Tuple[Any, int]:
try: try:
target_chunk_index = int(chunk_index_str) target_chunk_index = int(chunk_index_str)
except ValueError: except ValueError:
return error_response("BAD_REQUEST", BAD_REQUEST, "chunk_index 必须为整数", http_status=400) return jsonify({"error": "chunk_index 必须为整数"}), 400
# 按 chunk_index 排序Chroma 返回顺序不保证有序) # 按 chunk_index 排序Chroma 返回顺序不保证有序)
all_chunks.sort(key=lambda c: c.get('metadata', {}).get('chunk_index', 0)) all_chunks.sort(key=lambda c: c.get('metadata', {}).get('chunk_index', 0))
@@ -969,7 +924,9 @@ def preview_document(doc_path: str) -> Tuple[Any, int]:
(c.get('metadata', {}).get('chunk_index', 0) for c in all_chunks), (c.get('metadata', {}).get('chunk_index', 0) for c in all_chunks),
default=total - 1 default=total - 1
) )
return error_response("BAD_REQUEST", BAD_REQUEST, f"chunk_index={target_chunk_index} 未找到对应切片 (可用范围 0-{max_idx})", http_status=400) return jsonify({
"error": f"chunk_index={target_chunk_index} 未找到对应切片 (可用范围 0-{max_idx})"
}), 400
# 截取上下文窗口 # 截取上下文窗口
start = max(0, target_pos - context_count) start = max(0, target_pos - context_count)
@@ -980,7 +937,8 @@ def preview_document(doc_path: str) -> Tuple[Any, int]:
for i, c in enumerate(window): for i, c in enumerate(window):
c['is_target'] = (start + i == target_pos) c['is_target'] = (start + i == target_pos)
return success_response(data={ return jsonify({
"success": True,
"collection": collection, "collection": collection,
"source": filename, "source": filename,
"total_chunks": total, "total_chunks": total,
@@ -1011,7 +969,7 @@ def create_chunk() -> Tuple[Any, int]:
""" """
kb_manager = _get_kb_manager() kb_manager = _get_kb_manager()
if not kb_manager: if not kb_manager:
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "知识库管理器未初始化", http_status=503) return jsonify({"error": "知识库管理器未初始化"}), 503
data = request.json or {} data = request.json or {}
collection = data.get('collection') collection = data.get('collection')
@@ -1019,13 +977,17 @@ def create_chunk() -> Tuple[Any, int]:
metadata = data.get('metadata', {}) metadata = data.get('metadata', {})
if not collection: if not collection:
return error_response("NO_COLLECTION", NO_COLLECTION, "请指定向量库 (collection)", http_status=400) return jsonify({"error": "请指定向量库 (collection)"}), 400
if not content: if not content:
return error_response("NO_CONTENT", NO_CONTENT, "切片内容不能为空", http_status=400) return jsonify({"error": "切片内容不能为空"}), 400
chunk_id = kb_manager.add_chunk(collection, content, metadata) chunk_id = kb_manager.add_chunk(collection, content, metadata)
return success_response(data={"chunk_id": chunk_id}, message="切片已添加") return jsonify({
"success": True,
"chunk_id": chunk_id,
"message": "切片已添加"
})
@document_bp.route('/chunks/<chunk_id>', methods=['PUT']) @document_bp.route('/chunks/<chunk_id>', methods=['PUT'])
@@ -1051,7 +1013,7 @@ def update_chunk(chunk_id: str) -> Tuple[Any, int]:
""" """
kb_manager = _get_kb_manager() kb_manager = _get_kb_manager()
if not kb_manager: if not kb_manager:
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "知识库管理器未初始化", http_status=503) return jsonify({"error": "知识库管理器未初始化"}), 503
data = request.json or {} data = request.json or {}
collection = data.get('collection') collection = data.get('collection')
@@ -1059,13 +1021,13 @@ def update_chunk(chunk_id: str) -> Tuple[Any, int]:
metadata = data.get('metadata') metadata = data.get('metadata')
if not collection: if not collection:
return error_response("NO_COLLECTION", NO_COLLECTION, "请指定向量库 (collection)", http_status=400) return jsonify({"error": "请指定向量库 (collection)"}), 400
success = kb_manager.update_chunk(collection, chunk_id, content=content, metadata=metadata) success = kb_manager.update_chunk(collection, chunk_id, content=content, metadata=metadata)
if success: if success:
return success_response(message="切片已更新") return jsonify({"success": True, "message": "切片已更新"})
return error_response("INTERNAL_ERROR", INTERNAL_ERROR, "更新失败", http_status=500) return jsonify({"error": "更新失败"}), 500
@document_bp.route('/chunks/<chunk_id>', methods=['DELETE']) @document_bp.route('/chunks/<chunk_id>', methods=['DELETE'])
@@ -1093,11 +1055,11 @@ def delete_chunk(chunk_id: str) -> Tuple[Any, int]:
collection = request.args.get('collection') collection = request.args.get('collection')
if not collection: if not collection:
return error_response("NO_COLLECTION", NO_COLLECTION, "请指定向量库 (collection)", http_status=400) return jsonify({"error": "请指定向量库 (collection)"}), 400
kb_manager = _get_kb_manager() kb_manager = _get_kb_manager()
if not kb_manager: if not kb_manager:
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "知识库管理器未初始化", http_status=503) return jsonify({"error": "知识库管理器未初始化"}), 503
# 删除切片,返回 (success, source_file) # 删除切片,返回 (success, source_file)
success, source_file = kb_manager.delete_chunk(collection, chunk_id) success, source_file = kb_manager.delete_chunk(collection, chunk_id)
@@ -1122,8 +1084,8 @@ def delete_chunk(chunk_id: str) -> Tuple[Any, int]:
logger.warning(f"清理哈希记录失败: {e}") logger.warning(f"清理哈希记录失败: {e}")
if success: if success:
return success_response(status_code=DELETE_SUCCESS, message="切片已删除") return jsonify({"success": True, "message": "切片已删除"})
return error_response("INTERNAL_ERROR", INTERNAL_ERROR, "删除失败", http_status=500) return jsonify({"error": "删除失败"}), 500
@document_bp.route('/chunks/batch', methods=['DELETE']) @document_bp.route('/chunks/batch', methods=['DELETE'])
@@ -1149,13 +1111,13 @@ def delete_chunks_by_source() -> Tuple[Any, int]:
source = data.get('source') source = data.get('source')
if not collection: if not collection:
return error_response("NO_COLLECTION", NO_COLLECTION, "请指定向量库 (collection)", http_status=400) return jsonify({"error": "请指定向量库 (collection)"}), 400
if not source: if not source:
return error_response("BAD_REQUEST", BAD_REQUEST, "请指定文件名 (source)", http_status=400) return jsonify({"error": "请指定文件名 (source)"}), 400
kb_manager = _get_kb_manager() kb_manager = _get_kb_manager()
if not kb_manager: if not kb_manager:
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "知识库管理器未初始化", http_status=503) return jsonify({"error": "知识库管理器未初始化"}), 503
# 批量删除该文件的所有切片 # 批量删除该文件的所有切片
try: try:
@@ -1173,7 +1135,11 @@ def delete_chunks_by_source() -> Tuple[Any, int]:
except Exception as e: except Exception as e:
logger.warning(f"清理哈希记录失败: {e}") logger.warning(f"清理哈希记录失败: {e}")
return success_response(data={"deleted_count": deleted_count}, message=f"已删除 {deleted_count} 个切片") return jsonify({
"success": True,
"deleted_count": deleted_count,
"message": f"已删除 {deleted_count} 个切片"
})
except Exception as e: except Exception as e:
logger.error(f"操作异常: {e}") logger.error(f"操作异常: {e}")
return error_response("INTERNAL_ERROR", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"error": "操作失败,请稍后重试"}), 500

View File

@@ -33,11 +33,6 @@
from flask import Blueprint, request, jsonify from flask import Blueprint, request, jsonify
from auth.gateway import require_gateway_auth, require_role 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 import logging
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -80,10 +75,10 @@ def submit_feedback():
user_id = data.get('user_id', '') user_id = data.get('user_id', '')
if not session_id or not query or rating is None: if not session_id or not query or rating is None:
return error_response("MISSING_PARAMS", BAD_REQUEST, "缺少必要参数", http_status=400) return jsonify({"error": "缺少必要参数"}), 400
if rating not in [1, -1]: if rating not in [1, -1]:
return error_response("INVALID_PARAMS", BAD_REQUEST, "rating 必须是 1 或 -1", http_status=400) return jsonify({"error": "rating 必须是 1 或 -1"}), 400
try: try:
feedback_service = _get_feedback_service() feedback_service = _get_feedback_service()
@@ -96,14 +91,15 @@ def submit_feedback():
reason=reason, reason=reason,
user_id=user_id user_id=user_id
) )
return success_response(data={ return jsonify({
"success": True,
"feedback_id": result['feedback_id'], "feedback_id": result['feedback_id'],
"faq_suggested": result.get('faq_suggested', False), "faq_suggested": result.get('faq_suggested', False),
"suggestion_id": result.get('suggestion_id') "suggestion_id": result.get('suggestion_id')
}) })
except Exception as e: except Exception as e:
logger.error(f"反馈操作异常: {e}") logger.error(f"反馈操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"error": "操作失败,请稍后重试"}), 500
@feedback_bp.route('/feedback/stats', methods=['GET']) @feedback_bp.route('/feedback/stats', methods=['GET'])
@@ -116,12 +112,13 @@ def get_feedback_stats():
try: try:
feedback_db = _get_feedback_db() feedback_db = _get_feedback_db()
stats = feedback_db.get_feedback_stats(start_date, end_date) stats = feedback_db.get_feedback_stats(start_date, end_date)
return success_response(data={ return jsonify({
"success": True,
"stats": stats "stats": stats
}) })
except Exception as e: except Exception as e:
logger.error(f"反馈操作异常: {e}") logger.error(f"反馈操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"error": "操作失败,请稍后重试"}), 500
@feedback_bp.route('/feedback/list', methods=['GET']) @feedback_bp.route('/feedback/list', methods=['GET'])
@@ -143,13 +140,14 @@ def get_feedback_list():
end_date=end_date, end_date=end_date,
limit=limit limit=limit
) )
return success_response(data={ return jsonify({
"success": True,
"feedbacks": feedbacks, "feedbacks": feedbacks,
"total": len(feedbacks) "total": len(feedbacks)
}) })
except Exception as e: except Exception as e:
logger.error(f"反馈操作异常: {e}") logger.error(f"反馈操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"error": "操作失败,请稍后重试"}), 500
@feedback_bp.route('/reports/weekly', methods=['GET']) @feedback_bp.route('/reports/weekly', methods=['GET'])
@@ -159,12 +157,13 @@ def get_weekly_report():
try: try:
feedback_service = _get_feedback_service() feedback_service = _get_feedback_service()
report = feedback_service.generate_report("weekly") report = feedback_service.generate_report("weekly")
return success_response(data={ return jsonify({
"success": True,
"report": report.to_dict() "report": report.to_dict()
}) })
except Exception as e: except Exception as e:
logger.error(f"反馈操作异常: {e}") logger.error(f"反馈操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"error": "操作失败,请稍后重试"}), 500
@feedback_bp.route('/reports/monthly', methods=['GET']) @feedback_bp.route('/reports/monthly', methods=['GET'])
@@ -174,12 +173,13 @@ def get_monthly_report():
try: try:
feedback_service = _get_feedback_service() feedback_service = _get_feedback_service()
report = feedback_service.generate_report("monthly") report = feedback_service.generate_report("monthly")
return success_response(data={ return jsonify({
"success": True,
"report": report.to_dict() "report": report.to_dict()
}) })
except Exception as e: except Exception as e:
logger.error(f"反馈操作异常: {e}") logger.error(f"反馈操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"error": "操作失败,请稍后重试"}), 500
@feedback_bp.route('/faq', methods=['GET']) @feedback_bp.route('/faq', methods=['GET'])
@@ -192,13 +192,14 @@ def get_faq_list():
try: try:
feedback_db = _get_feedback_db() feedback_db = _get_feedback_db()
faqs = feedback_db.get_faqs(status=status, limit=limit) faqs = feedback_db.get_faqs(status=status, limit=limit)
return success_response(data={ return jsonify({
"success": True,
"faqs": faqs, "faqs": faqs,
"total": len(faqs) "total": len(faqs)
}) })
except Exception as e: except Exception as e:
logger.error(f"反馈操作异常: {e}") logger.error(f"反馈操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"error": "操作失败,请稍后重试"}), 500
@feedback_bp.route('/faq', methods=['POST']) @feedback_bp.route('/faq', methods=['POST'])
@@ -219,7 +220,7 @@ def create_faq():
answer = data.get('answer') answer = data.get('answer')
if not question or not answer: if not question or not answer:
return error_response("MISSING_PARAMS", BAD_REQUEST, "缺少问题或答案", http_status=400) return jsonify({"error": "缺少问题或答案"}), 400
try: try:
from services.feedback import FAQ from services.feedback import FAQ
@@ -234,14 +235,15 @@ def create_faq():
) )
faq_id = feedback_db.add_faq(faq) faq_id = feedback_db.add_faq(faq)
return success_response(data={ return jsonify({
"success": True,
"faq_id": faq_id, "faq_id": faq_id,
"status": "draft", "status": "draft",
"message": "FAQ已创建请通过 /faq/<id>/approve 接口确认后生效" "message": "FAQ已创建请通过 /faq/<id>/approve 接口确认后生效"
}) })
except Exception as e: except Exception as e:
logger.error(f"反馈操作异常: {e}") logger.error(f"反馈操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"error": "操作失败,请稍后重试"}), 500
@feedback_bp.route('/faq/<int:faq_id>/approve', methods=['POST']) @feedback_bp.route('/faq/<int:faq_id>/approve', methods=['POST'])
@@ -261,10 +263,10 @@ def approve_faq(faq_id):
# 检查 FAQ 状态 # 检查 FAQ 状态
faq = feedback_db.get_faq(faq_id) faq = feedback_db.get_faq(faq_id)
if not faq: if not faq:
return error_response("NOT_FOUND", NOT_FOUND, "FAQ不存在", http_status=404) return jsonify({"error": "FAQ不存在"}), 404
if faq.get('status') == 'approved': if faq.get('status') == 'approved':
return success_response(message="FAQ已经是批准状态") return jsonify({"success": True, "message": "FAQ已经是批准状态"})
# 更新状态为 approved # 更新状态为 approved
feedback_db.update_faq(faq_id, {"status": "approved"}) feedback_db.update_faq(faq_id, {"status": "approved"})
@@ -277,14 +279,15 @@ def approve_faq(faq_id):
answer=faq['answer'] answer=faq['answer']
) )
return success_response(data={ return jsonify({
"success": True,
"faq_id": faq_id, "faq_id": faq_id,
"sync_status": "synced" if sync_success else "sync_failed", "sync_status": "synced" if sync_success else "sync_failed",
"message": "FAQ已批准并同步到知识库" "message": "FAQ已批准并同步到知识库"
}) })
except Exception as e: except Exception as e:
logger.error(f"反馈操作异常: {e}") logger.error(f"反馈操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"error": "操作失败,请稍后重试"}), 500
@feedback_bp.route('/faq/<int:faq_id>', methods=['PUT']) @feedback_bp.route('/faq/<int:faq_id>', methods=['PUT'])
@@ -300,11 +303,11 @@ def update_faq(faq_id):
# 获取更新前的 FAQ 信息(用于判断是否需要重新同步) # 获取更新前的 FAQ 信息(用于判断是否需要重新同步)
old_faq = feedback_db.get_faq(faq_id) old_faq = feedback_db.get_faq(faq_id)
if not old_faq: if not old_faq:
return error_response("NOT_FOUND", NOT_FOUND, "FAQ不存在", http_status=404) return jsonify({"error": "FAQ不存在"}), 404
updated = feedback_db.update_faq(faq_id, data) updated = feedback_db.update_faq(faq_id, data)
if not updated: if not updated:
return error_response("UPDATE_FAILED", INTERNAL_ERROR, "FAQ更新失败", http_status=500) return jsonify({"error": "FAQ更新失败"}), 500
# 检查是否需要重新同步向量库question 或 answer 变更时) # 检查是否需要重新同步向量库question 或 answer 变更时)
need_sync = False need_sync = False
@@ -328,13 +331,14 @@ def update_faq(faq_id):
) )
sync_status = "synced" if sync_success else "sync_failed" sync_status = "synced" if sync_success else "sync_failed"
return success_response(data={ return jsonify({
"success": True,
"message": "FAQ更新成功", "message": "FAQ更新成功",
"sync_status": sync_status "sync_status": sync_status
}) })
except Exception as e: except Exception as e:
logger.error(f"反馈操作异常: {e}") logger.error(f"反馈操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"error": "操作失败,请稍后重试"}), 500
@feedback_bp.route('/faq/<int:faq_id>', methods=['DELETE']) @feedback_bp.route('/faq/<int:faq_id>', methods=['DELETE'])
@@ -361,10 +365,13 @@ def delete_faq(faq_id):
feedback_service = _get_feedback_service() feedback_service = _get_feedback_service()
feedback_service._delete_faq_vectors(faq_id) feedback_service._delete_faq_vectors(faq_id)
return success_response(data={"deleted": deleted}, message="FAQ删除成功" if deleted else "FAQ不存在") return jsonify({
"success": deleted,
"message": "FAQ删除成功" if deleted else "FAQ不存在"
})
except Exception as e: except Exception as e:
logger.error(f"反馈操作异常: {e}") logger.error(f"反馈操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"error": "操作失败,请稍后重试"}), 500
@feedback_bp.route('/faq/suggestions', methods=['GET']) @feedback_bp.route('/faq/suggestions', methods=['GET'])
@@ -378,13 +385,14 @@ def get_faq_suggestions():
try: try:
feedback_db = _get_feedback_db() feedback_db = _get_feedback_db()
suggestions = feedback_db.get_faq_suggestions(status=status, limit=limit) suggestions = feedback_db.get_faq_suggestions(status=status, limit=limit)
return success_response(data={ return jsonify({
"success": True,
"suggestions": suggestions, "suggestions": suggestions,
"total": len(suggestions) "total": len(suggestions)
}) })
except Exception as e: except Exception as e:
logger.error(f"反馈操作异常: {e}") logger.error(f"反馈操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"error": "操作失败,请稍后重试"}), 500
@feedback_bp.route('/faq/suggestions/<int:suggestion_id>/approve', methods=['POST']) @feedback_bp.route('/faq/suggestions/<int:suggestion_id>/approve', methods=['POST'])
@@ -410,16 +418,17 @@ def approve_faq_suggestion(suggestion_id):
) )
if result.get('success'): if result.get('success'):
return success_response(data={ return jsonify({
"success": True,
"faq_id": result['faq_id'], "faq_id": result['faq_id'],
"sync_status": result.get('sync_status'), "sync_status": result.get('sync_status'),
"message": "FAQ建议已批准并同步到知识库" "message": "FAQ建议已批准并同步到知识库"
}) })
else: else:
return error_response("APPROVE_FAILED", BAD_REQUEST, result.get('error', '批准失败'), http_status=400) return jsonify({"error": result.get('error', '批准失败')}), 400
except Exception as e: except Exception as e:
logger.error(f"反馈操作异常: {e}") logger.error(f"反馈操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"error": "操作失败,请稍后重试"}), 500
@feedback_bp.route('/faq/suggestions/<int:suggestion_id>/reject', methods=['POST']) @feedback_bp.route('/faq/suggestions/<int:suggestion_id>/reject', methods=['POST'])
@@ -430,13 +439,13 @@ def reject_faq_suggestion(suggestion_id):
try: try:
feedback_db = _get_feedback_db() feedback_db = _get_feedback_db()
rejected = feedback_db.reject_faq_suggestion(suggestion_id) rejected = feedback_db.reject_faq_suggestion(suggestion_id)
return success_response(data={ return jsonify({
"rejected": rejected, "success": rejected,
"message": "FAQ建议已拒绝" if rejected else "建议不存在" "message": "FAQ建议已拒绝" if rejected else "建议不存在"
}) })
except Exception as e: except Exception as e:
logger.error(f"反馈操作异常: {e}") logger.error(f"反馈操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"error": "操作失败,请稍后重试"}), 500
# ==================== Bad Case 分析接口 ==================== # ==================== Bad Case 分析接口 ====================
@@ -468,7 +477,8 @@ def get_bad_cases():
for case in bad_cases: for case in bad_cases:
case['status'] = 'pending' # pending/resolved/ignored case['status'] = 'pending' # pending/resolved/ignored
return success_response(data={ return jsonify({
"success": True,
"bad_cases": bad_cases, "bad_cases": bad_cases,
"blacklisted_sources": blacklisted_sources, "blacklisted_sources": blacklisted_sources,
"suggestions": [ "suggestions": [
@@ -479,7 +489,7 @@ def get_bad_cases():
}) })
except Exception as e: except Exception as e:
logger.error(f"反馈操作异常: {e}") logger.error(f"反馈操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"error": "操作失败,请稍后重试"}), 500
@feedback_bp.route('/feedback/blacklist', methods=['GET']) @feedback_bp.route('/feedback/blacklist', methods=['GET'])
@@ -497,11 +507,12 @@ def get_chunk_blacklist():
feedback_service = _get_feedback_service() feedback_service = _get_feedback_service()
blacklist = feedback_service.get_chunk_blacklist(min_dislikes=min_dislikes) blacklist = feedback_service.get_chunk_blacklist(min_dislikes=min_dislikes)
return success_response(data={ return jsonify({
"success": True,
"blacklist": list(blacklist), "blacklist": list(blacklist),
"count": len(blacklist), "count": len(blacklist),
"usage": "在检索时过滤这些来源以提升回答质量" "usage": "在检索时过滤这些来源以提升回答质量"
}) })
except Exception as e: except Exception as e:
logger.error(f"反馈操作异常: {e}") logger.error(f"反馈操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"error": "操作失败,请稍后重试"}), 500

View File

@@ -11,8 +11,6 @@
import os import os
import logging import logging
from flask import Blueprint, send_file, jsonify, current_app 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__) logger = logging.getLogger(__name__)
@@ -53,7 +51,7 @@ def get_image(image_id: str):
# 安全检查:防止路径遍历攻击 # 安全检查:防止路径遍历攻击
if '..' in image_id or '/' in image_id or '\\' in image_id: if '..' in image_id or '/' in image_id or '\\' in image_id:
return error_response("BAD_REQUEST", BAD_REQUEST, "无效的图片 ID", http_status=400) return jsonify({"error": "无效的图片 ID"}), 400
images_path = get_images_base_path() images_path = get_images_base_path()
@@ -76,9 +74,9 @@ def get_image(image_id: str):
return send_file(image_path, mimetype=mimetype) return send_file(image_path, mimetype=mimetype)
except Exception as e: except Exception as e:
logger.error(f"读取图片异常: {e}") logger.error(f"读取图片异常: {e}")
return error_response("READ_ERROR", INTERNAL_ERROR, "读取图片失败", http_status=500) return jsonify({"error": "读取图片失败"}), 500
return error_response("NOT_FOUND", NOT_FOUND, "图片不存在", http_status=404, image_id=image_id) return jsonify({"error": "图片不存在", "image_id": image_id}), 404
@image_bp.route('/images/<image_id>/info', methods=['GET']) @image_bp.route('/images/<image_id>/info', methods=['GET'])
@@ -98,7 +96,7 @@ def get_image_info(image_id: str):
# 安全检查 # 安全检查
if '..' in image_id or '/' in image_id or '\\' in image_id: if '..' in image_id or '/' in image_id or '\\' in image_id:
return error_response("BAD_REQUEST", BAD_REQUEST, "无效的图片 ID", http_status=400) return jsonify({"error": "无效的图片 ID"}), 400
images_path = get_images_base_path() images_path = get_images_base_path()
supported_formats = ['.png', '.jpg', '.jpeg', '.gif', '.bmp', '.webp'] supported_formats = ['.png', '.jpg', '.jpeg', '.gif', '.bmp', '.webp']
@@ -111,7 +109,7 @@ def get_image_info(image_id: str):
from PIL import Image from PIL import Image
with Image.open(image_path) as img: with Image.open(image_path) as img:
return success_response(data={ return jsonify({
"image_id": image_id, "image_id": image_id,
"width": img.width, "width": img.width,
"height": img.height, "height": img.height,
@@ -122,16 +120,16 @@ def get_image_info(image_id: str):
}) })
except ImportError: except ImportError:
# PIL 未安装,返回基本信息 # PIL 未安装,返回基本信息
return success_response(data={ return jsonify({
"image_id": image_id, "image_id": image_id,
"size_bytes": os.path.getsize(image_path), "size_bytes": os.path.getsize(image_path),
"url": f"/images/{image_id}" "url": f"/images/{image_id}"
}) })
except Exception as e: except Exception as e:
logger.error(f"读取图片信息异常: {e}") logger.error(f"读取图片信息异常: {e}")
return error_response("READ_ERROR", INTERNAL_ERROR, "读取图片信息失败", http_status=500) return jsonify({"error": "读取图片信息失败"}), 500
return error_response("NOT_FOUND", NOT_FOUND, "图片不存在", http_status=404, image_id=image_id) return jsonify({"error": "图片不存在", "image_id": image_id}), 404
@image_bp.route('/images/list', methods=['GET']) @image_bp.route('/images/list', methods=['GET'])
@@ -154,7 +152,7 @@ def list_images():
images_path = get_images_base_path() images_path = get_images_base_path()
if not os.path.exists(images_path): if not os.path.exists(images_path):
return success_response(data={"images": [], "total": 0}) return jsonify({"images": [], "total": 0})
supported_extensions = {'.png', '.jpg', '.jpeg', '.gif', '.bmp', '.webp'} supported_extensions = {'.png', '.jpg', '.jpeg', '.gif', '.bmp', '.webp'}
images = [] images = []
@@ -178,7 +176,7 @@ def list_images():
total = len(images) total = len(images)
images = images[offset:offset + limit] images = images[offset:offset + limit]
return success_response(data={ return jsonify({
"images": images, "images": images,
"total": total, "total": total,
"limit": limit, "limit": limit,
@@ -187,7 +185,7 @@ def list_images():
except Exception as e: except Exception as e:
logger.error(f"列出图片异常: {e}") logger.error(f"列出图片异常: {e}")
return error_response("LIST_ERROR", INTERNAL_ERROR, "列出图片失败", http_status=500) return jsonify({"error": "列出图片失败"}), 500
@image_bp.route('/images/stats', methods=['GET']) @image_bp.route('/images/stats', methods=['GET'])
@@ -201,7 +199,7 @@ def image_stats():
images_path = get_images_base_path() images_path = get_images_base_path()
if not os.path.exists(images_path): if not os.path.exists(images_path):
return success_response(data={ return jsonify({
"total_images": 0, "total_images": 0,
"total_size_bytes": 0, "total_size_bytes": 0,
"supported_formats": ['.png', '.jpg', '.jpeg', '.gif', '.bmp', '.webp'] "supported_formats": ['.png', '.jpg', '.jpeg', '.gif', '.bmp', '.webp']
@@ -221,7 +219,7 @@ def image_stats():
total_size += os.path.getsize(filepath) total_size += os.path.getsize(filepath)
format_counts[ext] = format_counts.get(ext, 0) + 1 format_counts[ext] = format_counts.get(ext, 0) + 1
return success_response(data={ return jsonify({
"total_images": total_count, "total_images": total_count,
"total_size_bytes": total_size, "total_size_bytes": total_size,
"total_size_mb": round(total_size / (1024 * 1024), 2), "total_size_mb": round(total_size / (1024 * 1024), 2),
@@ -231,4 +229,4 @@ def image_stats():
except Exception as e: except Exception as e:
logger.error(f"获取统计信息异常: {e}") logger.error(f"获取统计信息异常: {e}")
return error_response("STATS_ERROR", INTERNAL_ERROR, "获取统计信息失败", http_status=500) return jsonify({"error": "获取统计信息失败"}), 500

View File

@@ -37,13 +37,6 @@ from typing import Tuple, Optional, Any
from flask import Blueprint, request, jsonify, current_app from flask import Blueprint, request, jsonify, current_app
import logging 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__) logger = logging.getLogger(__name__)
from auth.gateway import require_gateway_auth from auth.gateway import require_gateway_auth
@@ -140,7 +133,10 @@ def list_collections() -> Tuple[Any, int]:
"description": coll.description "description": coll.description
}) })
return success_response(data={"collections": result, "total": len(result)}) return jsonify({
"collections": result,
"total": len(result)
})
@kb_bp.route('/collections', methods=['POST']) @kb_bp.route('/collections', methods=['POST'])
@@ -176,19 +172,22 @@ def create_collection() -> Tuple[Any, int]:
description = data.get('description', '') description = data.get('description', '')
if not name: if not name:
return error_response("INVALID_NAME", BAD_REQUEST, "向量库名称不能为空", http_status=400) return jsonify({"error": "向量库名称不能为空"}), 400
# 验证名称格式ChromaDB 限制) # 验证名称格式ChromaDB 限制)
if not name.replace('_', '').replace('-', '').isalnum(): if not name.replace('_', '').replace('-', '').isalnum():
return error_response("INVALID_NAME_FORMAT", BAD_REQUEST, "向量库名称只能包含字母、数字、下划线和连字符", http_status=400) return jsonify({
"error": "名称格式错误",
"message": "向量库名称只能包含字母、数字、下划线和连字符"
}), 400
success, message = kb_manager.create_collection( success, message = kb_manager.create_collection(
name, display_name, department, description name, display_name, department, description
) )
if success: if success:
return success_response(data={"name": name}, status_code=CREATED, message=message, http_status=201) return jsonify({"success": True, "message": message, "name": name}), 201
return error_response("CREATE_FAILED", BAD_REQUEST, message, http_status=400) return jsonify({"error": message}), 400
@kb_bp.route('/collections/<kb_name>', methods=['PUT']) @kb_bp.route('/collections/<kb_name>', methods=['PUT'])
@@ -221,7 +220,7 @@ def update_collection(kb_name: str) -> Tuple[Any, int]:
# 检查向量库是否存在 # 检查向量库是否存在
collections = kb_manager.list_collections() collections = kb_manager.list_collections()
if not any(c.name == kb_name for c in collections): if not any(c.name == kb_name for c in collections):
return error_response("COLLECTION_NOT_FOUND", COLLECTION_NOT_FOUND, f"向量库 '{kb_name}' 不存在", http_status=404) return jsonify({"error": f"向量库 '{kb_name}' 不存在"}), 404
# 更新元数据 # 更新元数据
success = kb_manager.update_collection_metadata( success = kb_manager.update_collection_metadata(
@@ -231,8 +230,8 @@ def update_collection(kb_name: str) -> Tuple[Any, int]:
) )
if success: if success:
return success_response(data=None, status_code=UPDATE_SUCCESS, message="向量库信息已更新") return jsonify({"success": True, "message": "向量库信息已更新"})
return error_response("UPDATE_FAILED", INTERNAL_ERROR, "更新失败", http_status=500) return jsonify({"error": "更新失败"}), 500
@kb_bp.route('/collections/<kb_name>', methods=['DELETE']) @kb_bp.route('/collections/<kb_name>', methods=['DELETE'])
@@ -261,8 +260,12 @@ def delete_collection(kb_name: str) -> Tuple[Any, int]:
success, message = kb_manager.delete_collection(kb_name, delete_documents) success, message = kb_manager.delete_collection(kb_name, delete_documents)
if success: if success:
return success_response(data={"deleted_documents": delete_documents}, status_code=DELETE_SUCCESS, message=message) return jsonify({
return error_response("DELETE_FAILED", BAD_REQUEST, message, http_status=400) "success": True,
"message": message,
"deleted_documents": delete_documents
})
return jsonify({"error": message}), 400
@kb_bp.route('/collections/<kb_name>/documents', methods=['GET']) @kb_bp.route('/collections/<kb_name>/documents', methods=['GET'])
@@ -283,7 +286,11 @@ def list_collection_documents(kb_name: str) -> Tuple[Any, int]:
documents = kb_manager.list_documents(kb_name) documents = kb_manager.list_documents(kb_name)
return success_response(data={"collection": kb_name, "documents": documents, "total": len(documents)}) return jsonify({
"collection": kb_name,
"documents": documents,
"total": len(documents)
})
@kb_bp.route('/collections/<kb_name>/chunks', methods=['GET']) @kb_bp.route('/collections/<kb_name>/chunks', methods=['GET'])
@@ -314,16 +321,21 @@ 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) chunks = kb_manager.list_chunks(kb_name, document_id=document_id, limit=limit, offset=offset)
return success_response(data={"collection": kb_name, "chunks": chunks, "total": len(chunks)}) return jsonify({
"collection": kb_name,
"chunks": chunks,
"total": len(chunks)
})
@kb_bp.route('/documents/sync', methods=['POST']) @kb_bp.route('/documents/sync', methods=['POST'])
@require_gateway_auth @require_gateway_auth
def sync_documents() -> Tuple[Any, int]: def sync_documents() -> Tuple[Any, int]:
""" """
触发文档向量化同步(异步任务) 触发文档向量化同步
立即返回 task_id后台线程执行同步。 扫描文档目录,检测新增、修改、删除的文件,
自动更新向量库索引。
请求体: 请求体:
{ {
@@ -331,7 +343,11 @@ def sync_documents() -> Tuple[Any, int]:
} }
Returns: Returns:
{"success": true, "data": {"task_id": "xxx", "message": "..."}} {
"success": true,
"results": [{"collection": "...", "status": "...", ...}],
"synced_count": N
}
""" """
kb_manager, _, err = _require_multi_kb() kb_manager, _, err = _require_multi_kb()
if err: if err:
@@ -339,6 +355,7 @@ def sync_documents() -> Tuple[Any, int]:
from config import DOCUMENTS_PATH from config import DOCUMENTS_PATH
user = request.current_user
data = request.json or {} data = request.json or {}
target_collection = data.get('collection') target_collection = data.get('collection')
@@ -346,67 +363,54 @@ def sync_documents() -> Tuple[Any, int]:
if target_collection: if target_collection:
collections_to_sync = [target_collection] collections_to_sync = [target_collection]
else: else:
# 同步所有向量库
all_collections = kb_manager.list_collections() all_collections = kb_manager.list_collections()
collections_to_sync = [c.name for c in all_collections] collections_to_sync = [c.name for c in all_collections]
if not collections_to_sync: if not collections_to_sync:
return error_response("NO_COLLECTION", NO_COLLECTION, "没有可同步的向量库", http_status=400) return jsonify({"error": "没有可同步的向量库"}), 400
# 执行同步
results = []
# 使用 sync_service 执行同步
sync_service = current_app.config.get('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
)
from core.task_registry import get_registry if sync_service:
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: try:
registry.update_progress(task.id, stage='扫描文档', message='正在检测变更...') sync_result = sync_service.sync_now()
sync_result = sync_svc.sync_now() results.append({
return { "collection": "all",
'collection': target_collection or 'all', "status": "success",
'status': 'success', "message": f"同步完成: 处理 {sync_result.documents_processed} 个文档",
'message': f"同步完成: 处理 {sync_result.documents_processed} 个文档", "details": {
'details': { "added": sync_result.documents_added,
'added': sync_result.documents_added, "modified": sync_result.documents_modified,
'modified': sync_result.documents_modified, "deleted": sync_result.documents_deleted,
'deleted': sync_result.documents_deleted, "errors": sync_result.errors
'errors': sync_result.errors,
} }
} })
finally: except Exception as e:
sync_svc.on_change_callback = old_callback 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 端点"
})
registry.start_task(task.id, _do_sync, sync_service) return jsonify({
"success": True,
return success_response( "results": results,
data={"task_id": task.id, "message": f"同步任务已启动,通过 GET /tasks/{task.id} 查询进度"}, "synced_count": len([r for r in results if r["status"] == "success"])
status_code=SYNC_SUCCESS, })
message="同步任务已启动"
)
@kb_bp.route('/debug/scan', methods=['GET']) @kb_bp.route('/debug/scan', methods=['GET'])
@@ -469,16 +473,22 @@ def debug_scan() -> Tuple[Any, int]:
@require_gateway_auth @require_gateway_auth
def reindex_collection(kb_name: str) -> Tuple[Any, int]: def reindex_collection(kb_name: str) -> Tuple[Any, int]:
""" """
强制重新向量化指定集合的所有文档(异步任务) 强制重新向量化指定集合的所有文档
清除该集合的文档哈希记录,触发完整重新索引。 清除该集合的文档哈希记录,触发完整重新索引。
立即返回 task_id后台线程执行重建 适用于文档内容更新后需要重建索引的场景
Args: Args:
kb_name: 向量库名称 kb_name: 向量库名称
Returns: Returns:
{"success": true, "data": {"task_id": "xxx", "message": "..."}} {
"success": true,
"message": "...",
"documents_processed": N,
"documents_added": N,
"errors": [...]
}
""" """
kb_manager, _, err = _require_multi_kb() kb_manager, _, err = _require_multi_kb()
if err: if err:
@@ -486,12 +496,14 @@ def reindex_collection(kb_name: str) -> Tuple[Any, int]:
from config import DOCUMENTS_PATH from config import DOCUMENTS_PATH
# 清除该集合的文档哈希记录(同步完成,很快) # 清除该集合的文档哈希记录
try: try:
from data.db import get_connection from data.db import get_connection
with get_connection("knowledge") as conn: with get_connection("knowledge") as conn:
cursor = conn.cursor() cursor = conn.cursor()
# 转义 LIKE 通配符,防止 kb_name 中的 % 或 _ 导致非预期匹配
escaped_kb = kb_name.replace('%', '\\%').replace('_', '\\_') 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 '\\'", cursor.execute("DELETE FROM document_hashes WHERE document_id LIKE ? ESCAPE '\\' OR document_id LIKE ? ESCAPE '\\'",
(f"{escaped_kb}/%", f"{escaped_kb}\\%")) (f"{escaped_kb}/%", f"{escaped_kb}\\%"))
deleted = cursor.rowcount deleted = cursor.rowcount
@@ -499,54 +511,23 @@ def reindex_collection(kb_name: str) -> Tuple[Any, int]:
except Exception as e: except Exception as e:
logger.warning(f"清除哈希记录失败: {e}") logger.warning(f"清除哈希记录失败: {e}")
# 触发同步
sync_service = current_app.config.get('SYNC_SERVICE') sync_service = current_app.config.get('SYNC_SERVICE')
if not sync_service: if 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: try:
registry.update_progress(task.id, stage='扫描文档', message='正在检测变更...') result = sync_service.sync_now()
result = sync_svc.sync_now() return jsonify({
return { "success": True,
'message': f"重新索引完成: 处理 {result.documents_processed} 个文档", "message": f"重新索引完成: 处理 {result.documents_processed} 个文档",
'documents_processed': result.documents_processed, "documents_processed": result.documents_processed,
'documents_added': result.documents_added, "documents_added": result.documents_added,
'errors': result.errors, "errors": result.errors
} })
finally: except Exception as e:
sync_svc.on_change_callback = old_callback logger.error(f"reindex 异常: {e}")
return jsonify({"error": "操作失败,请稍后重试"}), 500
registry.start_task(task.id, _do_reindex, sync_service, kb_name) else:
return jsonify({"error": "同步服务不可用"}), 503
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']) @kb_bp.route('/kb/route', methods=['POST'])
@@ -586,7 +567,7 @@ def test_routing() -> Tuple[Any, int]:
query = data.get('query', '') query = data.get('query', '')
if not query: if not query:
return error_response("MISSING_PARAMS", BAD_REQUEST, "请提供查询内容", http_status=400) return jsonify({"error": "请提供查询内容"}), 400
# 获取路由结果 # 获取路由结果
target_kbs = route_query( target_kbs = route_query(
@@ -598,7 +579,7 @@ def test_routing() -> Tuple[Any, int]:
# 获取意图分析 # 获取意图分析
intent = kb_router.analyze_intent(query) intent = kb_router.analyze_intent(query)
return success_response(data={ return jsonify({
"query": query, "query": query,
"user_role": user.get("role"), "user_role": user.get("role"),
"user_department": user.get("department", ""), "user_department": user.get("department", ""),
@@ -655,10 +636,10 @@ def deprecate_document(kb_name: str, filename: str) -> Tuple[Any, int]:
reason, reason,
deprecated_by=user.get('user_id', 'unknown') deprecated_by=user.get('user_id', 'unknown')
) )
return success_response(data=result) return jsonify(result)
except Exception as e: except Exception as e:
logger.error(f"操作异常: {e}") logger.error(f"操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"success": False, "error": "操作失败,请稍后重试"}), 500
@kb_bp.route('/collections/<kb_name>/documents/<path:filename>/restore', methods=['POST']) @kb_bp.route('/collections/<kb_name>/documents/<path:filename>/restore', methods=['POST'])
@@ -687,10 +668,10 @@ def restore_document(kb_name: str, filename: str) -> Tuple[Any, int]:
try: try:
result = kb_manager.restore_document(kb_name, filename) result = kb_manager.restore_document(kb_name, filename)
return success_response(data=result) return jsonify(result)
except Exception as e: except Exception as e:
logger.error(f"操作异常: {e}") logger.error(f"操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"success": False, "error": "操作失败,请稍后重试"}), 500
@kb_bp.route('/collections/<kb_name>/documents/<path:filename>/versions', methods=['GET']) @kb_bp.route('/collections/<kb_name>/documents/<path:filename>/versions', methods=['GET'])
@@ -738,7 +719,8 @@ def get_document_versions(kb_name: str, filename: str) -> Tuple[Any, int]:
versions = version_query.get_document_history(kb_name, filename, limit) versions = version_query.get_document_history(kb_name, filename, limit)
versions_data = [v.to_dict() for v in versions] versions_data = [v.to_dict() for v in versions]
return success_response(data={ return jsonify({
"success": True,
"document_id": filename, "document_id": filename,
"collection": kb_name, "collection": kb_name,
"versions": versions_data, "versions": versions_data,
@@ -746,7 +728,7 @@ def get_document_versions(kb_name: str, filename: str) -> Tuple[Any, int]:
}) })
except Exception as e: except Exception as e:
logger.error(f"操作异常: {e}") logger.error(f"操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"success": False, "error": "操作失败,请稍后重试"}), 500
@kb_bp.route('/collections/<kb_name>/update-image-descriptions', methods=['POST']) @kb_bp.route('/collections/<kb_name>/update-image-descriptions', methods=['POST'])
@@ -778,10 +760,10 @@ def update_image_descriptions(kb_name: str) -> Tuple[Any, int]:
try: try:
result = kb_manager.update_image_descriptions(kb_name) result = kb_manager.update_image_descriptions(kb_name)
return success_response(data=result) return jsonify(result)
except Exception as e: except Exception as e:
logger.error(f"操作异常: {e}") logger.error(f"操作异常: {e}")
return error_response("OPERATION_FAILED", INTERNAL_ERROR, "操作失败,请稍后重试", http_status=500) return jsonify({"success": False, "error": "操作失败,请稍后重试"}), 500
@kb_bp.route('/collections/sync-vlm-cache', methods=['POST']) @kb_bp.route('/collections/sync-vlm-cache', methods=['POST'])
@@ -822,12 +804,12 @@ def sync_vlm_cache() -> Tuple[Any, int]:
images_dir = Path(".data/images") images_dir = Path(".data/images")
if not vlm_cache_dir.exists(): if not vlm_cache_dir.exists():
return error_response("VLM_CACHE_NOT_FOUND", BAD_REQUEST, "VLM 缓存目录不存在", http_status=400) return jsonify({"success": False, "error": "VLM 缓存目录不存在"}), 400
# 获取所有 VLM 缓存文件 # 获取所有 VLM 缓存文件
cache_files = list(vlm_cache_dir.glob("*.txt")) cache_files = list(vlm_cache_dir.glob("*.txt"))
if not cache_files: if not cache_files:
return success_response(data={"total_cache_files": 0, "synced_count": 0}, message="无 VLM 缓存文件") return jsonify({"success": True, "total_cache_files": 0, "synced_count": 0, "message": "无 VLM 缓存文件"})
# 构建图片 MD5 → 文件名 的映射 # 构建图片 MD5 → 文件名 的映射
image_hash_map = {} image_hash_map = {}
@@ -897,7 +879,8 @@ def sync_vlm_cache() -> Tuple[Any, int]:
skipped_count += 1 skipped_count += 1
details.append({"cache": cache_file.name, "image": image_filename, "status": "skipped", "reason": "向量库中未找到对应切片"}) details.append({"cache": cache_file.name, "image": image_filename, "status": "skipped", "reason": "向量库中未找到对应切片"})
return success_response(data={ return jsonify({
"success": True,
"total_cache_files": len(cache_files), "total_cache_files": len(cache_files),
"synced_count": synced_count, "synced_count": synced_count,
"skipped_count": skipped_count, "skipped_count": skipped_count,

View File

@@ -10,8 +10,6 @@
from flask import Blueprint, request, jsonify, current_app from flask import Blueprint, request, jsonify, current_app
from auth.gateway import require_gateway_auth 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__) session_bp = Blueprint('session', __name__)
@@ -36,7 +34,7 @@ def get_sessions():
""" """
session_manager = current_app.config['SESSION_MANAGER'] session_manager = current_app.config['SESSION_MANAGER']
if session_manager is None: if session_manager is None:
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "会话服务不可用", http_status=503) return jsonify({"error": "会话服务不可用"}), 503
user_id = request.current_user["user_id"] user_id = request.current_user["user_id"]
sessions = session_manager.get_user_sessions(user_id, limit=20) sessions = session_manager.get_user_sessions(user_id, limit=20)
@@ -49,7 +47,7 @@ def get_sessions():
else: else:
s["preview"] = "空会话" s["preview"] = "空会话"
return success_response(data={"sessions": sessions}) return jsonify({"sessions": sessions})
@session_bp.route('/history/<session_id>', methods=['GET']) @session_bp.route('/history/<session_id>', methods=['GET'])
@@ -67,19 +65,19 @@ def get_history(session_id):
""" """
session_manager = current_app.config['SESSION_MANAGER'] session_manager = current_app.config['SESSION_MANAGER']
if session_manager is None: if session_manager is None:
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "会话服务不可用", http_status=503) return jsonify({"error": "会话服务不可用"}), 503
user_id = request.current_user["user_id"] user_id = request.current_user["user_id"]
# 验证会话归属 # 验证会话归属
sessions = session_manager.get_user_sessions(user_id) sessions = session_manager.get_user_sessions(user_id, limit=20)
session_ids = [s["session_id"] for s in sessions] session_ids = [s["session_id"] for s in sessions]
if session_id not in session_ids: if session_id not in session_ids:
return error_response("FORBIDDEN", FORBIDDEN, "无权访问此会话", http_status=403) return jsonify({"error": "无权访问此会话"}), 403
history = session_manager.get_history(session_id, limit=100) history = session_manager.get_history(session_id, limit=100)
return success_response(data={"history": history}) return jsonify({"history": history})
@session_bp.route('/session/<session_id>', methods=['DELETE']) @session_bp.route('/session/<session_id>', methods=['DELETE'])
@@ -88,19 +86,19 @@ def delete_session(session_id):
"""删除会话""" """删除会话"""
session_manager = current_app.config['SESSION_MANAGER'] session_manager = current_app.config['SESSION_MANAGER']
if session_manager is None: if session_manager is None:
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "会话服务不可用", http_status=503) return jsonify({"error": "会话服务不可用"}), 503
user_id = request.current_user["user_id"] user_id = request.current_user["user_id"]
# 验证会话归属 # 验证会话归属
sessions = session_manager.get_user_sessions(user_id) sessions = session_manager.get_user_sessions(user_id, limit=20)
session_ids = [s["session_id"] for s in sessions] session_ids = [s["session_id"] for s in sessions]
if session_id not in session_ids: if session_id not in session_ids:
return error_response("FORBIDDEN", FORBIDDEN, "无权访问此会话", http_status=403) return jsonify({"error": "无权删除此会话"}), 403
session_manager.delete_session(session_id) session_manager.delete_session(session_id)
return success_response(status_code=DELETE_SUCCESS, message="会话已删除") return jsonify({"success": True, "message": "会话已删除"})
@session_bp.route('/clear/<session_id>', methods=['POST']) @session_bp.route('/clear/<session_id>', methods=['POST'])
@@ -109,16 +107,16 @@ def clear_history(session_id):
"""清空会话历史(保留会话)""" """清空会话历史(保留会话)"""
session_manager = current_app.config['SESSION_MANAGER'] session_manager = current_app.config['SESSION_MANAGER']
if session_manager is None: if session_manager is None:
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "会话服务不可用", http_status=503) return jsonify({"error": "会话服务不可用"}), 503
user_id = request.current_user["user_id"] user_id = request.current_user["user_id"]
# 验证会话归属 # 验证会话归属
sessions = session_manager.get_user_sessions(user_id) sessions = session_manager.get_user_sessions(user_id, limit=20)
session_ids = [s["session_id"] for s in sessions] session_ids = [s["session_id"] for s in sessions]
if session_id not in session_ids: if session_id not in session_ids:
return error_response("FORBIDDEN", FORBIDDEN, "无权访问此会话", http_status=403) return jsonify({"error": "无权操作此会话"}), 403
session_manager.clear_history(session_id) session_manager.clear_history(session_id)
return success_response(message="历史已清空") return jsonify({"success": True, "message": "历史已清空"})

View File

@@ -31,12 +31,12 @@ Example:
""" """
from typing import Optional, Tuple, Any from typing import Optional, Tuple, Any
from flask import Blueprint, request, current_app from flask import Blueprint, request, jsonify, current_app
import logging import logging
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
from auth.gateway import require_gateway_auth from auth.gateway import require_gateway_auth
from core.status_codes import SUCCESS, SYNC_SUCCESS, SYNC_ERROR, INTERNAL_ERROR, SERVICE_UNAVAILABLE from core.status_codes import SYNC_SUCCESS, SYNC_ERROR, INTERNAL_ERROR
from api.response_utils import success_response, error_response from api.response_utils import success_response, error_response
sync_bp = Blueprint('sync', __name__) sync_bp = Blueprint('sync', __name__)
@@ -68,7 +68,9 @@ def _require_sync_service() -> Tuple[Optional[Any], Optional[Tuple]]:
service = _get_sync_service() service = _get_sync_service()
if not service: if not service:
return None, error_response( return None, error_response(
"SERVICE_UNAVAILABLE", INTERNAL_ERROR, "同步服务未启用", error="SERVICE_UNAVAILABLE",
error_code=INTERNAL_ERROR,
message="同步服务未启用",
http_status=503 http_status=503
) )
return service, None return service, None
@@ -80,10 +82,9 @@ def _require_sync_service() -> Tuple[Optional[Any], Optional[Tuple]]:
@require_gateway_auth @require_gateway_auth
def trigger_sync() -> Tuple[Any, int]: def trigger_sync() -> Tuple[Any, int]:
""" """
手动触发知识库同步(异步任务) 手动触发知识库同步
立即返回 task_id后台线程执行同步 扫描文档目录,检测变更并执行向量化处理
客户端通过 GET /tasks/<task_id> 轮询进度。
请求体 (可选): 请求体 (可选):
{ {
@@ -92,7 +93,7 @@ def trigger_sync() -> Tuple[Any, int]:
} }
Returns: Returns:
成功: {"success": true, "data": {"task_id": "xxx", "message": "同步任务已启动"}} 成功: {"success": true, "data": {"result": {...}}}
失败: {"error": "...", "error_code": "..."} 失败: {"error": "...", "error_code": "..."}
Example: Example:
@@ -103,63 +104,20 @@ def trigger_sync() -> Tuple[Any, int]:
if err: if err:
return err return err
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", 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: try:
registry.update_progress(task.id, stage='扫描文档', message='正在检测变更...') result = service.sync_now()
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( return success_response(
data={ data={"result": result.to_dict() if hasattr(result, 'to_dict') else result},
'task_id': task.id,
'message': '同步任务已启动,通过 GET /tasks/' + task.id + ' 查询进度'
},
status_code=SYNC_SUCCESS, status_code=SYNC_SUCCESS,
message="同步任务已启动" message="同步完成"
)
except Exception as e:
logger.error(f"同步操作异常: {e}")
return error_response(
error="SYNC_ERROR",
error_code=SYNC_ERROR,
message="操作失败",
http_status=500
) )
@@ -185,8 +143,12 @@ def get_sync_status() -> Tuple[Any, int]:
""" """
service, err = _require_sync_service() service, err = _require_sync_service()
if err: if err:
return error_response("SERVICE_UNAVAILABLE", SERVICE_UNAVAILABLE, "同步服务未启用", http_status=503, return jsonify({
enabled=False) "status": "failed",
"status_code": INTERNAL_ERROR,
"enabled": False,
"message": "同步服务未启用"
})
try: try:
# 获取状态信息 # 获取状态信息
@@ -201,11 +163,15 @@ def get_sync_status() -> Tuple[Any, int]:
if hasattr(service, 'get_status'): if hasattr(service, 'get_status'):
status.update(service.get_status()) status.update(service.get_status())
return success_response(data=status) return jsonify(status)
except Exception as e: except Exception as e:
logger.error(f"状态查询异常: {e}") logger.error(f"状态查询异常: {e}")
return error_response("INTERNAL_ERROR", INTERNAL_ERROR, "操作失败", http_status=500, return jsonify({
enabled=True) "status": "failed",
"status_code": INTERNAL_ERROR,
"enabled": True,
"error": "操作失败"
})
@sync_bp.route('/sync/history', methods=['GET']) @sync_bp.route('/sync/history', methods=['GET'])
@@ -230,11 +196,13 @@ def get_sync_history() -> Tuple[Any, int]:
try: try:
history = service.get_sync_history(limit=limit) if hasattr(service, 'get_sync_history') else [] history = service.get_sync_history(limit=limit) if hasattr(service, 'get_sync_history') else []
return success_response(data={"history": history}) return jsonify({"history": history})
except Exception as e: except Exception as e:
logger.error(f"同步操作异常: {e}") logger.error(f"同步操作异常: {e}")
return error_response( return error_response(
"SYNC_ERROR", SYNC_ERROR, "操作失败", error="SYNC_ERROR",
error_code=SYNC_ERROR,
message="操作失败",
http_status=500 http_status=500
) )
@@ -263,11 +231,13 @@ def get_change_logs() -> Tuple[Any, int]:
try: try:
changes = service.get_change_logs(limit=limit, collection=collection) if hasattr(service, 'get_change_logs') else [] changes = service.get_change_logs(limit=limit, collection=collection) if hasattr(service, 'get_change_logs') else []
return success_response(data={"changes": changes}) return jsonify({"changes": changes})
except Exception as e: except Exception as e:
logger.error(f"同步操作异常: {e}") logger.error(f"同步操作异常: {e}")
return error_response( return error_response(
"SYNC_ERROR", SYNC_ERROR, "操作失败", error="SYNC_ERROR",
error_code=SYNC_ERROR,
message="操作失败",
http_status=500 http_status=500
) )
@@ -293,23 +263,27 @@ def start_sync_monitor() -> Tuple[Any, int]:
try: try:
if hasattr(service, 'is_running') and service.is_running(): if hasattr(service, 'is_running') and service.is_running():
return success_response(status_code=SYNC_SUCCESS, message="文件监控已在运行") return jsonify({"status": "success", "status_code": SYNC_SUCCESS, "message": "文件监控已在运行"})
if hasattr(service, 'start'): if hasattr(service, 'start'):
success = service.start() success = service.start()
if success: if success:
return success_response(status_code=SYNC_SUCCESS, message="文件监控已启动") return jsonify({"status": "success", "status_code": SYNC_SUCCESS, "message": "文件监控已启动"})
else: else:
return error_response( return error_response(
"SYNC_ERROR", SYNC_ERROR, "启动文件监控失败", error="SYNC_ERROR",
error_code=SYNC_ERROR,
message="启动文件监控失败",
http_status=500 http_status=500
) )
else: else:
return success_response(message="文件监控功能不可用") return jsonify({"status": "success", "message": "文件监控功能不可用"})
except Exception as e: except Exception as e:
logger.error(f"同步操作异常: {e}") logger.error(f"同步操作异常: {e}")
return error_response( return error_response(
"SYNC_ERROR", SYNC_ERROR, "操作失败", error="SYNC_ERROR",
error_code=SYNC_ERROR,
message="操作失败",
http_status=500 http_status=500
) )
@@ -332,10 +306,12 @@ def stop_sync_monitor() -> Tuple[Any, int]:
try: try:
if hasattr(service, 'stop'): if hasattr(service, 'stop'):
service.stop() service.stop()
return success_response(status_code=SYNC_SUCCESS, message="文件监控已停止") return jsonify({"status": "success", "status_code": SYNC_SUCCESS, "message": "文件监控已停止"})
except Exception as e: except Exception as e:
logger.error(f"同步操作异常: {e}") logger.error(f"同步操作异常: {e}")
return error_response( return error_response(
"SYNC_ERROR", SYNC_ERROR, "操作失败", error="SYNC_ERROR",
error_code=SYNC_ERROR,
message="操作失败",
http_status=500 http_status=500
) )

View File

@@ -1,191 +0,0 @@
"""
异步任务查询 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'
}
)

View File

@@ -20,12 +20,12 @@
## 模式说明 ## 模式说明
开发模式 (DEV_MODE=true,默认): 开发模式 (DEV_MODE=true):
- 支持 mock token 模拟用户Authorization: Bearer mock-token-admin - 支持 mock token 模拟用户Authorization: Bearer mock-token-admin
- 无 Header 时自动使用开发测试用户 - 无 Header 时自动使用开发测试用户
- 适用于前端测试和开发调试 - 适用于前端测试和开发调试
生产模式 (DEV_MODE=false): 生产模式 (DEV_MODE=false,默认):
- 不需要 Header直接放行 - 不需要 Header直接放行
- 权限由后端完全控制,通过 collections 参数传入 - 权限由后端完全控制,通过 collections 参数传入
- RAG 服务完全无状态,只负责问答检索 - RAG 服务完全无状态,只负责问答检索
@@ -34,17 +34,11 @@
from functools import wraps from functools import wraps
from flask import request, jsonify from flask import request, jsonify
from typing import Dict, Optional from typing import Dict, Optional
import os from config import DEV_MODE
from pathlib import Path
from dotenv import load_dotenv
# 加载 .env 文件(从项目根目录)
env_path = Path(__file__).parent.parent / '.env'
load_dotenv(env_path)
# ==================== 模拟用户数据(开发环境)==================== # ==================== 模拟用户数据(开发环境)====================
# 用于前端模拟登录测试,仅 DEV_MODE=true 时生效 # 用于前端模拟登录测试,仅 DEV_MODE=true 时生效(需在 .env 中显式开启)
MOCK_USERS = { MOCK_USERS = {
'admin': { 'admin': {
'user_id': 'admin001', 'user_id': 'admin001',
@@ -83,19 +77,19 @@ def require_gateway_auth(f):
""" """
网关认证装饰器 - 从 Header 读取用户信息 网关认证装饰器 - 从 Header 读取用户信息
开发模式 (DEV_MODE=true,默认): 开发模式 (DEV_MODE=true):
- 支持 mock token: Authorization: Bearer mock-token-admin - 支持 mock token: Authorization: Bearer mock-token-admin
- 无 Header 时自动使用开发测试用户admin 角色) - 无 Header 时自动使用开发测试用户admin 角色)
生产模式 (DEV_MODE=false): 生产模式 (DEV_MODE=false,默认):
- 不需要 Header直接放行 - 不需要 Header直接放行
- 用户信息设为默认值 - 用户信息设为默认值
- 权限由后端通过 collections 参数控制 - 权限由后端通过 collections 参数控制
""" """
@wraps(f) @wraps(f)
def decorated(*args, **kwargs): def decorated(*args, **kwargs):
# 开发模式开关(默认开启,生产环境设置 DEV_MODE=false # 开发模式开关(统一由 config.py 管理
dev_mode = os.environ.get('DEV_MODE', 'true').lower() != 'false' dev_mode = DEV_MODE
# 开发模式:支持 mock token # 开发模式:支持 mock token
if dev_mode: if dev_mode:
@@ -202,11 +196,24 @@ def can_delete_collection(role: str) -> bool:
def require_role(*roles): def require_role(*roles):
""" """
兼容旧代码 - 权限由后端管理,此装饰器不再执行权限验证 角色验证装饰器(开发和生产环境均生效)
需搭配 @require_gateway_auth 使用(先设置 current_user再验证角色
开发模式: 检查 mock token 对应用户的角色
生产模式: 检查网关注入的 X-User-Role Header
""" """
def decorator(f): def decorator(f):
@wraps(f) @wraps(f)
def decorated(*args, **kwargs): 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 f(*args, **kwargs)
return decorated return decorated
return decorator return decorator
@@ -214,11 +221,12 @@ def require_role(*roles):
def require_collection_permission(operation: str): def require_collection_permission(operation: str):
""" """
兼容旧代码 - 权限由后端管理,此装饰器不再执行权限验证 集合权限验证装饰器(开发模式下为占位实现,生产环境权限由网关控制)
""" """
def decorator(f): def decorator(f):
@wraps(f) @wraps(f)
def decorated(*args, **kwargs): def decorated(*args, **kwargs):
# 生产环境下权限由网关/后端统一管控,此处放行
return f(*args, **kwargs) return f(*args, **kwargs)
return decorated return decorated
return decorator return decorator

View File

@@ -201,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_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_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_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 CHUNK_SIZE = 1000

View File

@@ -29,6 +29,7 @@ RAG 核心引擎
import os import os
import gc import gc
import re
import time import time
import logging import logging
import threading import threading
@@ -112,7 +113,7 @@ except ImportError:
RERANK_DEVICE = "auto" RERANK_DEVICE = "auto"
RERANK_USE_ONNX = False RERANK_USE_ONNX = False
RERANK_BACKEND = "local" RERANK_BACKEND = "local"
RERANK_CLOUD_MODEL = "qwen3-rerank" RERANK_CLOUD_MODEL = "xop3qwen8breranker"
RERANK_CLOUD_API_KEY = "" RERANK_CLOUD_API_KEY = ""
RERANK_CLOUD_BASE_URL = "https://dashscope.aliyuncs.com/compatible-api/v1/reranks" RERANK_CLOUD_BASE_URL = "https://dashscope.aliyuncs.com/compatible-api/v1/reranks"
RERANK_CLOUD_TIMEOUT = 15 RERANK_CLOUD_TIMEOUT = 15
@@ -272,8 +273,7 @@ class CloudReranker:
body = { body = {
"model": self.model, "model": self.model,
"query": query, "query": query,
"documents": documents, "documents": documents
"top_n": len(documents)
} }
session = self._get_session() session = self._get_session()
@@ -1231,6 +1231,7 @@ class RAGEngine:
logger.warning(f"扩展连续切片失败: {e}") logger.warning(f"扩展连续切片失败: {e}")
return {'ids': [], 'documents': [], 'metadatas': []} return {'ids': [], 'documents': [], 'metadatas': []}
# 扩展同 section 的 text 邻居
where_filter = {"$and": [{"source": source}, {"chunk_type": "text"}]} where_filter = {"$and": [{"source": source}, {"chunk_type": "text"}]}
if section: if section:
where_filter["$and"].append({"section": section}) where_filter["$and"].append({"section": section})
@@ -1241,6 +1242,20 @@ class RAGEngine:
if not neighbors.get('ids') or len(neighbors.get('ids', [])) <= 1: if not neighbors.get('ids') or len(neighbors.get('ids', [])) <= 1:
neighbors = _get_neighbors({"$and": [{"source": source}, {"chunk_type": "text"}]}) 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 = [] neighbor_rows = []
for n_id, n_doc, n_meta in zip( for n_id, n_doc, n_meta in zip(
neighbors.get('ids', []), neighbors.get('ids', []),
@@ -1253,6 +1268,18 @@ class RAGEngine:
if seed_index - CONTEXT_EXPANSION_BEFORE <= n_index <= seed_index + CONTEXT_EXPANSION_AFTER: if seed_index - CONTEXT_EXPANSION_BEFORE <= n_index <= seed_index + CONTEXT_EXPANSION_AFTER:
neighbor_rows.append((n_index, n_id, n_doc, n_meta)) 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 seed_neighbors_added = 0
for n_index, n_id, n_doc, n_meta in sorted(neighbor_rows, key=lambda row: row[0]): for n_index, n_id, n_doc, n_meta in sorted(neighbor_rows, key=lambda row: row[0]):
if len(items) >= max_chunks: if len(items) >= max_chunks:
@@ -2041,21 +2068,33 @@ class RAGEngine:
"content": ( "content": (
"你是一个严谨的知识库问答助手。" "你是一个严谨的知识库问答助手。"
"你必须且只能根据用户提供的【参考资料】回答问题。" "你必须且只能根据用户提供的【参考资料】回答问题。"
"参考资料中每段内容前标有章节路径(━格式),请注意区分不同章节的内容,"
"特别当不同章节标题相似或包含相同关键词时,务必根据章节路径准确定位,不要混淆。"
"如果参考资料中有答案,必须引用对应内容回答,并在回答末尾标注引用编号(如[1]、[2])。" "如果参考资料中有答案,必须引用对应内容回答,并在回答末尾标注引用编号(如[1]、[2])。"
"如果参考资料中确实没有相关信息,简短说明即可,不要编造或补充资料外的内容。" "如果参考资料中确实没有相关信息,简短说明即可,不要编造或补充资料外的内容。"
"禁止使用参考资料以外的知识进行补充或推测。" "禁止使用参考资料以外的知识进行补充或推测。"
"【重要-表格处理规则】当用户询问表格、要求展示表格内容时,你必须将参考资料中的 Markdown 表格原样输出(保留 | 分隔符和表格结构),"
"不要仅用文字描述表格存在或仅列出章节名称。如果参考资料中多个章节都有表格,"
"优先展示与用户问题最相关的表格完整内容。"
) )
}) })
# 添加当前问题(带上下文)- 强化指令 # 添加当前问题(带上下文)- 强化指令
if context: 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"""【参考资料】 user_message = f"""【参考资料】
{context} {context}
【用户问题】 【用户问题】
{query} {query}
请仔细阅读以上全部参考资料后回答。如果参考资料中包含相关内容,必须引用回答并标注编号。如果资料中没有相关信息,请明确说明。""" 请仔细阅读以上全部参考资料后回答。注意参考资料中标有章节路径,请根据章节路径准确定位相关内容。如果参考资料中包含相关内容,必须引用回答并标注编号。如果资料中没有相关信息,请明确说明。{_table_hint}"""
else: else:
user_message = query user_message = query

View File

@@ -83,9 +83,16 @@ class IntentAnalyzer:
根据对话历史和当前用户消息,输出一个 JSON 对象,包含以下字段: 根据对话历史和当前用户消息,输出一个 JSON 对象,包含以下字段:
1. **rewritten_query**: 改写后的完整问题 1. **rewritten_query**: 改写后的完整问题
- 如果问题包含指代(如"这两张图片""继续说"),将其改写为完整、独立的问题 - **指代消解**如果问题包含指代(如"这两张图片""继续说"),将其改写为完整、独立的问题
- 例如:"分析一下这两张图片""分析一下对话历史中提到的图片" - 例如:"分析一下这两张图片""分析一下对话历史中提到的图片"
- 如果问题本身已经完整,直接返回原文 - **追问补全**:如果问题是省略式追问(省略了上一轮讨论的主题实体),必须补全为完整问题
- 判断方法:当前问题缺少主语/宾语,且对话历史中可以推断出省略的实体
- 补全方法:从上一轮用户问题中提取主题实体,与追问组合成完整问题
- 例如:
- 上一轮问"吸烟点C1类是什么区",追问"有完整表格吗?""吸烟点C1类有完整表格吗"
- 上一轮问"三峡工程的投资情况",追问"建设地点在哪?""三峡工程的建设地点在哪?"
- 上一轮问"货源投放有哪些原则?",追问"具体内容是什么?""货源投放原则的具体内容是什么?"
- 如果问题本身已经完整且独立,直接返回原文
2. **use_context**: 布尔值 2. **use_context**: 布尔值
- true: 问题依赖历史对话中的信息,答案已经在历史回答中 - true: 问题依赖历史对话中的信息,答案已经在历史回答中
@@ -105,7 +112,9 @@ class IntentAnalyzer:
- 推理类intent="reasoning"生成最多2个子查询 - 推理类intent="reasoning"生成最多2个子查询
* 原问题的检索查询 * 原问题的检索查询
* 一个补充角度的检索查询(如原因、背景、影响等),帮助获取更全面的上下文 * 一个补充角度的检索查询(如原因、背景、影响等),帮助获取更全面的上下文
- 其他类factual/instruction/other严格只生成1个子查询(原问题) - 其他类factual/instruction/other严格只生成1个子查询
* 子查询应基于 rewritten_query改写后的完整问题而非用户原始输入
* 例如:追问"有完整表格吗?"改写为"吸烟点C1类有完整表格吗"后,子查询应为"吸烟点C1类的完整表格内容"
- 不要为同一实体生成语义重叠的查询 - 不要为同一实体生成语义重叠的查询
- 子查询应保持原问题的关键词长度20-60字符为宜 - 子查询应保持原问题的关键词长度20-60字符为宜
@@ -307,7 +316,8 @@ class IntentAnalyzer:
if cache_emb is not None: if cache_emb is not None:
cached = cache.get(cache_emb) cached = cache.get(cache_emb)
if cached: # 确保缓存条目是意图分析结果(非 RAG 回答缓存)
if cached and cached.get("cache_type") != "rag_answer":
logger.info(f"意图分析缓存命中: {cached.get('reason', '')[:50]}") logger.info(f"意图分析缓存命中: {cached.get('reason', '')[:50]}")
return IntentAnalysis.from_dict(cached) return IntentAnalysis.from_dict(cached)
else: else:
@@ -377,9 +387,11 @@ class IntentAnalyzer:
intent=intent_type intent=intent_type
) )
# 存入语义缓存 # 存入语义缓存(标记类型,避免与 RAG 回答缓存混淆)
if cache and cache_emb is not None: 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: if len(self._exact_cache) < self._exact_cache_max:
@@ -437,6 +449,7 @@ class IntentAnalyzer:
if not history: if not history:
return "(无历史对话)" return "(无历史对话)"
import re
parts = [] parts = []
# 提取最近 3 轮对话 # 提取最近 3 轮对话
@@ -445,6 +458,7 @@ class IntentAnalyzer:
for msg in recent_history: for msg in recent_history:
role = "用户" if msg.get("role") == "user" else "助手" role = "用户" if msg.get("role") == "user" else "助手"
content = msg.get("content", "") content = msg.get("content", "")
original_content = content # 保留原始内容用于结构化提取
# 截断过长的内容 # 截断过长的内容
if len(content) > 500: if len(content) > 500:
@@ -452,9 +466,10 @@ class IntentAnalyzer:
parts.append(f"{role}{content}") parts.append(f"{role}{content}")
# 提取图片信息 # 提取结构化信息(从 assistant 消息中提取章节、表格、来源等)
metadata = msg.get("metadata", {}) metadata = msg.get("metadata", {})
if isinstance(metadata, dict): if isinstance(metadata, dict):
# 已有的图片提取
images = metadata.get("images", []) images = metadata.get("images", [])
if images: if images:
for img in images[:3]: for img in images[:3]:
@@ -463,6 +478,43 @@ class IntentAnalyzer:
img_type = img.get("type", "图片") img_type = img.get("type", "图片")
parts.append(f" └─ {img_type}: {desc}") 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: if context_images:
parts.append("\n【上下文中的图片】") parts.append("\n【上下文中的图片】")

View File

@@ -352,7 +352,7 @@ def quick_yes_no(
if keywords is None: if keywords is None:
keywords = ["", "需要", "yes", "true"] 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: if result is None:
return False return False

View File

@@ -49,8 +49,6 @@ _STATUS_MESSAGES: Dict[int, str] = {
4011: "向量库不存在", 4011: "向量库不存在",
4012: "文件内容为空", 4012: "文件内容为空",
4013: "权限不足", 4013: "权限不足",
4014: "任务不存在",
4015: "任务冲突",
# 服务端错误 (50xx) # 服务端错误 (50xx)
5000: "服务器内部错误", 5000: "服务器内部错误",
@@ -61,7 +59,6 @@ _STATUS_MESSAGES: Dict[int, str] = {
5020: "出题失败", 5020: "出题失败",
5021: "批阅失败", 5021: "批阅失败",
5030: "图片描述更新失败", 5030: "图片描述更新失败",
5040: "重建索引失败",
} }
@@ -116,8 +113,6 @@ FILE_NOT_FOUND = 4010
COLLECTION_NOT_FOUND = 4011 COLLECTION_NOT_FOUND = 4011
NO_CONTENT = 4012 NO_CONTENT = 4012
PERMISSION_DENIED = 4013 PERMISSION_DENIED = 4013
TASK_NOT_FOUND = 4014
TASK_CONFLICT = 4015
# 服务端错误 (50xx) # 服务端错误 (50xx)
INTERNAL_ERROR = 5000 INTERNAL_ERROR = 5000
@@ -128,4 +123,3 @@ SYNC_ERROR = 5010
EXAM_ERROR = 5020 EXAM_ERROR = 5020
GRADE_ERROR = 5021 GRADE_ERROR = 5021
IMAGE_DESC_ERROR = 5030 IMAGE_DESC_ERROR = 5030
REINDEX_ERROR = 5040

View File

@@ -1,311 +0,0 @@
"""
异步任务注册表
为长时间运行的操作(同步、重建索引、上传向量化、出题、批阅)提供统一的
任务状态跟踪机制。支持 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

View File

@@ -22,7 +22,6 @@
- [12. 图片服务](#12-图片服务) - [12. 图片服务](#12-图片服务)
- [13. 报告服务](#13-报告服务) - [13. 报告服务](#13-报告服务)
- [14. 知识库路由](#14-知识库路由) - [14. 知识库路由](#14-知识库路由)
- [15. 异步任务查询](#15-异步任务查询)
- [附录. 已知问题](#附录-已知问题) - [附录. 已知问题](#附录-已知问题)
--- ---
@@ -588,18 +587,15 @@ curl -s -X POST http://localhost:5001/documents/upload \
"size": 18, "size": 18,
"replaced": false "replaced": false
}, },
"sync_status": "已保存,向量化任务已启动", "sync_status": "已保存并添加到向量库"
"task_id": "a1b2c3d4e5f6"
}, },
"message": "文件上传成功,已保存,向量化任务已启动", "message": "文件上传成功,已保存并添加到向量库",
"status": "success", "status": "success",
"status_code": 2002, "status_code": 2002,
"success": true "success": true
} }
``` ```
> **异步说明**:文件保存为同步操作,向量化在后台线程异步执行。响应中的 `task_id` 可用于轮询向量化进度(`GET /tasks/<task_id>`)。
>
> **同名文件处理**:上传同名文件时,旧版本的切片会被自动清理后覆盖(`replaced: true`),不会生成时间戳后缀文件。 > **同名文件处理**:上传同名文件时,旧版本的切片会被自动清理后覆盖(`replaced: true`),不会生成时间戳后缀文件。
**验证结果**:✅ 通过 **验证结果**:✅ 通过
@@ -628,8 +624,7 @@ curl -s -X POST http://localhost:5001/documents/batch-upload \
{"filename": "file2.txt", "path": "public_kb/file2.txt", "status": "success", "replaced": false} {"filename": "file2.txt", "path": "public_kb/file2.txt", "status": "success", "replaced": false}
], ],
"success_count": 2, "success_count": 2,
"total": 2, "total": 2
"task_id": "b2c3d4e5f6a1"
}, },
"message": "批量上传完成,成功 2/2 个文件", "message": "批量上传完成,成功 2/2 个文件",
"status": "success", "status": "success",
@@ -638,8 +633,6 @@ curl -s -X POST http://localhost:5001/documents/batch-upload \
} }
``` ```
> **异步说明**:批量上传成功后自动触发后台向量化任务。`task_id` 可用于轮询向量化进度(`GET /tasks/<task_id>`)。若无成功上传的文件则不返回 `task_id`。
>
> **同名文件处理**:与单文件上传相同,批量上传中遇到同名文件也会自动覆盖旧版本(`replaced: true`)。 > **同名文件处理**:与单文件上传相同,批量上传中遇到同名文件也会自动覆盖旧版本(`replaced: true`)。
**验证结果**:✅ 通过 **验证结果**:✅ 通过
@@ -889,7 +882,7 @@ curl -s -X DELETE "http://localhost:5001/chunks/人员名册.txt_text_0?collecti
### POST /sync ### POST /sync
触发文档同步(异步任务) 触发文档同步。
```bash ```bash
curl -s -X POST http://localhost:5001/sync \ curl -s -X POST http://localhost:5001/sync \
@@ -900,23 +893,24 @@ curl -s -X POST http://localhost:5001/sync \
```json ```json
{ {
"data": { "data": {
"task_id": "c3d4e5f6a1b2", "result": {
"message": "同步任务已启动,通过 GET /tasks/c3d4e5f6a1b2 查询进度" "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"
}
}, },
"message": "同步任务已启动", "message": "同步完成",
"status": "success", "status": "success",
"status_code": 2010, "status_code": 2010,
"success": true "success": true
} }
``` ```
> **⚠️ 异步变更**:此接口已从同步改为异步。不再直接返回同步结果,而是返回 `task_id`。后端需通过 `GET /tasks/<task_id>` 轮询任务状态,直到 `status` 为 `completed` 或 `failed`。
>
> **冲突检测**:如果已有同步任务正在运行,返回 HTTP 409
> ```json
> {"error": "TASK_RUNNING", "message": "同步任务正在执行中 (task_id: xxx),请等待完成"}
> ```
**验证结果**:✅ 通过 **验证结果**:✅ 通过
--- ---
@@ -1462,21 +1456,6 @@ curl -s -X POST http://localhost:5001/exam/generate \
```json ```json
{ {
"data": { "data": {
"task_id": "d4e5f6a1b2c3",
"message": "出题任务已启动 (10题),通过 GET /tasks/d4e5f6a1b2c3 查询结果"
},
"message": "出题任务已启动",
"status": "success",
"status_code": 2020,
"success": true
}
```
> **⚠️ 异步变更**:此接口已从同步改为异步。响应仅返回 `task_id`,后端需通过 `GET /tasks/<task_id>` 轮询任务状态。任务完成后,`result` 字段包含完整出题结果(格式见下方说明)。
**轮询结果GET /tasks/\<task_id\> 完成后的 result 字段)**
```json
{
"questions": [ "questions": [
{ {
"content": { "content": {
@@ -1507,6 +1486,11 @@ curl -s -X POST http://localhost:5001/exam/generate \
"requested_types": {"single_choice": 5, "true_false": 3, "fill_blank": 2}, "requested_types": {"single_choice": 5, "true_false": 3, "fill_blank": 2},
"actual_types": {"single_choice": 5, "true_false": 3, "fill_blank": 2}, "actual_types": {"single_choice": 5, "true_false": 3, "fill_blank": 2},
"warnings": [] "warnings": []
},
"message": "出题成功",
"status": "success",
"status_code": 2020,
"success": true
} }
``` ```
@@ -1543,21 +1527,6 @@ curl -s -X POST http://localhost:5001/exam/generate-smart \
```json ```json
{ {
"data": { "data": {
"task_id": "e5f6a1b2c3d4",
"message": "AI 智能出题任务已启动,通过 GET /tasks/e5f6a1b2c3d4 查询结果"
},
"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": { "ai_analysis": {
"total_knowledge_points": 21, "total_knowledge_points": 21,
"suitable_types": ["single_choice", "multiple_choice", "true_false", "subjective"], "suitable_types": ["single_choice", "multiple_choice", "true_false", "subjective"],
@@ -1575,6 +1544,11 @@ curl -s -X POST http://localhost:5001/exam/generate-smart \
"source_chunks_used": 15, "source_chunks_used": 15,
"success": true, "success": true,
"total": 16 "total": 16
},
"message": "AI 智能出题成功",
"status": "success",
"status_code": 2020,
"success": true
} }
``` ```
@@ -1627,21 +1601,6 @@ curl -s -X POST http://localhost:5001/exam/grade \
```json ```json
{ {
"data": { "data": {
"task_id": "f6a1b2c3d4e5",
"message": "批阅任务已启动 (1题),通过 GET /tasks/f6a1b2c3d4e5 查询结果"
},
"message": "批阅任务已启动",
"status": "success",
"status_code": 2021,
"success": true
}
```
> **⚠️ 异步变更**:此接口已从同步改为异步。响应仅返回 `task_id`,后端需通过 `GET /tasks/<task_id>` 轮询任务状态。任务完成后,`result` 字段包含完整批阅结果(格式见下方说明)。
**轮询结果GET /tasks/\<task_id\> 完成后的 result 字段)**
```json
{
"request_id": null, "request_id": null,
"results": [ "results": [
{ {
@@ -1661,6 +1620,11 @@ curl -s -X POST http://localhost:5001/exam/grade \
"success": true, "success": true,
"total_max_score": 2, "total_max_score": 2,
"total_score": 2 "total_score": 2
},
"message": "批阅完成",
"status": "success",
"status_code": 2021,
"success": true
} }
``` ```
@@ -1859,190 +1823,6 @@ 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 参数 ### 1. /rag 接口 collections 参数
@@ -2171,11 +1951,11 @@ curl -s http://localhost:5001/tasks/stats
当代码升级涉及元数据字段变更时(如新增 `chunk_index``doc_type`),需要重构向量库: 当代码升级涉及元数据字段变更时(如新增 `chunk_index``doc_type`),需要重构向量库:
```bash ```bash
# 重构指定知识库(清除哈希记录 → 触发全量同步,异步任务 # 重构指定知识库(清除哈希记录 → 触发全量同步)
curl -s -X POST http://127.0.0.1:5001/collections/<kb_name>/reindex curl -s -X POST http://127.0.0.1:5001/collections/<kb_name>/reindex
``` ```
> **⚠️ 注意**reindex 已改为异步任务,立即返回 `task_id`。通过 `GET /tasks/<task_id>` 轮询进度。不再阻塞 gunicorn worker搜索和问答接口不受影响 > **⚠️ 注意**reindex 会调用 `sync_now()` 全局同步,期间 gunicorn worker 被阻塞,搜索和问答接口将暂时无响应。建议在低峰期执行
### Reranker 配置 ### Reranker 配置

View File

@@ -4,33 +4,9 @@
## 📋 变更记录2026-06-05 更新) ## 📋 变更记录2026-06-05 更新)
> **本次更新内容**:新增 AI 智能出题端口、更新生产环境测试结果、**长操作改为异步任务** > **本次更新内容**:新增 AI 智能出题端口、更新生产环境测试结果
> >
> **2026-06-05 更新**:出题批卷接口格式优化与输入校验增强 > **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 | 任务统计 | 按状态和类型分组统计 |
### 新增端口 ### 新增端口
@@ -147,22 +123,10 @@ RAG服务负责
### 2.4 出题系统(可选) ### 2.4 出题系统(可选)
| 端点 | 方法 | 功能 | 说明 | | 端点 | 方法 | 功能 |
|-----|------|------|------| |-----|------|------|
| `/exam/generate` | POST | 生成题目 | 异步任务,返回 task_id | | `/exam/generate` | POST | 生成题目 |
| `/exam/generate-smart` | POST | AI 智能出题 | 异步任务,返回 task_id | | `/exam/grade` | POST | 批阅答案 |
| `/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 | 任务统计 | 按状态和类型分组统计 |
--- ---
@@ -283,7 +247,7 @@ ENABLE_DIFY_WORKFLOW=false
--- ---
**最后更新**: 2026-06-05 **最后更新**: 2026-04-29
> 本文档供后端开发人员参考,用于对接 RAG 知识库服务。 > 本文档供后端开发人员参考,用于对接 RAG 知识库服务。
@@ -1125,24 +1089,17 @@ Content-Type: multipart/form-data
```json ```json
{ {
"success": true, "success": true,
"status_code": 2002, "message": "文件上传成功,已保存并添加到向量库",
"message": "文件上传成功,已保存,向量化任务已启动",
"data": {
"file": { "file": {
"filename": "document.pdf", "filename": "document.pdf",
"collection": "public_kb", "collection": "public_kb",
"path": "public_kb/document.pdf", "path": "public_kb/document.pdf",
"size": 1024000, "size": 1024000,
"replaced": false "replaced": false
},
"sync_status": "已保存,向量化任务已启动",
"task_id": "a1b2c3d4e5f6"
} }
} }
``` ```
> **异步说明**:文件保存为同步操作,向量化在后台线程异步执行。响应中的 `task_id` 可用于轮询向量化进度(`GET /tasks/<task_id>`)。当同步服务不可用时,`task_id` 为 `null``sync_status` 为 `"已保存,等待手动同步"`。
**同名文件处理**:上传同名文件时,旧版本的切片会被自动清理后覆盖(`replaced: true`),不会生成时间戳后缀文件。这确保了向量库中不会出现同一文档的新旧切片共存的情况。 **同名文件处理**:上传同名文件时,旧版本的切片会被自动清理后覆盖(`replaced: true`),不会生成时间戳后缀文件。这确保了向量库中不会出现同一文档的新旧切片共存的情况。
### 5.2 批量上传 ### 5.2 批量上传
@@ -1300,49 +1257,24 @@ POST /sync
**请求体:** 无需传递参数(同步所有知识库) **请求体:** 无需传递参数(同步所有知识库)
**响应(异步任务)** **响应:**
```json ```json
{ {
"success": true, "success": true,
"status": "success",
"status_code": 2010, "status_code": 2010,
"message": "同步任务已启动", "message": "同步完成",
"data": { "data": {
"task_id": "c3d4e5f6a1b2", "result": {
"message": "同步任务已启动,通过 GET /tasks/c3d4e5f6a1b2 查询进度" "documents_added": 1,
"documents_deleted": 1,
"documents_modified": 0,
"documents_processed": 2,
"errors": [],
"status": "completed"
}
} }
} }
```
> **⚠️ 异步变更**:此接口已从同步改为异步。不再直接返回同步结果,而是返回 `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 同步状态 ### 6.2 同步状态
@@ -1427,7 +1359,7 @@ POST /sync/stop
```json ```json
{ {
"status": "success", "status": "success",
"status_code": 2010, "status_code": 3001,
"message": "文件监控已启动" "message": "文件监控已启动"
} }
``` ```
@@ -1485,24 +1417,26 @@ POST /exam/generate
| difficulty | 必须为 1-5 的整数 | HTTP 400 INVALID_PARAMS | | difficulty | 必须为 1-5 的整数 | HTTP 400 INVALID_PARAMS |
| **总题数上限** | **所有题型数量之和不能超过 20** | HTTP 400 INVALID_PARAMS | | **总题数上限** | **所有题型数量之和不能超过 20** | HTTP 400 INVALID_PARAMS |
**响应(异步任务)** **响应:**
**完整响应格式**(包含外层包装):
```json ```json
{ {
"success": true, "success": true,
"status_code": 2020, "status": "success",
"message": "出题任务已启动", "status_code": 2011,
"message": "出题完成",
"data": { "data": {
"task_id": "d4e5f6a1b2c3", "success": true,
"message": "出题任务已启动 (10题),通过 GET /tasks/d4e5f6a1b2c3 查询结果" "request_id": "xxx",
"total": 10,
"source_chunks_used": 15,
"questions": [...]
} }
} }
``` ```
> **⚠️ 异步变更**:此接口已从同步改为异步。响应仅返回 `task_id`,后端需通过 `GET /tasks/<task_id>` 轮询任务状态。任务完成后,`result` 字段包含完整出题结果(格式见下方说明)。 **data 内部结构**
**轮询结果GET /tasks/\<task_id\> 完成后的 result 字段)**
```json ```json
{ {
"success": true, "success": true,
@@ -1701,22 +1635,25 @@ POST /exam/grade
#### 响应 #### 响应
**响应(异步任务):** **完整响应格式**(包含外层包装):
```json ```json
{ {
"success": true, "success": true,
"status": "success",
"status_code": 2021, "status_code": 2021,
"message": "批阅任务已启动", "message": "批阅完成",
"data": { "data": {
"task_id": "f6a1b2c3d4e5", "request_id": "可选,原样返回",
"message": "批阅任务已启动 (5题),通过 GET /tasks/f6a1b2c3d4e5 查询结果" "success": true,
"total_score": 12.5,
"total_max_score": 22.0,
"score_rate": 56.8,
"results": [...]
} }
} }
``` ```
> **⚠️ 异步变更**:此接口已从同步改为异步。响应仅返回 `task_id`,后端需通过 `GET /tasks/<task_id>` 轮询任务状态。任务完成后,`result` 字段包含完整批阅结果(格式见下方说明)。 **data 内部结构**
**轮询结果GET /tasks/\<task_id\> 完成后的 result 字段)**
```json ```json
{ {
@@ -1880,30 +1817,16 @@ def build_grade_request(answer_list, questions_map):
return {"answers": grade_answers} return {"answers": grade_answers}
``` ```
**Step 4: 调用 RAG 批卷接口(异步任务)** **Step 4: 调用 RAG 批卷接口**
```python ```python
import time
def call_rag_grade(grade_request): def call_rag_grade(grade_request):
"""调用 RAG 批卷接口并轮询等待结果""" """调用 RAG 批卷接口"""
# 1. 提交批阅任务
response = requests.post( response = requests.post(
'http://rag-service:5001/exam/grade', 'http://rag-service:5001/exam/grade',
json=grade_request json=grade_request
) )
task_id = response.json()['data']['task_id'] return response.json()
# 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: 更新学生成绩** **Step 5: 更新学生成绩**
@@ -2507,22 +2430,19 @@ ChromaDB 集合名称限制:
| 端点 | 方法 | 说明 | | 端点 | 方法 | 说明 |
|------|------|------| |------|------|------|
| `/sync` | POST | 触发同步**异步任务,返回 task_id** | | `/sync` | POST | 触发同步 |
| `/sync/status` | GET | 同步状态 | | `/sync/status` | GET | 同步状态 |
| `/sync/history` | GET | 同步历史 | | `/sync/history` | GET | 同步历史 |
| `/sync/changes` | GET | 变更日志 | | `/sync/changes` | GET | 变更日志 |
| `/sync/start` | POST | 启动文件监控 | | `/sync/start` | POST | 启动文件监控 |
| `/sync/stop` | POST | 停止文件监控 | | `/sync/stop` | POST | 停止文件监控 |
| `/documents/sync` | POST | 触发文档同步(**异步任务,返回 task_id** |
| `/collections/<kb_name>/reindex` | POST | 重建索引(**异步任务,返回 task_id** |
### 13.7 出题系统 ### 13.7 出题系统
| 端点 | 方法 | 说明 | | 端点 | 方法 | 说明 |
|------|------|------| |------|------|------|
| `/exam/generate` | POST | 生成题目**异步任务,返回 task_id** | | `/exam/generate` | POST | 生成题目 |
| `/exam/generate-smart` | POST | AI 智能出题(**异步任务,返回 task_id** | | `/exam/grade` | POST | 批改答案 |
| `/exam/grade` | POST | 批改答案(**异步任务,返回 task_id** |
| `/exam/health` | GET | 出题服务健康检查 | | `/exam/health` | GET | 出题服务健康检查 |
### 13.8 反馈与 FAQ 管理 ### 13.8 反馈与 FAQ 管理
@@ -2577,77 +2497,6 @@ ChromaDB 集合名称限制:
|------|------|------| |------|------|------|
| `/documents/<path>/preview` | GET | 文档预览,按 `chunk_index` 跳转到具体切片dev-ui 引用溯源用) | | `/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)")
```
--- ---
## 十四、文件管理服务(后端负责) ## 十四、文件管理服务(后端负责)

View File

@@ -1,554 +0,0 @@
# 异步任务接口变更说明
> **变更日期**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` 字段包含错误描述。

View File

@@ -2,8 +2,8 @@
出题与批题系统 API 蓝图 出题与批题系统 API 蓝图
提供 REST API 接口: 提供 REST API 接口:
- 出题:生成题目(异步任务,返回 task_id - 出题:生成题目(返回 JSON 给后端
- 批题:批阅答案(异步任务,返回 task_id - 批题:批阅答案(返回结果给后端
职责边界: 职责边界:
- RAG 服务负责:生成题目 + 批阅答案 - RAG 服务负责:生成题目 + 批阅答案
@@ -12,11 +12,6 @@
使用方式: 使用方式:
from exam_pkg.api import exam_bp from exam_pkg.api import exam_bp
app.register_blueprint(exam_bp, url_prefix='/exam') 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 from flask import Blueprint, request, jsonify
@@ -34,7 +29,7 @@ from auth.gateway import (
) )
# 导入统一响应格式 # 导入统一响应格式
from core.status_codes import EXAM_SUCCESS, GRADE_SUCCESS, EXAM_ERROR, GRADE_ERROR, BAD_REQUEST, UNAUTHORIZED, FORBIDDEN, NO_CONTENT, LLM_ERROR from core.status_codes import EXAM_SUCCESS, GRADE_SUCCESS, EXAM_ERROR, GRADE_ERROR, BAD_REQUEST, NO_CONTENT, LLM_ERROR
from api.response_utils import success_response, error_response from api.response_utils import success_response, error_response
# 合法题型 # 合法题型
@@ -174,7 +169,7 @@ def api_generate_questions():
# 获取当前用户 # 获取当前用户
user = get_current_user() user = get_current_user()
if not user: if not user:
return error_response("UNAUTHORIZED", UNAUTHORIZED, "未认证", http_status=401) return error_response("UNAUTHORIZED", BAD_REQUEST, "未认证", http_status=401)
# 检查向量库访问权限 # 检查向量库访问权限
if not check_collection_permission( if not check_collection_permission(
@@ -183,51 +178,20 @@ def api_generate_questions():
collection_name=collection, collection_name=collection,
operation="read" operation="read"
): ):
return error_response("FORBIDDEN", FORBIDDEN, "权限不足", http_status=403) return error_response("FORBIDDEN", BAD_REQUEST, "权限不足", http_status=403)
# 调用新版出题接口(异步任务) # 调用新版出题接口
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
)
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( result = generate_questions_from_file(
file_path=fp, file_path=file_path,
collection=coll, collection=collection,
question_types=q_types, question_types=question_types,
difficulty=diff, difficulty=data.get('difficulty', 3),
options=opts, options=data.get('options', {}),
request_id=req_id, request_id=data.get('request_id'),
exclude_stems=excl exclude_stems=data.get('exclude_stems')
)
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( return success_response(data=result, status_code=EXAM_SUCCESS, message="出题成功")
data={
'task_id': task.id,
'message': f'出题任务已启动 ({total_questions}题),通过 GET /tasks/{task.id} 查询结果'
},
status_code=EXAM_SUCCESS,
message="出题任务已启动"
)
except Exception as e: except Exception as e:
return error_response("EXAM_ERROR", EXAM_ERROR, str(e), http_status=500) return error_response("EXAM_ERROR", EXAM_ERROR, str(e), http_status=500)
@@ -276,7 +240,7 @@ def api_generate_smart():
# 获取当前用户 # 获取当前用户
user = get_current_user() user = get_current_user()
if not user: if not user:
return error_response("UNAUTHORIZED", UNAUTHORIZED, "未认证", http_status=401) return error_response("UNAUTHORIZED", BAD_REQUEST, "未认证", http_status=401)
# 检查向量库访问权限 # 检查向量库访问权限
if not check_collection_permission( if not check_collection_permission(
@@ -285,57 +249,34 @@ def api_generate_smart():
collection_name=collection, collection_name=collection,
operation="read" operation="read"
): ):
return error_response("FORBIDDEN", FORBIDDEN, "权限不足", http_status=403) return error_response("FORBIDDEN", BAD_REQUEST, "权限不足", http_status=403)
# AI 智能出题(异步任务) # 1. 调用 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)}"
)
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 from exam_pkg.manager import analyze_file_for_exam
ai_analysis = analyze_file_for_exam(file_path=fp, collection=coll) ai_analysis = analyze_file_for_exam(
file_path=file_path,
collection=collection
)
q_types = ai_analysis.get('question_types', {}) question_types = ai_analysis.get('question_types', {})
if not q_types or sum(q_types.values()) == 0: if not question_types or sum(question_types.values()) == 0:
raise ValueError("AI 分析后未生成有效题型配置") return error_response("EXAM_ERROR", EXAM_ERROR, "AI 分析后未生成有效题型配置", http_status=500)
total = sum(q_types.values())
registry.update_progress(task.id, total=total, stage='生成题目',
message=f"AI 推荐 {total} 道题,正在生成...")
# 2. 使用 AI 推荐的题型调用出题接口
result = generate_questions_from_file( result = generate_questions_from_file(
file_path=fp, collection=coll, file_path=file_path,
question_types=q_types, difficulty=diff, collection=collection,
options=opts, request_id=req_id, exclude_stems=excl 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')
) )
# 3. 在返回结果中添加 AI 分析信息
result['ai_analysis'] = ai_analysis result['ai_analysis'] = ai_analysis
registry.update_progress(task.id, stage='完成', message=f"生成 {result.get('total', 0)} 道题")
return result
registry.start_task( return success_response(data=result, status_code=EXAM_SUCCESS, message="AI 智能出题成功")
task.id, _do_smart_generate,
file_path, collection,
data.get('difficulty', 3), data.get('options', {}),
data.get('request_id'), data.get('exclude_stems')
)
return success_response(
data={
'task_id': task.id,
'message': 'AI 智能出题任务已启动,通过 GET /tasks/' + task.id + ' 查询结果'
},
status_code=EXAM_SUCCESS,
message="AI 智能出题任务已启动"
)
except Exception as e: except Exception as e:
return error_response("EXAM_ERROR", EXAM_ERROR, str(e), http_status=500) return error_response("EXAM_ERROR", EXAM_ERROR, str(e), http_status=500)
@@ -426,34 +367,13 @@ def api_grade_answers():
f"{i+1} 题的 question_type 无效: {q_type},合法值: {', '.join(sorted(VALID_QUESTION_TYPES))}", f"{i+1} 题的 question_type 无效: {q_type},合法值: {', '.join(sorted(VALID_QUESTION_TYPES))}",
http_status=400) http_status=400)
# 调用批题接口(异步任务) # 调用新版批题接口
from core.task_registry import get_registry result = grade_answers(
registry = get_registry() answers=answers,
request_id=data.get('request_id')
task = registry.create_task(
'exam_grade',
f"批阅: {len(answers)} 道题",
total=len(answers)
) )
def _do_grade(task, ans_list, req_id): return success_response(data=result, status_code=GRADE_SUCCESS, message="批阅完成")
"""后台执行批阅"""
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: except Exception as e:
return error_response("GRADE_ERROR", GRADE_ERROR, str(e), http_status=500) return error_response("GRADE_ERROR", GRADE_ERROR, str(e), http_status=500)

View File

@@ -393,6 +393,11 @@ class KnowledgeBaseManager(
if len(chunks) < 2: if len(chunks) < 2:
return chunks 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 = [] merged_chunks = []
i = 0 i = 0
merge_count = 0 merge_count = 0
@@ -405,6 +410,7 @@ class KnowledgeBaseManager(
# 查找下一个表格(跳过中间的"续表"文本) # 查找下一个表格(跳过中间的"续表"文本)
next_table_idx = None next_table_idx = None
next_chunk = None next_chunk = None
intermediate_texts = [] # 收集中间文本用于降级判断
for j in range(i + 1, min(i + 4, len(chunks))): # 最多向前看3个切片 for j in range(i + 1, min(i + 4, len(chunks))): # 最多向前看3个切片
candidate = chunks[j] candidate = chunks[j]
@@ -419,7 +425,11 @@ class KnowledgeBaseManager(
elif candidate_type == 'text' and ('续表' in candidate_title or '续表' in candidate_content): elif candidate_type == 'text' and ('续表' in candidate_title or '续表' in candidate_content):
# 遇到"续表"文本,继续查找下一个表格 # 遇到"续表"文本,继续查找下一个表格
continue continue
elif candidate_type not in ('text',): elif candidate_type == 'text':
# 非"续表"文本,收集后停止查找
intermediate_texts.append(candidate)
break
else:
# 遇到非文本类型,停止查找 # 遇到非文本类型,停止查找
break break
@@ -439,12 +449,22 @@ class KnowledgeBaseManager(
# 获取内容(用于检测"续表" # 获取内容(用于检测"续表"
next_content = getattr(next_chunk, 'content', '') next_content = getattr(next_chunk, 'content', '')
# 通用/无意义标题集合,这些标题不能用于"标题相似"判定
_GENERIC_TITLES = {'表格', 'table', '表格', ''}
# 判断是否为跨页表格 # 判断是否为跨页表格
is_cross_page = False is_cross_page = False
# 规则1: 页码连续(如果页码有效) # 规则1: 页码连续(如果页码有效)
page_valid = curr_page_end > 0 and next_page_start > 0 page_valid = curr_page_end > 0 and next_page_start > 0
if page_valid and curr_page_end + 1 == next_page_start: if page_valid and curr_page_end + 1 == next_page_start:
# 页码连续时,还需标题匹配或为通用标题才合并
# 避免把不同页面上不相关的表格错误合并
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 is_cross_page = True
# 规则2: 第二个表格标题或内容包含"续表" # 规则2: 第二个表格标题或内容包含"续表"
@@ -452,11 +472,36 @@ class KnowledgeBaseManager(
is_cross_page = True is_cross_page = True
# 规则3: 标题相似(去掉"续表"后比较) # 规则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() 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 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: if is_cross_page:
# 执行合并 # 执行合并
merge_count += 1 merge_count += 1
@@ -469,14 +514,27 @@ class KnowledgeBaseManager(
# 合并两个表格的 HTML # 合并两个表格的 HTML
current.table_html = curr_html + '\n' + next_html current.table_html = curr_html + '\n' + next_html
# 合并 image_path 到 images # 合并 image_path 和嵌入图片到 images
curr_img = getattr(current, 'image_path', None) curr_img = getattr(current, 'image_path', None)
next_img = getattr(next_chunk, 'image_path', None) next_img = getattr(next_chunk, 'image_path', None)
merged_images = [] curr_images = getattr(current, 'images', None) or []
if curr_img: 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}) 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}) merged_images.append({'id': next_img, 'page': next_page_start})
if merged_images: if merged_images:
current.images = merged_images current.images = merged_images
# 保留第一个图片作为主 image_path # 保留第一个图片作为主 image_path
@@ -525,7 +583,7 @@ class KnowledgeBaseManager(
try: try:
from config import get_llm_client, DASHSCOPE_MODEL from config import get_llm_client, DASHSCOPE_MODEL
client = get_llm_client() 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 "" return summary.strip() if summary else ""
except Exception as e: except Exception as e:
logger.warning(f"生成表格摘要失败: {e}") logger.warning(f"生成表格摘要失败: {e}")
@@ -584,7 +642,7 @@ class KnowledgeBaseManager(
] ]
} }
], ],
max_tokens=200 max_tokens=512
) )
description = response.choices[0].message.content description = response.choices[0].message.content

View File

@@ -283,7 +283,7 @@ class KnowledgeBaseRouter:
content = call_llm( content = call_llm(
self.llm_client, prompt, MODEL, self.llm_client, prompt, MODEL,
temperature=0.1, temperature=0.1,
max_tokens=100 max_tokens=512
) )
if content is None: if content is None:

347
parsers/heading_rules.py Normal file
View 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

View File

@@ -286,14 +286,138 @@ def parse_with_mineru_online(
raise RuntimeError(f"MinerU 在线 API 调用失败: {e}") raise RuntimeError(f"MinerU 在线 API 调用失败: {e}")
def _parse_v2_content_list(v2_data: list) -> list:
"""
将 MinerU v2 嵌套格式转换为 v1 兼容的扁平 content_list
v2 格式: [[page0_items], [page1_items], ...]
v1 格式: [item, item, ...] 扁平列表
转换规则:
- paragraph → text拼接 paragraph_content提取 style 信息)
- table → table保留 table_body/html
- title → text提取 level 信息)
- page_header / page_footer / page_number → 过滤掉
Args:
v2_data: v2 格式的嵌套列表
Returns:
v1 兼容的扁平列表,额外包含 _v2_styles / _v2_table_type 字段
"""
flat_list: List[Dict] = []
for page_idx, page_items in enumerate(v2_data):
if not isinstance(page_items, list):
continue
for item in page_items:
if not isinstance(item, dict):
continue
v2_type: str = item.get('type', '')
content = item.get('content', {})
# 过滤噪音类型(页眉、页脚、页码)
if v2_type in ('page_header', 'page_footer', 'page_number'):
continue
if v2_type == 'paragraph':
# 拼接段落文本,收集样式信息
text_parts: List[str] = []
has_bold = False
para_content = content.get('paragraph_content', []) if isinstance(content, dict) else []
for part in para_content:
if not isinstance(part, dict):
continue
part_text = part.get('content', '')
text_parts.append(part_text)
if 'bold' in part.get('style', []):
has_bold = True
full_text = ''.join(text_parts).strip()
if not full_text:
continue
flat_item: Dict[str, Any] = {
'type': 'text',
'text': full_text,
'content': full_text,
'page_idx': page_idx,
'bbox': item.get('bbox', []),
'text_level': 0,
'_v2_styles': ['bold'] if has_bold else [],
}
flat_list.append(flat_item)
elif v2_type == 'table':
flat_item = {
'type': 'table',
'page_idx': page_idx,
'bbox': item.get('bbox', []),
'html': content.get('html', '') if isinstance(content, dict) else '',
'table_body': content.get('html', '') if isinstance(content, dict) else '',
'caption': content.get('caption', '') if isinstance(content, dict) else '',
'_v2_table_type': content.get('table_type', '') if isinstance(content, dict) else '',
}
flat_list.append(flat_item)
elif v2_type == 'title':
# v2 title 项有 level 信息
level = content.get('level', 0) if isinstance(content, dict) else 0
text_parts = []
para_content = content.get('paragraph_content', []) if isinstance(content, dict) else []
for part in para_content:
if isinstance(part, dict):
text_parts.append(part.get('content', ''))
full_text = ''.join(text_parts).strip()
if full_text:
flat_item = {
'type': 'text',
'text': full_text,
'content': full_text,
'page_idx': page_idx,
'bbox': item.get('bbox', []),
'text_level': level,
'_v2_styles': ['bold'],
}
flat_list.append(flat_item)
elif v2_type in ('image', 'chart'):
flat_item = {
'type': v2_type,
'page_idx': page_idx,
'bbox': item.get('bbox', []),
'img_path': content.get('img_path', '') if isinstance(content, dict) else '',
'image_path': content.get('img_path', '') if isinstance(content, dict) else '',
'caption': content.get('caption', '') if isinstance(content, dict) else '',
}
flat_list.append(flat_item)
elif v2_type == 'equation':
flat_item = {
'type': 'equation',
'page_idx': page_idx,
'bbox': item.get('bbox', []),
'content': content.get('latex', '') if isinstance(content, dict) else '',
'text': content.get('latex', '') if isinstance(content, dict) else '',
'latex': content.get('latex', '') if isinstance(content, dict) else '',
'img_path': content.get('img_path', '') if isinstance(content, dict) else '',
}
flat_list.append(flat_item)
logger.info(f"v2 格式转换: {len(v2_data)} 页 → {len(flat_list)} 项(已过滤噪音类型)")
return flat_list
def _parse_mineru_online_zip(zip_content: bytes, file_path: Path) -> Dict[str, Any]: def _parse_mineru_online_zip(zip_content: bytes, file_path: Path) -> Dict[str, Any]:
""" """
解析 MinerU 在线 API 返回的 zip 包 解析 MinerU 在线 API 返回的 zip 包
zip 包结构 zip 包结构:
- full.md - Markdown 解析结果 - full.md - Markdown 解析结果
- *_content_list.json - 内容列表(扁平格式,优先使用 - *_content_list.json - 内容列表(v1 扁平格式)
- *_content_list_v2.json - 内容列表(嵌套格式) - *_content_list_v2.json - 内容列表(v2 嵌套格式,含 style 信息
- images/ - 图片目录 - images/ - 图片目录
Args: Args:
@@ -306,23 +430,44 @@ def _parse_mineru_online_zip(zip_content: bytes, file_path: Path) -> Dict[str, A
import zipfile import zipfile
import io import io
# 读取格式偏好配置
try:
from config import MINERU_PREFER_V2
except ImportError:
MINERU_PREFER_V2 = True # 默认优先 v2
with zipfile.ZipFile(io.BytesIO(zip_content), 'r') as zf: with zipfile.ZipFile(io.BytesIO(zip_content), 'r') as zf:
# 列出所有文件 # 列出所有文件
file_list = zf.namelist() file_list = zf.namelist()
logger.debug(f"zip 包内容: {file_list}") logger.debug(f"zip 包内容: {file_list}")
# 查找 content_list.json优先使用扁平格式 # 查找 content_list 文件v2 含 style 信息,优先使用
content_list_path = None content_list_path = None
is_v2 = False
if MINERU_PREFER_V2:
# 优先使用 v2 格式(含 style 信息)
for f in file_list:
if f.endswith('_content_list_v2.json'):
content_list_path = f
is_v2 = True
break
if not content_list_path:
for f in file_list:
if f.endswith('_content_list.json') and not f.endswith('_v2.json'):
content_list_path = f
break
else:
# 优先使用 v1 格式
for f in file_list: for f in file_list:
# 优先使用 content_list.json扁平格式
if f.endswith('_content_list.json') and not f.endswith('_v2.json'): if f.endswith('_content_list.json') and not f.endswith('_v2.json'):
content_list_path = f content_list_path = f
break break
# 如果没有找到,再尝试 v2 格式
if not content_list_path: if not content_list_path:
for f in file_list: for f in file_list:
if f.endswith('_content_list_v2.json'): if f.endswith('_content_list_v2.json'):
content_list_path = f content_list_path = f
is_v2 = True
break break
# 查找 markdown 文件 # 查找 markdown 文件
@@ -337,7 +482,13 @@ def _parse_mineru_online_zip(zip_content: bytes, file_path: Path) -> Dict[str, A
if content_list_path: if content_list_path:
with zf.open(content_list_path) as f: with zf.open(content_list_path) as f:
content_list = json.load(f) content_list = json.load(f)
logger.info(f"读取 content_list: {len(content_list)} 项, 来源: {content_list_path}")
# v2 格式需要转换为 v1 兼容的扁平列表
if is_v2 and isinstance(content_list, list) and content_list and isinstance(content_list[0], list):
content_list = _parse_v2_content_list(content_list)
logger.info(f"v2 格式已转换为扁平列表,共 {len(content_list)}")
logger.info(f"读取 content_list: {len(content_list)} 项, 格式={'v2' if is_v2 else 'v1'}, 来源: {content_list_path}")
# 读取 markdown # 读取 markdown
markdown_content = "" markdown_content = ""
@@ -424,10 +575,14 @@ def _parse_mineru_online_result(result: Dict, file_path: Path) -> Dict[str, Any]
if item_type == "text": if item_type == "text":
text = item.get("content", "") or item.get("text", "") text = item.get("content", "") or item.get("text", "")
v2_styles = item.get("_v2_styles", [])
# 启发式标题识别 # 启发式标题识别
if text_level == 0: if text_level == 0:
text_level = _detect_heading_level(text) from parsers.heading_rules import get_heading_engine
engine = get_heading_engine()
detected_level, _ = engine.detect(text, style=v2_styles)
text_level = detected_level
title = "" title = ""
if text_level > 0: if text_level > 0:
@@ -471,7 +626,8 @@ def _parse_mineru_online_result(result: Dict, file_path: Path) -> Dict[str, Any]
else: else:
md_table = "" md_table = ""
table_images = extract_images_from_markdown(md_table) if md_table else [] # 从原始 HTML 提取嵌入图片md_table 经 get_text 转换后已丢失 <img> 标签)
table_images = extract_images_from_markdown(table_body) if table_body else []
chunk = MinerUChunk( chunk = MinerUChunk(
content=table_caption or "表格", content=table_caption or "表格",
@@ -558,8 +714,14 @@ def _parse_mineru_online_result(result: Dict, file_path: Path) -> Dict[str, Any]
min_merge = 100 min_merge = 100
max_size = 1200 max_size = 1200
# 表单类型二次校正(在 _post_process_chunks 之前,因为 table 不参与合并)
_reclassify_text_chunks(chunks)
chunks = _post_process_chunks(chunks, min_merge_size=min_merge, max_chunk_size=max_size) chunks = _post_process_chunks(chunks, min_merge_size=min_merge, max_chunk_size=max_size)
# 验证分类标题是否被规则引擎正确识别(安全网,仅告警不修改)
_validate_category_section_paths(chunks)
return { return {
'markdown': "\n".join(markdown_parts) if markdown_parts else markdown_content, 'markdown': "\n".join(markdown_parts) if markdown_parts else markdown_content,
'chunks': chunks, 'chunks': chunks,
@@ -697,16 +859,18 @@ def parse_with_mineru(
logger.error(f"MinerU 解析失败: {e}") logger.error(f"MinerU 解析失败: {e}")
raise raise
finally: finally:
# 清理临时目录 清理临时目录
if cleanup_output and os.path.exists(output_dir): if cleanup_output and os.path.exists(output_dir):
shutil.rmtree(output_dir, ignore_errors=True) shutil.rmtree(output_dir, ignore_errors=True)
def _detect_heading_level(text: str) -> int: def _detect_heading_level(text: str) -> int:
""" """
启发式标题识别 启发式标题识别(规则引擎版)
用于 MinerU 解析 DOCX 等 Office 格式时不提供 text_level 的情况 MinerU 解析 DOCX 等 Office 格式时不提供 text_level 时使用
规则按优先级从高到低匹配,第一个命中即返回。
规则定义见 parsers/heading_rules.py可通过 config.py 覆盖。
Args: Args:
text: 文本内容 text: 文本内容
@@ -714,57 +878,91 @@ def _detect_heading_level(text: str) -> int:
Returns: Returns:
标题级别 (0=正文, 1=h1, 2=h2, 3=h3) 标题级别 (0=正文, 1=h1, 2=h2, 3=h3)
""" """
import re from parsers.heading_rules import get_heading_engine
engine = get_heading_engine()
level, _ = engine.detect(text)
return level
text = text.strip() def _reclassify_text_chunks(chunks: List[MinerUChunk]) -> None:
"""
二次校正:检测被标记为 text 但实际是表格/表单的 chunk
# 空文本 MinerU 解析 Word 文档时,某些带下划线填空项的表单
被标记为 text 类型,需要根据内容特征修正为 table。
就地修改 chunks 列表中的 chunk_type 字段。
检测依据(基于实测数据设计):
- 连续下划线 ___ (3个以上) —— Word 表单填空项
- 冒号后跟下划线 如 "日期____" —— 键值对式表单
需 >= min_indicators 个指标同时命中才校正,避免误判。
"""
try:
from config import FORM_RECLASSIFY_ENABLED
if not FORM_RECLASSIFY_ENABLED:
return
except ImportError:
pass
try:
from config import FORM_RECLASSIFY_MIN_INDICATORS
min_indicators = FORM_RECLASSIFY_MIN_INDICATORS
except ImportError:
min_indicators = 2
# 表单特征指标(编译一次,避免循环内重复编译)
# 注意MinerU 输出的下划线是 Markdown 转义格式 \_\_\_
# 需要同时匹配纯下划线 ___ 和转义下划线 \_\_\_
_FORM_INDICATORS = [
re.compile(r'(?:\\_|_){3,}'), # 连续下划线3个以上含转义格式
re.compile(r'[:]\s*(?:\\_|_){2,}'), # 冒号后跟下划线(含转义格式)
]
reclassified = 0
for chunk in chunks:
if chunk.chunk_type != 'text':
continue
text = (chunk.content or '').strip()
if not text: if not text:
return 0 continue
# 中文章节标题模式 indicator_count = sum(1 for p in _FORM_INDICATORS if p.search(text))
# 第一章、第二章、... -> h1 if indicator_count >= min_indicators:
if re.match(r'^第[一二三四五六七八九十百千万]+[章节篇部]', text): chunk.chunk_type = 'table'
return 1 reclassified += 1
logger.info(f"表单检测: text -> table, content='{text[:80]}'")
# 第一条、第二条、... -> h2 (条文编号) if reclassified > 0:
if re.match(r'^第[一二三四五六七八九十百千万]+[条款]', text): logger.info(f"表单类型二次校正: 共 {reclassified} 个 text -> table")
return 2
# 数字章节: 1. 2. 3. 或 1、2、3、
# 一级标题: 1. 2. 3. (单数字)
if re.match(r'^\d+[\.、\s]', text):
# 短文本可能是标题
if len(text) < 50:
return 1
# 二级标题: 1.1 1.2 2.1 等 def _validate_category_section_paths(chunks: List[MinerUChunk]) -> None:
if re.match(r'^\d+\.\d+[\.、\s]', text): """
if len(text) < 80: 验证分类标题是否被规则引擎正确识别(安全网)
return 2
# 三级标题: 1.1.1 1.1.2 等 当规则引擎正确识别分类标题后section_stack 自然会更新,
if re.match(r'^\d+\.\d+\.\d+[\.、\s]', text): 不再需要后处理修改 section_path。此函数仅做验证和告警
if len(text) < 100: 便于发现规则引擎的遗漏。
return 3
# 英文章节标题 支持的模式A1类、B2类、C1类含可选 ** 粗体标记)。
# Chapter 1, Section 2, etc. """
if re.match(r'^(Chapter|Section|Part|Chapter\s+\d+|Section\s+\d+)', text, re.IGNORECASE): cat_pattern = re.compile(r'^\*{0,2}[A-Z]\d+[类類]\*{0,2}[:]')
return 1
# 短文本 + 加粗标记 (**xxx**) 可能是标题 missed_count = 0
if re.match(r'^\*\*.+\*\*$', text) and len(text) < 50: for chunk in chunks:
return 2 if chunk.chunk_type == 'text':
text = (chunk.content or '').strip()
if cat_pattern.match(text) and chunk.text_level == 0:
missed_count += 1
logger.warning(
f"分类标题未被识别为标题: '{text[:50]}', "
f"section_path='{chunk.section_path}'"
)
# 非常短的文本 (< 20 字符) 可能是标题 if missed_count > 0:
# 但需要排除常见的非标题短文本 logger.warning(
if len(text) < 20 and not re.match(r'^[\d\s\.,;:!?,。;:!?、]+$', text): f"发现 {missed_count} 个分类标题未被规则引擎识别,"
# 排除纯数字、纯标点 f"请检查 heading_rules 配置"
if re.search(r'[\u4e00-\u9fff]', text): # 包含中文 )
return 2
return 0
def _parse_mineru_output(file_path: Path, output_dir) -> Dict[str, Any]: def _parse_mineru_output(file_path: Path, output_dir) -> Dict[str, Any]:
@@ -803,16 +1001,45 @@ def _parse_mineru_output(file_path: Path, output_dir) -> Dict[str, Any]:
else: else:
raise RuntimeError(f"MinerU 输出目录不存在: {auto_dir}{office_dir}") raise RuntimeError(f"MinerU 输出目录不存在: {auto_dir}{office_dir}")
# 读取 content_list.json # 读取 content_listv2 含 style 信息,优先使用)
content_list_path = output_subdir / f"{doc_name}_content_list.json" try:
if not content_list_path.exists(): from config import MINERU_PREFER_V2
content_list_path = output_subdir / f"{doc_name}_content_list_v2.json" except ImportError:
MINERU_PREFER_V2 = True
v1_path = output_subdir / f"{doc_name}_content_list.json"
v2_path = output_subdir / f"{doc_name}_content_list_v2.json"
content_list_path = None
is_v2 = False
if MINERU_PREFER_V2:
# 优先使用 v2 格式
if v2_path.exists():
content_list_path = v2_path
is_v2 = True
elif v1_path.exists():
content_list_path = v1_path
else:
# 优先使用 v1 格式
if v1_path.exists():
content_list_path = v1_path
elif v2_path.exists():
content_list_path = v2_path
is_v2 = True
content_list = [] content_list = []
if content_list_path.exists(): if content_list_path and content_list_path.exists():
with open(content_list_path, 'r', encoding='utf-8') as f: with open(content_list_path, 'r', encoding='utf-8') as f:
content_list = json.load(f) content_list = json.load(f)
# v2 格式需要转换为 v1 兼容的扁平列表
if is_v2 and isinstance(content_list, list) and content_list and isinstance(content_list[0], list):
content_list = _parse_v2_content_list(content_list)
logger.info(f"v2 格式已转换为扁平列表,共 {len(content_list)}")
logger.info(f"读取 content_list: {len(content_list)} 项, 格式={'v2' if is_v2 else 'v1'}")
# 读取 Markdown # 读取 Markdown
md_path = output_subdir / f"{doc_name}.md" md_path = output_subdir / f"{doc_name}.md"
markdown_content = "" markdown_content = ""
@@ -857,10 +1084,14 @@ def _parse_mineru_output(file_path: Path, output_dir) -> Dict[str, Any]:
if item_type == "text": if item_type == "text":
text = item.get("text", "") text = item.get("text", "")
v2_styles = item.get("_v2_styles", [])
# 启发式标题识别(当 text_level 为 0 时) # 启发式标题识别(当 text_level 为 0 时)
if text_level == 0: if text_level == 0:
text_level = _detect_heading_level(text) from parsers.heading_rules import get_heading_engine
engine = get_heading_engine()
detected_level, _ = engine.detect(text, style=v2_styles)
text_level = detected_level
# 处理标题 # 处理标题
title = "" title = ""
@@ -909,8 +1140,8 @@ def _parse_mineru_output(file_path: Path, output_dir) -> Dict[str, Any]:
else: else:
md_table = "" md_table = ""
# 提取表格中的嵌入图片 # 从原始 HTML 提取嵌入图片md_table 经 get_text 转换后已丢失 <img> 标签)
table_images = extract_images_from_markdown(md_table) if md_table else [] table_images = extract_images_from_markdown(table_body) if table_body else []
chunk = MinerUChunk( chunk = MinerUChunk(
content=table_caption or "表格", content=table_caption or "表格",
@@ -1003,8 +1234,14 @@ def _parse_mineru_output(file_path: Path, output_dir) -> Dict[str, Any]:
min_merge = 100 min_merge = 100
max_size = 1200 max_size = 1200
# 表单类型二次校正(在 _post_process_chunks 之前,因为 table 不参与合并)
_reclassify_text_chunks(chunks)
chunks = _post_process_chunks(chunks, min_merge_size=min_merge, max_chunk_size=max_size) chunks = _post_process_chunks(chunks, min_merge_size=min_merge, max_chunk_size=max_size)
# 验证分类标题是否被规则引擎正确识别(安全网,仅告警不修改)
_validate_category_section_paths(chunks)
return { return {
'markdown': "\n".join(markdown_parts), 'markdown': "\n".join(markdown_parts),
'chunks': chunks, 'chunks': chunks,
@@ -1349,7 +1586,17 @@ def html_table_to_markdown(html_table: str) -> str:
del rowspan_tracker[col_idx] del rowspan_tracker[col_idx]
col_idx += 1 col_idx += 1
# 提取单元格内容 # 提取单元格内容(保留图片引用信息)
img_tags = cell.find_all('img')
if img_tags:
# 单元格包含图片,生成占位标记供 LLM 感知
text_part = cell.get_text(strip=True)
img_count = len(img_tags)
if text_part:
content = f"{text_part} [{'图片' if img_count == 1 else f'{img_count}张图片'}]"
else:
content = f"[{'图片' if img_count == 1 else f'{img_count}张图片'}]"
else:
content = cell.get_text(strip=True) content = cell.get_text(strip=True)
# 处理 rowspan # 处理 rowspan
@@ -1630,6 +1877,21 @@ def parse_with_mineru_persistent(
chunk.image_path = new_name chunk.image_path = new_name
break break
# 更新表格嵌入图片的路径映射images 字段)
if hasattr(chunk, 'images') and chunk.images:
for img_info in chunk.images:
if isinstance(img_info, dict) and 'id' in img_info:
old_id = img_info['id']
if old_id in image_path_map:
img_info['id'] = image_path_map[old_id]
else:
# 兼容:尝试用文件名匹配映射
old_basename = os.path.basename(old_id)
for old_path, new_name in image_path_map.items():
if old_basename == os.path.basename(old_path):
img_info['id'] = new_name
break
# 更新结果中的图片路径列表(供外部使用) # 更新结果中的图片路径列表(供外部使用)
result['images'] = list(image_path_map.values()) result['images'] = list(image_path_map.values())

View File

@@ -1,179 +0,0 @@
# 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 阈值调优和知识库版本联动失效。

View File

@@ -1,342 +0,0 @@
## 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` 淘汰策略。按前面估算,四层缓存满载约 95MB128MB 足够且留有余量。
### 九、迁移步骤建议
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 不一致 | 全局一致 |

View File

@@ -616,7 +616,7 @@ class FeedbackService:
prompt=prompt, prompt=prompt,
model=self.model, model=self.model,
temperature=0.7, temperature=0.7,
max_tokens=200 max_tokens=512
) )
if not response: if not response:

View File

@@ -159,7 +159,7 @@ class SessionManager:
SELECT id, role, content, metadata, created_at SELECT id, role, content, metadata, created_at
FROM messages FROM messages
WHERE session_id = ? WHERE session_id = ?
ORDER BY created_at DESC ORDER BY created_at DESC, id DESC
LIMIT ? LIMIT ?
''', (session_id, limit)) ''', (session_id, limit))