diff --git a/.gitignore b/.gitignore index 381dd6b..cfb8adc 100644 --- a/.gitignore +++ b/.gitignore @@ -59,8 +59,8 @@ CLAUDE.md # 会话数据库(含用户数据) sessions.db -# 测试文件 -tests/ +# pytest / 测试运行缓存(测试源码需要纳入版本库) +.pytest_cache/ # 文档源文件(可能含敏感内容) documents/ @@ -81,13 +81,20 @@ test_output/ # 非必要文档 docs/superpowers/ -# 以下为v5.0.0新增忽略:多向量库存储及SQLite数据 -data/ +# 多向量库存储及 SQLite 运行数据 +# data/db.py 与 data/__init__.py 是应用源码,必须纳入版本库;仅忽略数据库文件。 +data/prod/ +data/dev/ +data/sqlite/ +data/**/*.db +data/**/*.db-shm +data/**/*.db-wal vector_store/ knowledge/vector_store/ # 缓存和临时数据目录 .data/ +.recovery/ # 测试JSON文件 exam_questions.json @@ -114,14 +121,13 @@ chat-ui/ # 生产环境配置(含 API 密钥) deploy/.env.production -# 根目录临时测试脚本 -test_*.py -test_*.json +# 根目录临时测试脚本(tests/ 下的正式回归测试不忽略) +/test_*.py +/test_*.json rag_response.json nul -# 调试脚本和临时计划(仅本地使用) -scripts/ +# 临时计划(正式 scripts/ 源码需要纳入版本库) plans/ # Qoder 工具目录 @@ -143,6 +149,8 @@ docs/环境分离说明.md # 临时脚本 scripts/evaluate_enum_context.py scripts/migrate_split_databases.py +/scripts/test_cache*.py +/scripts/test_intent*.py # 上传清单(仅本地使用) UPLOAD_CHECKLIST.md diff --git a/config.example.py b/config.example.py index d69de81..608529d 100644 --- a/config.example.py +++ b/config.example.py @@ -9,11 +9,13 @@ import os # 一、API 密钥与模型 # ============================================================================== -# 通义千问 LLM 服务 +# OpenAI 兼容 LLM 服务(当前默认使用小米 MiMo) DASHSCOPE_API_KEY = os.getenv("DASHSCOPE_API_KEY", "") -DASHSCOPE_BASE_URL = os.getenv("DASHSCOPE_BASE_URL", "https://dashscope.aliyuncs.com/compatible-mode/v1") -DASHSCOPE_MODEL = os.getenv("DASHSCOPE_MODEL", "qwen-flash") # 文本生成模型 -RAG_CHAT_MODEL = os.getenv("RAG_CHAT_MODEL", "qwen-flash") # RAG 对话模型 +DASHSCOPE_BASE_URL = os.getenv("DASHSCOPE_BASE_URL", "https://token-plan-cn.xiaomimimo.com/v1") +DASHSCOPE_MODEL = os.getenv("DASHSCOPE_MODEL", "mimo-v2.5") # 文本生成模型 +RAG_CHAT_MODEL = os.getenv("RAG_CHAT_MODEL", "mimo-v2.5") # RAG 对话模型 +INTENT_MODEL = os.getenv("INTENT_MODEL", "mimo-v2.5") # 意图分析模型 +VLM_MODEL = os.getenv("VLM_MODEL", "mimo-v2.5") # 图片理解模型 # 兼容旧变量名(逐步迁移到 DASHSCOPE_* 命名) API_KEY = DASHSCOPE_API_KEY @@ -27,6 +29,7 @@ MODEL = DASHSCOPE_MODEL APP_ENV = os.getenv("APP_ENV", "dev") # dev / prod IS_DEV = APP_ENV == "dev" IS_PROD = APP_ENV == "prod" +DEV_MODE = os.getenv("DEV_MODE", "true").lower() != "false" # 开发/生产环境自动切换 ENABLE_SESSION = IS_DEV # 会话存储(仅开发环境) @@ -85,10 +88,11 @@ RERANK_DEVICE = os.getenv("RERANK_DEVICE", os.getenv("DEVICE", "auto")) # ----- 通用问答 ----- LLM_TEMPERATURE = 0.7 # 生成温度(0=确定性,1=随机性) LLM_MAX_TOKENS = 3000 # 最大输出 token 数 +LLM_DISABLE_THINKING = os.getenv("LLM_DISABLE_THINKING", "true").lower() != "false" # ----- 意图分析(轻量、确定性高)----- INTENT_TEMPERATURE = 0.1 -INTENT_MAX_TOKENS = 300 +INTENT_MAX_TOKENS = 2048 INTENT_HISTORY_WINDOW = 6 # 分析时取最近几条历史消息 # ============================================================================== @@ -201,6 +205,9 @@ MAX_QUERY_REWRITES = 1 # MinerU 解析器 MINERU_DEVICE_MODE = os.getenv("MINERU_DEVICE_MODE", "cpu") # cpu / cuda +MINERU_MODEL_VERSION = os.getenv("MINERU_MODEL_VERSION", "vlm") +MINERU_LOCAL_BACKEND = os.getenv("MINERU_LOCAL_BACKEND", "pipeline") +MINERU_ONLINE_TIMEOUT = int(os.getenv("MINERU_ONLINE_TIMEOUT", "300")) # MinerU 在线 API(作为本地解析失败时的备选) MINERU_API_TOKEN = os.getenv("MINERU_API_TOKEN", "") # 在 https://mineru.net/apiManage/token 申请 @@ -246,6 +253,20 @@ CONFIDENCE_CAUTION_THRESHOLD = 0.30 # top-3 均分低于此值时,提示 LLM ENUM_QUERY_DISABLE_TOPK_SHRINK = True ENUM_QUERY_MMR_LAMBDA = 0.85 +# 章节聚类提升与 BM25/CrossEncoder 分歧救援 +SECTION_CLUSTER_BOOST_ENABLED = True +SECTION_CLUSTER_RESCUE_ENABLED = True +BM25_DIVERGENCE_RESCUE_ENABLED = True +BM25_DIVERGENCE_MAX_RANK = 3 +CLUSTER_MIN_MEMBERS = 3 +CLUSTER_MIN_TYPES = 2 +CLUSTER_SEED_FLOOR = 0.35 +CLUSTER_RESCUE_FLOOR = 0.06 +CLUSTER_MAX_BOOST_PER_SECTION = 8 +CLUSTER_MAX_SECTIONS = 3 +CLUSTER_MAX_RESCUE_PER_SECTION = 6 +CLUSTER_SECTION_PREFIX_LEVELS = 2 + # ============================================================================== # 十一、可选功能配置 # ============================================================================== @@ -261,3 +282,20 @@ def get_llm_client(): """获取 LLM 客户端实例""" from openai import OpenAI return OpenAI(api_key=DASHSCOPE_API_KEY, base_url=DASHSCOPE_BASE_URL) + + +_intent_client = None + + +def get_intent_client(): + """获取意图分析专用 LLM 客户端(进程内复用连接池)。""" + global _intent_client + if _intent_client is None: + if not DASHSCOPE_API_KEY: + raise ValueError("DASHSCOPE_API_KEY 未配置,请在环境变量中设置") + from openai import OpenAI + _intent_client = OpenAI( + api_key=DASHSCOPE_API_KEY, + base_url=DASHSCOPE_BASE_URL, + ) + return _intent_client diff --git a/core/cache.py b/core/cache.py index e9c2beb..f7a3b0c 100644 --- a/core/cache.py +++ b/core/cache.py @@ -189,6 +189,9 @@ class RAGCacheManager: # 失效旧版本缓存 self.query_cache.invalidate_by_version(old_version) self.embedding_cache.invalidate_by_version(old_version) + # Rerank 缓存未携带知识库版本;文档变更后必须全量清空, + # 否则相同 doc_id 可能继续复用旧文档内容对应的分数。 + self.rerank_cache.clear() logger.info(f"知识库 {kb_name} 版本更新: {old_version} -> {new_version}") return new_version diff --git a/core/intent_analyzer.py b/core/intent_analyzer.py index b2ab6dd..080f6da 100644 --- a/core/intent_analyzer.py +++ b/core/intent_analyzer.py @@ -316,8 +316,8 @@ class IntentAnalyzer: if cache_emb is not None: cached = cache.get(cache_emb) - # 确保缓存条目是意图分析结果(非 RAG 回答缓存) - if cached and cached.get("cache_type") != "rag_answer": + # 仅接受明确标记的意图缓存,避免未来新增缓存类型时交叉命中。 + if cached and cached.get("cache_type") == "intent_analysis": logger.info(f"意图分析缓存命中: {cached.get('reason', '')[:50]}") return IntentAnalysis.from_dict(cached) else: diff --git a/data/__init__.py b/data/__init__.py new file mode 100644 index 0000000..cbfade4 --- /dev/null +++ b/data/__init__.py @@ -0,0 +1,46 @@ +""" +数据目录 + +统一数据库结构(按业务域和运行环境划分): +- data/prod/feedback.db: 反馈、FAQ 与质量报告 +- data/prod/knowledge.db: 同步、纲要与文档版本 +- data/dev/session.db: 本地开发会话与审计日志 +- data/dev/exam.db: 本地出题、试卷与批阅数据 + +使用方式: + from data.db import get_connection, init_databases + + # 首次运行时初始化 + init_databases() + + # 使用连接 + with get_connection("session") as conn: + cursor = conn.cursor() + cursor.execute("SELECT * FROM sessions") +""" + +from data.db import ( + get_connection, + get_raw_connection, + init_databases, + get_db_stats, + row_to_dict, + rows_to_list, + json_parse, + json_stringify, + DB_PATHS, + DATA_DIR, +) + +__all__ = [ + "get_connection", + "get_raw_connection", + "init_databases", + "get_db_stats", + "row_to_dict", + "rows_to_list", + "json_parse", + "json_stringify", + "DB_PATHS", + "DATA_DIR", +] diff --git a/data/db.py b/data/db.py new file mode 100644 index 0000000..988a793 --- /dev/null +++ b/data/db.py @@ -0,0 +1,821 @@ +""" +统一数据访问层 - 集中管理所有数据库连接 + +功能: +1. 统一数据库路径配置(按环境分离:prod/dev) +2. 连接池管理(上下文管理器) +3. WAL 模式 + 外键约束 +4. 自动事务管理(commit/rollback) +5. 初始化所有表结构 + +数据库分离: +- feedback.db: 反馈系统(生产模式也使用) +- knowledge.db: 知识管理(生产模式也使用) +- session.db: 会话管理(仅开发模式) +- exam.db: 出题系统(仅开发模式) + +使用方式: + from data.db import get_connection, init_databases + + # 初始化数据库(首次运行时调用) + init_databases() + + # 使用连接 + with get_connection("feedback") as conn: + cursor = conn.cursor() + cursor.execute("SELECT * FROM feedbacks WHERE user_id = ?", (user_id,)) + rows = cursor.fetchall() +""" + +import sqlite3 +import os +import json +import logging +import threading +from contextlib import contextmanager +from typing import Generator, Dict, List, Optional + +logger = logging.getLogger(__name__) + +# ==================== 路径配置 ==================== + +DATA_DIR = os.path.dirname(os.path.abspath(__file__)) + +DB_PATHS = { + # 生产模式数据库(开发模式也使用) + "feedback": os.path.join(DATA_DIR, "prod", "feedback.db"), + "knowledge": os.path.join(DATA_DIR, "prod", "knowledge.db"), + + # 仅开发模式数据库 + "session": os.path.join(DATA_DIR, "dev", "session.db"), + "exam": os.path.join(DATA_DIR, "dev", "exam.db"), +} + +# 数据库名称映射(用于日志) +DB_NAMES = { + "feedback": "反馈系统数据库(反馈、FAQ、质量报告)", + "knowledge": "知识管理数据库(同步、大纲、版本)", + "session": "会话管理数据库(会话、消息、审计日志)", + "exam": "出题系统数据库(题库、试卷、批卷)", +} + +_init_lock = threading.Lock() +_initialized = False + + +# ==================== 连接管理 ==================== + +@contextmanager +def get_connection(db_name: str, row_factory: bool = True) -> Generator[sqlite3.Connection, None, None]: + """ + 获取数据库连接(上下文管理器) + + Args: + db_name: 数据库名称 ("feedback", "knowledge", "session", "exam") + row_factory: 是否启用行工厂(返回字典格式) + + Yields: + sqlite3.Connection: 数据库连接 + + Example: + with get_connection("feedback") as conn: + cursor = conn.cursor() + cursor.execute("SELECT * FROM feedbacks") + rows = cursor.fetchall() + """ + if db_name not in DB_PATHS: + raise ValueError(f"未知的数据库: {db_name},可用: {list(DB_PATHS.keys())}") + + db_path = DB_PATHS[db_name] + + # fresh clone 中 prod/dev 目录尚不存在;连接前统一创建父目录。 + os.makedirs(os.path.dirname(db_path), exist_ok=True) + + # 关键:check_same_thread=False 支持多worker + conn = sqlite3.connect(db_path, timeout=30, check_same_thread=False) + + # 启用 WAL 模式(提升并发性能) + conn.execute("PRAGMA journal_mode=WAL") + + # 设置繁忙超时(5秒) + conn.execute("PRAGMA busy_timeout=5000") + + # 启用外键约束 + conn.execute("PRAGMA foreign_keys=ON") + + # 设置行工厂(返回字典格式) + if row_factory: + conn.row_factory = sqlite3.Row + + try: + yield conn + conn.commit() + except Exception: + conn.rollback() + raise + finally: + conn.close() + + +def get_raw_connection(db_name: str) -> sqlite3.Connection: + """ + 获取原始数据库连接(不使用上下文管理器) + + 注意:调用者需要手动关闭连接 + + Args: + db_name: 数据库名称 + + Returns: + sqlite3.Connection: 数据库连接 + """ + if db_name not in DB_PATHS: + raise ValueError(f"未知的数据库: {db_name}") + + db_path = DB_PATHS[db_name] + os.makedirs(os.path.dirname(db_path), exist_ok=True) + conn = sqlite3.connect(db_path, timeout=30, check_same_thread=False) + conn.execute("PRAGMA journal_mode=WAL") + conn.execute("PRAGMA busy_timeout=5000") + conn.execute("PRAGMA foreign_keys=ON") + conn.row_factory = sqlite3.Row + return conn + + +# ==================== 数据库初始化 ==================== + +def init_databases(): + """ + 初始化所有数据库表结构 + + 按环境分离为: + - 生产模式数据库(开发模式也使用):feedback.db, knowledge.db + - 仅开发模式数据库:session.db, exam.db + """ + global _initialized + if _initialized: + return + + from config import IS_DEV + + with _init_lock: + if _initialized: + return + + # 始终初始化生产模式数据库 + _init_feedback_db() + _init_knowledge_db() + + # 仅开发模式初始化会话和出题数据库 + if IS_DEV: + _init_session_db() + _init_exam_db() + logger.info("所有数据库初始化完成(开发模式)") + else: + logger.info("生产模式数据库初始化完成") + + _initialized = True + + +def _init_feedback_db(): + """初始化 feedback.db(反馈系统数据库 - 生产模式也使用)""" + with get_connection("feedback", row_factory=False) as conn: + cursor = conn.cursor() + + # ========== 反馈表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS feedbacks ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + session_id TEXT NOT NULL, + query TEXT NOT NULL, + answer TEXT, + sources TEXT, + rating INTEGER NOT NULL, + reason TEXT, + user_id TEXT, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_feedback_session + ON feedbacks(session_id) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_feedback_rating + ON feedbacks(rating) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_feedback_created + ON feedbacks(created_at) + ''') + + # ========== FAQ 表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS faqs ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + question TEXT NOT NULL, + answer TEXT NOT NULL, + source_documents TEXT, + frequency INTEGER DEFAULT 1, + avg_rating REAL DEFAULT 0, + status TEXT DEFAULT 'draft', + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_faq_status + ON faqs(status) + ''') + + # ========== FAQ 问题变体表(Multi-Query Indexing)========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS faq_variants ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + faq_id INTEGER NOT NULL, + variant_question TEXT NOT NULL, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (faq_id) REFERENCES faqs(id) ON DELETE CASCADE + ) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_faq_variant_faq + ON faq_variants(faq_id) + ''') + + # ========== 质量报告表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS quality_reports ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + report_type TEXT NOT NULL, + start_date DATE NOT NULL, + end_date DATE NOT NULL, + total_queries INTEGER DEFAULT 0, + total_feedback INTEGER DEFAULT 0, + positive_count INTEGER DEFAULT 0, + negative_count INTEGER DEFAULT 0, + avg_rating REAL DEFAULT 0, + satisfaction_rate REAL DEFAULT 0, + high_freq_queries TEXT, + low_rating_queries TEXT, + improvement_suggestions TEXT, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + # ========== FAQ 建议表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS faq_suggestions ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + query TEXT NOT NULL, + answer TEXT, + frequency INTEGER DEFAULT 1, + avg_rating REAL DEFAULT 0, + status TEXT DEFAULT 'pending', + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_faq_suggestion_status + ON faq_suggestions(status) + ''') + + logger.info(f"初始化数据库: {DB_PATHS['feedback']} (feedback.db)") + + +def _init_session_db(): + """初始化 session.db(会话管理数据库 - 仅开发模式使用)""" + with get_connection("session", row_factory=False) as conn: + cursor = conn.cursor() + + # ========== 会话表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS sessions ( + session_id TEXT PRIMARY KEY, + user_id TEXT NOT NULL, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + last_active TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + metadata TEXT + ) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_sessions_user + ON sessions(user_id) + ''') + + # ========== 消息表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS messages ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + session_id TEXT NOT NULL, + role TEXT NOT NULL, + content TEXT NOT NULL, + metadata TEXT, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (session_id) REFERENCES sessions(session_id) ON DELETE CASCADE + ) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_messages_session + ON messages(session_id, created_at) + ''') + + # ========== 审计日志表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS audit_logs ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id TEXT NOT NULL, + username TEXT DEFAULT '', + action TEXT NOT NULL, + query TEXT, + result_summary TEXT, + sources TEXT, + role TEXT, + department TEXT, + ip_address TEXT, + duration_ms INTEGER, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_audit_user + ON audit_logs(user_id, created_at) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_audit_action + ON audit_logs(action, created_at) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_audit_created + ON audit_logs(created_at) + ''') + + logger.info(f"初始化数据库: {DB_PATHS['session']} (session.db)") + + +def _init_knowledge_db(): + """初始化 knowledge.db(知识管理数据库)""" + with get_connection("knowledge", row_factory=False) as conn: + cursor = conn.cursor() + + # ========== 文档哈希表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS document_hashes ( + document_id TEXT PRIMARY KEY, + document_name TEXT, + content_hash TEXT, + file_size INTEGER, + last_modified TIMESTAMP, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + # ========== 变更日志表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS change_logs ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + document_id TEXT, + document_name TEXT, + change_type TEXT, + old_hash TEXT, + new_hash TEXT, + change_time TIMESTAMP, + processed INTEGER DEFAULT 0, + error_message TEXT, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_change_logs_time + ON change_logs(change_time) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_change_logs_processed + ON change_logs(processed) + ''') + + # ========== 同步状态表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS sync_status ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + sync_type TEXT, + status TEXT, + start_time TIMESTAMP, + end_time TIMESTAMP, + documents_processed INTEGER, + documents_added INTEGER, + documents_modified INTEGER, + documents_deleted INTEGER, + error_message TEXT, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + # ========== 纲要缓存表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS outline_cache ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + document_id TEXT NOT NULL UNIQUE, + document_name TEXT, + total_pages INTEGER DEFAULT 0, + content_hash TEXT NOT NULL, + outline_json TEXT NOT NULL, + generated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_outline_doc + ON outline_cache(document_id) + ''') + + # ========== 文档向量缓存表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS document_vectors ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + document_id TEXT NOT NULL UNIQUE, + document_name TEXT, + vector_hash TEXT, + vector_json TEXT NOT NULL, + tags_json TEXT, + updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_vector_doc + ON document_vectors(document_id) + ''') + + # ========== 推荐缓存表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS recommendation_cache ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + document_id TEXT NOT NULL, + recommendations_json TEXT NOT NULL, + generated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + # ========== 文档版本表(统一结构) ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS document_versions ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + document_id TEXT NOT NULL, + collection TEXT, + version TEXT NOT NULL DEFAULT 'v1', + content_hash TEXT, + status TEXT NOT NULL DEFAULT 'active', + effective_date DATE, + expiry_date DATE, + deprecated_date DATETIME, + deprecated_reason TEXT, + deprecated_by TEXT, + change_summary TEXT, + changed_sections TEXT, + supersedes TEXT, + chunk_count INTEGER DEFAULT 0, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + created_by TEXT, + UNIQUE(document_id, collection, version) + ) + ''') + + # ========== 版本变更日志表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS version_change_logs ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + document_id TEXT NOT NULL, + collection TEXT, + old_version TEXT, + new_version TEXT, + old_status TEXT, + new_status TEXT, + change_type TEXT NOT NULL, + reason TEXT, + changed_by TEXT, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + # 添加性能优化索引 + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_document_versions_status + ON document_versions(document_id, collection, status) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_version_change_logs_document + ON version_change_logs(document_id, collection, created_at DESC) + ''') + + logger.info(f"初始化数据库: {DB_PATHS['knowledge']} (knowledge.db)") + + +def _init_exam_db(): + """初始化 exam.db(出题系统数据库)""" + with get_connection("exam", row_factory=False) as conn: + cursor = conn.cursor() + + # ========== 题目表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS questions ( + id TEXT PRIMARY KEY, + question_type TEXT NOT NULL, + content TEXT NOT NULL, + options TEXT, + correct_answer TEXT NOT NULL, + analysis TEXT, + knowledge_points TEXT, + difficulty INTEGER DEFAULT 3, + score INTEGER NOT NULL, + source_file TEXT NOT NULL, + source_collection TEXT NOT NULL, + source_snippet TEXT, + source_hash TEXT, + status TEXT DEFAULT 'approved', + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + created_by TEXT, + updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_questions_source + ON questions(source_file) + ''') + + # ========== 试卷表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS exams ( + id TEXT PRIMARY KEY, + name TEXT NOT NULL, + description TEXT, + total_score INTEGER NOT NULL, + total_count INTEGER NOT NULL, + duration INTEGER DEFAULT 60, + status TEXT DEFAULT 'published', + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + created_by TEXT + ) + ''') + + # ========== 试卷题目关联表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS exam_questions ( + exam_id TEXT NOT NULL, + question_id TEXT NOT NULL, + question_order INTEGER NOT NULL, + PRIMARY KEY (exam_id, question_id), + FOREIGN KEY (exam_id) REFERENCES exams(id) ON DELETE CASCADE, + FOREIGN KEY (question_id) REFERENCES questions(id) ON DELETE CASCADE + ) + ''') + + # ========== 学生答卷表(只存 student_id) ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS student_answers ( + id TEXT PRIMARY KEY, + exam_id TEXT NOT NULL, + student_id TEXT NOT NULL, + question_id TEXT NOT NULL, + question_type TEXT NOT NULL, + student_answer TEXT NOT NULL, + score REAL DEFAULT 0, + max_score INTEGER NOT NULL, + feedback TEXT, + score_details TEXT, + submitted_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + graded_at TIMESTAMP, + FOREIGN KEY (exam_id) REFERENCES exams(id) ON DELETE CASCADE, + FOREIGN KEY (question_id) REFERENCES questions(id) ON DELETE CASCADE + ) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_student_answers_exam + ON student_answers(exam_id, student_id) + ''') + + # ========== 批阅报告表(只存 student_id) ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS grade_reports ( + id TEXT PRIMARY KEY, + exam_id TEXT NOT NULL, + student_id TEXT NOT NULL, + total_score REAL NOT NULL, + max_score REAL NOT NULL, + score_rate REAL, + analysis TEXT, + graded_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (exam_id) REFERENCES exams(id) ON DELETE CASCADE + ) + ''') + + # ========== 题目-制度关联表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS question_document_links ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + question_id TEXT NOT NULL, + question_type TEXT NOT NULL, + exam_id TEXT NOT NULL, + document_id TEXT NOT NULL, + document_name TEXT, + chapter TEXT, + key_points TEXT, + relevance_score REAL DEFAULT 0, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_qdl_question + ON question_document_links(question_id) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_qdl_document + ON question_document_links(document_id) + ''') + + # ========== 知识点表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS knowledge_points ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL UNIQUE, + category TEXT, + description TEXT, + parent_id INTEGER, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (parent_id) REFERENCES knowledge_points(id) + ) + ''') + + # ========== 题目-知识点关联表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS question_knowledge_links ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + question_id TEXT NOT NULL, + question_type TEXT NOT NULL, + exam_id TEXT NOT NULL, + knowledge_point_id INTEGER NOT NULL, + knowledge_point_name TEXT, + weight REAL DEFAULT 1.0, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (knowledge_point_id) REFERENCES knowledge_points(id) + ) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_qkl_question + ON question_knowledge_links(question_id) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_qkl_knowledge + ON question_knowledge_links(knowledge_point_id) + ''') + + # ========== 题目状态表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS question_status ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + question_id TEXT NOT NULL UNIQUE, + question_type TEXT NOT NULL, + exam_id TEXT NOT NULL, + status TEXT DEFAULT 'approved', + affected_by TEXT, + affect_reason TEXT, + updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_qs_status + ON question_status(status) + ''') + + # ========== 整卷分析报告表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS exam_analysis_reports ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + report_id TEXT NOT NULL UNIQUE, + exam_id TEXT, + exam_name TEXT, + student_id TEXT, + total_score REAL, + max_score REAL, + score_rate REAL, + type_scores TEXT, + knowledge_analysis TEXT, + weak_points TEXT, + strong_points TEXT, + ai_comment TEXT, + study_suggestions TEXT, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + # ========== 新题建议表 ========== + cursor.execute(''' + CREATE TABLE IF NOT EXISTS question_suggestions ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + document_id TEXT NOT NULL, + suggestion TEXT, + status TEXT DEFAULT 'pending', + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + ''') + + logger.info(f"初始化数据库: {DB_PATHS['exam']} (exam.db)") + + +# ==================== 工具函数 ==================== + +def row_to_dict(row: sqlite3.Row) -> Dict: + """将 sqlite3.Row 转换为字典""" + if row is None: + return {} + return dict(row) + + +def rows_to_list(rows: List[sqlite3.Row]) -> List[Dict]: + """将 sqlite3.Row 列表转换为字典列表""" + return [row_to_dict(row) for row in rows] + + +def json_parse(value: str, default=None): + """安全解析 JSON 字符串""" + if not value: + return default + try: + return json.loads(value) + except (json.JSONDecodeError, TypeError): + return default + + +def json_stringify(value, ensure_ascii: bool = False) -> str: + """将对象转换为 JSON 字符串""" + if value is None: + return None + return json.dumps(value, ensure_ascii=ensure_ascii) + + +def get_db_stats() -> Dict: + """获取所有数据库的统计信息""" + stats = {} + + for db_name, db_path in DB_PATHS.items(): + if os.path.exists(db_path): + file_size = os.path.getsize(db_path) + + with get_connection(db_name, row_factory=False) as conn: + cursor = conn.cursor() + + # 获取所有表名 + cursor.execute( + "SELECT name FROM sqlite_master WHERE type='table' ORDER BY name" + ) + tables = [row[0] for row in cursor.fetchall()] + + # 统计每个表的行数 + table_counts = {} + for table in tables: + cursor.execute(f"SELECT COUNT(*) FROM {table}") + table_counts[table] = cursor.fetchone()[0] + + stats[db_name] = { + "path": db_path, + "name": DB_NAMES.get(db_name, db_name), + "file_size": file_size, + "tables": table_counts, + "total_rows": sum(table_counts.values()) + } + else: + stats[db_name] = { + "path": db_path, + "name": DB_NAMES.get(db_name, db_name), + "exists": False + } + + return stats + + +# ==================== 自动初始化 ==================== + +# 模块加载时不自动初始化,由调用方决定何时初始化 +# 这样可以避免在 import 时创建数据库文件 diff --git a/exam_pkg/generator.py b/exam_pkg/generator.py index eda54c6..dbbf2bf 100644 --- a/exam_pkg/generator.py +++ b/exam_pkg/generator.py @@ -1032,6 +1032,7 @@ def analyze_document_for_exam(chunks: List[Dict], max_total: int = None) -> Dict ) prompt += f"\n### {section_name[:30]}\n{content_preview[:300]}\n" + question_limit = min(total_knowledge_points * 2, max_total if max_total else 20) prompt += """ ## 要求 根据文档内容特点,决定: @@ -1055,10 +1056,10 @@ def analyze_document_for_exam(chunks: List[Dict], max_total: int = None) -> Dict 注意: - 不适合的题型数量设为 0 -- 所有数量之和不要超过 {min(total_knowledge_points * 2, 20)} +- 所有数量之和不要超过 %d - 必须返回合法 JSON,不要有其他内容 -请直接输出 JSON:""" +请直接输出 JSON:""" % question_limit try: response = generator._call_llm(prompt) diff --git a/exam_pkg/grader.py b/exam_pkg/grader.py index 18f0e1b..48aa994 100644 --- a/exam_pkg/grader.py +++ b/exam_pkg/grader.py @@ -219,6 +219,41 @@ def grade_fill_blank(answer: Dict) -> Dict: student_answers = answer.get('student_answer', []) max_score = answer.get('max_score', 4.0) + if not isinstance(student_answers, list): + logger.warning( + "填空题学生答案格式错误: 期望列表,实际为 %s", + type(student_answers).__name__, + ) + return { + "question_id": answer.get('question_id'), + "score": 0, + "max_score": max_score, + "grading_status": "failed", + "details": { + "error": f"学生答案格式错误,期望列表,实际为 {type(student_answers).__name__}" + }, + } + + for index, student_answer in enumerate(student_answers): + if not isinstance(student_answer, str): + logger.warning( + "填空题学生答案第 %s 项格式错误: 期望字符串,实际为 %s", + index + 1, + type(student_answer).__name__, + ) + return { + "question_id": answer.get('question_id'), + "score": 0, + "max_score": max_score, + "grading_status": "failed", + "details": { + "error": ( + f"填空题学生答案第 {index + 1} 项格式错误,期望字符串," + f"实际为 {type(student_answer).__name__}" + ) + }, + } + # 归一化答案格式(修复 LLM 生成的扁平数组问题) blank_count = question_content.get('data', {}).get('blank_count', 0) correct_answers = _normalize_fill_blank_answer(correct_answers, blank_count) diff --git a/exam_pkg/local_db.py b/exam_pkg/local_db.py index bcab583..0e0768f 100644 --- a/exam_pkg/local_db.py +++ b/exam_pkg/local_db.py @@ -170,8 +170,9 @@ class ExamLocalDB: VALUES (?, ?, ?, ?, ?, ?, ?, ?) ''', (exam_id, name, description, total_score, len(questions), duration, 'published', created_by)) - # 关联题目 - for order, qid in enumerate(question_ids): + # 仅关联实际存在的题目,避免无效 ID 触发外键约束。 + valid_question_ids = [q['id'] for q in questions] + for order, qid in enumerate(valid_question_ids): cursor.execute(''' INSERT INTO exam_questions (exam_id, question_id, question_order) VALUES (?, ?, ?) @@ -183,7 +184,7 @@ class ExamLocalDB: 'total_score': total_score, 'total_count': len(questions), 'duration': duration, - 'question_ids': question_ids + 'question_ids': valid_question_ids } def get_exam(self, exam_id: str) -> Optional[Dict]: diff --git a/knowledge/sync.py b/knowledge/sync.py index a669458..620b0ef 100644 --- a/knowledge/sync.py +++ b/knowledge/sync.py @@ -664,6 +664,17 @@ class KnowledgeSyncService: except Exception as e: logger.warning(f"递增缓存版本号失败: {e}") + # 语义缓存没有知识库版本字段,文档增删改后必须清空;否则回答缓存 + # 可能继续返回已过时的来源、引用或图片路径。 + try: + from core.semantic_cache import get_semantic_cache + semantic_cache = get_semantic_cache() + if semantic_cache: + semantic_cache.clear() + logger.debug(f"已清空语义缓存(文档变更触发): {kb_name}") + except Exception as e: + logger.warning(f"清空语义缓存失败: {e}") + return True except Exception as e: diff --git a/scripts/analyze_rag_performance.py b/scripts/analyze_rag_performance.py new file mode 100644 index 0000000..6cc3763 --- /dev/null +++ b/scripts/analyze_rag_performance.py @@ -0,0 +1,273 @@ +# -*- coding: utf-8 -*- +""" +RAG 性能分析脚本 + +分析 RAG 流程各阶段耗时,帮助定位性能瓶颈 +""" + +import sys +import os +import json +import time +import argparse +from datetime import datetime +from typing import Dict + +# 添加项目根目录到路径 +PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +sys.path.insert(0, PROJECT_ROOT) + + +def call_rag_stream_api( + question: str, + kb_name: str = "public_kb", + base_url: str = "http://localhost:5001", +) -> Dict: + """ + 调用 RAG 流式 API 并收集各阶段事件 + + Returns: + 包含各阶段耗时信息的字典 + """ + import requests + + url = f"{base_url.rstrip('/')}/rag" + headers = { + "Content-Type": "application/json", + "Authorization": "Bearer mock-token-admin" + } + data = { + "message": question, + "collections": [kb_name] if kb_name else [] + } + + events = [] + timing = { + 'total_duration_ms': 0, + 'search_time_ms': 0, + 'rerank_time_ms': 0, + 'llm_time_ms': 0, + 'stages': [] + } + + try: + start_time = time.time() + response = requests.post(url, json=data, headers=headers, stream=True, timeout=300) + response.raise_for_status() + + for line in response.iter_lines(): + if line: + line_str = line.decode('utf-8') + if line_str.startswith('data: '): + try: + event = json.loads(line_str[6:]) + event['_received_at'] = time.time() + events.append(event) + + # 分析 finish 事件中的耗时信息 + if event.get('type') == 'finish': + timing['total_duration_ms'] = event.get('duration_ms', 0) + if 'timing' in event: + timing['search_time_ms'] = event['timing'].get('total_search_ms', 0) + timing['rerank_time_ms'] = event['timing'].get('rerank_ms', 0) + timing['rerank_cached'] = event['timing'].get('rerank_cached', False) + except json.JSONDecodeError: + continue + + end_time = time.time() + timing['actual_elapsed_ms'] = int((end_time - start_time) * 1000) + + # 计算 LLM 生成时间 + if timing['total_duration_ms'] > 0 and timing['search_time_ms'] > 0: + timing['llm_time_ms'] = timing['total_duration_ms'] - timing['search_time_ms'] + + finish_event = next((event for event in events if event.get('type') == 'finish'), None) + error_event = next((event for event in events if event.get('type') == 'error'), None) + return { + 'success': finish_event is not None and error_event is None, + 'error': error_event.get('message', '') if error_event else ( + '' if finish_event else 'SSE 流结束但未收到 finish 事件' + ), + 'events': events, + 'timing': timing, + 'answer': finish_event.get('answer', '') if finish_event else '', + } + + except requests.exceptions.RequestException as e: + return { + 'success': False, + 'error': str(e), + 'events': events, + 'timing': timing + } + + +def analyze_performance(result: Dict, question: str) -> None: + """分析并打印性能数据""" + print("\n" + "=" * 80) + print(f"问题: {question}") + print("=" * 80) + + if not result['success']: + print(f"[X] 请求失败: {result.get('error', '未知错误')}") + return + + timing = result['timing'] + events = result['events'] + + # 打印各阶段事件时间线 + print("\n[事件时间线]") + print("-" * 80) + first_event_time = None + for event in events: + event_type = event.get('type', 'unknown') + received_at = event.get('_received_at', 0) + + if first_event_time is None: + first_event_time = received_at + relative_time = 0 + else: + relative_time = (received_at - first_event_time) * 1000 + + if event_type == 'finish': + print(f" {relative_time:>8.0f}ms | {event_type:20s} | 总耗时: {event.get('duration_ms', 0)}ms") + elif event_type == 'sources': + sources = event.get('sources', []) + print(f" {relative_time:>8.0f}ms | {event_type:20s} | 找到 {len(sources)} 个来源") + elif event_type == 'chunks_retrieved': + chunks = event.get('data', {}).get('chunks', []) + print(f" {relative_time:>8.0f}ms | {event_type:20s} | 召回 {len(chunks)} 个切片") + elif event_type == 'chunk': + # 流式输出,只显示第一个 + if not hasattr(analyze_performance, '_chunk_printed'): + print(f" {relative_time:>8.0f}ms | {event_type:20s} | 开始流式输出...") + analyze_performance._chunk_printed = True + else: + print(f" {relative_time:>8.0f}ms | {event_type:20s}") + + if hasattr(analyze_performance, '_chunk_printed'): + delattr(analyze_performance, '_chunk_printed') + + # 打印耗时统计 + print("\n[耗时统计]") + print("-" * 80) + total = timing['total_duration_ms'] + search = timing['search_time_ms'] + rerank = timing['rerank_time_ms'] + llm = timing['llm_time_ms'] + actual = timing.get('actual_elapsed_ms', total) + + if total > 0: + print(f" 总耗时 (API报告): {total:>8}ms ({total/1000:.2f}s)") + print(f" 实际耗时 (本地测量): {actual:>8}ms ({actual/1000:.2f}s)") + print() + print(f" 检索阶段: {search:>8}ms ({search/total*100:.1f}%)") + if rerank > 0: + cached_flag = " [缓存]" if timing.get('rerank_cached') else "" + print(f" 重排序阶段: {rerank:>8}ms ({rerank/total*100:.1f}%){cached_flag}") + print(f" LLM生成阶段: {llm:>8}ms ({llm/total*100:.1f}%)") + + # 性能诊断 + print("\n[性能诊断]") + print("-" * 80) + if total > 10000: + print(" [!] 总耗时超过 10 秒,需要优化") + + if search > 3000: + print(f" [!] 检索耗时过长 ({search}ms),可能原因:") + print(" - 向量库数据量过大") + print(" - 未命中查询缓存") + print(" - BM25 索引加载慢") + + if rerank > 2000 and not timing.get('rerank_cached'): + print(f" [!] 重排序耗时过长 ({rerank}ms),可能原因:") + print(" - Rerank 模型计算量大") + print(" - 候选切片数量过多") + + if llm > 5000: + print(f" [!] LLM生成耗时过长 ({llm}ms),可能原因:") + print(" - 模型生成速度慢") + print(" - 输出 token 数量多") + print(" - 网络延迟") + print(" 建议: 检查 MiMo thinking 是否关闭、上下文长度及 max_tokens") + + if total < 5000: + print(" [OK] 性能良好") + else: + print(" [!] 未能获取耗时数据") + + +def run_performance_tests( + base_url: str = "http://localhost:5001", + kb_name: str = "public_kb", + question: str = None, +): + """执行性能测试""" + # 测试问题(不同类型) + test_questions = [ + {"question": "智启科技成立于哪一年?", "type": "精确匹配"}, + {"question": "公司的愿景是什么?", "type": "语义理解"}, + {"question": "公司有哪些分公司,分别在哪些城市?", "type": "跨文档关联"}, + {"question": "一个入职3年的员工,累计病假1个月,能拿到多少病假工资?", "type": "复杂推理"}, + {"question": "P4级工程师的年薪范围是多少?", "type": "表格数据"}, + ] + if question: + test_questions = [{"question": question, "type": "自定义问题"}] + + print("=" * 80) + print("RAG Performance Analysis") + print(f"Time: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}") + print("=" * 80) + + results = [] + for i, test in enumerate(test_questions): + print(f"\n\n[{i+1}/{len(test_questions)}] 类型: {test['type']}") + result = call_rag_stream_api(test['question'], kb_name=kb_name, base_url=base_url) + analyze_performance(result, test['question']) + results.append({ + 'question': test['question'], + 'type': test['type'], + 'success': result['success'], + 'timing': result['timing'] + }) + + # 汇总统计 + print("\n\n" + "=" * 80) + print("[Performance Summary]") + print("=" * 80) + + successful = [r for r in results if r['success']] + if successful: + total_times = [r['timing']['total_duration_ms'] for r in successful] + search_times = [r['timing']['search_time_ms'] for r in successful] + llm_times = [r['timing']['llm_time_ms'] for r in successful] + + print(f"\n成功请求数: {len(successful)}/{len(results)}") + print(f"\n平均总耗时: {sum(total_times)/len(total_times):.0f}ms") + print(f"平均检索耗时: {sum(search_times)/len(search_times):.0f}ms") + print(f"平均LLM耗时: {sum(llm_times)/len(llm_times):.0f}ms") + print(f"\n最快总耗时: {min(total_times)}ms") + print(f"最慢总耗时: {max(total_times)}ms") + + # 保存详细报告 + report_dir = os.path.join(PROJECT_ROOT, ".data", "performance") + os.makedirs(report_dir, exist_ok=True) + report_file = os.path.join( + report_dir, + f"performance_report_{datetime.now().strftime('%Y%m%d_%H%M%S')}.json", + ) + with open(report_file, 'w', encoding='utf-8') as f: + json.dump({ + 'test_time': datetime.now().isoformat(), + 'results': results + }, f, ensure_ascii=False, indent=2) + print(f"\n详细报告已保存: {report_file}") + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description="RAG SSE 链路耗时分析") + parser.add_argument("--base-url", default="http://localhost:5001") + parser.add_argument("--kb", default="public_kb", help="目标知识库名称") + parser.add_argument("--question", help="只测试一个自定义问题") + args = parser.parse_args() + run_performance_tests(args.base_url, args.kb, args.question) diff --git a/scripts/download_models.py b/scripts/download_models.py new file mode 100644 index 0000000..dff6eb6 --- /dev/null +++ b/scripts/download_models.py @@ -0,0 +1,82 @@ +""" +下载 BGE embedding 和 reranker 模型到本地 models/ 目录 + +首次配置时运行: + python scripts/download_models.py +""" + +import os +import sys + +PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +MODELS_DIR = os.path.join(PROJECT_ROOT, "models") +EMBEDDING_MODEL_PATH = os.path.join(MODELS_DIR, "bge-base-zh-v1.5") +RERANK_MODEL_PATH = os.path.join(MODELS_DIR, "bge-reranker-base") + +os.makedirs(MODELS_DIR, exist_ok=True) + + +def download_embedding_model(): + """下载 BGE 中文向量模型""" + if os.path.exists(os.path.join(EMBEDDING_MODEL_PATH, "config.json")): + print(f"[跳过] Embedding 模型已存在: {EMBEDDING_MODEL_PATH}") + return True + + print(f"[下载] BGE 中文向量模型 (bge-base-zh-v1.5)...") + print(f" 目标路径: {EMBEDDING_MODEL_PATH}") + try: + from sentence_transformers import SentenceTransformer + model = SentenceTransformer("BAAI/bge-base-zh-v1.5") + model.save(EMBEDDING_MODEL_PATH) + print(f"[完成] Embedding 模型下载成功!") + return True + except Exception as e: + print(f"[失败] Embedding 模型下载失败: {e}") + return False + + +def download_reranker_model(): + """下载 BGE reranker 模型""" + if os.path.exists(os.path.join(RERANK_MODEL_PATH, "config.json")): + print(f"[跳过] Reranker 模型已存在: {RERANK_MODEL_PATH}") + return True + + print(f"[下载] BGE Reranker 模型 (bge-reranker-base)...") + print(f" 目标路径: {RERANK_MODEL_PATH}") + try: + from transformers import AutoModelForSequenceClassification, AutoTokenizer + os.makedirs(RERANK_MODEL_PATH, exist_ok=True) + model = AutoModelForSequenceClassification.from_pretrained("BAAI/bge-reranker-base") + tokenizer = AutoTokenizer.from_pretrained("BAAI/bge-reranker-base") + model.save_pretrained(RERANK_MODEL_PATH) + tokenizer.save_pretrained(RERANK_MODEL_PATH) + print(f"[完成] Reranker 模型下载成功!") + return True + except Exception as e: + print(f"[失败] Reranker 模型下载失败: {e}") + print(f" (项目启动时会自动尝试下载,或可稍后手动重试)") + return False + + +if __name__ == "__main__": + print("=" * 60) + print("RAG 项目模型下载脚本") + print("=" * 60) + + ok1 = download_embedding_model() + print() + ok2 = download_reranker_model() + + print() + print("=" * 60) + if ok1: + print("Embedding 模型: OK (项目可以启动)") + else: + print("Embedding 模型: FAILED (项目无法启动,请检查网络)") + + if ok2: + print("Reranker 模型: OK") + else: + print("Reranker 模型: FAILED (不影响启动,首次使用 rerank 时自动下载)") + + print("=" * 60) diff --git a/tests/test_data_db.py b/tests/test_data_db.py new file mode 100644 index 0000000..c6b0337 --- /dev/null +++ b/tests/test_data_db.py @@ -0,0 +1,101 @@ +"""data.db fresh-clone and transaction regression tests.""" + +import os +import shutil +import unittest +import uuid + +import config +import data.db as db +from exam_pkg.local_db import ExamLocalDB + + +class DataDbTests(unittest.TestCase): + def setUp(self): + project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + test_root = os.path.join(project_root, ".data", "test-runs") + os.makedirs(test_root, exist_ok=True) + self.temp_dir = os.path.join(test_root, f"run-{uuid.uuid4().hex}") + os.makedirs(self.temp_dir) + self.original_paths = db.DB_PATHS + self.original_is_dev = config.IS_DEV + db.DB_PATHS = { + "feedback": os.path.join(self.temp_dir, "prod", "feedback.db"), + "knowledge": os.path.join(self.temp_dir, "prod", "knowledge.db"), + "session": os.path.join(self.temp_dir, "dev", "session.db"), + "exam": os.path.join(self.temp_dir, "dev", "exam.db"), + } + config.IS_DEV = True + db._initialized = False + + def tearDown(self): + db._initialized = False + db.DB_PATHS = self.original_paths + config.IS_DEV = self.original_is_dev + test_root = os.path.realpath(os.path.join(os.path.dirname(self.temp_dir))) + target = os.path.realpath(self.temp_dir) + if os.path.commonpath([test_root, target]) != test_root: + raise RuntimeError(f"拒绝清理测试目录之外的路径: {target}") + shutil.rmtree(target, ignore_errors=True) + + def _table_names(self, db_name): + with db.get_connection(db_name) as conn: + rows = conn.execute( + "SELECT name FROM sqlite_master WHERE type = 'table'" + ).fetchall() + return {row[0] for row in rows} + + def test_fresh_clone_initializes_all_dev_databases(self): + db.init_databases() + + for path in db.DB_PATHS.values(): + self.assertTrue(os.path.isfile(path), path) + + self.assertTrue({"feedbacks", "faqs", "faq_suggestions"} <= self._table_names("feedback")) + self.assertTrue({"document_hashes", "document_versions", "sync_status"} <= self._table_names("knowledge")) + self.assertTrue({"sessions", "messages", "audit_logs"} <= self._table_names("session")) + self.assertTrue({"questions", "exams", "exam_questions"} <= self._table_names("exam")) + + with db.get_connection("knowledge") as conn: + columns = { + row[1] for row in conn.execute("PRAGMA table_info(document_versions)").fetchall() + } + self.assertTrue( + {"content_hash", "expiry_date", "deprecated_date", "deprecated_by", "changed_sections"} + <= columns + ) + + # Initialization must remain safe when several services call it. + db.init_databases() + + def test_connection_commits_and_rolls_back(self): + with db.get_connection("session") as conn: + conn.execute("CREATE TABLE tx_test (value TEXT)") + conn.execute("INSERT INTO tx_test VALUES ('committed')") + + with self.assertRaisesRegex(RuntimeError, "rollback"): + with db.get_connection("session") as conn: + conn.execute("INSERT INTO tx_test VALUES ('rolled-back')") + raise RuntimeError("rollback") + + with db.get_connection("session") as conn: + values = [row[0] for row in conn.execute("SELECT value FROM tx_test").fetchall()] + self.assertEqual(["committed"], values) + + def test_exam_creation_ignores_missing_question_ids(self): + exam_db = ExamLocalDB() + question_id = exam_db.add_question( + {"content": "示例题", "answer": "答案", "score": 2}, + source_file="example.txt", + source_collection="public", + ) + + exam = exam_db.create_exam("回归测试", [question_id, "missing-question"]) + + self.assertEqual(1, exam["total_count"]) + self.assertEqual([question_id], exam["question_ids"]) + self.assertEqual(1, len(exam_db.get_exam(exam["id"])["questions"])) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_shared_regressions.py b/tests/test_shared_regressions.py new file mode 100644 index 0000000..82e732d --- /dev/null +++ b/tests/test_shared_regressions.py @@ -0,0 +1,197 @@ +"""Regression tests for branch-shared RAG behavior.""" + +import ast +import json +import os +import pathlib +import unittest +from datetime import datetime +from unittest.mock import patch + +from core.cache import RAGCacheManager +from core.intent_analyzer import IntentAnalyzer +from exam_pkg.generator import QuestionGenerator, analyze_document_for_exam +from exam_pkg.grader import grade_fill_blank +from knowledge.sync import ChangeType, DocumentChange, KnowledgeSyncService + + +class CacheInvalidationTests(unittest.TestCase): + def test_kb_version_change_clears_rerank_scores(self): + cache = RAGCacheManager() + cache.set_rerank_scores("问题", ["doc-1"], [0.91]) + self.assertEqual({"doc-1": 0.91}, cache.get_rerank_scores("问题", ["doc-1"])) + + cache.increment_kb_version("public") + + self.assertIsNone(cache.get_rerank_scores("问题", ["doc-1"])) + + +class IntentCacheTests(unittest.TestCase): + def test_unknown_semantic_cache_type_is_not_used_as_intent(self): + class FakeEmbeddingModel: + def encode(self, _text): + return [1.0, 0.0] + + class FakeCache: + def __init__(self): + self.saved = None + + def get(self, _embedding): + return { + "cache_type": "future_cache_type", + "rewritten_query": "错误的缓存结果", + } + + def set(self, _embedding, value): + self.saved = value + + def get_stats(self): + return {} + + analyzer = IntentAnalyzer(model="mimo-v2.5") + analyzer._embedding_model = FakeEmbeddingModel() + analyzer._cache = FakeCache() + analyzer._client = object() + + llm_result = json.dumps( + { + "rewritten_query": "来自 LLM 的新结果", + "use_context": False, + "need_retrieval": True, + "reason": "新问题", + "sub_queries": ["来自 LLM 的新结果"], + "intent": "factual", + }, + ensure_ascii=False, + ) + with patch("core.llm_utils.call_llm", return_value=llm_result): + result = analyzer.analyze("当前问题", []) + + self.assertEqual("来自 LLM 的新结果", result.rewritten_query) + self.assertEqual("intent_analysis", analyzer._cache.saved["cache_type"]) + + +class ExamRegressionTests(unittest.TestCase): + def test_analysis_prompt_respects_requested_max_total(self): + captured = {} + + def fake_call(_self, prompt): + captured["prompt"] = prompt + return json.dumps( + { + "total_knowledge_points": 2, + "suitable_types": ["single_choice"], + "question_types": { + "single_choice": 3, + "multiple_choice": 0, + "true_false": 0, + "fill_blank": 0, + "subjective": 0, + }, + "reason": "测试", + }, + ensure_ascii=False, + ) + + chunks = [{"content": "知识点" * 80, "section": "第一章"}] + with patch.object(QuestionGenerator, "_call_llm", fake_call): + result = analyze_document_for_exam(chunks, max_total=3) + + self.assertIn("所有数量之和不要超过 3", captured["prompt"]) + self.assertLessEqual(sum(result["question_types"].values()), 3) + + def test_fill_blank_rejects_non_list_student_answer(self): + result = grade_fill_blank( + { + "question_id": "q1", + "content": {"answer": [["正确答案"]], "data": {"blank_count": 1}}, + "student_answer": "正确答案", + "max_score": 2, + } + ) + self.assertEqual("failed", result["grading_status"]) + self.assertEqual(0, result["score"]) + + def test_fill_blank_rejects_non_string_items(self): + result = grade_fill_blank( + { + "question_id": "q1", + "content": {"answer": [["正确答案"]], "data": {"blank_count": 1}}, + "student_answer": [{"text": "正确答案"}], + "max_score": 2, + } + ) + self.assertEqual("failed", result["grading_status"]) + self.assertIn("第 1 项", result["details"]["error"]) + + +class SyncInvalidationTests(unittest.TestCase): + def test_document_delete_clears_semantic_cache(self): + class FakeDb: + def delete_document_hash(self, _document_id): + return None + + class FakeKbManager: + def delete_document(self, _kb_name, _document_name): + return 1 + + class FakeSemanticCache: + def __init__(self): + self.cleared = False + + def clear(self): + self.cleared = True + + service = object.__new__(KnowledgeSyncService) + service.documents_path = "unused" + service.db = FakeDb() + semantic_cache = FakeSemanticCache() + change = DocumentChange( + document_id="public/example.pdf", + document_name="example.pdf", + change_type=ChangeType.DELETED, + old_hash="old", + new_hash=None, + change_time=datetime.now(), + ) + + with ( + patch("knowledge.manager.get_kb_manager", return_value=FakeKbManager()), + patch("knowledge.sync.CACHE_AVAILABLE", False), + patch("core.semantic_cache.get_semantic_cache", return_value=semantic_cache), + ): + self.assertTrue(service.process_change(change)) + + self.assertTrue(semantic_cache.cleared) + + +class ConfigTemplateTests(unittest.TestCase): + def test_config_example_defines_all_direct_imports(self): + root = pathlib.Path(__file__).resolve().parents[1] + template_tree = ast.parse((root / "config.example.py").read_text(encoding="utf-8")) + defined = { + node.name + for node in template_tree.body + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) + } + for node in template_tree.body: + if isinstance(node, ast.Assign): + defined.update(target.id for target in node.targets if isinstance(target, ast.Name)) + elif isinstance(node, ast.AnnAssign) and isinstance(node.target, ast.Name): + defined.add(node.target.id) + + imported = set() + excluded = {"venv", ".git", ".data", ".recovery", "__pycache__"} + for path in root.rglob("*.py"): + if any(part in excluded for part in path.parts): + continue + tree = ast.parse(path.read_text(encoding="utf-8-sig")) + for node in ast.walk(tree): + if isinstance(node, ast.ImportFrom) and node.module == "config": + imported.update(alias.name for alias in node.names) + + self.assertEqual([], sorted(imported - defined)) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_upload_dedup.py b/tests/test_upload_dedup.py index f04322c..0f34ca8 100644 --- a/tests/test_upload_dedup.py +++ b/tests/test_upload_dedup.py @@ -9,12 +9,30 @@ Phase 3 验证测试:重复上传旧切片残留修复 import os import sys -import tempfile import shutil +import uuid from pathlib import Path # 添加项目根目录到 Python 路径 -sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +TEST_TEMP_ROOT = os.path.join(PROJECT_ROOT, ".data", "test-runs") +sys.path.insert(0, PROJECT_ROOT) + + +def _make_test_dir(): + """在工作区内创建测试目录,兼容 Codex 的 Windows 文件权限。""" + os.makedirs(TEST_TEMP_ROOT, exist_ok=True) + path = os.path.join(TEST_TEMP_ROOT, f"upload-{uuid.uuid4().hex}") + os.makedirs(path) + return path + + +def _cleanup_test_dir(path): + root = os.path.realpath(TEST_TEMP_ROOT) + target = os.path.realpath(path) + if os.path.commonpath([root, target]) != root: + raise RuntimeError(f"拒绝清理测试目录之外的路径: {target}") + shutil.rmtree(target, ignore_errors=True) def test_add_file_to_kb_dedup(): @@ -143,7 +161,7 @@ def test_upload_document_overwrite(): print("\n=== 测试 2:upload_document 同名文件覆盖 ===") # 创建临时目录 - temp_dir = tempfile.mkdtemp() + temp_dir = _make_test_dir() try: target_dir = os.path.join(temp_dir, 'public_kb') os.makedirs(target_dir, exist_ok=True) @@ -184,14 +202,14 @@ def test_upload_document_overwrite(): print(f" [PASS] replaced={replaced}") finally: - shutil.rmtree(temp_dir) + _cleanup_test_dir(temp_dir) def test_batch_upload_overwrite(): """测试批量上传中的同名文件覆盖逻辑""" print("\n=== 测试 3:批量上传同名文件覆盖 ===") - temp_dir = tempfile.mkdtemp() + temp_dir = _make_test_dir() try: target_dir = os.path.join(temp_dir, 'public_kb') os.makedirs(target_dir, exist_ok=True) @@ -240,14 +258,14 @@ def test_batch_upload_overwrite(): print(f" [PASS] 文件内容已更新") finally: - shutil.rmtree(temp_dir) + _cleanup_test_dir(temp_dir) def test_docstore_cleanup(): """测试 DocStore 文件清理逻辑""" print("\n=== 测试 4:DocStore 文件清理 ===") - temp_dir = tempfile.mkdtemp() + temp_dir = _make_test_dir() try: docstore_dir = Path(temp_dir) / 'docstore' docstore_dir.mkdir() @@ -289,7 +307,7 @@ def test_docstore_cleanup(): print(f" [PASS] 剩余 {len(remaining)} 个无关文件未被清理") finally: - shutil.rmtree(temp_dir) + _cleanup_test_dir(temp_dir) def test_sync_hash_cleanup():