init: RAG 知识库服务初始提交
- 后端 API(Flask + Gunicorn) - RAG 引擎(混合检索 + 云端 Reranker + 引用溯源) - 文档解析(MinerU + 多格式支持) - Docker 生产部署配置 - 排除前端项目、敏感配置、模型文件
This commit is contained in:
40
knowledge/__init__.py
Normal file
40
knowledge/__init__.py
Normal file
@@ -0,0 +1,40 @@
|
||||
"""
|
||||
知识库管理模块
|
||||
|
||||
包含:
|
||||
- manager: 多向量库管理器 (KnowledgeBaseManager)
|
||||
- router: 知识库路由器 (KnowledgeBaseRouter)
|
||||
- sync: 知识库同步服务 (KnowledgeSyncService)
|
||||
- base: 基础类和常量
|
||||
- collection: 向量库管理 Mixin
|
||||
- document: 文档管理 Mixin
|
||||
- chunk: 切片管理 Mixin
|
||||
- index: BM25 索引管理 Mixin
|
||||
- search: 检索功能 Mixin
|
||||
- processing: 图片/表格处理 Mixin
|
||||
- permission: 权限控制 Mixin
|
||||
"""
|
||||
|
||||
from .manager import KnowledgeBaseManager
|
||||
from .router import KnowledgeBaseRouter
|
||||
from .base import (
|
||||
BM25Index,
|
||||
CollectionInfo,
|
||||
SearchResult,
|
||||
PUBLIC_KB_NAME,
|
||||
)
|
||||
|
||||
try:
|
||||
from .sync import KnowledgeSyncService
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
__all__ = [
|
||||
'KnowledgeBaseManager',
|
||||
'KnowledgeBaseRouter',
|
||||
'KnowledgeSyncService',
|
||||
'BM25Index',
|
||||
'CollectionInfo',
|
||||
'SearchResult',
|
||||
'PUBLIC_KB_NAME',
|
||||
]
|
||||
469
knowledge/base.py
Normal file
469
knowledge/base.py
Normal file
@@ -0,0 +1,469 @@
|
||||
"""
|
||||
知识库管理器 - 基础模块
|
||||
|
||||
包含:
|
||||
- 配置常量
|
||||
- 数据类定义
|
||||
- 辅助函数
|
||||
- BM25Index 类
|
||||
"""
|
||||
|
||||
import os
|
||||
import json
|
||||
import pickle
|
||||
import logging
|
||||
from typing import List, Dict, Optional, Tuple
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
from rank_bm25 import BM25Okapi
|
||||
import jieba
|
||||
|
||||
# 设置日志
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ==================== 配置常量 ====================
|
||||
|
||||
# 向量存储基础路径(位于 knowledge/vector_store/)
|
||||
VECTOR_STORE_BASE_PATH = os.path.join(
|
||||
os.path.dirname(os.path.abspath(__file__)),
|
||||
"vector_store"
|
||||
)
|
||||
|
||||
# 向量库基础路径(ChromaDB 数据存储)
|
||||
CHROMA_DB_BASE_PATH = os.path.join(VECTOR_STORE_BASE_PATH, "chroma")
|
||||
|
||||
# BM25 索引基础路径
|
||||
BM25_INDEX_BASE_PATH = os.path.join(VECTOR_STORE_BASE_PATH, "bm25")
|
||||
|
||||
# 向量库元数据文件
|
||||
KB_METADATA_FILE = "kb_metadata.json"
|
||||
|
||||
# 预定义的公开知识库名称
|
||||
PUBLIC_KB_NAME = "public_kb"
|
||||
|
||||
# 默认部门列表
|
||||
DEFAULT_DEPARTMENTS = ["finance", "hr", "tech", "operation", "marketing"]
|
||||
|
||||
# 部门名称映射
|
||||
DEPARTMENT_NAME_MAP = {
|
||||
"财务部": "finance", "财务": "finance",
|
||||
"人事部": "hr", "人事": "hr", "人力资源部": "hr", "人力资源": "hr",
|
||||
"技术部": "tech", "技术": "tech", "研发部": "tech", "研发": "tech",
|
||||
"运营部": "operation", "运营": "operation",
|
||||
"市场部": "marketing", "市场": "marketing",
|
||||
"法务部": "legal", "法务": "legal",
|
||||
"行政部": "admin", "行政": "admin",
|
||||
"finance": "finance", "hr": "hr", "tech": "tech",
|
||||
"operation": "operation", "marketing": "marketing", "legal": "legal", "admin": "admin",
|
||||
}
|
||||
|
||||
|
||||
# ==================== 数据结构 ====================
|
||||
|
||||
@dataclass
|
||||
class CollectionInfo:
|
||||
"""向量库信息"""
|
||||
name: str
|
||||
display_name: str
|
||||
document_count: int = 0
|
||||
created_at: str = ""
|
||||
department: str = ""
|
||||
description: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class SearchResult:
|
||||
"""检索结果"""
|
||||
ids: List[str]
|
||||
documents: List[str]
|
||||
metadatas: List[dict]
|
||||
distances: List[float]
|
||||
collection_name: str = ""
|
||||
|
||||
|
||||
# ==================== 辅助函数 ====================
|
||||
|
||||
def _get_doc_type(filename: str) -> str:
|
||||
"""
|
||||
根据文件扩展名判断文档类型
|
||||
|
||||
Args:
|
||||
filename: 文件名,包含扩展名
|
||||
|
||||
Returns:
|
||||
文档类型字符串: 'pdf' | 'word' | 'excel' | 'ppt' | 'other'
|
||||
|
||||
Example:
|
||||
>>> _get_doc_type("report.pdf")
|
||||
'pdf'
|
||||
>>> _get_doc_type("data.xlsx")
|
||||
'excel'
|
||||
"""
|
||||
ext = Path(filename).suffix.lower()
|
||||
type_map = {
|
||||
'.pdf': 'pdf',
|
||||
'.docx': 'word', '.doc': 'word',
|
||||
'.xlsx': 'excel', '.xls': 'excel',
|
||||
'.pptx': 'ppt', '.ppt': 'ppt',
|
||||
}
|
||||
return type_map.get(ext, 'other')
|
||||
|
||||
|
||||
def _extract_figure_number(caption: str, section: str = '') -> str:
|
||||
"""
|
||||
从 caption 或 section 中提取图号(增强版)
|
||||
|
||||
支持格式:
|
||||
- 图2.4, 图2-4
|
||||
- Fig.2.4, Fig 2.4, Figure 2.4
|
||||
- (图2)
|
||||
|
||||
Args:
|
||||
caption: 图片标题/说明文字
|
||||
section: 章节信息(可选)
|
||||
|
||||
Returns:
|
||||
图号字符串,如 "2.4";未找到返回空字符串
|
||||
|
||||
Example:
|
||||
>>> _extract_figure_number("图2.4 系统架构图")
|
||||
'2.4'
|
||||
>>> _extract_figure_number("参见Figure 3.1所示")
|
||||
'3.1'
|
||||
"""
|
||||
import re
|
||||
text = f"{caption} {section}"
|
||||
patterns = [
|
||||
r'图\s*(\d+[\.\-]\d+)',
|
||||
r'Fig\.?\s*(\d+[\.\-]\d+)',
|
||||
r'Figure\s*(\d+[\.\-]\d+)',
|
||||
r'[((]\s*图\s*(\d+)\s*[))]',
|
||||
]
|
||||
for pattern in patterns:
|
||||
match = re.search(pattern, text, re.IGNORECASE)
|
||||
if match:
|
||||
return match.group(1).replace('-', '.')
|
||||
return ""
|
||||
|
||||
|
||||
def normalize_department_name(department: str) -> str:
|
||||
"""
|
||||
将部门名称标准化为英文标识
|
||||
|
||||
支持中文部门名(如"财务部")和英文标识(如"finance"),
|
||||
返回符合 ChromaDB 命名规范的英文标识。
|
||||
|
||||
Args:
|
||||
department: 原始部门名称(中文或英文)
|
||||
|
||||
Returns:
|
||||
标准化的英文标识;无法识别时返回空字符串
|
||||
|
||||
Example:
|
||||
>>> normalize_department_name("财务部")
|
||||
'finance'
|
||||
>>> normalize_department_name("tech")
|
||||
'tech'
|
||||
"""
|
||||
if not department:
|
||||
return ""
|
||||
if department in DEPARTMENT_NAME_MAP:
|
||||
return DEPARTMENT_NAME_MAP[department]
|
||||
if department.replace("_", "").replace("-", "").isalnum() and department.isascii():
|
||||
return department.lower()
|
||||
logger.warning(f"无法识别的部门名称: {department}")
|
||||
return ""
|
||||
|
||||
|
||||
def _extract_section(section_path: str, max_levels: int = 3) -> str:
|
||||
"""
|
||||
动态截断章节路径,保留核心+末尾层级
|
||||
|
||||
用于在向量检索时提供简洁的章节上下文,
|
||||
避免过长的章节路径影响语义匹配。
|
||||
|
||||
Args:
|
||||
section_path: 完整章节路径,用 '>' 分隔
|
||||
max_levels: 最多保留的层级数(默认3级)
|
||||
|
||||
Returns:
|
||||
截断后的章节路径
|
||||
|
||||
Example:
|
||||
>>> _extract_section("第一章 > 1.1 概述 > 1.1.1 背景 > 1.1.1.1 详细说明")
|
||||
'1.1 概述 > 1.1.1 背景 > 1.1.1.1 详细说明'
|
||||
"""
|
||||
parts = [p.strip() for p in section_path.split('>') if p.strip()]
|
||||
if len(parts) <= max_levels:
|
||||
return ' > '.join(parts)
|
||||
return ' > '.join(parts[-max_levels:])
|
||||
|
||||
|
||||
def _build_semantic_content_for_text(chunk, page_info: dict, doc_type: str) -> str:
|
||||
"""
|
||||
构建语义增强内容(文本类型)
|
||||
|
||||
将原始文本切片转换为适合向量检索的语义增强格式,
|
||||
包含标题、章节上下文和正文内容。
|
||||
|
||||
PDF 和 Word 差异化处理:
|
||||
- PDF: 有 text_level、bbox,标题识别准确
|
||||
- Word: 无 text_level、bbox,依赖启发式识别
|
||||
|
||||
Args:
|
||||
chunk: 文档切片对象,含 title、content、text_level 属性
|
||||
page_info: 页面信息字典,含 section_path、section 等
|
||||
doc_type: 文档类型 ('pdf' | 'word' | 'excel' | 'ppt')
|
||||
|
||||
Returns:
|
||||
语义增强后的内容字符串,格式为:
|
||||
标题
|
||||
主题:章节路径
|
||||
正文内容
|
||||
"""
|
||||
parts = []
|
||||
title = getattr(chunk, 'title', '') or ''
|
||||
if isinstance(title, list):
|
||||
title = ' '.join(str(t) for t in title if t) or ''
|
||||
if not isinstance(title, str):
|
||||
title = ''
|
||||
text_level = getattr(chunk, 'text_level', 0)
|
||||
|
||||
if title and title.strip() and text_level > 0:
|
||||
parts.append(title.strip())
|
||||
|
||||
section = page_info.get('section_path', '') or page_info.get('section', '')
|
||||
if isinstance(section, list):
|
||||
section = ' > '.join(str(s) for s in section if s) or ''
|
||||
if not isinstance(section, str):
|
||||
section = ''
|
||||
if section and section.strip():
|
||||
section = _extract_section(section, max_levels=3)
|
||||
parts.append(f"主题:{section}")
|
||||
|
||||
content = chunk.content if hasattr(chunk, 'content') else page_info.get('text', '')
|
||||
if isinstance(content, list):
|
||||
content = '\n'.join(str(item) for item in content)
|
||||
parts.append(content)
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _build_semantic_content_for_table(table_md: str, page_info: dict, chunk, doc_type: str) -> str:
|
||||
"""
|
||||
构建语义增强内容(表格类型)
|
||||
|
||||
将 Markdown 表格转换为包含语义摘要的增强格式,
|
||||
提取表头、字段、行数和示例数据,提升向量检索命中率。
|
||||
|
||||
Args:
|
||||
table_md: 表格的 Markdown 内容
|
||||
page_info: 页面信息字典,含 section_path 等
|
||||
chunk: 文档切片对象,含 title 属性
|
||||
doc_type: 文档类型 ('pdf' | 'word' | 'excel' | 'ppt')
|
||||
|
||||
Returns:
|
||||
语义增强后的内容字符串,包含:
|
||||
主题:章节路径
|
||||
表格:标题
|
||||
字段:表头列表
|
||||
描述:行数和字段概要
|
||||
示例:首行数据示例
|
||||
"""
|
||||
parts = []
|
||||
section = page_info.get('section_path', '') or page_info.get('section', '')
|
||||
if isinstance(section, list):
|
||||
section = ' > '.join(str(s) for s in section if s) or ''
|
||||
if not isinstance(section, str):
|
||||
section = ''
|
||||
if section and section.strip():
|
||||
section = _extract_section(section, max_levels=3)
|
||||
parts.append(f"主题:{section}")
|
||||
|
||||
caption = getattr(chunk, 'title', '') or ''
|
||||
if isinstance(caption, list):
|
||||
caption = ' '.join(str(c) for c in caption if c) or ''
|
||||
if caption and caption.strip() and caption != "表格":
|
||||
parts.append(f"表格:{caption.strip()}")
|
||||
|
||||
# 检测是否是 HTML 格式,如果是则转换为 Markdown
|
||||
if '<table' in table_md.lower():
|
||||
try:
|
||||
from parsers.mineru_parser import html_table_to_markdown
|
||||
table_md = html_table_to_markdown(table_md)
|
||||
except Exception as e:
|
||||
logger.warning(f"HTML 表格转换失败: {e}")
|
||||
|
||||
lines = table_md.split('\n')
|
||||
headers = []
|
||||
for line in lines:
|
||||
if line.startswith('|') and '---' not in line:
|
||||
headers = [h.strip() for h in line.split('|') if h.strip()]
|
||||
if headers and len(headers) > 1:
|
||||
parts.append(f"字段:{', '.join(headers)}")
|
||||
break
|
||||
|
||||
row_count = len([l for l in lines if l.startswith('|') and '---' not in l])
|
||||
if headers:
|
||||
parts.append(f"描述:该表包含{row_count}行数据,记录各{', '.join(headers[:3])}信息")
|
||||
else:
|
||||
parts.append(f"描述:该表包含{row_count}行数据")
|
||||
|
||||
for line in lines:
|
||||
if line.startswith('|') and '---' not in line and headers:
|
||||
cells = [c.strip() for c in line.split('|') if c.strip()]
|
||||
if cells and cells != headers:
|
||||
example_parts = []
|
||||
for i, h in enumerate(headers[:2]):
|
||||
if i < len(cells):
|
||||
example_parts.append(f"{h}={cells[i]}")
|
||||
if example_parts:
|
||||
parts.append(f"示例:{', '.join(example_parts)}")
|
||||
break
|
||||
|
||||
# 添加完整表格内容(Markdown 格式)
|
||||
if lines and any(l.startswith('|') for l in lines):
|
||||
parts.append("") # 空行分隔
|
||||
parts.append("表格内容:")
|
||||
parts.extend(lines)
|
||||
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
# ==================== BM25 索引管理 ====================
|
||||
|
||||
class BM25Index:
|
||||
"""
|
||||
BM25 关键词检索索引
|
||||
|
||||
基于 rank_bm25.BM25Okapi 实现,使用 jieba 进行中文分词。
|
||||
支持文档的添加、检索、持久化和加载。
|
||||
|
||||
Attributes:
|
||||
bm25: BM25Okapi 索引实例
|
||||
ids: 文档 ID 列表
|
||||
documents: 文档内容列表
|
||||
metadatas: 文档元数据列表
|
||||
|
||||
Example:
|
||||
>>> index = BM25Index()
|
||||
>>> index.add_documents(
|
||||
... ids=["doc1", "doc2"],
|
||||
... documents=["财务报销流程", "请假审批制度"],
|
||||
... metadatas=[{"source": "a.pdf"}, {"source": "b.pdf"}]
|
||||
... )
|
||||
>>> ids, docs, metas, scores = index.search("报销", top_k=5)
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""初始化空的 BM25 索引"""
|
||||
self.bm25: Optional[BM25Okapi] = None
|
||||
self.ids: List[str] = []
|
||||
self.documents: List[str] = []
|
||||
self.metadatas: List[dict] = []
|
||||
|
||||
def tokenize(self, text: str) -> List[str]:
|
||||
"""
|
||||
使用 jieba 对文本进行分词
|
||||
|
||||
Args:
|
||||
text: 待分词的文本
|
||||
|
||||
Returns:
|
||||
分词结果列表
|
||||
"""
|
||||
return list(jieba.cut(text))
|
||||
|
||||
def add_documents(self, ids: List[str], documents: List[str], metadatas: List[dict]) -> None:
|
||||
"""
|
||||
添加文档到索引(会覆盖原有索引)
|
||||
|
||||
Args:
|
||||
ids: 文档 ID 列表
|
||||
documents: 文档内容列表
|
||||
metadatas: 文档元数据列表
|
||||
"""
|
||||
self.ids = ids
|
||||
self.documents = documents
|
||||
self.metadatas = metadatas
|
||||
if documents:
|
||||
tokenized = [self.tokenize(doc) for doc in documents]
|
||||
self.bm25 = BM25Okapi(tokenized)
|
||||
|
||||
def search(self, query: str, top_k: int = 10) -> Tuple[List[str], List[str], List[dict], List[float]]:
|
||||
"""
|
||||
检索与查询最相关的文档
|
||||
|
||||
Args:
|
||||
query: 查询文本
|
||||
top_k: 返回的最大文档数
|
||||
|
||||
Returns:
|
||||
元组 (ids, documents, metadatas, scores):
|
||||
- ids: 文档 ID 列表
|
||||
- documents: 文档内容列表
|
||||
- metadatas: 文档元数据列表
|
||||
- scores: BM25 分数列表
|
||||
"""
|
||||
if not self.bm25 or not self.documents:
|
||||
return [], [], [], []
|
||||
tokenized_query = self.tokenize(query)
|
||||
scores = self.bm25.get_scores(tokenized_query)
|
||||
top_indices = np.argsort(scores)[::-1][:top_k]
|
||||
return (
|
||||
[self.ids[i] for i in top_indices],
|
||||
[self.documents[i] for i in top_indices],
|
||||
[self.metadatas[i] for i in top_indices],
|
||||
[float(scores[i]) for i in top_indices]
|
||||
)
|
||||
|
||||
def save(self, filepath: str) -> None:
|
||||
"""
|
||||
持久化索引到文件
|
||||
|
||||
Args:
|
||||
filepath: 保存路径(.pkl 文件)
|
||||
"""
|
||||
data = {'ids': self.ids, 'documents': self.documents, 'metadatas': self.metadatas}
|
||||
os.makedirs(os.path.dirname(filepath), exist_ok=True)
|
||||
with open(filepath, 'wb') as f:
|
||||
pickle.dump(data, f)
|
||||
|
||||
def load(self, filepath: str) -> bool:
|
||||
"""
|
||||
从文件加载索引
|
||||
|
||||
Args:
|
||||
filepath: 索引文件路径(.pkl 文件)
|
||||
|
||||
Returns:
|
||||
加载成功返回 True,失败返回 False
|
||||
"""
|
||||
if not os.path.exists(filepath):
|
||||
return False
|
||||
try:
|
||||
with open(filepath, 'rb') as f:
|
||||
data = pickle.load(f)
|
||||
self.ids = data.get('ids', [])
|
||||
self.documents = data.get('documents', [])
|
||||
self.metadatas = data.get('metadatas', [])
|
||||
if self.documents:
|
||||
tokenized = [self.tokenize(doc) for doc in self.documents]
|
||||
self.bm25 = BM25Okapi(tokenized)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"加载 BM25 索引失败: {e}")
|
||||
return False
|
||||
|
||||
def clear(self) -> None:
|
||||
"""清空索引数据"""
|
||||
self.bm25 = None
|
||||
self.ids = []
|
||||
self.documents = []
|
||||
self.metadatas = []
|
||||
369
knowledge/chunk.py
Normal file
369
knowledge/chunk.py
Normal file
@@ -0,0 +1,369 @@
|
||||
"""
|
||||
知识库管理器 - 切片管理 Mixin
|
||||
|
||||
提供文档切片(Chunk)的 CRUD 操作,支持:
|
||||
- 新增切片:手动添加文本切片到向量库
|
||||
- 修改切片:更新切片内容和元数据
|
||||
- 删除切片:从向量库移除切片
|
||||
- 查询切片:分页获取切片列表
|
||||
|
||||
切片是向量检索的基本单位,每个切片包含:
|
||||
- 文档内容(用于向量化)
|
||||
- 元数据(来源、页码、状态等)
|
||||
- 向量嵌入(由 embedding 模型生成)
|
||||
|
||||
主要方法:
|
||||
- add_chunk: 新增切片
|
||||
- update_chunk: 修改切片
|
||||
- delete_chunk: 删除切片
|
||||
- list_chunks: 查询切片列表
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from typing import List, Dict, Optional, Tuple
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ChunkMixin:
|
||||
"""
|
||||
切片管理 Mixin
|
||||
|
||||
提供切片级别的 CRUD 操作,同时维护向量索引和 BM25 索引。
|
||||
|
||||
依赖属性(需由主类提供):
|
||||
- self.get_collection: 获取向量库集合的方法
|
||||
- self.get_bm25_index: 获取 BM25 索引的方法
|
||||
- self.save_bm25_index: 保存 BM25 索引的方法
|
||||
"""
|
||||
|
||||
def add_chunk(
|
||||
self,
|
||||
kb_name: str,
|
||||
content: str,
|
||||
metadata: dict = None
|
||||
) -> str:
|
||||
"""
|
||||
新增单个切片
|
||||
|
||||
将文本内容切片添加到指定向量库,自动生成向量嵌入和 BM25 索引。
|
||||
|
||||
Args:
|
||||
kb_name: 目标向量库名称
|
||||
content: 切片文本内容
|
||||
metadata: 可选的元数据字典,可包含:
|
||||
- source: 来源文件名
|
||||
- page: 页码
|
||||
- chunk_type: 切片类型 ('text' | 'table' | 'image')
|
||||
- status: 状态 ('active' | 'deprecated')
|
||||
- 其他自定义字段
|
||||
|
||||
Returns:
|
||||
新创建的切片 ID,格式为: manual_{时间戳}_{内容哈希}
|
||||
|
||||
Raises:
|
||||
ValueError: 向量库不存在
|
||||
Exception: 向量生成失败
|
||||
|
||||
Example:
|
||||
>>> chunk_id = kb_manager.add_chunk(
|
||||
... "public_kb",
|
||||
... "这是要添加的文本内容",
|
||||
... {"source": "manual", "page": 1}
|
||||
... )
|
||||
"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
raise ValueError(f"向量库 '{kb_name}' 不存在")
|
||||
|
||||
content_hash = hashlib.md5(content.encode('utf-8')).hexdigest()[:8]
|
||||
chunk_id = f"manual_{datetime.now().strftime('%Y%m%d%H%M%S')}_{content_hash}"
|
||||
|
||||
chunk_metadata = {
|
||||
"chunk_id": chunk_id,
|
||||
"chunk_type": "text",
|
||||
"collection": kb_name,
|
||||
"source": "manual",
|
||||
"status": "active",
|
||||
"version": "v1",
|
||||
"change_time": datetime.now().isoformat(),
|
||||
}
|
||||
if metadata:
|
||||
chunk_metadata.update(metadata)
|
||||
|
||||
try:
|
||||
from core.engine import get_engine
|
||||
engine = get_engine()
|
||||
if not engine._initialized:
|
||||
engine.initialize()
|
||||
embedding = engine.embedding_model.encode(content).tolist()
|
||||
except Exception as e:
|
||||
logger.error(f"生成向量失败: {e}")
|
||||
raise
|
||||
|
||||
collection.add(
|
||||
ids=[chunk_id],
|
||||
documents=[content],
|
||||
metadatas=[chunk_metadata],
|
||||
embeddings=[embedding]
|
||||
)
|
||||
|
||||
bm25 = self.get_bm25_index(kb_name)
|
||||
if bm25:
|
||||
bm25.add_documents([chunk_id], [content], [chunk_metadata])
|
||||
self.save_bm25_index(kb_name)
|
||||
|
||||
logger.info(f"新增切片: {chunk_id} -> {kb_name}")
|
||||
return chunk_id
|
||||
|
||||
def update_chunk(
|
||||
self,
|
||||
kb_name: str,
|
||||
chunk_id: str,
|
||||
content: str = None,
|
||||
metadata: dict = None
|
||||
) -> bool:
|
||||
"""
|
||||
修改切片
|
||||
|
||||
更新切片的内容和/或元数据。如果更新内容,会重新生成向量嵌入。
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称
|
||||
chunk_id: 切片 ID
|
||||
content: 新的文本内容(None 表示不更新)
|
||||
metadata: 要更新的元数据字段(None 表示不更新)
|
||||
|
||||
Returns:
|
||||
更新成功返回 True,切片不存在或更新失败返回 False
|
||||
|
||||
Example:
|
||||
>>> kb_manager.update_chunk(
|
||||
... "public_kb",
|
||||
... "manual_20260517_abc123",
|
||||
... content="更新后的内容",
|
||||
... metadata={"status": "updated"}
|
||||
... )
|
||||
True
|
||||
"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return False
|
||||
|
||||
existing = collection.get(ids=[chunk_id])
|
||||
if not existing['ids']:
|
||||
logger.warning(f"切片不存在: {chunk_id}")
|
||||
return False
|
||||
|
||||
update_kwargs = {"ids": [chunk_id]}
|
||||
|
||||
if content is not None:
|
||||
update_kwargs["documents"] = [content]
|
||||
try:
|
||||
from core.engine import get_engine
|
||||
engine = get_engine()
|
||||
if not engine._initialized:
|
||||
engine.initialize()
|
||||
embedding = engine.embedding_model.encode(content).tolist()
|
||||
update_kwargs["embeddings"] = [embedding]
|
||||
except Exception as e:
|
||||
logger.error(f"生成向量失败: {e}")
|
||||
return False
|
||||
|
||||
if metadata is not None:
|
||||
existing_metadata = existing['metadatas'][0] if existing['metadatas'] else {}
|
||||
existing_metadata.update(metadata)
|
||||
update_kwargs["metadatas"] = [existing_metadata]
|
||||
|
||||
collection.update(**update_kwargs)
|
||||
|
||||
if content is not None:
|
||||
bm25 = self.get_bm25_index(kb_name)
|
||||
if bm25:
|
||||
bm25.add_documents([chunk_id], [content], [metadata or {}])
|
||||
self.save_bm25_index(kb_name)
|
||||
|
||||
logger.info(f"更新切片: {chunk_id}")
|
||||
return True
|
||||
|
||||
def get_chunk(self, kb_name: str, chunk_id: str) -> Optional[Dict]:
|
||||
"""
|
||||
获取单个切片信息
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称
|
||||
chunk_id: 切片 ID
|
||||
|
||||
Returns:
|
||||
切片信息字典,包含 id, document, metadata 等字段
|
||||
切片不存在时返回 None
|
||||
"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return None
|
||||
|
||||
result = collection.get(ids=[chunk_id], include=["documents", "metadatas"])
|
||||
if not result['ids']:
|
||||
return None
|
||||
|
||||
return {
|
||||
'id': result['ids'][0],
|
||||
'document': result['documents'][0] if result['documents'] else '',
|
||||
'metadata': result['metadatas'][0] if result['metadatas'] else {}
|
||||
}
|
||||
|
||||
def delete_chunk(self, kb_name: str, chunk_id: str) -> Tuple[bool, Optional[str]]:
|
||||
"""
|
||||
删除切片
|
||||
|
||||
从向量库中移除指定的切片。
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称
|
||||
chunk_id: 切片 ID
|
||||
|
||||
Returns:
|
||||
元组 (success, source_file):
|
||||
- success: 删除是否成功
|
||||
- source_file: 被删除切片的来源文件名(用于清理哈希记录)
|
||||
|
||||
Note:
|
||||
删除后 BM25 索引不会立即更新,需要手动调用 rebuild_bm25_index。
|
||||
"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return False, None
|
||||
|
||||
# 获取切片信息(删除前)
|
||||
existing = collection.get(ids=[chunk_id], include=["metadatas"])
|
||||
if not existing['ids']:
|
||||
logger.warning(f"切片不存在: {chunk_id}")
|
||||
return False, None
|
||||
|
||||
source_file = None
|
||||
if existing['metadatas'] and existing['metadatas'][0]:
|
||||
source_file = existing['metadatas'][0].get('source')
|
||||
|
||||
collection.delete(ids=[chunk_id])
|
||||
logger.info(f"删除切片: {chunk_id}")
|
||||
return True, source_file
|
||||
|
||||
def delete_chunks_by_source(self, kb_name: str, source: str) -> int:
|
||||
"""
|
||||
批量删除指定文件的所有切片
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称
|
||||
source: 文件名(source 字段值)
|
||||
|
||||
Returns:
|
||||
删除的切片数量
|
||||
|
||||
Note:
|
||||
删除后会重建 BM25 索引。
|
||||
"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return 0
|
||||
|
||||
# 获取所有切片,找到匹配 source 的
|
||||
all_chunks = collection.get(include=["metadatas"])
|
||||
|
||||
# 找到要删除的切片 ID
|
||||
chunk_ids_to_delete = []
|
||||
for i, chunk_id in enumerate(all_chunks['ids']):
|
||||
metadata = all_chunks['metadatas'][i] if all_chunks['metadatas'] else {}
|
||||
if metadata.get('source') == source:
|
||||
chunk_ids_to_delete.append(chunk_id)
|
||||
|
||||
if not chunk_ids_to_delete:
|
||||
logger.warning(f"未找到文件 {source} 的切片")
|
||||
return 0
|
||||
|
||||
# 批量删除
|
||||
collection.delete(ids=chunk_ids_to_delete)
|
||||
logger.info(f"批量删除切片: {source}, 共 {len(chunk_ids_to_delete)} 个")
|
||||
|
||||
# 重建 BM25 索引
|
||||
try:
|
||||
self._bm25_indexes.pop(kb_name, None)
|
||||
bm25 = self.get_bm25_index(kb_name)
|
||||
if bm25:
|
||||
remaining = collection.get(include=["documents", "metadatas"])
|
||||
if remaining['ids']:
|
||||
bm25.add_documents(
|
||||
remaining['ids'],
|
||||
remaining['documents'] or [],
|
||||
remaining['metadatas'] or []
|
||||
)
|
||||
self.save_bm25_index(kb_name)
|
||||
except Exception as e:
|
||||
logger.warning(f"重建 BM25 索引失败: {e}")
|
||||
|
||||
return len(chunk_ids_to_delete)
|
||||
|
||||
def list_chunks(
|
||||
self,
|
||||
kb_name: str,
|
||||
document_id: str = None,
|
||||
limit: int = 100,
|
||||
offset: int = 0
|
||||
) -> List[Dict]:
|
||||
"""
|
||||
获取向量库中的切片列表
|
||||
|
||||
支持分页和按文档过滤。
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称
|
||||
document_id: 文档 ID(source 字段),用于过滤特定文档的切片
|
||||
limit: 返回数量限制(默认 100)
|
||||
offset: 偏移量(默认 0)
|
||||
|
||||
Returns:
|
||||
切片字典列表,每个字典包含:
|
||||
- id: 切片 ID
|
||||
- document: 切片内容
|
||||
- metadata: 元数据字典
|
||||
- status: 状态
|
||||
- version: 版本号
|
||||
|
||||
Example:
|
||||
>>> chunks = kb_manager.list_chunks("public_kb", limit=10)
|
||||
>>> for chunk in chunks:
|
||||
... print(chunk["id"], chunk["status"])
|
||||
"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return []
|
||||
|
||||
where_filter = {}
|
||||
if document_id:
|
||||
where_filter["source"] = document_id
|
||||
|
||||
try:
|
||||
result = collection.get(
|
||||
where=where_filter if where_filter else None,
|
||||
limit=limit,
|
||||
offset=offset
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"获取切片失败: {e}")
|
||||
return []
|
||||
|
||||
return [
|
||||
{
|
||||
"id": id,
|
||||
"document": doc,
|
||||
"metadata": meta,
|
||||
"status": meta.get("status", "active") if meta else "active",
|
||||
"version": meta.get("version", "v1") if meta else "v1"
|
||||
}
|
||||
for id, doc, meta in zip(
|
||||
result['ids'],
|
||||
result.get('documents', []),
|
||||
result.get('metadatas', [])
|
||||
)
|
||||
]
|
||||
241
knowledge/cleanup.py
Normal file
241
knowledge/cleanup.py
Normal file
@@ -0,0 +1,241 @@
|
||||
"""
|
||||
文档版本自动清理任务
|
||||
|
||||
定期清理 superseded 状态的旧版本,控制存储成本。
|
||||
|
||||
使用方式:
|
||||
from knowledge.cleanup import cleanup_superseded_versions
|
||||
|
||||
# 清理超过 7 天的 superseded 版本
|
||||
cleaned = cleanup_superseded_versions(days_to_keep=7)
|
||||
"""
|
||||
|
||||
import logging
|
||||
from datetime import datetime, timedelta
|
||||
from typing import List
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def cleanup_superseded_versions(days_to_keep: int = 7) -> int:
|
||||
"""
|
||||
清理超过指定天数的 superseded 版本
|
||||
|
||||
Args:
|
||||
days_to_keep: 保留天数,默认7天
|
||||
|
||||
Returns:
|
||||
清理的 chunk 数量
|
||||
"""
|
||||
from knowledge.manager import get_kb_manager
|
||||
|
||||
kb_manager = get_kb_manager()
|
||||
cutoff_date = (datetime.now() - timedelta(days=days_to_keep)).isoformat()
|
||||
|
||||
logger.info(f"开始清理 superseded 版本(保留 {days_to_keep} 天内的)")
|
||||
|
||||
# 获取所有向量库
|
||||
try:
|
||||
kb_names = kb_manager.list_collections()
|
||||
except Exception as e:
|
||||
logger.error(f"获取向量库列表失败: {e}")
|
||||
return 0
|
||||
|
||||
total_cleaned = 0
|
||||
|
||||
for kb_name in kb_names:
|
||||
try:
|
||||
collection = kb_manager.get_collection(kb_name)
|
||||
if not collection:
|
||||
continue
|
||||
|
||||
# 查询超过保留期的 superseded chunks
|
||||
# 注意:ChromaDB 的 where 过滤可能不支持 $lt 操作符
|
||||
# 所以我们先获取所有 superseded chunks,然后在 Python 中过滤
|
||||
result = collection.get(
|
||||
where={"status": "superseded"}
|
||||
)
|
||||
|
||||
if not result['ids']:
|
||||
continue
|
||||
|
||||
# 在 Python 中过滤超过保留期的 chunks
|
||||
ids_to_delete = []
|
||||
for i, meta in enumerate(result['metadatas']):
|
||||
superseded_time = meta.get('superseded_time', '')
|
||||
if superseded_time and superseded_time < cutoff_date:
|
||||
ids_to_delete.append(result['ids'][i])
|
||||
|
||||
if ids_to_delete:
|
||||
# 删除这些 chunks
|
||||
collection.delete(ids=ids_to_delete)
|
||||
total_cleaned += len(ids_to_delete)
|
||||
logger.info(f"清理 {kb_name}: {len(ids_to_delete)} chunks")
|
||||
|
||||
# 重建 BM25 索引
|
||||
kb_manager.rebuild_bm25_index(kb_name)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"清理 {kb_name} 失败: {e}")
|
||||
continue
|
||||
|
||||
logger.info(f"清理完成,共删除 {total_cleaned} 个 superseded chunks")
|
||||
return total_cleaned
|
||||
|
||||
|
||||
def cleanup_deprecated_versions(days_to_keep: int = 30) -> int:
|
||||
"""
|
||||
清理超过指定天数的 deprecated 版本
|
||||
|
||||
Args:
|
||||
days_to_keep: 保留天数,默认30天
|
||||
|
||||
Returns:
|
||||
清理的 chunk 数量
|
||||
"""
|
||||
from knowledge.manager import get_kb_manager
|
||||
|
||||
kb_manager = get_kb_manager()
|
||||
cutoff_date = (datetime.now() - timedelta(days=days_to_keep)).isoformat()
|
||||
|
||||
logger.info(f"开始清理 deprecated 版本(保留 {days_to_keep} 天内的)")
|
||||
|
||||
try:
|
||||
kb_names = kb_manager.list_collections()
|
||||
except Exception as e:
|
||||
logger.error(f"获取向量库列表失败: {e}")
|
||||
return 0
|
||||
|
||||
total_cleaned = 0
|
||||
|
||||
for kb_name in kb_names:
|
||||
try:
|
||||
collection = kb_manager.get_collection(kb_name)
|
||||
if not collection:
|
||||
continue
|
||||
|
||||
# 获取所有 deprecated chunks
|
||||
result = collection.get(
|
||||
where={"status": "deprecated"}
|
||||
)
|
||||
|
||||
if not result['ids']:
|
||||
continue
|
||||
|
||||
# 在 Python 中过滤超过保留期的 chunks
|
||||
ids_to_delete = []
|
||||
for i, meta in enumerate(result['metadatas']):
|
||||
deprecated_date = meta.get('deprecated_date', '')
|
||||
if deprecated_date and deprecated_date < cutoff_date:
|
||||
ids_to_delete.append(result['ids'][i])
|
||||
|
||||
if ids_to_delete:
|
||||
# 删除这些 chunks
|
||||
collection.delete(ids=ids_to_delete)
|
||||
total_cleaned += len(ids_to_delete)
|
||||
logger.info(f"清理 {kb_name}: {len(ids_to_delete)} deprecated chunks")
|
||||
|
||||
# 重建 BM25 索引
|
||||
kb_manager.rebuild_bm25_index(kb_name)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"清理 {kb_name} 失败: {e}")
|
||||
continue
|
||||
|
||||
logger.info(f"清理完成,共删除 {total_cleaned} 个 deprecated chunks")
|
||||
return total_cleaned
|
||||
|
||||
|
||||
def cleanup_all_old_versions(
|
||||
superseded_days: int = 7,
|
||||
deprecated_days: int = 30
|
||||
) -> dict:
|
||||
"""
|
||||
清理所有旧版本
|
||||
|
||||
Args:
|
||||
superseded_days: superseded 版本保留天数
|
||||
deprecated_days: deprecated 版本保留天数
|
||||
|
||||
Returns:
|
||||
清理统计信息
|
||||
"""
|
||||
logger.info("开始清理所有旧版本")
|
||||
|
||||
superseded_cleaned = cleanup_superseded_versions(superseded_days)
|
||||
deprecated_cleaned = cleanup_deprecated_versions(deprecated_days)
|
||||
|
||||
result = {
|
||||
"superseded_cleaned": superseded_cleaned,
|
||||
"deprecated_cleaned": deprecated_cleaned,
|
||||
"total_cleaned": superseded_cleaned + deprecated_cleaned
|
||||
}
|
||||
|
||||
logger.info(f"清理完成: {result}")
|
||||
return result
|
||||
|
||||
|
||||
# ==================== 定时任务(可选) ====================
|
||||
|
||||
def start_cleanup_scheduler(
|
||||
superseded_days: int = 7,
|
||||
deprecated_days: int = 30,
|
||||
schedule_time: str = "03:00"
|
||||
):
|
||||
"""
|
||||
启动清理调度器(每天定时执行)
|
||||
|
||||
Args:
|
||||
superseded_days: superseded 版本保留天数
|
||||
deprecated_days: deprecated 版本保留天数
|
||||
schedule_time: 执行时间(24小时制,如 "03:00")
|
||||
|
||||
注意:需要安装 schedule 库:pip install schedule
|
||||
"""
|
||||
try:
|
||||
import schedule
|
||||
import threading
|
||||
import time
|
||||
except ImportError:
|
||||
logger.error("schedule 库未安装,无法启动定时任务。请运行: pip install schedule")
|
||||
return
|
||||
|
||||
def job():
|
||||
"""清理任务"""
|
||||
try:
|
||||
cleanup_all_old_versions(superseded_days, deprecated_days)
|
||||
except Exception as e:
|
||||
logger.error(f"定时清理任务失败: {e}")
|
||||
|
||||
# 设置定时任务
|
||||
schedule.every().day.at(schedule_time).do(job)
|
||||
logger.info(f"清理调度器已启动,每天 {schedule_time} 执行")
|
||||
|
||||
def run_scheduler():
|
||||
"""调度器运行循环"""
|
||||
while True:
|
||||
schedule.run_pending()
|
||||
time.sleep(3600) # 每小时检查一次
|
||||
|
||||
# 在后台线程中运行
|
||||
thread = threading.Thread(target=run_scheduler, daemon=True)
|
||||
thread.start()
|
||||
logger.info("清理调度器后台线程已启动")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# 测试清理功能
|
||||
import sys
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
|
||||
)
|
||||
|
||||
if len(sys.argv) > 1:
|
||||
days = int(sys.argv[1])
|
||||
else:
|
||||
days = 7
|
||||
|
||||
print(f"清理超过 {days} 天的 superseded 版本...")
|
||||
result = cleanup_all_old_versions(superseded_days=days)
|
||||
print(f"清理完成: {result}")
|
||||
385
knowledge/collection.py
Normal file
385
knowledge/collection.py
Normal file
@@ -0,0 +1,385 @@
|
||||
"""
|
||||
知识库管理器 - 集合管理 Mixin
|
||||
|
||||
提供向量库(ChromaDB Collection)的创建、删除、查询等管理功能。
|
||||
每个向量库对应一个独立的 ChromaDB 集合,支持按部门隔离。
|
||||
|
||||
主要方法:
|
||||
- get_collection: 获取或创建向量库集合
|
||||
- create_collection: 创建新向量库
|
||||
- delete_collection: 删除向量库
|
||||
- list_collections: 列出所有向量库
|
||||
- collection_exists: 检查向量库是否存在
|
||||
"""
|
||||
|
||||
import os
|
||||
import gc
|
||||
import time
|
||||
import logging
|
||||
from typing import List, Optional, Tuple
|
||||
from datetime import datetime
|
||||
|
||||
import chromadb
|
||||
from chromadb import Collection
|
||||
|
||||
from .base import (
|
||||
CollectionInfo, PUBLIC_KB_NAME
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class CollectionMixin:
|
||||
"""
|
||||
集合管理 Mixin
|
||||
|
||||
提供 ChromaDB 向量库集合的管理功能,包括:
|
||||
- 集合的创建与删除
|
||||
- 集合元数据管理
|
||||
- 集合列表查询
|
||||
|
||||
依赖属性(需由主类提供):
|
||||
- self.base_path: 向量库存储根路径
|
||||
- self.bm25_base_path: BM25 索引存储路径
|
||||
- self._collections: 集合缓存字典
|
||||
- self._clients: 客户端缓存字典
|
||||
- self._bm25_indexes: BM25 索引缓存字典
|
||||
- self._metadata: 元数据字典
|
||||
- self._lock: 线程锁
|
||||
"""
|
||||
|
||||
def _get_client(self, kb_name: str) -> chromadb.PersistentClient:
|
||||
"""
|
||||
获取或创建 ChromaDB 客户端
|
||||
|
||||
每个向量库使用独立的客户端实例,客户端会被缓存以避免重复创建。
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称
|
||||
|
||||
Returns:
|
||||
ChromaDB PersistentClient 实例
|
||||
"""
|
||||
if kb_name not in self._clients:
|
||||
db_path = os.path.join(self.base_path, kb_name)
|
||||
os.makedirs(db_path, exist_ok=True)
|
||||
self._clients[kb_name] = chromadb.PersistentClient(path=db_path)
|
||||
return self._clients[kb_name]
|
||||
|
||||
def get_collection(self, kb_name: str) -> Optional[Collection]:
|
||||
"""
|
||||
获取或创建向量库集合
|
||||
|
||||
如果集合已存在于缓存中,直接返回缓存实例。
|
||||
否则从 ChromaDB 获取或创建新集合。
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称(只能包含字母、数字和下划线)
|
||||
|
||||
Returns:
|
||||
ChromaDB Collection 对象;失败时返回 None
|
||||
|
||||
Note:
|
||||
使用 cosine 相似度作为向量距离度量。
|
||||
sync_threshold 设置为 100000 以支持大批量写入。
|
||||
"""
|
||||
with self._lock:
|
||||
if kb_name in self._collections:
|
||||
return self._collections[kb_name]
|
||||
try:
|
||||
client = self._get_client(kb_name)
|
||||
collection = client.get_or_create_collection(
|
||||
name=kb_name,
|
||||
metadata={
|
||||
"hnsw:space": "cosine",
|
||||
"hnsw:sync_threshold": 100000
|
||||
}
|
||||
)
|
||||
self._collections[kb_name] = collection
|
||||
logger.info(f"获取向量库: {kb_name}, 文档数: {collection.count()}")
|
||||
return collection
|
||||
except Exception as e:
|
||||
logger.error(f"获取向量库失败: {kb_name}, 错误: {e}")
|
||||
return None
|
||||
|
||||
def create_collection(
|
||||
self,
|
||||
kb_name: str,
|
||||
display_name: str = "",
|
||||
department: str = "",
|
||||
description: str = ""
|
||||
) -> Tuple[bool, str]:
|
||||
"""
|
||||
创建新向量库
|
||||
|
||||
创建一个新的 ChromaDB 集合,同时初始化对应的 BM25 索引。
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称(只能包含字母、数字和下划线)
|
||||
display_name: 显示名称(用于 UI 展示)
|
||||
department: 所属部门(用于权限控制)
|
||||
description: 向量库描述
|
||||
|
||||
Returns:
|
||||
元组 (success, message):
|
||||
- success: 创建是否成功
|
||||
- message: 成功或失败的描述信息
|
||||
|
||||
Example:
|
||||
>>> success, msg = kb_manager.create_collection(
|
||||
... "dept_finance",
|
||||
... display_name="财务部知识库",
|
||||
... department="finance"
|
||||
... )
|
||||
"""
|
||||
from .base import BM25Index
|
||||
|
||||
if not kb_name or not kb_name.replace('_', '').isalnum():
|
||||
return False, "向量库名称只能包含字母、数字和下划线"
|
||||
|
||||
if kb_name in self._metadata.get("collections", {}):
|
||||
return False, f"向量库 '{kb_name}' 已存在"
|
||||
|
||||
chroma_dir = os.path.join(self.base_path, kb_name)
|
||||
if os.path.exists(chroma_dir):
|
||||
import shutil
|
||||
shutil.rmtree(chroma_dir)
|
||||
logger.warning(f"清理残留向量库文件夹: {chroma_dir}")
|
||||
|
||||
try:
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return False, "创建向量库失败"
|
||||
|
||||
self._bm25_indexes[kb_name] = BM25Index()
|
||||
|
||||
if "collections" not in self._metadata:
|
||||
self._metadata["collections"] = {}
|
||||
|
||||
self._metadata["collections"][kb_name] = {
|
||||
"display_name": display_name or kb_name,
|
||||
"department": department,
|
||||
"description": description,
|
||||
"created_at": datetime.now().isoformat()
|
||||
}
|
||||
self._save_metadata()
|
||||
|
||||
logger.info(f"创建向量库: {kb_name}")
|
||||
return True, f"向量库 '{kb_name}' 创建成功"
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"创建向量库失败: {e}")
|
||||
return False, f"创建失败: {str(e)}"
|
||||
|
||||
def update_collection_metadata(
|
||||
self,
|
||||
kb_name: str,
|
||||
display_name: str = None,
|
||||
description: str = None
|
||||
) -> bool:
|
||||
"""
|
||||
更新向量库元数据
|
||||
|
||||
仅更新 display_name 和 description,不影响集合中的数据。
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称
|
||||
display_name: 新的显示名称(None 表示不更新)
|
||||
description: 新的描述(None 表示不更新)
|
||||
|
||||
Returns:
|
||||
更新成功返回 True,向量库不存在返回 False
|
||||
"""
|
||||
collections = self._metadata.get("collections", {})
|
||||
if kb_name not in collections:
|
||||
return False
|
||||
|
||||
if display_name is not None:
|
||||
collections[kb_name]["display_name"] = display_name
|
||||
if description is not None:
|
||||
collections[kb_name]["description"] = description
|
||||
|
||||
self._save_metadata()
|
||||
logger.info(f"更新向量库元数据: {kb_name}")
|
||||
return True
|
||||
|
||||
def delete_collection(self, kb_name: str, delete_documents: bool = False) -> Tuple[bool, str]:
|
||||
"""
|
||||
删除向量库
|
||||
|
||||
删除向量库及其所有数据,包括:
|
||||
- ChromaDB 集合
|
||||
- BM25 索引文件
|
||||
- 向量库文件夹
|
||||
- 可选:原始文档文件
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称
|
||||
delete_documents: 是否同时删除原始文档文件
|
||||
|
||||
Returns:
|
||||
元组 (success, message):
|
||||
- success: 删除是否成功
|
||||
- message: 成功或失败的描述信息
|
||||
|
||||
Warning:
|
||||
公开知识库 (public_kb) 不能被删除。
|
||||
删除操作不可逆,请谨慎使用。
|
||||
"""
|
||||
import shutil
|
||||
|
||||
if kb_name == PUBLIC_KB_NAME:
|
||||
return False, "公开知识库不能删除"
|
||||
|
||||
if kb_name not in self._metadata.get("collections", {}):
|
||||
return False, f"向量库 '{kb_name}' 不存在"
|
||||
|
||||
try:
|
||||
client = self._clients.get(kb_name)
|
||||
if client:
|
||||
try:
|
||||
client.delete_collection(kb_name)
|
||||
except Exception as e:
|
||||
logger.debug(f"删除集合失败: {e}")
|
||||
try:
|
||||
client.close()
|
||||
except Exception as e:
|
||||
logger.debug(f"关闭客户端失败: {e}")
|
||||
|
||||
if kb_name in self._collections:
|
||||
del self._collections[kb_name]
|
||||
if kb_name in self._bm25_indexes:
|
||||
del self._bm25_indexes[kb_name]
|
||||
if kb_name in self._clients:
|
||||
del self._clients[kb_name]
|
||||
|
||||
gc.collect()
|
||||
time.sleep(1)
|
||||
|
||||
bm25_path = os.path.join(self.bm25_base_path, f"{kb_name}.pkl")
|
||||
if os.path.exists(bm25_path):
|
||||
os.remove(bm25_path)
|
||||
|
||||
chroma_dir = os.path.join(self.base_path, kb_name)
|
||||
if os.path.exists(chroma_dir):
|
||||
for attempt in range(3):
|
||||
try:
|
||||
shutil.rmtree(chroma_dir)
|
||||
logger.info(f"删除向量库文件夹: {chroma_dir}")
|
||||
break
|
||||
except PermissionError:
|
||||
if attempt < 2:
|
||||
logger.warning(f"删除文件夹失败,重试 {attempt + 1}/3")
|
||||
gc.collect()
|
||||
time.sleep(1)
|
||||
else:
|
||||
logger.error(f"删除文件夹失败,文件被占用: {chroma_dir}")
|
||||
raise
|
||||
|
||||
if delete_documents:
|
||||
from config import DOCUMENTS_PATH
|
||||
docs_dir = os.path.join(DOCUMENTS_PATH, kb_name)
|
||||
if os.path.exists(docs_dir):
|
||||
shutil.rmtree(docs_dir)
|
||||
logger.info(f"删除文档文件夹: {docs_dir}")
|
||||
|
||||
# 删除该向量库所有文档的哈希记录
|
||||
try:
|
||||
from knowledge.sync import SyncDatabase
|
||||
sync_db = SyncDatabase()
|
||||
# 获取所有以 kb_name/ 开头的哈希记录
|
||||
all_hashes = sync_db.get_all_document_hashes()
|
||||
deleted_count = 0
|
||||
for doc_id in list(all_hashes.keys()):
|
||||
if doc_id.startswith(f"{kb_name}/"):
|
||||
sync_db.delete_document_hash(doc_id)
|
||||
deleted_count += 1
|
||||
if deleted_count > 0:
|
||||
logger.info(f"清理向量库哈希记录: {kb_name}, 共 {deleted_count} 条")
|
||||
except Exception as e:
|
||||
logger.warning(f"清理哈希记录失败: {e}")
|
||||
|
||||
if kb_name in self._metadata.get("collections", {}):
|
||||
del self._metadata["collections"][kb_name]
|
||||
self._save_metadata()
|
||||
|
||||
logger.info(f"删除向量库: {kb_name}")
|
||||
return True, f"向量库 '{kb_name}' 已删除"
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"删除向量库失败: {e}")
|
||||
return False, f"删除失败: {str(e)}"
|
||||
|
||||
def list_collections(self) -> List[CollectionInfo]:
|
||||
"""
|
||||
列出所有向量库
|
||||
|
||||
返回所有已创建的向量库信息,包括名称、文档数量、创建时间等。
|
||||
|
||||
Returns:
|
||||
CollectionInfo 对象列表,每个对象包含:
|
||||
- name: 向量库名称
|
||||
- display_name: 显示名称
|
||||
- document_count: 文档数量
|
||||
- created_at: 创建时间
|
||||
- department: 所属部门
|
||||
- description: 描述
|
||||
"""
|
||||
result = []
|
||||
|
||||
# 扫描 base_path 下的所有子目录作为向量库
|
||||
# 每个向量库使用独立目录:base_path/my_ky, base_path/public_kb 等
|
||||
actual_collections = []
|
||||
try:
|
||||
if os.path.exists(self.base_path):
|
||||
for item in os.listdir(self.base_path):
|
||||
item_path = os.path.join(self.base_path, item)
|
||||
if os.path.isdir(item_path) and not item.startswith('.'):
|
||||
# 检查是否包含 chroma.sqlite3(有效的向量库目录)
|
||||
if os.path.exists(os.path.join(item_path, 'chroma.sqlite3')):
|
||||
actual_collections.append(item)
|
||||
except Exception as e:
|
||||
logger.warning(f"扫描向量库目录失败: {e}")
|
||||
|
||||
# 如果扫描失败,回退到元数据中的集合列表
|
||||
if not actual_collections:
|
||||
actual_collections = list(self._metadata.get("collections", {}).keys())
|
||||
|
||||
for name in actual_collections:
|
||||
if name not in self._metadata.get("collections", {}):
|
||||
if "collections" not in self._metadata:
|
||||
self._metadata["collections"] = {}
|
||||
self._metadata["collections"][name] = {
|
||||
"display_name": name,
|
||||
"department": "",
|
||||
"description": "",
|
||||
"created_at": datetime.now().isoformat()
|
||||
}
|
||||
logger.info(f"自动补充向量库元数据: {name}")
|
||||
|
||||
self._save_metadata()
|
||||
|
||||
for name, info in self._metadata.get("collections", {}).items():
|
||||
collection = self.get_collection(name)
|
||||
result.append(CollectionInfo(
|
||||
name=name,
|
||||
display_name=info.get("display_name", name),
|
||||
document_count=collection.count() if collection else 0,
|
||||
created_at=info.get("created_at", ""),
|
||||
department=info.get("department", ""),
|
||||
description=info.get("description", "")
|
||||
))
|
||||
|
||||
return result
|
||||
|
||||
def collection_exists(self, kb_name: str) -> bool:
|
||||
"""
|
||||
检查向量库是否存在
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称
|
||||
|
||||
Returns:
|
||||
存在返回 True,不存在返回 False
|
||||
"""
|
||||
return kb_name in self._metadata.get("collections", {})
|
||||
266
knowledge/document.py
Normal file
266
knowledge/document.py
Normal file
@@ -0,0 +1,266 @@
|
||||
"""
|
||||
知识库管理器 - 文档管理 Mixin
|
||||
|
||||
包含文档级别的管理方法
|
||||
"""
|
||||
|
||||
import os
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from typing import List, Dict, Optional
|
||||
|
||||
from .base import _get_doc_type
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class DocumentMixin:
|
||||
"""文档管理方法"""
|
||||
|
||||
def get_document_count(self, kb_name: str) -> int:
|
||||
"""获取向量库中的文档数量"""
|
||||
collection = self.get_collection(kb_name)
|
||||
return collection.count() if collection else 0
|
||||
|
||||
def list_documents(self, kb_name: str) -> List[dict]:
|
||||
"""列出向量库中的文档"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return []
|
||||
|
||||
result = collection.get()
|
||||
|
||||
from collections import Counter
|
||||
file_chunks = Counter()
|
||||
|
||||
for meta in result.get('metadatas', []):
|
||||
source = meta.get('source', 'unknown')
|
||||
file_chunks[source] += 1
|
||||
|
||||
return [
|
||||
{"source": source, "chunks": count}
|
||||
for source, count in file_chunks.items()
|
||||
]
|
||||
|
||||
def delete_document(self, kb_name: str, filename: str) -> int:
|
||||
"""从向量库删除文档"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return 0
|
||||
|
||||
result = collection.get(where={"source": filename})
|
||||
|
||||
if not result['ids']:
|
||||
return 0
|
||||
|
||||
collection.delete(ids=result['ids'])
|
||||
deleted = len(result['ids'])
|
||||
|
||||
logger.info(f"从 {kb_name} 删除文档: {filename}, 片段数: {deleted}")
|
||||
return deleted
|
||||
|
||||
def deprecate_document(
|
||||
self,
|
||||
kb_name: str,
|
||||
filename: str,
|
||||
reason: str = "制度废止",
|
||||
deprecated_by: str = ""
|
||||
) -> Dict:
|
||||
"""软删除文档 - 将chunks状态标记为deprecated"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return {"success": False, "error": "向量库不存在"}
|
||||
|
||||
result = collection.get(where={"source": filename})
|
||||
|
||||
if not result['ids']:
|
||||
return {"success": False, "error": "文档不存在"}
|
||||
|
||||
deprecated_date = datetime.now().isoformat()
|
||||
updated_metadatas = [
|
||||
{
|
||||
**m,
|
||||
"status": "deprecated",
|
||||
"deprecated_date": deprecated_date,
|
||||
"deprecated_reason": reason,
|
||||
"deprecated_by": deprecated_by
|
||||
}
|
||||
for m in result['metadatas']
|
||||
]
|
||||
|
||||
collection.update(
|
||||
ids=result['ids'],
|
||||
metadatas=updated_metadatas
|
||||
)
|
||||
|
||||
self.rebuild_bm25_index(kb_name)
|
||||
|
||||
logger.info(f"软删除文档: {kb_name}/{filename}, chunks: {len(result['ids'])}, 原因: {reason}")
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"deprecated_chunks": len(result['ids']),
|
||||
"document_id": filename,
|
||||
"collection": kb_name,
|
||||
"deprecated_date": deprecated_date
|
||||
}
|
||||
|
||||
def restore_document(self, kb_name: str, filename: str) -> Dict:
|
||||
"""恢复已废止的文档"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return {"success": False, "error": "向量库不存在"}
|
||||
|
||||
result = collection.get(
|
||||
where={
|
||||
"$and": [
|
||||
{"source": filename},
|
||||
{"status": "deprecated"}
|
||||
]
|
||||
}
|
||||
)
|
||||
|
||||
if not result['ids']:
|
||||
return {"success": False, "error": "未找到已废止的文档"}
|
||||
|
||||
updated_metadatas = [
|
||||
{
|
||||
**m,
|
||||
"status": "active",
|
||||
"deprecated_date": None,
|
||||
"deprecated_reason": None
|
||||
}
|
||||
for m in result['metadatas']
|
||||
]
|
||||
|
||||
collection.update(
|
||||
ids=result['ids'],
|
||||
metadatas=updated_metadatas
|
||||
)
|
||||
|
||||
self.rebuild_bm25_index(kb_name)
|
||||
|
||||
logger.info(f"恢复文档: {kb_name}/{filename}, chunks: {len(result['ids'])}")
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"restored_chunks": len(result['ids']),
|
||||
"document_id": filename,
|
||||
"collection": kb_name
|
||||
}
|
||||
|
||||
def get_document_chunks(
|
||||
self,
|
||||
kb_name: str,
|
||||
filename: str,
|
||||
status: str = None
|
||||
) -> List[Dict]:
|
||||
"""获取文档的chunks列表"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return []
|
||||
|
||||
where_filter = {"source": filename}
|
||||
if status:
|
||||
where_filter["status"] = status
|
||||
|
||||
result = collection.get(where=where_filter)
|
||||
|
||||
return [
|
||||
{
|
||||
"id": id,
|
||||
"document": doc,
|
||||
"metadata": meta,
|
||||
"status": meta.get("status", "active"),
|
||||
"version": meta.get("version", "v1")
|
||||
}
|
||||
for id, doc, meta in zip(
|
||||
result['ids'],
|
||||
result['documents'],
|
||||
result['metadatas']
|
||||
)
|
||||
]
|
||||
|
||||
def get_document_info(self, kb_name: str, filename: str) -> Optional[Dict]:
|
||||
"""获取文档基本信息"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return None
|
||||
|
||||
result = collection.get(where={"source": filename})
|
||||
|
||||
if not result['ids']:
|
||||
return None
|
||||
|
||||
status_counts = {}
|
||||
for meta in result['metadatas']:
|
||||
status = meta.get("status", "active")
|
||||
status_counts[status] = status_counts.get(status, 0) + 1
|
||||
|
||||
main_status = "active"
|
||||
if status_counts.get("deprecated", 0) > status_counts.get("active", 0):
|
||||
main_status = "deprecated"
|
||||
elif status_counts.get("superseded", 0) > 0:
|
||||
main_status = "superseded"
|
||||
|
||||
first_meta = result['metadatas'][0] if result['metadatas'] else {}
|
||||
|
||||
return {
|
||||
"document_id": filename,
|
||||
"collection": kb_name,
|
||||
"total_chunks": len(result['ids']),
|
||||
"status": main_status,
|
||||
"status_counts": status_counts,
|
||||
"version": first_meta.get("version", "v1"),
|
||||
"effective_date": first_meta.get("effective_date"),
|
||||
"deprecated_date": first_meta.get("deprecated_date"),
|
||||
"deprecated_reason": first_meta.get("deprecated_reason"),
|
||||
"security_level": first_meta.get("security_level", "public")
|
||||
}
|
||||
|
||||
def list_documents_by_status(
|
||||
self,
|
||||
kb_name: str,
|
||||
status: str = None
|
||||
) -> List[Dict]:
|
||||
"""按状态列出文档"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return []
|
||||
|
||||
result = collection.get()
|
||||
|
||||
doc_info = {}
|
||||
for meta in result.get('metadatas', []):
|
||||
source = meta.get('source', 'unknown')
|
||||
chunk_status = meta.get('status', 'active')
|
||||
|
||||
if source not in doc_info:
|
||||
doc_info[source] = {
|
||||
"source": source,
|
||||
"chunks": 0,
|
||||
"status_counts": {},
|
||||
"collection": kb_name
|
||||
}
|
||||
|
||||
doc_info[source]["chunks"] += 1
|
||||
doc_info[source]["status_counts"][chunk_status] = \
|
||||
doc_info[source]["status_counts"].get(chunk_status, 0) + 1
|
||||
|
||||
result_list = []
|
||||
for doc in doc_info.values():
|
||||
counts = doc["status_counts"]
|
||||
if counts.get("deprecated", 0) > counts.get("active", 0):
|
||||
doc["status"] = "deprecated"
|
||||
elif counts.get("superseded", 0) > 0:
|
||||
doc["status"] = "superseded"
|
||||
else:
|
||||
doc["status"] = "active"
|
||||
|
||||
if status is None or doc["status"] == status:
|
||||
result_list.append(doc)
|
||||
|
||||
return result_list
|
||||
|
||||
# add_file_to_kb 方法较长,暂时保留在 manager.py 中
|
||||
# 后续可以拆分到单独的 document_processing.py
|
||||
361
knowledge/document_versions.py
Normal file
361
knowledge/document_versions.py
Normal file
@@ -0,0 +1,361 @@
|
||||
"""
|
||||
文档版本查询模块(简化版)
|
||||
|
||||
保留核心的版本查询功能,删除未使用的生命周期管理功能。
|
||||
|
||||
功能:
|
||||
1. 查询文档版本历史
|
||||
2. 获取当前生效版本
|
||||
3. 记录版本变更日志
|
||||
|
||||
使用方式:
|
||||
from knowledge.document_versions import DocumentVersionQuery
|
||||
|
||||
query = DocumentVersionQuery()
|
||||
|
||||
# 获取版本历史
|
||||
history = query.get_document_history("public_kb", "报销制度.pdf")
|
||||
|
||||
# 获取生效版本
|
||||
active = query.get_active_version("public_kb", "报销制度.pdf")
|
||||
"""
|
||||
|
||||
import logging
|
||||
from enum import Enum
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Optional, List, Dict
|
||||
|
||||
from data.db import get_connection
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ==================== 枚举与数据类 ====================
|
||||
|
||||
class DocumentStatus(Enum):
|
||||
"""文档状态"""
|
||||
DRAFT = "draft" # 草稿
|
||||
ACTIVE = "active" # 生效中
|
||||
DEPRECATED = "deprecated" # 已废止
|
||||
SUPERSEDED = "superseded" # 被替代
|
||||
|
||||
|
||||
@dataclass
|
||||
class DocumentVersionInfo:
|
||||
"""文档版本信息"""
|
||||
document_id: str
|
||||
collection: str
|
||||
version: str
|
||||
status: DocumentStatus
|
||||
effective_date: Optional[str] = None
|
||||
deprecated_date: Optional[str] = None
|
||||
deprecated_reason: Optional[str] = None
|
||||
change_summary: Optional[str] = None
|
||||
supersedes: Optional[str] = None
|
||||
created_at: Optional[str] = None
|
||||
created_by: Optional[str] = None
|
||||
chunk_count: int = 0
|
||||
|
||||
def to_dict(self) -> Dict:
|
||||
"""转换为字典"""
|
||||
return {
|
||||
"document_id": self.document_id,
|
||||
"collection": self.collection,
|
||||
"version": self.version,
|
||||
"status": self.status.value if isinstance(self.status, DocumentStatus) else self.status,
|
||||
"effective_date": self.effective_date,
|
||||
"deprecated_date": self.deprecated_date,
|
||||
"deprecated_reason": self.deprecated_reason,
|
||||
"change_summary": self.change_summary,
|
||||
"supersedes": self.supersedes,
|
||||
"created_at": self.created_at,
|
||||
"created_by": self.created_by,
|
||||
"chunk_count": self.chunk_count
|
||||
}
|
||||
|
||||
|
||||
# ==================== 文档版本查询 ====================
|
||||
|
||||
class DocumentVersionQuery:
|
||||
"""文档版本查询(简化版)"""
|
||||
|
||||
def __init__(self):
|
||||
"""初始化"""
|
||||
pass
|
||||
|
||||
def get_document_history(
|
||||
self,
|
||||
collection: str,
|
||||
document_id: str,
|
||||
limit: int = 10
|
||||
) -> List[DocumentVersionInfo]:
|
||||
"""
|
||||
获取文档版本历史
|
||||
|
||||
Args:
|
||||
collection: 向量库名称
|
||||
document_id: 文档ID
|
||||
limit: 返回数量限制
|
||||
|
||||
Returns:
|
||||
版本信息列表(按时间倒序)
|
||||
"""
|
||||
try:
|
||||
with get_connection("knowledge") as conn:
|
||||
cursor = conn.execute(
|
||||
"""
|
||||
SELECT
|
||||
document_id, collection, version, status,
|
||||
effective_date, deprecated_date, deprecated_reason,
|
||||
change_summary, supersedes, created_at, created_by,
|
||||
chunk_count
|
||||
FROM document_versions
|
||||
WHERE collection = ? AND document_id = ?
|
||||
ORDER BY created_at DESC
|
||||
LIMIT ?
|
||||
""",
|
||||
(collection, document_id, limit)
|
||||
)
|
||||
|
||||
versions = []
|
||||
for row in cursor.fetchall():
|
||||
versions.append(DocumentVersionInfo(
|
||||
document_id=row[0],
|
||||
collection=row[1],
|
||||
version=row[2],
|
||||
status=DocumentStatus(row[3]) if row[3] else DocumentStatus.ACTIVE,
|
||||
effective_date=row[4],
|
||||
deprecated_date=row[5],
|
||||
deprecated_reason=row[6],
|
||||
change_summary=row[7],
|
||||
supersedes=row[8],
|
||||
created_at=row[9],
|
||||
created_by=row[10],
|
||||
chunk_count=row[11] or 0
|
||||
))
|
||||
|
||||
return versions
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取文档历史失败: {e}")
|
||||
return []
|
||||
|
||||
def get_active_version(
|
||||
self,
|
||||
collection: str,
|
||||
document_id: str
|
||||
) -> Optional[DocumentVersionInfo]:
|
||||
"""
|
||||
获取当前生效版本
|
||||
|
||||
Args:
|
||||
collection: 向量库名称
|
||||
document_id: 文档ID
|
||||
|
||||
Returns:
|
||||
生效版本信息,不存在则返回 None
|
||||
"""
|
||||
try:
|
||||
with get_connection("knowledge") as conn:
|
||||
cursor = conn.execute(
|
||||
"""
|
||||
SELECT
|
||||
document_id, collection, version, status,
|
||||
effective_date, deprecated_date, deprecated_reason,
|
||||
change_summary, supersedes, created_at, created_by,
|
||||
chunk_count
|
||||
FROM document_versions
|
||||
WHERE collection = ? AND document_id = ? AND status = 'active'
|
||||
ORDER BY created_at DESC
|
||||
LIMIT 1
|
||||
""",
|
||||
(collection, document_id)
|
||||
)
|
||||
|
||||
row = cursor.fetchone()
|
||||
if row:
|
||||
return DocumentVersionInfo(
|
||||
document_id=row[0],
|
||||
collection=row[1],
|
||||
version=row[2],
|
||||
status=DocumentStatus(row[3]) if row[3] else DocumentStatus.ACTIVE,
|
||||
effective_date=row[4],
|
||||
deprecated_date=row[5],
|
||||
deprecated_reason=row[6],
|
||||
change_summary=row[7],
|
||||
supersedes=row[8],
|
||||
created_at=row[9],
|
||||
created_by=row[10],
|
||||
chunk_count=row[11] or 0
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取生效版本失败: {e}")
|
||||
return None
|
||||
|
||||
def get_next_version(self, collection: str, document_id: str) -> str:
|
||||
"""
|
||||
获取下一个版本号
|
||||
|
||||
Args:
|
||||
collection: 向量库名称
|
||||
document_id: 文档ID
|
||||
|
||||
Returns:
|
||||
版本号字符串,如 'v1', 'v2', 'v3'
|
||||
"""
|
||||
try:
|
||||
with get_connection("knowledge") as conn:
|
||||
cursor = conn.execute(
|
||||
"""
|
||||
SELECT version FROM document_versions
|
||||
WHERE collection = ? AND document_id = ?
|
||||
ORDER BY created_at DESC
|
||||
LIMIT 1
|
||||
""",
|
||||
(collection, document_id)
|
||||
)
|
||||
row = cursor.fetchone()
|
||||
|
||||
if row and row[0]:
|
||||
# 解析现有版本号
|
||||
current = row[0]
|
||||
if current.startswith('v') and current[1:].isdigit():
|
||||
next_num = int(current[1:]) + 1
|
||||
else:
|
||||
# 非标准版本号,从1开始
|
||||
next_num = 1
|
||||
else:
|
||||
next_num = 1
|
||||
|
||||
return f"v{next_num}"
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"获取版本号失败: {e},使用默认 v1")
|
||||
return "v1"
|
||||
|
||||
def create_version_record(
|
||||
self,
|
||||
collection: str,
|
||||
document_id: str,
|
||||
version: str = None,
|
||||
status: str = "active",
|
||||
change_summary: str = "",
|
||||
supersedes: str = None,
|
||||
created_by: str = "",
|
||||
chunk_count: int = 0
|
||||
):
|
||||
"""
|
||||
创建版本记录
|
||||
|
||||
Args:
|
||||
collection: 向量库名称
|
||||
document_id: 文档ID
|
||||
version: 版本号(可选,不传则自动生成 v1, v2, v3...)
|
||||
status: 状态
|
||||
change_summary: 变更摘要
|
||||
supersedes: 替代的旧版本
|
||||
created_by: 创建者
|
||||
chunk_count: chunk数量
|
||||
"""
|
||||
# 自动生成版本号
|
||||
if version is None:
|
||||
version = self.get_next_version(collection, document_id)
|
||||
|
||||
try:
|
||||
with get_connection("knowledge") as conn:
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT OR REPLACE INTO document_versions
|
||||
(document_id, collection, version, status, effective_date,
|
||||
change_summary, supersedes, created_at, created_by, chunk_count)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
document_id,
|
||||
collection,
|
||||
version,
|
||||
status,
|
||||
datetime.now().isoformat(),
|
||||
change_summary,
|
||||
supersedes,
|
||||
datetime.now().isoformat(),
|
||||
created_by,
|
||||
chunk_count
|
||||
)
|
||||
)
|
||||
conn.commit()
|
||||
logger.info(f"创建版本记录: {collection}/{document_id} {version}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"创建版本记录失败: {e}")
|
||||
raise
|
||||
|
||||
def log_version_change(
|
||||
self,
|
||||
collection: str,
|
||||
document_id: str,
|
||||
change_type: str,
|
||||
old_version: str = None,
|
||||
new_version: str = None,
|
||||
old_status: str = None,
|
||||
new_status: str = None,
|
||||
reason: str = "",
|
||||
changed_by: str = ""
|
||||
):
|
||||
"""
|
||||
记录版本变更日志
|
||||
|
||||
Args:
|
||||
collection: 向量库名称
|
||||
document_id: 文档ID
|
||||
change_type: 变更类型(update/deprecate/restore)
|
||||
old_version: 旧版本号
|
||||
new_version: 新版本号
|
||||
old_status: 旧状态
|
||||
new_status: 新状态
|
||||
reason: 变更原因
|
||||
changed_by: 操作者
|
||||
"""
|
||||
try:
|
||||
with get_connection("knowledge") as conn:
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT INTO version_change_logs
|
||||
(document_id, collection, old_version, new_version,
|
||||
old_status, new_status, change_type, reason, changed_by, created_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
document_id,
|
||||
collection,
|
||||
old_version,
|
||||
new_version,
|
||||
old_status,
|
||||
new_status,
|
||||
change_type,
|
||||
reason,
|
||||
changed_by,
|
||||
datetime.now().isoformat()
|
||||
)
|
||||
)
|
||||
conn.commit()
|
||||
logger.info(f"记录版本变更: {collection}/{document_id} {change_type}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"记录版本变更失败: {e}")
|
||||
|
||||
|
||||
# ==================== 工厂函数 ====================
|
||||
|
||||
_version_query_instance = None
|
||||
|
||||
def get_version_query() -> DocumentVersionQuery:
|
||||
"""获取文档版本查询实例(单例)"""
|
||||
global _version_query_instance
|
||||
if _version_query_instance is None:
|
||||
_version_query_instance = DocumentVersionQuery()
|
||||
return _version_query_instance
|
||||
125
knowledge/index.py
Normal file
125
knowledge/index.py
Normal file
@@ -0,0 +1,125 @@
|
||||
"""
|
||||
知识库管理器 - BM25索引管理 Mixin
|
||||
|
||||
提供 BM25 关键词检索索引的管理功能,包括索引的加载、保存和重建。
|
||||
|
||||
BM25 是一种基于概率检索模型的排序函数,适合中文关键词检索。
|
||||
每个向量库对应一个独立的 BM25 索引,使用 jieba 进行中文分词。
|
||||
|
||||
主要方法:
|
||||
- get_bm25_index: 获取或加载 BM25 索引
|
||||
- save_bm25_index: 保存 BM25 索引到磁盘
|
||||
- rebuild_bm25_index: 从向量库重建 BM25 索引
|
||||
"""
|
||||
|
||||
import os
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
from .base import BM25Index
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class IndexMixin:
|
||||
"""
|
||||
BM25索引管理 Mixin
|
||||
|
||||
提供 BM25 索引的生命周期管理,支持:
|
||||
- 索引的懒加载(首次访问时从磁盘加载)
|
||||
- 索引的持久化(保存到 .pkl 文件)
|
||||
- 索引的重建(从向量库同步)
|
||||
|
||||
依赖属性(需由主类提供):
|
||||
- self.bm25_base_path: BM25 索引存储路径
|
||||
- self._bm25_indexes: BM25 索引缓存字典
|
||||
"""
|
||||
|
||||
def get_bm25_index(self, kb_name: str) -> BM25Index:
|
||||
"""
|
||||
获取或加载 BM25 索引
|
||||
|
||||
如果索引已在内存中,直接返回。
|
||||
否则尝试从磁盘加载,加载失败则创建空索引。
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称
|
||||
|
||||
Returns:
|
||||
BM25Index 实例
|
||||
|
||||
Note:
|
||||
索引文件路径: {bm25_base_path}/{kb_name}.pkl
|
||||
"""
|
||||
if kb_name not in self._bm25_indexes:
|
||||
self._bm25_indexes[kb_name] = BM25Index()
|
||||
bm25_path = os.path.join(self.bm25_base_path, f"{kb_name}.pkl")
|
||||
if os.path.exists(bm25_path):
|
||||
self._bm25_indexes[kb_name].load(bm25_path)
|
||||
return self._bm25_indexes[kb_name]
|
||||
|
||||
def save_bm25_index(self, kb_name: str) -> None:
|
||||
"""
|
||||
保存 BM25 索引到磁盘
|
||||
|
||||
将指定向量库的 BM25 索引序列化保存到 .pkl 文件。
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称
|
||||
"""
|
||||
if kb_name in self._bm25_indexes:
|
||||
bm25_path = os.path.join(self.bm25_base_path, f"{kb_name}.pkl")
|
||||
self._bm25_indexes[kb_name].save(bm25_path)
|
||||
logger.info(f"保存 BM25 索引: {kb_name}")
|
||||
|
||||
def rebuild_bm25_index(self, kb_name: str) -> bool:
|
||||
"""
|
||||
重建 BM25 索引
|
||||
|
||||
从 ChromaDB 向量库中读取所有文档,重新构建 BM25 索引。
|
||||
用于文档变更后同步索引状态。
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称
|
||||
|
||||
Returns:
|
||||
重建成功返回 True,失败返回 False
|
||||
|
||||
Warning:
|
||||
大量文档时重建可能耗时较长。
|
||||
"""
|
||||
try:
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return False
|
||||
|
||||
result = collection.get()
|
||||
|
||||
bm25_index = BM25Index()
|
||||
if result['ids']:
|
||||
bm25_index.add_documents(
|
||||
ids=result['ids'],
|
||||
documents=result['documents'],
|
||||
metadatas=result['metadatas']
|
||||
)
|
||||
|
||||
self._bm25_indexes[kb_name] = bm25_index
|
||||
self.save_bm25_index(kb_name)
|
||||
|
||||
logger.info(f"重建 BM25 索引: {kb_name}, 文档数: {len(result['ids'])}")
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"重建 BM25 索引失败: {e}")
|
||||
return False
|
||||
|
||||
def _rebuild_bm25_index(self, kb_name: str) -> None:
|
||||
"""
|
||||
重建指定向量库的 BM25 索引(内部方法)
|
||||
|
||||
兼容旧代码的内部方法,直接调用 rebuild_bm25_index。
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称
|
||||
"""
|
||||
self.rebuild_bm25_index(kb_name)
|
||||
287
knowledge/lazy_enhance.py
Normal file
287
knowledge/lazy_enhance.py
Normal file
@@ -0,0 +1,287 @@
|
||||
"""
|
||||
懒加载增强模块(Phase 4)
|
||||
|
||||
按需调用 LLM/VLM 生成表格摘要和图片描述
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 缓存目录(扁平化)
|
||||
VLM_CACHE_DIR = Path(".data/cache/vlm")
|
||||
LLM_CACHE_DIR = Path(".data/cache/llm")
|
||||
|
||||
|
||||
def compute_file_hash(file_path: str) -> str:
|
||||
"""计算文件哈希"""
|
||||
try:
|
||||
with open(file_path, 'rb') as f:
|
||||
return hashlib.md5(f.read()).hexdigest()
|
||||
except Exception as e:
|
||||
logger.warning(f"计算文件哈希失败: {e}")
|
||||
return hashlib.md5(file_path.encode()).hexdigest()
|
||||
|
||||
|
||||
async def lazy_vlm_description(chunk_id: str, image_path: str, kb_name: str, metadata: dict = None) -> str:
|
||||
"""
|
||||
懒加载 VLM 描述
|
||||
|
||||
触发条件:图片切片被检索命中
|
||||
|
||||
Args:
|
||||
chunk_id: 切片 ID
|
||||
image_path: 图片路径(相对路径或绝对路径)
|
||||
kb_name: 知识库名称
|
||||
metadata: 图片元数据(包含 section、page、caption、上下文等)
|
||||
|
||||
Returns:
|
||||
VLM 生成的图片描述
|
||||
"""
|
||||
import os
|
||||
from knowledge.manager import get_kb_manager
|
||||
|
||||
# 构建完整图片路径
|
||||
if not os.path.isabs(image_path):
|
||||
full_image_path = os.path.join('.data/images', image_path)
|
||||
else:
|
||||
full_image_path = image_path
|
||||
|
||||
# 1. 检查缓存
|
||||
img_hash = compute_file_hash(full_image_path)
|
||||
cache_file = VLM_CACHE_DIR / f"{img_hash}.txt"
|
||||
if cache_file.exists():
|
||||
logger.info(f"VLM 缓存命中: {image_path}")
|
||||
return cache_file.read_text(encoding='utf-8')
|
||||
|
||||
# 2. 调用 VLM(传入元数据)
|
||||
logger.info(f"VLM 懒加载: {image_path}")
|
||||
kb_manager = get_kb_manager()
|
||||
description = kb_manager._generate_image_description(full_image_path, metadata=metadata)
|
||||
|
||||
# 3. 写入缓存
|
||||
VLM_CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
cache_file.write_text(description, encoding='utf-8')
|
||||
|
||||
# 4. 更新向量库(metadata + embedding)
|
||||
try:
|
||||
collection = kb_manager.get_collection(kb_name)
|
||||
result = collection.get(ids=[chunk_id], include=['metadatas'])
|
||||
if result['metadatas']:
|
||||
# 更新 metadata
|
||||
new_metadata = {
|
||||
**result['metadatas'][0],
|
||||
'has_vlm_desc': True,
|
||||
'vlm_desc': description
|
||||
}
|
||||
|
||||
# 更新 embedding(使用 VLM 描述重新计算向量)
|
||||
# 这样 VLM 描述中的关键词(如"发电量")才能参与相似度检索
|
||||
embedding_model = kb_manager.embedding_model
|
||||
if embedding_model:
|
||||
new_vector = embedding_model.encode(description).tolist()
|
||||
if isinstance(new_vector[0], list):
|
||||
new_vector = new_vector[0]
|
||||
|
||||
collection.update(
|
||||
ids=[chunk_id],
|
||||
metadatas=[new_metadata],
|
||||
embeddings=[new_vector],
|
||||
documents=[description] # 同时更新 document 字段
|
||||
)
|
||||
logger.info(f"已更新向量库 embedding: {chunk_id}")
|
||||
else:
|
||||
# 无 embedding 模型时只更新 metadata
|
||||
collection.update(
|
||||
ids=[chunk_id],
|
||||
metadatas=[new_metadata]
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"更新向量库失败: {e}")
|
||||
|
||||
return description
|
||||
|
||||
|
||||
async def lazy_table_summary(chunk_id: str, table_md: str, kb_name: str) -> str:
|
||||
"""
|
||||
懒加载表格摘要
|
||||
|
||||
触发条件:表格切片被检索命中且相关性 > 0.7
|
||||
|
||||
Args:
|
||||
chunk_id: 切片 ID
|
||||
table_md: 表格 Markdown 内容
|
||||
kb_name: 知识库名称
|
||||
|
||||
Returns:
|
||||
LLM 生成的表格摘要
|
||||
"""
|
||||
from knowledge.manager import get_kb_manager
|
||||
|
||||
# 1. 检查缓存
|
||||
table_hash = hashlib.md5(table_md.encode()).hexdigest()
|
||||
cache_file = LLM_CACHE_DIR / f"{table_hash}.txt"
|
||||
if cache_file.exists():
|
||||
logger.info(f"LLM 缓存命中: {chunk_id}")
|
||||
return cache_file.read_text(encoding='utf-8')
|
||||
|
||||
# 2. 调用 LLM
|
||||
logger.info(f"LLM 懒加载: {chunk_id}")
|
||||
kb_manager = get_kb_manager()
|
||||
summary = kb_manager._generate_table_summary(table_md, None)
|
||||
|
||||
# 3. 写入缓存
|
||||
LLM_CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
cache_file.write_text(summary, encoding='utf-8')
|
||||
|
||||
# 4. 更新向量库(可选)
|
||||
try:
|
||||
collection = kb_manager.get_collection(kb_name)
|
||||
result = collection.get(ids=[chunk_id], include=['metadatas'])
|
||||
if result['metadatas']:
|
||||
# 新增摘要切片
|
||||
embedding_model = kb_manager.embedding_model
|
||||
vector = embedding_model.encode(summary).tolist()
|
||||
if isinstance(vector[0], list):
|
||||
vector = vector[0]
|
||||
|
||||
collection.add(
|
||||
ids=[f"{chunk_id}_summary"],
|
||||
embeddings=[vector],
|
||||
documents=[summary],
|
||||
metadatas=[{
|
||||
**result['metadatas'][0],
|
||||
'is_summary': True,
|
||||
'original_doc_id': chunk_id
|
||||
}]
|
||||
)
|
||||
# 更新原切片标记
|
||||
collection.update(
|
||||
ids=[chunk_id],
|
||||
metadatas=[{**result['metadatas'][0], 'has_summary': True}]
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"更新向量库失败: {e}")
|
||||
|
||||
return summary
|
||||
|
||||
|
||||
async def enhance_retrieved_chunks(contexts: list, query: str, kb_name: str):
|
||||
"""
|
||||
检索后增强:按需调用 LLM/VLM
|
||||
|
||||
Args:
|
||||
contexts: 检索上下文列表
|
||||
query: 用户查询
|
||||
kb_name: 知识库名称
|
||||
"""
|
||||
for ctx in contexts:
|
||||
meta = ctx.get('meta', {})
|
||||
chunk_type = meta.get('chunk_type', 'text')
|
||||
image_path = meta.get('image_path', '')
|
||||
|
||||
# 图片切片:懒加载 VLM 描述
|
||||
if chunk_type in ('image', 'chart') and not meta.get('has_vlm_desc'):
|
||||
if image_path:
|
||||
try:
|
||||
# 从 doc 字段中提取图号(上下文可能包含"见图2.5"等)
|
||||
doc_text = ctx.get('doc', '')
|
||||
import re
|
||||
|
||||
# 提取图号(从前文/后文中)
|
||||
figure_number = ""
|
||||
# 匹配 "见图2.5"、"图2.5"、"见图 2.5" 等
|
||||
fig_match = re.search(r'[见如]?图\s*(\d+\.?\d*)', doc_text)
|
||||
if fig_match:
|
||||
figure_number = fig_match.group(1)
|
||||
|
||||
# 如果 doc 中没有,尝试从 section 中提取
|
||||
section = meta.get('section') or meta.get('section_path', '')
|
||||
if not figure_number and section:
|
||||
fig_match = re.search(r'[见如]?图\s*(\d+\.?\d*)', section)
|
||||
if fig_match:
|
||||
figure_number = fig_match.group(1)
|
||||
|
||||
# 传入图片元数据,增强 VLM 描述
|
||||
image_metadata = {
|
||||
'section': section,
|
||||
'page': meta.get('page'),
|
||||
'caption': meta.get('caption', ''),
|
||||
'source': meta.get('source', ''),
|
||||
'figure_number': figure_number, # 添加提取的图号
|
||||
'doc_text': doc_text # 添加完整文档文本
|
||||
}
|
||||
vlm_desc = await lazy_vlm_description(
|
||||
meta.get('id', ''),
|
||||
image_path,
|
||||
kb_name,
|
||||
metadata=image_metadata
|
||||
)
|
||||
ctx['doc'] = vlm_desc
|
||||
ctx['vlm_enhanced'] = True
|
||||
except Exception as e:
|
||||
logger.warning(f"VLM 懒加载失败: {e}")
|
||||
|
||||
# 表格切片:同时处理摘要和关联图片的 VLM 描述
|
||||
elif chunk_type == 'table':
|
||||
doc_text = ctx.get('doc', '')
|
||||
|
||||
# 1. 懒加载表格摘要(高分切片)
|
||||
if not meta.get('has_summary'):
|
||||
score = meta.get('score', 0)
|
||||
if score > 0.7: # 只对高相关表格生成摘要
|
||||
try:
|
||||
summary = await lazy_table_summary(
|
||||
meta.get('id', ''),
|
||||
doc_text,
|
||||
kb_name
|
||||
)
|
||||
# 摘要作为补充信息
|
||||
ctx['summary'] = summary
|
||||
ctx['llm_enhanced'] = True
|
||||
except Exception as e:
|
||||
logger.warning(f"表格摘要懒加载失败: {e}")
|
||||
|
||||
# 2. 表格有关联图片时,懒加载 VLM 描述
|
||||
if image_path and not meta.get('has_vlm_desc'):
|
||||
try:
|
||||
import re
|
||||
|
||||
# 提取表号(如 "表2.2"、"见表2.1")
|
||||
table_number = ""
|
||||
# 匹配 "表2.2"、"见表2.2"、"见表 2.2" 等
|
||||
table_match = re.search(r'[见如]?表\s*(\d+\.?\d*)', doc_text)
|
||||
if table_match:
|
||||
table_number = table_match.group(1)
|
||||
|
||||
# 如果 doc 中没有,尝试从 section 中提取
|
||||
section = meta.get('section') or meta.get('section_path', '')
|
||||
if not table_number and section:
|
||||
table_match = re.search(r'[见如]?表\s*(\d+\.?\d*)', section)
|
||||
if table_match:
|
||||
table_number = table_match.group(1)
|
||||
|
||||
# 构建表格图片元数据
|
||||
table_image_metadata = {
|
||||
'section': section,
|
||||
'page': meta.get('page'),
|
||||
'caption': meta.get('caption', ''),
|
||||
'source': meta.get('source', ''),
|
||||
'table_number': table_number, # 表号
|
||||
'figure_number': table_number, # 兼容字段
|
||||
'doc_text': doc_text,
|
||||
'is_table': True # 标记为表格图片
|
||||
}
|
||||
vlm_desc = await lazy_vlm_description(
|
||||
meta.get('id', ''),
|
||||
image_path,
|
||||
kb_name,
|
||||
metadata=table_image_metadata
|
||||
)
|
||||
# 表格图片描述作为补充信息
|
||||
ctx['image_description'] = vlm_desc
|
||||
ctx['vlm_enhanced'] = True
|
||||
except Exception as e:
|
||||
logger.warning(f"表格图片 VLM 懒加载失败: {e}")
|
||||
717
knowledge/manager.py
Normal file
717
knowledge/manager.py
Normal file
@@ -0,0 +1,717 @@
|
||||
"""
|
||||
多向量库管理器 - 支持按部门/权限隔离的向量知识库
|
||||
|
||||
功能:
|
||||
1. 多向量库管理 - 创建、删除、列举向量库
|
||||
2. 多 BM25 索引管理 - 每个向量库独立的 BM25 索引
|
||||
3. 权限过滤 - 根据用户角色和部门返回可访问的向量库
|
||||
4. 并行检索 - 支持同时检索多个向量库
|
||||
|
||||
向量库命名规范:
|
||||
- public_kb: 公开知识库,所有人可访问
|
||||
- dept_{部门名}: 部门知识库,如 dept_finance, dept_hr, dept_tech
|
||||
|
||||
使用方式:
|
||||
from knowledge.manager import KnowledgeBaseManager
|
||||
|
||||
kb_manager = KnowledgeBaseManager()
|
||||
|
||||
# 获取向量库
|
||||
collection = kb_manager.get_collection("dept_finance")
|
||||
|
||||
# 列出所有向量库
|
||||
collections = kb_manager.list_collections()
|
||||
|
||||
# 获取用户可访问的向量库
|
||||
accessible = kb_manager.get_accessible_collections("manager", "finance")
|
||||
"""
|
||||
|
||||
import os
|
||||
import json
|
||||
import threading
|
||||
from typing import List, Dict, Optional, Tuple
|
||||
from pathlib import Path
|
||||
import logging
|
||||
|
||||
import chromadb
|
||||
|
||||
# 从 base.py 导入基础类和常量
|
||||
from .base import (
|
||||
BM25Index,
|
||||
CollectionInfo,
|
||||
SearchResult,
|
||||
_get_doc_type,
|
||||
_extract_figure_number,
|
||||
_extract_section,
|
||||
_build_semantic_content_for_text,
|
||||
_build_semantic_content_for_table,
|
||||
CHROMA_DB_BASE_PATH,
|
||||
BM25_INDEX_BASE_PATH,
|
||||
KB_METADATA_FILE,
|
||||
PUBLIC_KB_NAME,
|
||||
DEFAULT_DEPARTMENTS,
|
||||
DEPARTMENT_NAME_MAP,
|
||||
normalize_department_name,
|
||||
)
|
||||
|
||||
# 导入 Mixin 类
|
||||
from .collection import CollectionMixin
|
||||
from .document import DocumentMixin
|
||||
from .chunk import ChunkMixin
|
||||
from .index import IndexMixin
|
||||
from .search import SearchMixin
|
||||
from .processing import ProcessingMixin
|
||||
from .permission import PermissionMixin
|
||||
|
||||
# 导入 LLM 工具函数
|
||||
from core.llm_utils import call_llm
|
||||
|
||||
# 设置日志
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ==================== 多向量库管理器 ====================
|
||||
|
||||
class KnowledgeBaseManager(
|
||||
CollectionMixin,
|
||||
DocumentMixin,
|
||||
ChunkMixin,
|
||||
IndexMixin,
|
||||
SearchMixin,
|
||||
ProcessingMixin,
|
||||
PermissionMixin
|
||||
):
|
||||
"""
|
||||
多向量库管理器
|
||||
|
||||
管理多个独立的 ChromaDB 集合,每个集合对应一个知识库。
|
||||
支持按部门隔离,每个部门有独立的向量库和 BM25 索引。
|
||||
|
||||
通过 Mixin 模式组合功能:
|
||||
- CollectionMixin: 向量库管理
|
||||
- DocumentMixin: 文档管理
|
||||
- ChunkMixin: 切片管理
|
||||
- IndexMixin: BM25 索引管理
|
||||
- SearchMixin: 检索功能
|
||||
- ProcessingMixin: 图片/表格处理
|
||||
- PermissionMixin: 权限控制
|
||||
"""
|
||||
|
||||
def __init__(self, base_path: str = None, bm25_base_path: str = None):
|
||||
"""
|
||||
初始化
|
||||
|
||||
Args:
|
||||
base_path: 向量库存储路径
|
||||
bm25_base_path: BM25 索引存储路径
|
||||
"""
|
||||
self.base_path = base_path or CHROMA_DB_BASE_PATH
|
||||
self.bm25_base_path = bm25_base_path or BM25_INDEX_BASE_PATH
|
||||
|
||||
# 缓存
|
||||
self._collections: Dict[str, chromadb.Collection] = {}
|
||||
self._bm25_indexes: Dict[str, BM25Index] = {}
|
||||
self._clients: Dict[str, chromadb.PersistentClient] = {}
|
||||
self._lock = threading.Lock()
|
||||
|
||||
# 确保目录存在
|
||||
os.makedirs(self.base_path, exist_ok=True)
|
||||
os.makedirs(self.bm25_base_path, exist_ok=True)
|
||||
|
||||
# 加载元数据
|
||||
self._metadata = self._load_metadata()
|
||||
|
||||
# 初始化公开知识库
|
||||
self._ensure_public_kb()
|
||||
|
||||
# 扫描并列出所有已存在的向量库
|
||||
existing_kbs = []
|
||||
if os.path.exists(self.base_path):
|
||||
for item in os.listdir(self.base_path):
|
||||
if os.path.isdir(os.path.join(self.base_path, item)) and not item.startswith('.'):
|
||||
if os.path.exists(os.path.join(self.base_path, item, 'chroma.sqlite3')):
|
||||
existing_kbs.append(item)
|
||||
logger.info(f"知识库管理器初始化完成,路径: {self.base_path},发现 {len(existing_kbs)} 个向量库: {existing_kbs}")
|
||||
|
||||
def _load_metadata(self) -> dict:
|
||||
"""加载元数据"""
|
||||
metadata_path = os.path.join(self.base_path, KB_METADATA_FILE)
|
||||
if os.path.exists(metadata_path):
|
||||
try:
|
||||
with open(metadata_path, 'r', encoding='utf-8') as f:
|
||||
return json.load(f)
|
||||
except Exception as e:
|
||||
logger.error(f"加载元数据失败: {e}")
|
||||
return {"collections": {}}
|
||||
|
||||
def _save_metadata(self):
|
||||
"""保存元数据"""
|
||||
metadata_path = os.path.join(self.base_path, KB_METADATA_FILE)
|
||||
try:
|
||||
with open(metadata_path, 'w', encoding='utf-8') as f:
|
||||
json.dump(self._metadata, f, ensure_ascii=False, indent=2)
|
||||
except Exception as e:
|
||||
logger.error(f"保存元数据失败: {e}")
|
||||
|
||||
def _ensure_public_kb(self):
|
||||
"""确保公开知识库存在"""
|
||||
if PUBLIC_KB_NAME not in self._metadata.get("collections", {}):
|
||||
self.create_collection(
|
||||
PUBLIC_KB_NAME,
|
||||
display_name="公开知识库",
|
||||
department="",
|
||||
description="所有人可访问的公开文档"
|
||||
)
|
||||
|
||||
# ==================== 文件添加 ====================
|
||||
|
||||
def add_file_to_kb(
|
||||
self,
|
||||
kb_name: str,
|
||||
filepath: str,
|
||||
embedding_model=None,
|
||||
extra_metadata: dict = None,
|
||||
enable_table_summary: bool = True,
|
||||
enable_image_description: bool = False,
|
||||
file_content: bytes = None
|
||||
) -> int:
|
||||
"""
|
||||
添加文件到指定向量库(v6 支持企业文件系统)
|
||||
|
||||
使用统一的 parse_document() 入口,支持:
|
||||
- PDF/DOCX/PPTX/图片 → MinerU 解析
|
||||
- XLSX/XLS → Pandas 解析
|
||||
- TXT → 文本解析
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称
|
||||
filepath: 文件路径(相对路径或绝对路径)
|
||||
embedding_model: 向量模型(可选,默认使用 engine 的)
|
||||
extra_metadata: 额外的元数据(如 status, version 等)
|
||||
enable_table_summary: 是否启用表格摘要管道(LLM 生成摘要)
|
||||
enable_image_description: 是否启用图片描述管道(VLM 生成描述)
|
||||
file_content: 文件二进制内容(可选,用于企业文件系统集成)
|
||||
|
||||
Returns:
|
||||
添加的片段数量
|
||||
"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
raise ValueError(f"向量库 '{kb_name}' 不存在")
|
||||
|
||||
# 解析文档
|
||||
from parsers import parse_document
|
||||
result = parse_document(filepath, file_content=file_content)
|
||||
|
||||
if not result:
|
||||
logger.warning(f"文档解析结果为空: {filepath}")
|
||||
return 0
|
||||
|
||||
chunks = result.get('chunks', [])
|
||||
if not chunks:
|
||||
logger.warning(f"文档解析后无有效切片: {filepath}")
|
||||
return 0
|
||||
|
||||
# 合并跨页表格
|
||||
chunks = self._merge_cross_page_tables(chunks)
|
||||
|
||||
# 准备向量模型
|
||||
if embedding_model is None:
|
||||
from core.engine import get_engine
|
||||
engine = get_engine()
|
||||
if not engine._initialized:
|
||||
engine.initialize()
|
||||
embedding_model = engine.embedding_model
|
||||
|
||||
# 处理切片
|
||||
ids = []
|
||||
documents = []
|
||||
metadatas = []
|
||||
embeddings = []
|
||||
|
||||
filename = Path(filepath).name
|
||||
doc_type = _get_doc_type(filename)
|
||||
|
||||
for i, chunk in enumerate(chunks):
|
||||
chunk_id = f"{filename}_{i}"
|
||||
|
||||
# 获取切片类型(优先从 chunk 直接获取,兼容 MinerUChunk 和其他格式)
|
||||
chunk_type = getattr(chunk, 'chunk_type', None)
|
||||
if not chunk_type:
|
||||
page_info = getattr(chunk, 'page_info', {}) or {}
|
||||
chunk_type = page_info.get('chunk_type', 'text')
|
||||
|
||||
# 获取页码信息(兼容 MinerUChunk 和 page_info 格式)
|
||||
page_start = getattr(chunk, 'page_start', None)
|
||||
page_end = getattr(chunk, 'page_end', None)
|
||||
if page_start is None:
|
||||
page_info = getattr(chunk, 'page_info', {}) or {}
|
||||
page_start = page_info.get('page', 0)
|
||||
page_end = page_info.get('page_end', page_start)
|
||||
page_start = page_start or 0
|
||||
page_end = page_end or page_start
|
||||
|
||||
# 获取章节信息
|
||||
section_path = getattr(chunk, 'section_path', '') or ''
|
||||
if not section_path:
|
||||
page_info = getattr(chunk, 'page_info', {}) or {}
|
||||
section_path = page_info.get('section_path', '') or page_info.get('section', '')
|
||||
|
||||
if chunk_type == 'table':
|
||||
# 表格内容优先使用 table_html(完整表格),fallback 到 content(标题)
|
||||
table_md = getattr(chunk, 'table_html', None) or chunk.content
|
||||
semantic_content = _build_semantic_content_for_table(
|
||||
table_md, {'page': page_start, 'page_end': page_end, 'section_path': section_path}, chunk, doc_type
|
||||
)
|
||||
elif chunk_type in ('image', 'chart'):
|
||||
semantic_content = self.generate_lightweight_image_description(
|
||||
chunk.content, chunk, {'page': page_start, 'page_end': page_end, 'section_path': section_path}
|
||||
)
|
||||
else:
|
||||
semantic_content = _build_semantic_content_for_text(
|
||||
chunk, {'page': page_start, 'page_end': page_end, 'section_path': section_path}, doc_type
|
||||
)
|
||||
|
||||
# 构建元数据
|
||||
metadata = {
|
||||
"chunk_id": chunk_id,
|
||||
"chunk_index": i,
|
||||
"chunk_type": chunk_type,
|
||||
"source": filename,
|
||||
"collection": kb_name,
|
||||
"doc_type": doc_type, # 文档类型(pdf/word/excel/ppt/other),驱动前端差异化溯源展示
|
||||
"page": page_start,
|
||||
"page_end": page_end,
|
||||
"section": section_path,
|
||||
"status": "active",
|
||||
"version": "v1",
|
||||
}
|
||||
|
||||
if extra_metadata:
|
||||
metadata.update(extra_metadata)
|
||||
|
||||
# 序列化图片信息(修复图片召回断链)
|
||||
if hasattr(chunk, 'images') and chunk.images:
|
||||
metadata['images_json'] = json.dumps(chunk.images, ensure_ascii=False)
|
||||
|
||||
if hasattr(chunk, 'image_path') and chunk.image_path:
|
||||
metadata['image_path'] = chunk.image_path
|
||||
|
||||
# 生成向量
|
||||
try:
|
||||
embedding = embedding_model.encode(semantic_content).tolist()
|
||||
except Exception as e:
|
||||
logger.error(f"生成向量失败: {chunk_id}, 错误: {e}")
|
||||
continue
|
||||
|
||||
ids.append(chunk_id)
|
||||
documents.append(semantic_content)
|
||||
metadatas.append(metadata)
|
||||
embeddings.append(embedding)
|
||||
|
||||
# 处理表格摘要
|
||||
if chunk_type == 'table' and enable_table_summary:
|
||||
table_md = chunk.content
|
||||
summary = self._generate_table_summary(table_md, chunk)
|
||||
if summary:
|
||||
# 存储原始表格到 DocStore
|
||||
self._store_original_table(chunk_id, table_md, metadata)
|
||||
|
||||
# 处理图片描述
|
||||
if chunk_type in ('image', 'chart') and enable_image_description:
|
||||
image_path = chunk.content
|
||||
if self.should_process_image(image_path, "", page_info.get('caption', '')):
|
||||
desc = self._generate_image_description(image_path, metadata)
|
||||
if desc:
|
||||
self._store_image_reference(chunk_id, image_path, metadata)
|
||||
|
||||
if not ids:
|
||||
return 0
|
||||
|
||||
# 批量添加到向量库
|
||||
collection.add(
|
||||
ids=ids,
|
||||
documents=documents,
|
||||
metadatas=metadatas,
|
||||
embeddings=embeddings
|
||||
)
|
||||
|
||||
# 更新 BM25 索引
|
||||
bm25 = self.get_bm25_index(kb_name)
|
||||
if bm25:
|
||||
bm25.add_documents(ids, documents, metadatas)
|
||||
self.save_bm25_index(kb_name)
|
||||
|
||||
logger.info(f"添加文件: {filename} -> {kb_name}, 片段数: {len(ids)}")
|
||||
return len(ids)
|
||||
|
||||
def _merge_cross_page_tables(self, chunks: list) -> list:
|
||||
"""
|
||||
合并跨页表格
|
||||
|
||||
检测规则:
|
||||
1. 表格切片后面跟着"续表"文本或另一个表格
|
||||
2. 页码连续或无法判断时,通过标题匹配
|
||||
3. 第二个表格标题包含"续表"或标题相似
|
||||
|
||||
合并操作:
|
||||
1. 合并 table_html
|
||||
2. 将两个 image_path 存入 images 字段
|
||||
3. 更新 page_end
|
||||
"""
|
||||
if len(chunks) < 2:
|
||||
return chunks
|
||||
|
||||
merged_chunks = []
|
||||
i = 0
|
||||
merge_count = 0
|
||||
|
||||
while i < len(chunks):
|
||||
current = chunks[i]
|
||||
|
||||
# 检查是否为表格类型
|
||||
if getattr(current, 'chunk_type', '') == 'table':
|
||||
# 查找下一个表格(跳过中间的"续表"文本)
|
||||
next_table_idx = None
|
||||
next_chunk = None
|
||||
|
||||
for j in range(i + 1, min(i + 4, len(chunks))): # 最多向前看3个切片
|
||||
candidate = chunks[j]
|
||||
candidate_type = getattr(candidate, 'chunk_type', '')
|
||||
candidate_title = getattr(candidate, 'title', '') or ''
|
||||
candidate_content = getattr(candidate, 'content', '')
|
||||
|
||||
if candidate_type == 'table':
|
||||
next_table_idx = j
|
||||
next_chunk = candidate
|
||||
break
|
||||
elif candidate_type == 'text' and ('续表' in candidate_title or '续表' in candidate_content):
|
||||
# 遇到"续表"文本,继续查找下一个表格
|
||||
continue
|
||||
elif candidate_type not in ('text',):
|
||||
# 遇到非文本类型,停止查找
|
||||
break
|
||||
|
||||
if next_chunk is not None:
|
||||
# 获取页码信息
|
||||
curr_page_end = getattr(current, 'page_end', getattr(current, 'page_start', 0))
|
||||
next_page_start = getattr(next_chunk, 'page_start', 0)
|
||||
|
||||
# 获取标题
|
||||
curr_title = getattr(current, 'title', '') or ''
|
||||
next_title = getattr(next_chunk, 'title', '') or ''
|
||||
if isinstance(curr_title, list):
|
||||
curr_title = curr_title[0] if curr_title else ''
|
||||
if isinstance(next_title, list):
|
||||
next_title = next_title[0] if next_title else ''
|
||||
|
||||
# 获取内容(用于检测"续表")
|
||||
next_content = getattr(next_chunk, 'content', '')
|
||||
|
||||
# 判断是否为跨页表格
|
||||
is_cross_page = False
|
||||
|
||||
# 规则1: 页码连续(如果页码有效)
|
||||
page_valid = curr_page_end > 0 and next_page_start > 0
|
||||
if page_valid and curr_page_end + 1 == next_page_start:
|
||||
is_cross_page = True
|
||||
|
||||
# 规则2: 第二个表格标题或内容包含"续表"
|
||||
elif '续表' in next_title or '续表' in next_content:
|
||||
is_cross_page = True
|
||||
|
||||
# 规则3: 标题相似(去掉"续表"后比较)
|
||||
elif curr_title and next_title:
|
||||
clean_next = next_title.replace('续表', '').strip()
|
||||
if curr_title in clean_next or clean_next in curr_title:
|
||||
is_cross_page = True
|
||||
|
||||
if is_cross_page:
|
||||
# 执行合并
|
||||
merge_count += 1
|
||||
logger.info(f"合并跨页表格: {curr_title} (页{curr_page_end}) + {next_title} (页{next_page_start})")
|
||||
|
||||
# 合并 table_html
|
||||
curr_html = getattr(current, 'table_html', '') or ''
|
||||
next_html = getattr(next_chunk, 'table_html', '') or ''
|
||||
if curr_html and next_html:
|
||||
# 合并两个表格的 HTML
|
||||
current.table_html = curr_html + '\n' + next_html
|
||||
|
||||
# 合并 image_path 到 images
|
||||
curr_img = getattr(current, 'image_path', None)
|
||||
next_img = getattr(next_chunk, 'image_path', None)
|
||||
merged_images = []
|
||||
if curr_img:
|
||||
merged_images.append({'id': curr_img, 'page': curr_page_end})
|
||||
if next_img:
|
||||
merged_images.append({'id': next_img, 'page': next_page_start})
|
||||
if merged_images:
|
||||
current.images = merged_images
|
||||
# 保留第一个图片作为主 image_path
|
||||
current.image_path = curr_img
|
||||
|
||||
# 更新页码范围
|
||||
current.page_end = getattr(next_chunk, 'page_end', next_page_start)
|
||||
|
||||
# 添加合并后的切片,跳过中间所有切片
|
||||
merged_chunks.append(current)
|
||||
i = next_table_idx + 1
|
||||
continue
|
||||
|
||||
# 不需要合并,直接添加
|
||||
merged_chunks.append(current)
|
||||
i += 1
|
||||
|
||||
if merge_count > 0:
|
||||
logger.info(f"跨页表格合并完成: {merge_count} 组表格被合并")
|
||||
|
||||
return merged_chunks
|
||||
|
||||
def _generate_table_summary(self, table_md: str, chunk) -> str:
|
||||
"""生成表格摘要(LLM)"""
|
||||
# 提取表头
|
||||
lines = table_md.split('\n')
|
||||
headers = []
|
||||
for line in lines:
|
||||
if line.startswith('|') and '---' not in line:
|
||||
headers = [h.strip() for h in line.split('|') if h.strip()]
|
||||
break
|
||||
|
||||
if not headers:
|
||||
return ""
|
||||
|
||||
# 构建提示词
|
||||
prompt = f"""请用一句话总结以下表格的主要内容,不超过50字。
|
||||
|
||||
表头:{', '.join(headers)}
|
||||
|
||||
表格内容(前5行):
|
||||
{chr(10).join(lines[:6])}
|
||||
|
||||
摘要:"""
|
||||
|
||||
try:
|
||||
from config import get_llm_client, DASHSCOPE_MODEL
|
||||
client = get_llm_client()
|
||||
summary = call_llm(client, prompt, DASHSCOPE_MODEL, max_tokens=100)
|
||||
return summary.strip() if summary else ""
|
||||
except Exception as e:
|
||||
logger.warning(f"生成表格摘要失败: {e}")
|
||||
return ""
|
||||
|
||||
@staticmethod
|
||||
def _extract_table_title(table_md: str) -> str:
|
||||
"""从表格 Markdown 中提取标题"""
|
||||
lines = table_md.split('\n')
|
||||
for line in lines[:3]:
|
||||
if line.startswith('#'):
|
||||
return line.lstrip('#').strip()
|
||||
return ""
|
||||
|
||||
def _generate_image_description(self, image_path: str, metadata: dict = None) -> str:
|
||||
"""生成图片描述(VLM)"""
|
||||
# 检查缓存
|
||||
cached = self._get_vlm_cache(image_path)
|
||||
if cached:
|
||||
return cached
|
||||
|
||||
# 调用 VLM
|
||||
try:
|
||||
import base64
|
||||
from pathlib import Path
|
||||
|
||||
# 读取图片
|
||||
img_path = Path(image_path)
|
||||
if not img_path.exists():
|
||||
return ""
|
||||
|
||||
img_data = base64.b64encode(img_path.read_bytes()).decode()
|
||||
|
||||
# 构建提示词
|
||||
prompt = """请描述这张图片的内容,包括:
|
||||
1. 图片类型(如流程图、架构图、数据图表等)
|
||||
2. 主要内容和关键信息
|
||||
3. 如果是图表,描述数据趋势或关键数值
|
||||
|
||||
描述应简洁,不超过100字。"""
|
||||
|
||||
# 调用 VLM(需要支持视觉的模型)
|
||||
from config import DASHSCOPE_API_KEY, DASHSCOPE_BASE_URL, VLM_MODEL
|
||||
from openai import OpenAI
|
||||
|
||||
client = OpenAI(api_key=DASHSCOPE_API_KEY, base_url=DASHSCOPE_BASE_URL)
|
||||
|
||||
response = client.chat.completions.create(
|
||||
model=VLM_MODEL,
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": prompt},
|
||||
{"type": "image_url", "image_url": {"url": f"data:image/png;base64,{img_data}"}}
|
||||
]
|
||||
}
|
||||
],
|
||||
max_tokens=200
|
||||
)
|
||||
|
||||
description = response.choices[0].message.content
|
||||
|
||||
# 缓存结果
|
||||
import hashlib
|
||||
img_hash = hashlib.md5(img_path.read_bytes()).hexdigest()
|
||||
cache_dir = Path('.data/cache/vlm')
|
||||
cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
(cache_dir / f'{img_hash}.txt').write_text(description, encoding='utf-8')
|
||||
|
||||
return description
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"生成图片描述失败: {e}")
|
||||
return ""
|
||||
|
||||
def update_image_descriptions(self, kb_name: str) -> dict:
|
||||
"""更新向量库中所有图片的描述"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return {"success": False, "error": "向量库不存在"}
|
||||
|
||||
result = collection.get(where={"chunk_type": {"$in": ["image", "chart"]}})
|
||||
|
||||
if not result['ids']:
|
||||
return {"success": True, "updated": 0, "message": "没有图片切片"}
|
||||
|
||||
updated = 0
|
||||
failed = 0
|
||||
|
||||
for chunk_id, doc, meta in zip(
|
||||
result['ids'],
|
||||
result['documents'],
|
||||
result['metadatas']
|
||||
):
|
||||
image_path = meta.get('image_path', '')
|
||||
if not image_path:
|
||||
continue
|
||||
|
||||
description = self._generate_image_description(image_path, meta)
|
||||
if description:
|
||||
try:
|
||||
collection.update(
|
||||
ids=[chunk_id],
|
||||
documents=[description],
|
||||
metadatas=[{**meta, 'image_description': description}]
|
||||
)
|
||||
updated += 1
|
||||
except Exception as e:
|
||||
logger.warning(f"更新图片描述失败: {chunk_id}, {e}")
|
||||
failed += 1
|
||||
|
||||
# 重建 BM25 索引
|
||||
self.rebuild_bm25_index(kb_name)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"updated": updated,
|
||||
"failed": failed,
|
||||
"total": len(result['ids'])
|
||||
}
|
||||
|
||||
# ==================== 文档生命周期 ====================
|
||||
|
||||
def mark_document_as_superseded(
|
||||
self,
|
||||
kb_name: str,
|
||||
old_filename: str,
|
||||
new_filename: str,
|
||||
reason: str = "版本更新"
|
||||
) -> Dict:
|
||||
"""标记文档为已替代版本"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return {"success": False, "error": "向量库不存在"}
|
||||
|
||||
result = collection.get(where={"source": old_filename})
|
||||
|
||||
if not result['ids']:
|
||||
return {"success": False, "error": "旧文档不存在"}
|
||||
|
||||
from datetime import datetime
|
||||
superseded_date = datetime.now().isoformat()
|
||||
|
||||
updated_metadatas = [
|
||||
{
|
||||
**m,
|
||||
"status": "superseded",
|
||||
"superseded_by": new_filename,
|
||||
"superseded_date": superseded_date,
|
||||
"superseded_reason": reason
|
||||
}
|
||||
for m in result['metadatas']
|
||||
]
|
||||
|
||||
collection.update(
|
||||
ids=result['ids'],
|
||||
metadatas=updated_metadatas
|
||||
)
|
||||
|
||||
self.rebuild_bm25_index(kb_name)
|
||||
|
||||
logger.info(f"标记文档替代: {old_filename} -> {new_filename}")
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"superseded_chunks": len(result['ids']),
|
||||
"old_document": old_filename,
|
||||
"new_document": new_filename,
|
||||
"collection": kb_name
|
||||
}
|
||||
|
||||
# ==================== 辅助检索方法 ====================
|
||||
|
||||
def search_with_status_filter(
|
||||
self,
|
||||
kb_name: str,
|
||||
query_vector: List[float],
|
||||
top_k: int = 5,
|
||||
status_filter: str = None
|
||||
) -> Optional[SearchResult]:
|
||||
"""带状态过滤的检索"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
return None
|
||||
|
||||
where_filter = None
|
||||
if status_filter:
|
||||
where_filter = {"status": status_filter}
|
||||
|
||||
result = collection.query(
|
||||
query_embeddings=[query_vector],
|
||||
n_results=top_k,
|
||||
where=where_filter
|
||||
)
|
||||
|
||||
return SearchResult(
|
||||
ids=result['ids'][0] if result['ids'] else [],
|
||||
documents=result['documents'][0] if result['documents'] else [],
|
||||
metadatas=result['metadatas'][0] if result['metadatas'] else [],
|
||||
distances=result['distances'][0] if result['distances'] else [],
|
||||
collection_name=kb_name
|
||||
)
|
||||
|
||||
def _rebuild_bm25_index(self, kb_name: str):
|
||||
"""重建 BM25 索引(私有方法,兼容旧代码)"""
|
||||
return self.rebuild_bm25_index(kb_name)
|
||||
|
||||
|
||||
# ==================== 全局实例 ====================
|
||||
|
||||
_kb_manager: Optional[KnowledgeBaseManager] = None
|
||||
|
||||
|
||||
def get_kb_manager() -> KnowledgeBaseManager:
|
||||
"""获取全局知识库管理器实例"""
|
||||
global _kb_manager
|
||||
if _kb_manager is None:
|
||||
_kb_manager = KnowledgeBaseManager()
|
||||
return _kb_manager
|
||||
119
knowledge/permission.py
Normal file
119
knowledge/permission.py
Normal file
@@ -0,0 +1,119 @@
|
||||
"""
|
||||
知识库管理器 - 权限管理 Mixin
|
||||
|
||||
提供基于角色和部门的向量库访问控制功能。
|
||||
|
||||
权限模型:
|
||||
- admin: 可访问所有向量库,拥有所有操作权限
|
||||
- manager: 可访问公开库和本部门库,拥有读写权限
|
||||
- user: 可访问公开库和本部门库,仅有只读权限
|
||||
|
||||
向量库命名规则:
|
||||
- public_kb: 公开知识库,所有人可读
|
||||
- dept_{部门名}: 部门知识库,仅对应部门可访问
|
||||
|
||||
主要方法:
|
||||
- get_accessible_collections: 获取用户可访问的向量库列表
|
||||
- check_permission: 检查用户对指定向量库的操作权限
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import List
|
||||
|
||||
from .base import PUBLIC_KB_NAME, normalize_department_name
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class PermissionMixin:
|
||||
"""
|
||||
权限管理 Mixin
|
||||
|
||||
提供向量库级别的访问控制,基于用户角色和部门判断权限。
|
||||
|
||||
依赖属性(需由主类提供):
|
||||
- self._metadata: 元数据字典(包含 collections 信息)
|
||||
"""
|
||||
|
||||
def get_accessible_collections(
|
||||
self,
|
||||
role: str,
|
||||
department: str,
|
||||
operation: str = "read"
|
||||
) -> List[str]:
|
||||
"""
|
||||
获取用户可访问的向量库列表
|
||||
|
||||
根据用户角色和部门,返回其有权限访问的向量库名称列表。
|
||||
|
||||
权限规则:
|
||||
- admin: 可访问所有向量库
|
||||
- manager: 可读公开库 + 本部门库,可写本部门库
|
||||
- user: 仅可读公开库 + 本部门库
|
||||
|
||||
Args:
|
||||
role: 用户角色,可选值: 'admin' | 'manager' | 'user'
|
||||
department: 用户部门(支持中文名如"财务部"或英文标识如"finance")
|
||||
operation: 操作类型,可选值: 'read' | 'write' | 'delete' | 'sync'
|
||||
|
||||
Returns:
|
||||
可访问的向量库名称列表
|
||||
|
||||
Example:
|
||||
>>> kb_manager.get_accessible_collections("admin", "finance")
|
||||
['public_kb', 'dept_finance', 'dept_hr', ...]
|
||||
>>> kb_manager.get_accessible_collections("user", "技术部", "read")
|
||||
['public_kb', 'dept_tech']
|
||||
"""
|
||||
result = []
|
||||
|
||||
if role == "admin":
|
||||
for info in self.list_collections():
|
||||
result.append(info.name)
|
||||
return result
|
||||
|
||||
if PUBLIC_KB_NAME in self._metadata.get("collections", {}):
|
||||
result.append(PUBLIC_KB_NAME)
|
||||
|
||||
if department:
|
||||
normalized_dept = normalize_department_name(department)
|
||||
if normalized_dept:
|
||||
dept_kb = f"dept_{normalized_dept}"
|
||||
if dept_kb in self._metadata.get("collections", {}):
|
||||
if operation == "read":
|
||||
result.append(dept_kb)
|
||||
elif operation in ("write", "delete", "sync"):
|
||||
if role == "manager":
|
||||
result.append(dept_kb)
|
||||
else:
|
||||
logger.warning(f"部门名称无法标准化: {department}")
|
||||
|
||||
return result
|
||||
|
||||
def check_permission(
|
||||
self,
|
||||
role: str,
|
||||
department: str,
|
||||
kb_name: str,
|
||||
operation: str = "read"
|
||||
) -> bool:
|
||||
"""
|
||||
检查用户对指定向量库的操作权限
|
||||
|
||||
Args:
|
||||
role: 用户角色,可选值: 'admin' | 'manager' | 'user'
|
||||
department: 用户部门
|
||||
kb_name: 目标向量库名称
|
||||
operation: 操作类型,可选值: 'read' | 'write' | 'delete' | 'sync'
|
||||
|
||||
Returns:
|
||||
有权限返回 True,无权限返回 False
|
||||
|
||||
Example:
|
||||
>>> kb_manager.check_permission("user", "finance", "dept_finance", "read")
|
||||
True
|
||||
>>> kb_manager.check_permission("user", "finance", "dept_finance", "write")
|
||||
False
|
||||
"""
|
||||
accessible = self.get_accessible_collections(role, department, operation)
|
||||
return kb_name in accessible
|
||||
379
knowledge/processing.py
Normal file
379
knowledge/processing.py
Normal file
@@ -0,0 +1,379 @@
|
||||
"""
|
||||
知识库管理器 - 图片/表格处理 Mixin
|
||||
|
||||
提供图片和表格的智能处理功能,包括:
|
||||
- 图片过滤:判断图片是否值得处理(过滤 logo、icon 等垃圾图片)
|
||||
- 图片描述生成:生成用于向量检索的轻量级描述
|
||||
- 表格摘要生成:生成表格的语义摘要
|
||||
- 原始数据存储:将表格/图片引用存储到 DocStore
|
||||
|
||||
处理策略:
|
||||
- 图片:通过 VLM(视觉语言模型)生成描述,用于语义检索
|
||||
- 表格:通过 LLM 生成摘要,提取关键字段和示例数据
|
||||
|
||||
主要方法:
|
||||
- should_process_image: 判断图片是否值得处理
|
||||
- generate_image_short_summary: 生成图片短摘要
|
||||
- generate_lightweight_image_description: 生成轻量级图片描述
|
||||
"""
|
||||
|
||||
import os
|
||||
import re
|
||||
import json
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import Optional, List
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ProcessingMixin:
|
||||
"""
|
||||
图片/表格处理 Mixin
|
||||
|
||||
提供图片和表格的智能处理能力,支持:
|
||||
- 垃圾图片过滤(logo、icon、二维码等)
|
||||
- 图片描述生成(用于向量检索)
|
||||
- 表格摘要生成(提取语义信息)
|
||||
"""
|
||||
|
||||
def should_process_image(self, image_path: str, context_text: str, caption: str = "") -> bool:
|
||||
"""
|
||||
判断图片是否值得处理
|
||||
|
||||
通过多维度判断过滤无意义的图片:
|
||||
1. 文件名过滤:排除 logo、icon、qr、watermark 等
|
||||
2. 尺寸过滤:排除小于 100x100 的图片
|
||||
3. 上下文判断:有 caption 或上下文文本时保留
|
||||
|
||||
Args:
|
||||
image_path: 图片文件路径
|
||||
context_text: 图片周围的上下文文本
|
||||
caption: 图片标题/说明
|
||||
|
||||
Returns:
|
||||
值得处理返回 True,应该跳过返回 False
|
||||
|
||||
Example:
|
||||
>>> should_process_image("/path/to/logo.png", "", "")
|
||||
False # 文件名包含 logo
|
||||
>>> should_process_image("/path/to/chart.png", "如图所示...", "图2.1 架构图")
|
||||
True # 有 caption
|
||||
"""
|
||||
filename = os.path.basename(image_path).lower()
|
||||
|
||||
junk_keywords = ["logo", "icon", "qr", "watermark", "banner", "button", "avatar"]
|
||||
if any(kw in filename for kw in junk_keywords):
|
||||
logger.debug(f"图片过滤:文件名包含垃圾关键词 - {filename}")
|
||||
return False
|
||||
|
||||
try:
|
||||
from PIL import Image
|
||||
with Image.open(image_path) as img:
|
||||
width, height = img.size
|
||||
if width < 100 or height < 100:
|
||||
logger.debug(f"图片过滤:尺寸过小 ({width}x{height}) - {filename}")
|
||||
return False
|
||||
except Exception as e:
|
||||
logger.debug(f"图片尺寸检查失败: {e}")
|
||||
pass
|
||||
|
||||
if len(caption) >= 3:
|
||||
return True
|
||||
if len(context_text) >= 10:
|
||||
return True
|
||||
|
||||
logger.debug(f"图片保留:{filename}")
|
||||
return True
|
||||
|
||||
def _get_vlm_cache(self, image_path: str) -> Optional[str]:
|
||||
"""
|
||||
检查是否有 VLM 缓存描述
|
||||
|
||||
通过图片 MD5 哈希查找缓存文件,避免重复调用 VLM。
|
||||
|
||||
Args:
|
||||
image_path: 图片文件路径
|
||||
|
||||
Returns:
|
||||
缓存的描述文本;无缓存返回 None
|
||||
"""
|
||||
import hashlib
|
||||
|
||||
try:
|
||||
if not os.path.exists(image_path):
|
||||
return None
|
||||
|
||||
img_hash = hashlib.md5(open(image_path, 'rb').read()).hexdigest()
|
||||
|
||||
cache_file = Path(f'.data/cache/vlm/{img_hash}.txt')
|
||||
if cache_file.exists():
|
||||
return cache_file.read_text(encoding='utf-8')
|
||||
except Exception as e:
|
||||
logger.debug(f"检查 VLM 缓存失败: {e}")
|
||||
|
||||
return None
|
||||
|
||||
def generate_image_short_summary(self, chunk, page_info: dict, full_description: str = "") -> str:
|
||||
"""
|
||||
生成图片短摘要(用于向量检索匹配)
|
||||
|
||||
提取关键信息构建简洁摘要,包括:
|
||||
- 图号/表号
|
||||
- 关键数值(年份、金额等)
|
||||
- 章节主题
|
||||
- 图表类型
|
||||
|
||||
Args:
|
||||
chunk: 文档切片对象
|
||||
page_info: 页面信息字典
|
||||
full_description: 完整描述(可选)
|
||||
|
||||
Returns:
|
||||
短摘要字符串,不超过 45 字符
|
||||
|
||||
Example:
|
||||
>>> generate_image_short_summary(chunk, {"section": "1.2 概述"}, "2020-2025年数据...")
|
||||
"图2.1:2020-2025年,柱状图"
|
||||
"""
|
||||
figure_number = ""
|
||||
table_number = ""
|
||||
|
||||
section = page_info.get('section_path', '') or page_info.get('section', '')
|
||||
title = chunk.title if hasattr(chunk, 'title') and chunk.title else ""
|
||||
caption = page_info.get('caption', '')
|
||||
|
||||
for source in [section, title, caption, full_description]:
|
||||
if not source:
|
||||
continue
|
||||
fig_match = re.search(r'图\s*(\d+\.?\d*)', str(source))
|
||||
if fig_match and not figure_number:
|
||||
figure_number = fig_match.group(1)
|
||||
table_match = re.search(r'表\s*(\d+\.?\d*)', str(source))
|
||||
if table_match and not table_number:
|
||||
table_number = table_match.group(1)
|
||||
|
||||
chunk_type = page_info.get('chunk_type', 'image')
|
||||
keywords = []
|
||||
|
||||
year_match = re.search(r'(\d{4})\s*[-至到]\s*(\d{4})', full_description)
|
||||
if year_match:
|
||||
keywords.append(f"{year_match.group(1)}-{year_match.group(2)}年")
|
||||
|
||||
num_match = re.search(r'(\d+\.?\d*)\s*(亿|万|千瓦时|吨|米)', full_description)
|
||||
if num_match:
|
||||
keywords.append(f"{num_match.group(1)}{num_match.group(2)}")
|
||||
|
||||
if section:
|
||||
section_parts = section.split('>')
|
||||
if section_parts:
|
||||
last_part = section_parts[-1].strip()
|
||||
theme_match = re.search(r'(\d+\.?\d*)\s*(.+)', last_part)
|
||||
if theme_match:
|
||||
keywords.append(theme_match.group(2).strip()[:10])
|
||||
|
||||
if chunk_type == 'chart':
|
||||
type_label = "图表"
|
||||
if '柱状图' in full_description:
|
||||
type_label = "柱状图"
|
||||
elif '折线图' in full_description or '曲线图' in full_description:
|
||||
type_label = "折线图"
|
||||
elif '饼图' in full_description:
|
||||
type_label = "饼图"
|
||||
elif chunk_type == 'table':
|
||||
type_label = "表格"
|
||||
else:
|
||||
type_label = "图片"
|
||||
|
||||
if figure_number:
|
||||
summary = f"图{figure_number}:"
|
||||
elif table_number:
|
||||
summary = f"表{table_number}:"
|
||||
else:
|
||||
summary = f"{type_label}:"
|
||||
|
||||
if keywords:
|
||||
summary += ",".join(keywords[:3])
|
||||
|
||||
if full_description and len(keywords) < 2:
|
||||
first_sentence = full_description.split('。')[0][:30]
|
||||
if first_sentence:
|
||||
summary += first_sentence
|
||||
|
||||
if len(summary) > 45:
|
||||
summary = summary[:42] + "..."
|
||||
|
||||
return summary
|
||||
|
||||
def _extract_figure_references(self, text: str) -> List[str]:
|
||||
"""
|
||||
提取文本中的图表引用
|
||||
|
||||
从文本中提取所有"见图X.X"、"见表X.X"等引用。
|
||||
|
||||
Args:
|
||||
text: 待提取的文本
|
||||
|
||||
Returns:
|
||||
图表编号列表(去重)
|
||||
"""
|
||||
references = []
|
||||
|
||||
fig_matches = re.findall(r'(?:[见如及和与])?图\s*(\d+\.?\d*)', text)
|
||||
references.extend(fig_matches)
|
||||
|
||||
table_matches = re.findall(r'(?:[见如及和与])?表\s*(\d+\.?\d*)', text)
|
||||
references.extend(table_matches)
|
||||
|
||||
return list(set(references))
|
||||
|
||||
def generate_lightweight_image_description(self, image_path: str, chunk, page_info: dict) -> str:
|
||||
"""
|
||||
生成轻量级图片描述(用于语义检索)
|
||||
|
||||
构建包含上下文信息的描述,不调用 VLM,
|
||||
仅利用已有的元数据信息。
|
||||
|
||||
Args:
|
||||
image_path: 图片路径(或图片标识)
|
||||
chunk: 文档切片对象
|
||||
page_info: 页面信息字典
|
||||
|
||||
Returns:
|
||||
多行描述字符串,包含:
|
||||
- 图号/表号
|
||||
- 标题/caption
|
||||
- 章节位置
|
||||
- 页码
|
||||
- 前后文上下文
|
||||
|
||||
Example:
|
||||
>>> generate_lightweight_image_description("/path/to/img", chunk, page_info)
|
||||
图2.1,系统架构图,位于「第一章 > 1.2 概述」,第5页
|
||||
前文:系统由三个模块组成...
|
||||
后文:如图所示,各模块之间...
|
||||
"""
|
||||
parts = []
|
||||
|
||||
chunk_type = page_info.get('chunk_type', 'image')
|
||||
is_chart = chunk_type == 'chart'
|
||||
type_label = "图表" if is_chart else "图片"
|
||||
|
||||
figure_number = ""
|
||||
table_number = ""
|
||||
|
||||
sources_to_check = []
|
||||
|
||||
section = page_info.get('section_path', '') or page_info.get('section', '')
|
||||
if section:
|
||||
sources_to_check.append(section)
|
||||
|
||||
title = chunk.title if hasattr(chunk, 'title') and chunk.title else ""
|
||||
if title:
|
||||
sources_to_check.append(title)
|
||||
|
||||
caption = page_info.get('caption', '')
|
||||
if caption:
|
||||
sources_to_check.append(caption)
|
||||
|
||||
context_before = ""
|
||||
context_after = ""
|
||||
if hasattr(chunk, 'context_before') and chunk.context_before:
|
||||
context_before = chunk.context_before[:500]
|
||||
sources_to_check.append(context_before)
|
||||
if hasattr(chunk, 'context_after') and chunk.context_after:
|
||||
context_after = chunk.context_after[:500]
|
||||
sources_to_check.append(context_after)
|
||||
|
||||
for source_text in sources_to_check:
|
||||
if not figure_number:
|
||||
fig_match = re.search(r'(?:[见如及和与])?图\s*(\d+\.?\d*)', source_text)
|
||||
if fig_match:
|
||||
figure_number = fig_match.group(1)
|
||||
if not table_number:
|
||||
table_match = re.search(r'(?:[见如及和与])?表\s*(\d+\.?\d*)', source_text)
|
||||
if table_match:
|
||||
table_number = table_match.group(1)
|
||||
|
||||
page = page_info.get('page', 0)
|
||||
|
||||
if figure_number:
|
||||
parts.append(f"图{figure_number}")
|
||||
if table_number:
|
||||
parts.append(f"表{table_number}")
|
||||
|
||||
if caption and caption not in ("图片", "图表"):
|
||||
parts.append(caption)
|
||||
elif title and title not in ("图片", "图表"):
|
||||
parts.append(title)
|
||||
|
||||
if section:
|
||||
parts.append(f"位于「{section}」")
|
||||
|
||||
parts.append(f"第{page}页")
|
||||
|
||||
description_parts = [f"{type_label}:{','.join(parts)}"]
|
||||
|
||||
if context_before:
|
||||
description_parts.append(f"前文:{context_before}")
|
||||
if context_after:
|
||||
description_parts.append(f"后文:{context_after}")
|
||||
|
||||
return "\n".join(description_parts)
|
||||
|
||||
def _store_original_table(self, doc_id: str, table_md: str, metadata: dict) -> None:
|
||||
"""
|
||||
存储原始表格到 DocStore
|
||||
|
||||
将表格的 Markdown 内容持久化存储,用于后续检索展示。
|
||||
|
||||
Args:
|
||||
doc_id: 文档 ID(切片 ID)
|
||||
table_md: 表格的 Markdown 内容
|
||||
metadata: 元数据字典
|
||||
"""
|
||||
try:
|
||||
docstore_dir = Path(".data/docstore")
|
||||
docstore_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
record = {
|
||||
"content_type": "table",
|
||||
"markdown": table_md,
|
||||
"meta": metadata
|
||||
}
|
||||
|
||||
doc_path = docstore_dir / f"{doc_id}.json"
|
||||
with open(doc_path, 'w', encoding='utf-8') as f:
|
||||
json.dump(record, f, ensure_ascii=False, indent=2)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"存储原始表格失败: {e}")
|
||||
|
||||
def _store_image_reference(self, doc_id: str, image_path: str, metadata: dict) -> None:
|
||||
"""
|
||||
存储图片引用到 DocStore
|
||||
|
||||
记录图片文件路径,用于后续访问和展示。
|
||||
|
||||
Args:
|
||||
doc_id: 文档 ID(切片 ID)
|
||||
image_path: 图片文件路径
|
||||
metadata: 元数据字典
|
||||
"""
|
||||
try:
|
||||
docstore_dir = Path(".data/docstore")
|
||||
docstore_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
record = {
|
||||
"content_type": "image",
|
||||
"storage_type": "file",
|
||||
"file_path": image_path,
|
||||
"meta": metadata
|
||||
}
|
||||
|
||||
doc_path = docstore_dir / f"{doc_id}.json"
|
||||
with open(doc_path, 'w', encoding='utf-8') as f:
|
||||
json.dump(record, f, ensure_ascii=False, indent=2)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"存储图片引用失败: {e}")
|
||||
711
knowledge/router.py
Normal file
711
knowledge/router.py
Normal file
@@ -0,0 +1,711 @@
|
||||
"""
|
||||
知识库路由器 - 智能选择查询目标
|
||||
|
||||
功能:
|
||||
1. 查询意图分析 - 判断查询是否涉及特定部门
|
||||
2. 知识库路由 - 根据意图和权限选择目标向量库
|
||||
3. 单库优化 - 如果只需查询单库,避免不必要的并行检索
|
||||
|
||||
使用方式:
|
||||
from knowledge.router import KnowledgeBaseRouter
|
||||
|
||||
router = KnowledgeBaseRouter()
|
||||
|
||||
# 获取目标向量库
|
||||
target_kbs = router.route(
|
||||
query="财务部的报销流程是什么",
|
||||
role="user",
|
||||
department="tech"
|
||||
)
|
||||
# 返回: ["public_kb", "dept_finance"] # 如果有权限
|
||||
"""
|
||||
|
||||
import os
|
||||
import re
|
||||
import json
|
||||
import logging
|
||||
from typing import List, Dict, Optional, Tuple
|
||||
from dataclasses import dataclass
|
||||
from openai import OpenAI
|
||||
|
||||
# 导入配置
|
||||
from config import API_KEY, BASE_URL, MODEL
|
||||
|
||||
# 导入 LLM 工具函数
|
||||
from core.llm_utils import call_llm, parse_json_from_response
|
||||
|
||||
# 导入权限管理
|
||||
from auth.gateway import get_accessible_collections, normalize_department_name
|
||||
|
||||
# 设置日志
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ==================== 部门关键词配置 ====================
|
||||
|
||||
# 部门关键词映射(可根据实际情况扩展)
|
||||
DEPARTMENT_KEYWORDS = {
|
||||
"finance": [
|
||||
"财务", "报销", "发票", "预算", "支出", "收入", "成本",
|
||||
"账目", "会计", "审计", "税务", "工资", "奖金", "补贴",
|
||||
"费用", "付款", "收款", "借款", "报销单", "财务部"
|
||||
],
|
||||
"hr": [
|
||||
"人事", "招聘", "入职", "离职", "考勤", "请假", "休假",
|
||||
"员工", "培训", "绩效", "晋升", "调岗", "合同", "档案",
|
||||
"社保", "公积金", "福利", "加班", "年假", "人事部", "人力资源"
|
||||
],
|
||||
"tech": [
|
||||
"技术", "开发", "代码", "系统", "服务器", "数据库", "API",
|
||||
"接口", "部署", "测试", "Bug", "需求", "架构", "运维",
|
||||
"网络安全", "服务器", "云服务", "技术部", "研发", "IT"
|
||||
],
|
||||
"operation": [
|
||||
"运营", "推广", "营销", "活动", "用户", "增长", "数据",
|
||||
"分析", "客服", "售后", "投诉", "反馈", "运营部", "运营中心"
|
||||
],
|
||||
"marketing": [
|
||||
"市场", "品牌", "宣传", "广告", "公关", "媒体", "推广",
|
||||
"展会", "活动策划", "市场部", "营销部"
|
||||
],
|
||||
"legal": [
|
||||
"法务", "合同", "法律", "诉讼", "合规", "风险", "版权",
|
||||
"知识产权", "协议", "法务部"
|
||||
],
|
||||
"admin": [
|
||||
"行政", "办公室", "会议室", "采购", "固定资产", "办公用品",
|
||||
"印章", "档案", "行政部", "总务"
|
||||
]
|
||||
}
|
||||
|
||||
# 通用关键词(查询 public_kb)
|
||||
GENERAL_KEYWORDS = [
|
||||
"公司", "企业", "组织", "介绍", "简介", "文化", "价值观",
|
||||
"制度", "规定", "流程", "政策", "手册", "指南", "帮助",
|
||||
"联系方式", "地址", "电话", "邮箱"
|
||||
]
|
||||
|
||||
|
||||
# ==================== 数据结构 ====================
|
||||
|
||||
@dataclass
|
||||
class QueryIntent:
|
||||
"""查询意图"""
|
||||
is_general: bool # 是否为通用问题
|
||||
department: Optional[str] # 涉及的部门(如果有)
|
||||
confidence: float # 置信度
|
||||
keywords: List[str] # 匹配到的关键词
|
||||
reason: str # 判断理由
|
||||
|
||||
|
||||
# ==================== 知识库路由器 ====================
|
||||
|
||||
class KnowledgeBaseRouter:
|
||||
"""
|
||||
知识库路由器
|
||||
|
||||
根据查询内容和用户权限,智能选择需要查询的向量库。
|
||||
支持规则匹配和 LLM 意图分析两种方式。
|
||||
"""
|
||||
|
||||
def __init__(self, use_llm: bool = True):
|
||||
"""
|
||||
初始化
|
||||
|
||||
Args:
|
||||
use_llm: 是否使用 LLM 进行意图分析(更准确但更慢)
|
||||
"""
|
||||
self.use_llm = use_llm
|
||||
self.llm_client = None
|
||||
|
||||
if use_llm:
|
||||
try:
|
||||
self.llm_client = OpenAI(api_key=API_KEY, base_url=BASE_URL)
|
||||
logger.info("LLM 客户端初始化成功,将使用 LLM 进行意图分析")
|
||||
except Exception as e:
|
||||
logger.warning(f"LLM 客户端初始化失败: {e},将使用规则匹配")
|
||||
self.use_llm = False
|
||||
|
||||
def route(
|
||||
self,
|
||||
query: str,
|
||||
role: str,
|
||||
department: str,
|
||||
accessible_collections: List[str] = None
|
||||
) -> List[str]:
|
||||
"""
|
||||
根据查询意图和用户权限,决定查询哪些向量库
|
||||
|
||||
Args:
|
||||
query: 用户查询
|
||||
role: 用户角色
|
||||
department: 用户部门
|
||||
accessible_collections: 可访问的向量库列表(可选)
|
||||
|
||||
Returns:
|
||||
需要查询的向量库名称列表
|
||||
"""
|
||||
# 1. 获取可访问的向量库
|
||||
if accessible_collections is None:
|
||||
accessible_collections = get_accessible_collections(role, department)
|
||||
|
||||
if not accessible_collections:
|
||||
logger.warning(f"用户无可访问的向量库: role={role}, dept={department}")
|
||||
return []
|
||||
|
||||
# 2. 分析查询意图
|
||||
intent = self.analyze_intent(query)
|
||||
|
||||
# 3. 根据意图选择目标库
|
||||
target_kbs = self._select_knowledge_bases(
|
||||
intent, accessible_collections, role, department
|
||||
)
|
||||
|
||||
logger.info(
|
||||
f"路由决策: query='{query[:30]}...', "
|
||||
f"intent={intent.department or 'general'}, "
|
||||
f"targets={target_kbs}"
|
||||
)
|
||||
|
||||
return target_kbs
|
||||
|
||||
def analyze_intent(self, query: str) -> QueryIntent:
|
||||
"""
|
||||
分析查询意图
|
||||
|
||||
Args:
|
||||
query: 用户查询
|
||||
|
||||
Returns:
|
||||
QueryIntent 对象
|
||||
"""
|
||||
# 先尝试规则匹配(快速)
|
||||
rule_intent = self._analyze_by_rules(query)
|
||||
|
||||
# 如果规则匹配置信度高,直接返回
|
||||
if rule_intent.confidence > 0.8:
|
||||
return rule_intent
|
||||
|
||||
# 否则使用 LLM 分析(更准确)
|
||||
if self.use_llm and self.llm_client:
|
||||
llm_intent = self._analyze_by_llm(query)
|
||||
if llm_intent:
|
||||
# 取两者中置信度高的
|
||||
return llm_intent if llm_intent.confidence > rule_intent.confidence else rule_intent
|
||||
|
||||
return rule_intent
|
||||
|
||||
def _analyze_by_rules(self, query: str) -> QueryIntent:
|
||||
"""基于规则的意图分析"""
|
||||
query_lower = query.lower()
|
||||
matched_departments = {}
|
||||
matched_general = []
|
||||
|
||||
# 检查部门关键词
|
||||
for dept, keywords in DEPARTMENT_KEYWORDS.items():
|
||||
for keyword in keywords:
|
||||
if keyword in query_lower:
|
||||
if dept not in matched_departments:
|
||||
matched_departments[dept] = []
|
||||
matched_departments[dept].append(keyword)
|
||||
|
||||
# 检查通用关键词
|
||||
for keyword in GENERAL_KEYWORDS:
|
||||
if keyword in query_lower:
|
||||
matched_general.append(keyword)
|
||||
|
||||
# 判断结果
|
||||
if matched_departments:
|
||||
# 找到匹配最多的部门
|
||||
best_dept = max(
|
||||
matched_departments.keys(),
|
||||
key=lambda d: len(matched_departments[d])
|
||||
)
|
||||
keywords = matched_departments[best_dept]
|
||||
confidence = min(0.9, 0.5 + len(keywords) * 0.1)
|
||||
|
||||
return QueryIntent(
|
||||
is_general=False,
|
||||
department=best_dept,
|
||||
confidence=confidence,
|
||||
keywords=keywords,
|
||||
reason=f"匹配到部门关键词: {', '.join(keywords)}"
|
||||
)
|
||||
|
||||
elif matched_general:
|
||||
return QueryIntent(
|
||||
is_general=True,
|
||||
department=None,
|
||||
confidence=0.7,
|
||||
keywords=matched_general,
|
||||
reason=f"匹配到通用关键词: {', '.join(matched_general)}"
|
||||
)
|
||||
|
||||
else:
|
||||
return QueryIntent(
|
||||
is_general=False,
|
||||
department=None,
|
||||
confidence=0.3,
|
||||
keywords=[],
|
||||
reason="未匹配到关键词,需要查询所有可访问的库"
|
||||
)
|
||||
|
||||
def _analyze_by_llm(self, query: str) -> Optional[QueryIntent]:
|
||||
"""使用 LLM 进行意图分析"""
|
||||
prompt = f"""分析以下问题的意图,判断:
|
||||
|
||||
1. 是否为通用问题(涉及公司整体、产品、文化等,不特指某部门)
|
||||
2. 是否涉及特定部门(财务、人事、技术等)
|
||||
|
||||
问题:{query}
|
||||
|
||||
请直接返回 JSON 格式(不要包含其他内容):
|
||||
{{"is_general": true/false, "department": "部门英文名或null", "confidence": 0.0-1.0}}
|
||||
|
||||
部门英文名对照:
|
||||
- finance: 财务
|
||||
- hr: 人事
|
||||
- tech: 技术
|
||||
- operation: 运营
|
||||
- marketing: 市场
|
||||
- legal: 法务
|
||||
- admin: 行政
|
||||
|
||||
注意:
|
||||
- 如果问题涉及多个部门,返回 null
|
||||
- 如果问题明显指向某个部门,返回对应英文名
|
||||
- confidence 表示判断置信度,0-1之间"""
|
||||
|
||||
content = call_llm(
|
||||
self.llm_client, prompt, MODEL,
|
||||
temperature=0.1,
|
||||
max_tokens=100
|
||||
)
|
||||
|
||||
if content is None:
|
||||
logger.warning("LLM 意图分析失败: 调用返回空")
|
||||
return None
|
||||
|
||||
# 使用 parse_json_from_response 解析 JSON
|
||||
result = parse_json_from_response(content)
|
||||
|
||||
if result is None:
|
||||
logger.warning(f"LLM 意图分析失败: JSON 解析失败,原始内容: {content[:100]}")
|
||||
return None
|
||||
|
||||
return QueryIntent(
|
||||
is_general=result.get("is_general", False),
|
||||
department=result.get("department"),
|
||||
confidence=result.get("confidence", 0.5),
|
||||
keywords=[],
|
||||
reason="LLM 意图分析"
|
||||
)
|
||||
|
||||
def _select_knowledge_bases(
|
||||
self,
|
||||
intent: QueryIntent,
|
||||
accessible_collections: List[str],
|
||||
role: str,
|
||||
department: str
|
||||
) -> List[str]:
|
||||
"""
|
||||
选择要查询的知识库
|
||||
|
||||
Args:
|
||||
intent: 查询意图
|
||||
accessible_collections: 可访问的向量库
|
||||
role: 用户角色
|
||||
department: 用户部门
|
||||
|
||||
Returns:
|
||||
目标向量库列表
|
||||
"""
|
||||
result = []
|
||||
public_kb = "public_kb"
|
||||
|
||||
# 通用问题:优先查 public_kb
|
||||
if intent.is_general:
|
||||
if public_kb in accessible_collections:
|
||||
result.append(public_kb)
|
||||
# 但也可能需要查其他库(取决于置信度)
|
||||
if intent.confidence < 0.7:
|
||||
result.extend([kb for kb in accessible_collections if kb not in result])
|
||||
|
||||
# 涉及特定部门
|
||||
elif intent.department:
|
||||
dept_kb = f"dept_{intent.department}"
|
||||
|
||||
# 检查是否有权限访问该部门
|
||||
if dept_kb in accessible_collections:
|
||||
result.append(dept_kb)
|
||||
# 也查 public_kb(可能有相关政策)
|
||||
if public_kb in accessible_collections and public_kb not in result:
|
||||
result.append(public_kb)
|
||||
else:
|
||||
# 没有权限访问目标部门,查 public_kb
|
||||
if public_kb in accessible_collections:
|
||||
result.append(public_kb)
|
||||
logger.info(
|
||||
f"用户无权访问部门 {intent.department} 的知识库,"
|
||||
f"只查 public_kb"
|
||||
)
|
||||
|
||||
# 未识别意图
|
||||
else:
|
||||
# admin 查所有
|
||||
if role == "admin":
|
||||
result = accessible_collections
|
||||
# 其他用户查 public 和本部门
|
||||
else:
|
||||
if public_kb in accessible_collections:
|
||||
result.append(public_kb)
|
||||
# 使用标准化的部门名称
|
||||
normalized_dept = normalize_department_name(department)
|
||||
if normalized_dept:
|
||||
user_dept_kb = f"dept_{normalized_dept}"
|
||||
if user_dept_kb in accessible_collections and user_dept_kb not in result:
|
||||
result.append(user_dept_kb)
|
||||
|
||||
# 去重并保持顺序
|
||||
seen = set()
|
||||
unique_result = []
|
||||
for kb in result:
|
||||
if kb not in seen:
|
||||
seen.add(kb)
|
||||
unique_result.append(kb)
|
||||
|
||||
return unique_result
|
||||
|
||||
def get_routing_stats(self) -> Dict:
|
||||
"""获取路由统计信息(用于监控)"""
|
||||
return {
|
||||
"use_llm": self.use_llm,
|
||||
"department_keywords": {
|
||||
dept: len(keywords)
|
||||
for dept, keywords in DEPARTMENT_KEYWORDS.items()
|
||||
},
|
||||
"general_keywords_count": len(GENERAL_KEYWORDS)
|
||||
}
|
||||
|
||||
# ==================== 版本感知检索 ====================
|
||||
|
||||
def route_with_version_awareness(
|
||||
self,
|
||||
query: str,
|
||||
role: str,
|
||||
department: str,
|
||||
accessible_collections: List[str] = None,
|
||||
include_deprecated: bool = False,
|
||||
top_k: int = 5
|
||||
) -> Dict:
|
||||
"""
|
||||
版本感知的路由
|
||||
|
||||
在普通路由基础上,额外查询已废止的相关文档,
|
||||
为用户提供版本提示。
|
||||
|
||||
Args:
|
||||
query: 用户查询
|
||||
role: 用户角色
|
||||
department: 用户部门
|
||||
accessible_collections: 可访问的向量库列表
|
||||
include_deprecated: 是否包含废止版本在结果中
|
||||
top_k: 返回数量
|
||||
|
||||
Returns:
|
||||
{
|
||||
"target_collections": ["public_kb", "dept_finance"],
|
||||
"version_hints": [
|
||||
{
|
||||
"document": "报销制度.pdf",
|
||||
"status": "deprecated",
|
||||
"message": "该文档已于2026-03-01废止"
|
||||
}
|
||||
]
|
||||
}
|
||||
"""
|
||||
# 1. 获取目标向量库(复用现有逻辑)
|
||||
target_kbs = self.route(query, role, department, accessible_collections)
|
||||
|
||||
if not target_kbs:
|
||||
return {
|
||||
"target_collections": [],
|
||||
"version_hints": []
|
||||
}
|
||||
|
||||
# 2. 查询是否有相关的废止版本
|
||||
version_hints = []
|
||||
if not include_deprecated:
|
||||
version_hints = self._find_deprecated_versions(query, target_kbs, top_k=3)
|
||||
|
||||
logger.info(
|
||||
f"版本感知路由: query='{query[:30]}...', "
|
||||
f"targets={target_kbs}, hints={len(version_hints)}"
|
||||
)
|
||||
|
||||
return {
|
||||
"target_collections": target_kbs,
|
||||
"version_hints": version_hints
|
||||
}
|
||||
|
||||
def _find_deprecated_versions(
|
||||
self,
|
||||
query: str,
|
||||
collections: List[str],
|
||||
top_k: int = 3
|
||||
) -> List[Dict]:
|
||||
"""
|
||||
查找与查询相关的已废止版本
|
||||
|
||||
Args:
|
||||
query: 用户查询
|
||||
collections: 目标向量库列表
|
||||
top_k: 每个库返回数量
|
||||
|
||||
Returns:
|
||||
已废止版本提示列表
|
||||
"""
|
||||
try:
|
||||
from knowledge.manager import get_kb_manager
|
||||
|
||||
kb_manager = get_kb_manager()
|
||||
|
||||
# 获取查询向量
|
||||
query_vector = self._get_query_vector(query)
|
||||
if query_vector is None:
|
||||
return []
|
||||
|
||||
# 使用知识库管理器查找废止版本
|
||||
hints = kb_manager.find_deprecated_versions(
|
||||
kb_names=collections,
|
||||
query_vector=query_vector,
|
||||
top_k=top_k
|
||||
)
|
||||
|
||||
# 去重(同一文档只提示一次)
|
||||
seen_docs = set()
|
||||
unique_hints = []
|
||||
for hint in hints:
|
||||
doc_key = f"{hint['collection']}/{hint['document']}"
|
||||
if doc_key not in seen_docs:
|
||||
seen_docs.add(doc_key)
|
||||
unique_hints.append(hint)
|
||||
|
||||
return unique_hints
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"查找废止版本失败: {e}")
|
||||
return []
|
||||
|
||||
def _get_query_vector(self, query: str) -> Optional[List[float]]:
|
||||
"""
|
||||
获取查询向量
|
||||
|
||||
Args:
|
||||
query: 查询文本
|
||||
|
||||
Returns:
|
||||
查询向量,失败返回None
|
||||
"""
|
||||
try:
|
||||
# 尝试使用 RAGEngine 的 embedding_model
|
||||
from core.engine import get_engine
|
||||
embedding_model = get_engine().embedding_model
|
||||
return embedding_model.encode(query).tolist()
|
||||
except Exception as e:
|
||||
logger.debug(f"无法从 RAGEngine 获取向量模型: {e}")
|
||||
|
||||
try:
|
||||
# 尝试使用 sentence-transformers
|
||||
from sentence_transformers import SentenceTransformer
|
||||
model = SentenceTransformer('BAAI/bge-base-zh-v1.5')
|
||||
return model.encode(query).tolist()
|
||||
except Exception as e:
|
||||
logger.debug(f"Embedding 编码失败: {e}")
|
||||
|
||||
logger.warning("无法加载向量模型,跳过废止版本检测")
|
||||
return None
|
||||
|
||||
def search_with_version_context(
|
||||
self,
|
||||
query: str,
|
||||
role: str,
|
||||
department: str,
|
||||
top_k: int = 5
|
||||
) -> Dict:
|
||||
"""
|
||||
带版本上下文的搜索
|
||||
|
||||
执行完整搜索流程:
|
||||
1. 版本感知路由
|
||||
2. 执行检索(只返回生效版本)
|
||||
3. 返回结果 + 废止版本提示
|
||||
|
||||
Args:
|
||||
query: 用户查询
|
||||
role: 用户角色
|
||||
department: 用户部门
|
||||
top_k: 返回数量
|
||||
|
||||
Returns:
|
||||
{
|
||||
"results": [...], # 生效版本的检索结果
|
||||
"version_hints": [...], # 废止版本提示
|
||||
"target_collections": [...]
|
||||
}
|
||||
"""
|
||||
from knowledge.manager import get_kb_manager
|
||||
|
||||
kb_manager = get_kb_manager()
|
||||
|
||||
# 1. 版本感知路由
|
||||
route_result = self.route_with_version_awareness(
|
||||
query, role, department, include_deprecated=False
|
||||
)
|
||||
|
||||
target_kbs = route_result["target_collections"]
|
||||
version_hints = route_result["version_hints"]
|
||||
|
||||
if not target_kbs:
|
||||
return {
|
||||
"results": [],
|
||||
"version_hints": version_hints,
|
||||
"target_collections": []
|
||||
}
|
||||
|
||||
# 2. 执行检索(只返回生效版本)
|
||||
query_vector = self._get_query_vector(query)
|
||||
if query_vector is None:
|
||||
return {
|
||||
"results": [],
|
||||
"version_hints": version_hints,
|
||||
"target_collections": target_kbs
|
||||
}
|
||||
|
||||
# 多库检索,只返回active状态的chunks
|
||||
search_result = kb_manager.search_multiple(
|
||||
kb_names=target_kbs,
|
||||
query_vector=query_vector,
|
||||
query_text=query,
|
||||
top_k=top_k,
|
||||
use_bm25=True
|
||||
)
|
||||
|
||||
# 过滤只返回active状态的chunks
|
||||
active_results = []
|
||||
if search_result.ids:
|
||||
for i, (doc_id, doc, meta, score) in enumerate(zip(
|
||||
search_result.ids,
|
||||
search_result.documents,
|
||||
search_result.metadatas,
|
||||
search_result.distances
|
||||
)):
|
||||
if meta.get("status", "active") == "active":
|
||||
active_results.append({
|
||||
"id": doc_id,
|
||||
"document": doc,
|
||||
"metadata": meta,
|
||||
"score": score
|
||||
})
|
||||
|
||||
return {
|
||||
"results": active_results[:top_k],
|
||||
"version_hints": version_hints,
|
||||
"target_collections": target_kbs
|
||||
}
|
||||
|
||||
|
||||
# ==================== 全局实例 ====================
|
||||
|
||||
_kb_router: Optional[KnowledgeBaseRouter] = None
|
||||
|
||||
|
||||
def get_kb_router() -> KnowledgeBaseRouter:
|
||||
"""获取全局知识库路由器实例"""
|
||||
global _kb_router
|
||||
if _kb_router is None:
|
||||
_kb_router = KnowledgeBaseRouter()
|
||||
return _kb_router
|
||||
|
||||
|
||||
# ==================== 便捷函数 ====================
|
||||
|
||||
def route_query(
|
||||
query: str,
|
||||
role: str,
|
||||
department: str,
|
||||
accessible_collections: List[str] = None
|
||||
) -> List[str]:
|
||||
"""
|
||||
路由查询到目标知识库(便捷函数)
|
||||
|
||||
Args:
|
||||
query: 用户查询
|
||||
role: 用户角色
|
||||
department: 用户部门
|
||||
accessible_collections: 可访问的向量库列表
|
||||
|
||||
Returns:
|
||||
目标向量库列表
|
||||
"""
|
||||
router = get_kb_router()
|
||||
return router.route(query, role, department, accessible_collections)
|
||||
|
||||
|
||||
def route_query_with_version(
|
||||
query: str,
|
||||
role: str,
|
||||
department: str,
|
||||
accessible_collections: List[str] = None,
|
||||
include_deprecated: bool = False
|
||||
) -> Dict:
|
||||
"""
|
||||
版本感知的路由(便捷函数)
|
||||
|
||||
Args:
|
||||
query: 用户查询
|
||||
role: 用户角色
|
||||
department: 用户部门
|
||||
accessible_collections: 可访问的向量库列表
|
||||
include_deprecated: 是否包含废止版本
|
||||
|
||||
Returns:
|
||||
{
|
||||
"target_collections": [...],
|
||||
"version_hints": [...]
|
||||
}
|
||||
"""
|
||||
router = get_kb_router()
|
||||
return router.route_with_version_awareness(
|
||||
query, role, department, accessible_collections, include_deprecated
|
||||
)
|
||||
|
||||
|
||||
def search_with_version_context(
|
||||
query: str,
|
||||
role: str,
|
||||
department: str,
|
||||
top_k: int = 5
|
||||
) -> Dict:
|
||||
"""
|
||||
带版本上下文的搜索(便捷函数)
|
||||
|
||||
Args:
|
||||
query: 用户查询
|
||||
role: 用户角色
|
||||
department: 用户部门
|
||||
top_k: 返回数量
|
||||
|
||||
Returns:
|
||||
{
|
||||
"results": [...],
|
||||
"version_hints": [...],
|
||||
"target_collections": [...]
|
||||
}
|
||||
"""
|
||||
router = get_kb_router()
|
||||
return router.search_with_version_context(query, role, department, top_k)
|
||||
382
knowledge/search.py
Normal file
382
knowledge/search.py
Normal file
@@ -0,0 +1,382 @@
|
||||
"""
|
||||
知识库管理器 - 检索功能 Mixin
|
||||
|
||||
提供多源融合检索功能,支持:
|
||||
- 单向量库检索:向量检索 + BM25 混合检索,RRF 融合排序
|
||||
- 多向量库并行检索:跨多个向量库并行检索并合并结果
|
||||
- 废止版本检测:查找与查询相关的已废止文档
|
||||
|
||||
检索流程:
|
||||
1. 向量检索:使用 cosine 相似度在 ChromaDB 中检索
|
||||
2. BM25 检索:使用关键词匹配在 BM25 索引中检索
|
||||
3. RRF 融合:使用 Reciprocal Rank Fusion 合并两路结果
|
||||
4. 过滤:排除已废止/已替代的文档
|
||||
|
||||
主要方法:
|
||||
- search_single: 单向量库检索
|
||||
- search_multiple: 多向量库并行检索
|
||||
- find_deprecated_versions: 查找已废止版本
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import List, Tuple, Dict, Optional
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
|
||||
from .base import SearchResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SearchMixin:
|
||||
"""
|
||||
检索功能 Mixin
|
||||
|
||||
提供混合检索(向量 + BM25)和多库并行检索能力。
|
||||
|
||||
依赖属性(需由主类提供):
|
||||
- self.get_collection: 获取向量库集合的方法
|
||||
- self.get_bm25_index: 获取 BM25 索引的方法
|
||||
"""
|
||||
|
||||
def search_single(
|
||||
self,
|
||||
kb_name: str,
|
||||
query_vector: List[float],
|
||||
query_text: str,
|
||||
top_k: int = 5,
|
||||
use_bm25: bool = True,
|
||||
include_deprecated: bool = False
|
||||
) -> Optional[SearchResult]:
|
||||
"""
|
||||
单向量库检索
|
||||
|
||||
对单个向量库执行混合检索:向量检索 + BM25 检索,
|
||||
使用 RRF (Reciprocal Rank Fusion) 融合排序。
|
||||
|
||||
Args:
|
||||
kb_name: 向量库名称
|
||||
query_vector: 查询向量(由 embedding 模型生成)
|
||||
query_text: 查询文本(用于 BM25 检索)
|
||||
top_k: 返回结果数量(默认 5)
|
||||
use_bm25: 是否启用 BM25 混合检索(默认 True)
|
||||
include_deprecated: 是否包含已废止/已替代的文档(默认 False)
|
||||
|
||||
Returns:
|
||||
SearchResult 对象,包含:
|
||||
- ids: 文档 ID 列表
|
||||
- documents: 文档内容列表
|
||||
- metadatas: 元数据列表
|
||||
- distances: 距离/分数列表
|
||||
- collection_name: 向量库名称
|
||||
|
||||
向量库为空时返回 None
|
||||
"""
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection or collection.count() == 0:
|
||||
return None
|
||||
|
||||
where_filter = None
|
||||
if not include_deprecated:
|
||||
where_filter = {"status": "active"}
|
||||
|
||||
vector_result = collection.query(
|
||||
query_embeddings=[query_vector],
|
||||
n_results=top_k,
|
||||
where=where_filter
|
||||
)
|
||||
|
||||
if not use_bm25:
|
||||
return SearchResult(
|
||||
ids=vector_result['ids'][0] if vector_result['ids'] else [],
|
||||
documents=vector_result['documents'][0] if vector_result['documents'] else [],
|
||||
metadatas=vector_result['metadatas'][0] if vector_result['metadatas'] else [],
|
||||
distances=vector_result['distances'][0] if vector_result['distances'] else [],
|
||||
collection_name=kb_name
|
||||
)
|
||||
|
||||
bm25_index = self.get_bm25_index(kb_name)
|
||||
bm25_ids, bm25_docs, bm25_metas, bm25_scores = bm25_index.search(
|
||||
query_text, top_k=min(top_k * 2, 20)
|
||||
)
|
||||
|
||||
if not include_deprecated and bm25_metas:
|
||||
filtered_bm25 = []
|
||||
for i, meta in enumerate(bm25_metas):
|
||||
if meta.get('status', 'active') == 'active':
|
||||
filtered_bm25.append((bm25_ids[i], bm25_docs[i], bm25_metas[i], bm25_scores[i]))
|
||||
|
||||
if filtered_bm25:
|
||||
bm25_ids, bm25_docs, bm25_metas, bm25_scores = zip(*filtered_bm25)
|
||||
else:
|
||||
bm25_ids, bm25_docs, bm25_metas, bm25_scores = [], [], [], []
|
||||
|
||||
return self._merge_results(
|
||||
vector_result,
|
||||
(list(bm25_ids), list(bm25_docs), list(bm25_metas), list(bm25_scores)),
|
||||
top_k=top_k,
|
||||
collection_name=kb_name
|
||||
)
|
||||
|
||||
def search_multiple(
|
||||
self,
|
||||
kb_names: List[str],
|
||||
query_vector: List[float],
|
||||
query_text: str,
|
||||
top_k: int = 5,
|
||||
use_bm25: bool = True
|
||||
) -> SearchResult:
|
||||
"""
|
||||
多向量库并行检索
|
||||
|
||||
同时在多个向量库中检索,使用线程池并行执行,
|
||||
最终合并去重并按分数排序。
|
||||
|
||||
Args:
|
||||
kb_names: 向量库名称列表
|
||||
query_vector: 查询向量
|
||||
query_text: 查询文本
|
||||
top_k: 每个库返回的数量
|
||||
use_bm25: 是否启用 BM25
|
||||
|
||||
Returns:
|
||||
合并后的 SearchResult 对象
|
||||
"""
|
||||
if not kb_names:
|
||||
return SearchResult(
|
||||
ids=[], documents=[], metadatas=[], distances=[]
|
||||
)
|
||||
|
||||
results = []
|
||||
with ThreadPoolExecutor(max_workers=len(kb_names)) as executor:
|
||||
futures = {
|
||||
executor.submit(
|
||||
self.search_single,
|
||||
kb_name,
|
||||
query_vector,
|
||||
query_text,
|
||||
top_k,
|
||||
use_bm25
|
||||
): kb_name for kb_name in kb_names
|
||||
}
|
||||
|
||||
for future in as_completed(futures):
|
||||
result = future.result()
|
||||
if result:
|
||||
results.append(result)
|
||||
|
||||
return self._merge_multiple_results(results, top_k)
|
||||
|
||||
def _merge_results(
|
||||
self,
|
||||
vector_result: dict,
|
||||
bm25_result: Tuple,
|
||||
top_k: int,
|
||||
collection_name: str
|
||||
) -> SearchResult:
|
||||
"""
|
||||
RRF 融合向量检索和 BM25 检索结果
|
||||
|
||||
使用 Reciprocal Rank Fusion 算法合并两路检索结果,
|
||||
综合考虑向量相似度和 BM25 分数进行排序。
|
||||
|
||||
Args:
|
||||
vector_result: ChromaDB 向量检索结果
|
||||
bm25_result: BM25 检索结果元组 (ids, docs, metas, scores)
|
||||
top_k: 返回数量
|
||||
collection_name: 向量库名称
|
||||
|
||||
Returns:
|
||||
融合后的 SearchResult 对象
|
||||
|
||||
Note:
|
||||
RRF 参数 k=60,向量权重 0.5,BM25 权重 0.5。
|
||||
"""
|
||||
k = 60 # RRF 参数
|
||||
|
||||
doc_scores = {}
|
||||
|
||||
# 向量检索结果
|
||||
if vector_result['ids'] and vector_result['ids'][0]:
|
||||
for rank, (doc_id, doc, meta, dist) in enumerate(zip(
|
||||
vector_result['ids'][0],
|
||||
vector_result['documents'][0],
|
||||
vector_result['metadatas'][0],
|
||||
vector_result['distances'][0]
|
||||
)):
|
||||
rrf_score = 1 / (k + rank + 1)
|
||||
sim_score = 1 - dist
|
||||
combined = rrf_score * 0.5 + sim_score * 0.5
|
||||
|
||||
doc_scores[doc_id] = {
|
||||
'score': combined,
|
||||
'doc': doc,
|
||||
'meta': meta
|
||||
}
|
||||
|
||||
# BM25 结果
|
||||
bm25_ids, bm25_docs, bm25_metas, bm25_scores = bm25_result
|
||||
for rank, (doc_id, doc, meta, score) in enumerate(zip(
|
||||
bm25_ids, bm25_docs, bm25_metas, bm25_scores
|
||||
)):
|
||||
rrf_score = 1 / (k + rank + 1)
|
||||
norm_score = score / 10.0 if score > 0 else 0
|
||||
combined = rrf_score * 0.5 + norm_score * 0.5
|
||||
|
||||
if doc_id in doc_scores:
|
||||
doc_scores[doc_id]['score'] += combined
|
||||
else:
|
||||
doc_scores[doc_id] = {
|
||||
'score': combined,
|
||||
'doc': doc,
|
||||
'meta': meta
|
||||
}
|
||||
|
||||
# 排序
|
||||
sorted_items = sorted(
|
||||
doc_scores.items(),
|
||||
key=lambda x: x[1]['score'],
|
||||
reverse=True
|
||||
)[:top_k]
|
||||
|
||||
return SearchResult(
|
||||
ids=[item[0] for item in sorted_items],
|
||||
documents=[item[1]['doc'] for item in sorted_items],
|
||||
metadatas=[item[1]['meta'] for item in sorted_items],
|
||||
distances=[item[1]['score'] for item in sorted_items],
|
||||
collection_name=collection_name
|
||||
)
|
||||
|
||||
def _merge_multiple_results(
|
||||
self,
|
||||
results: List[SearchResult],
|
||||
top_k: int
|
||||
) -> SearchResult:
|
||||
"""
|
||||
合并多个向量库的检索结果
|
||||
|
||||
将多个向量库的检索结果合并、去重、排序。
|
||||
|
||||
Args:
|
||||
results: 各向量库的检索结果列表
|
||||
top_k: 最终返回数量
|
||||
|
||||
Returns:
|
||||
合并后的 SearchResult 对象
|
||||
"""
|
||||
if not results:
|
||||
return SearchResult(
|
||||
ids=[], documents=[], metadatas=[], distances=[]
|
||||
)
|
||||
|
||||
if len(results) == 1:
|
||||
return results[0]
|
||||
|
||||
all_items = []
|
||||
for result in results:
|
||||
for i, doc_id in enumerate(result.ids):
|
||||
all_items.append({
|
||||
'id': doc_id,
|
||||
'doc': result.documents[i],
|
||||
'meta': result.metadatas[i],
|
||||
'score': result.distances[i],
|
||||
'collection': result.collection_name
|
||||
})
|
||||
|
||||
all_items.sort(key=lambda x: x['score'], reverse=True)
|
||||
|
||||
seen = set()
|
||||
unique_items = []
|
||||
for item in all_items:
|
||||
if item['id'] not in seen:
|
||||
seen.add(item['id'])
|
||||
unique_items.append(item)
|
||||
|
||||
unique_items = unique_items[:top_k]
|
||||
|
||||
return SearchResult(
|
||||
ids=[item['id'] for item in unique_items],
|
||||
documents=[item['doc'] for item in unique_items],
|
||||
metadatas=[item['meta'] for item in unique_items],
|
||||
distances=[item['score'] for item in unique_items],
|
||||
collection_name="multiple"
|
||||
)
|
||||
|
||||
def find_deprecated_versions(
|
||||
self,
|
||||
kb_names: List[str],
|
||||
query_vector: List[float],
|
||||
top_k: int = 3
|
||||
) -> List[Dict]:
|
||||
"""
|
||||
查找与查询相关的已废止版本
|
||||
|
||||
在指定向量库中搜索已废止的文档,
|
||||
当相似度 >= 0.7 时返回废止提示信息。
|
||||
|
||||
Args:
|
||||
kb_names: 向量库名称列表
|
||||
query_vector: 查询向量
|
||||
top_k: 每个库返回的数量
|
||||
|
||||
Returns:
|
||||
废止提示列表,每个元素包含:
|
||||
- document: 文档来源
|
||||
- collection: 向量库名称
|
||||
- status: 状态("deprecated")
|
||||
- deprecated_date: 废止日期
|
||||
- deprecated_reason: 废止原因
|
||||
- similarity: 相似度分数
|
||||
- snippet: 内容摘要
|
||||
- message: 废止提示消息
|
||||
"""
|
||||
hints = []
|
||||
|
||||
for kb_name in kb_names:
|
||||
collection = self.get_collection(kb_name)
|
||||
if not collection:
|
||||
continue
|
||||
|
||||
result = collection.query(
|
||||
query_embeddings=[query_vector],
|
||||
n_results=top_k,
|
||||
where={"status": "deprecated"}
|
||||
)
|
||||
|
||||
if result['ids'] and result['ids'][0]:
|
||||
for doc, meta, score in zip(
|
||||
result['documents'][0],
|
||||
result['metadatas'][0],
|
||||
result['distances'][0]
|
||||
):
|
||||
sim_score = 1 - score
|
||||
if sim_score >= 0.7:
|
||||
hints.append({
|
||||
"document": meta.get("source", ""),
|
||||
"collection": kb_name,
|
||||
"status": "deprecated",
|
||||
"deprecated_date": meta.get("deprecated_date", ""),
|
||||
"deprecated_reason": meta.get("deprecated_reason", ""),
|
||||
"similarity": sim_score,
|
||||
"snippet": doc[:100] + "..." if len(doc) > 100 else doc,
|
||||
"message": self._build_deprecation_hint(meta)
|
||||
})
|
||||
|
||||
return hints
|
||||
|
||||
def _build_deprecation_hint(self, metadata: Dict) -> str:
|
||||
"""
|
||||
构建废止提示消息
|
||||
|
||||
Args:
|
||||
metadata: 切片元数据,包含 deprecated_date 和 deprecated_reason
|
||||
|
||||
Returns:
|
||||
格式化的废止提示消息
|
||||
"""
|
||||
deprecated_date = metadata.get("deprecated_date", "")
|
||||
deprecated_reason = metadata.get("deprecated_reason", "")
|
||||
|
||||
date_str = deprecated_date[:10] if deprecated_date else "未知日期"
|
||||
reason_str = f",原因:{deprecated_reason}" if deprecated_reason else ""
|
||||
|
||||
return f"⚠️ 该文档已于 {date_str} 废止{reason_str},内容不再有效"
|
||||
891
knowledge/sync.py
Normal file
891
knowledge/sync.py
Normal file
@@ -0,0 +1,891 @@
|
||||
"""
|
||||
知识库同步服务 - 自动检测文档变更并触发增量更新
|
||||
|
||||
功能:
|
||||
1. 文件变更监控 - 使用 watchdog 监控 documents 目录
|
||||
2. 哈希比对 - 识别文件具体变更类型(新增/修改/删除)
|
||||
3. 增量向量化 - 仅处理变更文件
|
||||
4. 变更日志 - 记录变更历史
|
||||
|
||||
使用方式:
|
||||
from knowledge.sync import KnowledgeSyncService
|
||||
|
||||
# 启动同步服务
|
||||
sync_service = KnowledgeSyncService()
|
||||
sync_service.start() # 启动后台监控
|
||||
|
||||
# 手动触发同步
|
||||
result = sync_service.sync_now()
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
import json
|
||||
import hashlib
|
||||
import threading
|
||||
import time
|
||||
from datetime import datetime
|
||||
from typing import Dict, List, Optional, Callable
|
||||
from dataclasses import dataclass, asdict
|
||||
from enum import Enum
|
||||
import logging
|
||||
|
||||
from data.db import get_connection, init_databases
|
||||
|
||||
# 缓存支持
|
||||
try:
|
||||
from core.cache import get_cache_manager
|
||||
CACHE_AVAILABLE = True
|
||||
except ImportError:
|
||||
CACHE_AVAILABLE = False
|
||||
|
||||
# 设置日志
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 尝试导入 watchdog
|
||||
try:
|
||||
from watchdog.observers import Observer
|
||||
from watchdog.events import FileSystemEventHandler, FileCreatedEvent, FileModifiedEvent, FileDeletedEvent
|
||||
HAS_WATCHDOG = True
|
||||
except ImportError:
|
||||
HAS_WATCHDOG = False
|
||||
logger.warning("watchdog 未安装,文件监控功能不可用。请运行: pip install watchdog")
|
||||
|
||||
|
||||
class ChangeType(Enum):
|
||||
"""变更类型"""
|
||||
ADDED = "added" # 新增
|
||||
MODIFIED = "modified" # 修改
|
||||
DELETED = "deleted" # 删除
|
||||
|
||||
|
||||
class SyncStatus(Enum):
|
||||
"""同步状态"""
|
||||
IDLE = "idle" # 空闲
|
||||
RUNNING = "running" # 运行中
|
||||
COMPLETED = "completed" # 已完成
|
||||
FAILED = "failed" # 失败
|
||||
|
||||
|
||||
@dataclass
|
||||
class DocumentChange:
|
||||
"""文档变更记录"""
|
||||
document_id: str # 文档ID(相对路径)
|
||||
document_name: str # 文件名
|
||||
change_type: ChangeType # 变更类型
|
||||
old_hash: Optional[str] # 旧哈希
|
||||
new_hash: Optional[str] # 新哈希
|
||||
change_time: datetime # 变更时间
|
||||
processed: bool = False # 是否已处理
|
||||
error_message: Optional[str] = None
|
||||
|
||||
def to_dict(self):
|
||||
return {
|
||||
"document_id": self.document_id,
|
||||
"document_name": self.document_name,
|
||||
"change_type": self.change_type.value,
|
||||
"old_hash": self.old_hash,
|
||||
"new_hash": self.new_hash,
|
||||
"change_time": self.change_time.isoformat(),
|
||||
"processed": self.processed,
|
||||
"error_message": self.error_message
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class SyncResult:
|
||||
"""同步结果"""
|
||||
status: SyncStatus
|
||||
start_time: datetime
|
||||
end_time: Optional[datetime]
|
||||
documents_processed: int
|
||||
documents_added: int
|
||||
documents_modified: int
|
||||
documents_deleted: int
|
||||
errors: List[str]
|
||||
|
||||
def to_dict(self):
|
||||
return {
|
||||
"status": self.status.value,
|
||||
"start_time": self.start_time.isoformat(),
|
||||
"end_time": self.end_time.isoformat() if self.end_time else None,
|
||||
"documents_processed": self.documents_processed,
|
||||
"documents_added": self.documents_added,
|
||||
"documents_modified": self.documents_modified,
|
||||
"documents_deleted": self.documents_deleted,
|
||||
"errors": self.errors
|
||||
}
|
||||
|
||||
|
||||
class SyncDatabase:
|
||||
"""同步数据库管理"""
|
||||
|
||||
def __init__(self):
|
||||
"""初始化数据库"""
|
||||
init_databases()
|
||||
|
||||
def get_document_hash(self, document_id: str) -> Optional[Dict]:
|
||||
"""获取文档的当前哈希"""
|
||||
with get_connection("knowledge") as conn:
|
||||
cursor = conn.cursor()
|
||||
cursor.execute('''
|
||||
SELECT document_id, document_name, content_hash, file_size, last_modified
|
||||
FROM document_hashes WHERE document_id = ?
|
||||
''', (document_id,))
|
||||
row = cursor.fetchone()
|
||||
|
||||
if row:
|
||||
return {
|
||||
"document_id": row[0],
|
||||
"document_name": row[1],
|
||||
"content_hash": row[2],
|
||||
"file_size": row[3],
|
||||
"last_modified": row[4]
|
||||
}
|
||||
return None
|
||||
|
||||
def set_document_hash(self, document_id: str, document_name: str,
|
||||
content_hash: str, file_size: int, last_modified: datetime):
|
||||
"""设置文档哈希"""
|
||||
with get_connection("knowledge") as conn:
|
||||
cursor = conn.cursor()
|
||||
cursor.execute('''
|
||||
INSERT OR REPLACE INTO document_hashes
|
||||
(document_id, document_name, content_hash, file_size, last_modified, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, CURRENT_TIMESTAMP)
|
||||
''', (document_id, document_name, content_hash, file_size, last_modified))
|
||||
|
||||
def delete_document_hash(self, document_id: str):
|
||||
"""删除文档哈希记录"""
|
||||
with get_connection("knowledge") as conn:
|
||||
cursor = conn.cursor()
|
||||
cursor.execute('DELETE FROM document_hashes WHERE document_id = ?', (document_id,))
|
||||
|
||||
def get_all_document_hashes(self) -> Dict[str, Dict]:
|
||||
"""获取所有文档哈希"""
|
||||
with get_connection("knowledge") as conn:
|
||||
cursor = conn.cursor()
|
||||
cursor.execute('''
|
||||
SELECT document_id, document_name, content_hash, file_size, last_modified
|
||||
FROM document_hashes
|
||||
''')
|
||||
rows = cursor.fetchall()
|
||||
|
||||
return {
|
||||
row[0]: {
|
||||
"document_id": row[0],
|
||||
"document_name": row[1],
|
||||
"content_hash": row[2],
|
||||
"file_size": row[3],
|
||||
"last_modified": row[4]
|
||||
}
|
||||
for row in rows
|
||||
}
|
||||
|
||||
def log_change(self, change: DocumentChange) -> int:
|
||||
"""记录变更"""
|
||||
with get_connection("knowledge") as conn:
|
||||
cursor = conn.cursor()
|
||||
cursor.execute('''
|
||||
INSERT INTO change_logs
|
||||
(document_id, document_name, change_type, old_hash, new_hash, change_time, processed, error_message)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
|
||||
''', (
|
||||
change.document_id,
|
||||
change.document_name,
|
||||
change.change_type.value,
|
||||
change.old_hash,
|
||||
change.new_hash,
|
||||
change.change_time,
|
||||
change.processed,
|
||||
change.error_message
|
||||
))
|
||||
return cursor.lastrowid
|
||||
|
||||
def get_change_logs(self, limit: int = 100, processed: Optional[bool] = None,
|
||||
days: int = 30) -> List[Dict]:
|
||||
"""获取变更日志"""
|
||||
with get_connection("knowledge") as conn:
|
||||
cursor = conn.cursor()
|
||||
|
||||
sql = '''
|
||||
SELECT id, document_id, document_name, change_type, old_hash, new_hash,
|
||||
change_time, processed, error_message
|
||||
FROM change_logs
|
||||
WHERE change_time >= datetime('now', ?)
|
||||
'''
|
||||
params = [f'-{days} days']
|
||||
|
||||
if processed is not None:
|
||||
sql += ' AND processed = ?'
|
||||
params.append(1 if processed else 0)
|
||||
|
||||
sql += ' ORDER BY change_time DESC LIMIT ?'
|
||||
params.append(limit)
|
||||
|
||||
cursor.execute(sql, params)
|
||||
rows = cursor.fetchall()
|
||||
|
||||
return [
|
||||
{
|
||||
"id": row[0],
|
||||
"document_id": row[1],
|
||||
"document_name": row[2],
|
||||
"change_type": row[3],
|
||||
"old_hash": row[4],
|
||||
"new_hash": row[5],
|
||||
"change_time": row[6],
|
||||
"processed": bool(row[7]),
|
||||
"error_message": row[8]
|
||||
}
|
||||
for row in rows
|
||||
]
|
||||
|
||||
def mark_change_processed(self, change_id: int, error_message: str = None):
|
||||
"""标记变更已处理"""
|
||||
with get_connection("knowledge") as conn:
|
||||
cursor = conn.cursor()
|
||||
cursor.execute('''
|
||||
UPDATE change_logs
|
||||
SET processed = 1, error_message = ?
|
||||
WHERE id = ?
|
||||
''', (error_message, change_id))
|
||||
|
||||
def log_sync_status(self, result: SyncResult) -> int:
|
||||
"""记录同步状态"""
|
||||
with get_connection("knowledge") as conn:
|
||||
cursor = conn.cursor()
|
||||
cursor.execute('''
|
||||
INSERT INTO sync_status
|
||||
(sync_type, status, start_time, end_time, documents_processed,
|
||||
documents_added, documents_modified, documents_deleted, error_message)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
''', (
|
||||
"incremental",
|
||||
result.status.value,
|
||||
result.start_time,
|
||||
result.end_time,
|
||||
result.documents_processed,
|
||||
result.documents_added,
|
||||
result.documents_modified,
|
||||
result.documents_deleted,
|
||||
"; ".join(result.errors) if result.errors else None
|
||||
))
|
||||
return cursor.lastrowid
|
||||
|
||||
def get_sync_history(self, limit: int = 20) -> List[Dict]:
|
||||
"""获取同步历史"""
|
||||
with get_connection("knowledge") as conn:
|
||||
cursor = conn.cursor()
|
||||
cursor.execute('''
|
||||
SELECT id, sync_type, status, start_time, end_time,
|
||||
documents_processed, documents_added, documents_modified, documents_deleted, error_message
|
||||
FROM sync_status
|
||||
ORDER BY start_time DESC
|
||||
LIMIT ?
|
||||
''', (limit,))
|
||||
rows = cursor.fetchall()
|
||||
|
||||
return [
|
||||
{
|
||||
"id": row[0],
|
||||
"sync_type": row[1],
|
||||
"status": row[2],
|
||||
"start_time": row[3],
|
||||
"end_time": row[4],
|
||||
"documents_processed": row[5],
|
||||
"documents_added": row[6],
|
||||
"documents_modified": row[7],
|
||||
"documents_deleted": row[8],
|
||||
"error_message": row[9]
|
||||
}
|
||||
for row in rows
|
||||
]
|
||||
|
||||
|
||||
class FileChangeHandler(FileSystemEventHandler if HAS_WATCHDOG else object):
|
||||
"""文件变更处理器"""
|
||||
|
||||
def __init__(self, sync_service: 'KnowledgeSyncService'):
|
||||
if HAS_WATCHDOG:
|
||||
super().__init__()
|
||||
self.sync_service = sync_service
|
||||
# v5 统一解析支持的所有格式
|
||||
self.supported_extensions = {'.pdf', '.docx', '.doc', '.xlsx', '.xls', '.pptx', '.txt', '.png', '.jpg', '.jpeg', '.bmp', '.tiff'}
|
||||
self._pending_changes = {} # 防抖:短时间内多次修改只记录一次
|
||||
self._debounce_seconds = 2
|
||||
|
||||
def _is_supported_file(self, file_path: str) -> bool:
|
||||
"""检查是否为支持的文件类型"""
|
||||
ext = os.path.splitext(file_path)[1].lower()
|
||||
return ext in self.supported_extensions
|
||||
|
||||
def _debounce_change(self, file_path: str, change_type: ChangeType):
|
||||
"""防抖处理:短时间内多次修改合并为一次"""
|
||||
current_time = time.time()
|
||||
|
||||
if file_path in self._pending_changes:
|
||||
last_time, last_type = self._pending_changes[file_path]
|
||||
# 如果是修改事件且距离上次事件很近,忽略
|
||||
if current_time - last_time < self._debounce_seconds:
|
||||
return
|
||||
|
||||
self._pending_changes[file_path] = (current_time, change_type)
|
||||
|
||||
# 延迟处理
|
||||
threading.Timer(self._debounce_seconds, self._process_change, args=[file_path, change_type]).start()
|
||||
|
||||
def _process_change(self, file_path: str, change_type: ChangeType):
|
||||
"""处理文件变更"""
|
||||
try:
|
||||
# 计算相对路径
|
||||
rel_path = os.path.relpath(file_path, self.sync_service.documents_path).replace(chr(92), "/")
|
||||
document_name = os.path.basename(file_path)
|
||||
|
||||
logger.info(f"检测到文件变更: {rel_path} ({change_type.value})")
|
||||
|
||||
# 创建变更记录
|
||||
change = DocumentChange(
|
||||
document_id=rel_path,
|
||||
document_name=document_name,
|
||||
change_type=change_type,
|
||||
old_hash=None,
|
||||
new_hash=None,
|
||||
change_time=datetime.now()
|
||||
)
|
||||
|
||||
# 获取旧哈希
|
||||
old_doc = self.sync_service.db.get_document_hash(rel_path)
|
||||
if old_doc:
|
||||
change.old_hash = old_doc['content_hash']
|
||||
|
||||
# 计算新哈希(如果不是删除)
|
||||
if change_type != ChangeType.DELETED and os.path.exists(file_path):
|
||||
change.new_hash = self.sync_service.calculate_file_hash(file_path)
|
||||
|
||||
# 记录变更
|
||||
self.sync_service.db.log_change(change)
|
||||
|
||||
# 触发回调
|
||||
if self.sync_service.on_change_callback:
|
||||
self.sync_service.on_change_callback(change)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"处理文件变更失败: {file_path}, 错误: {e}")
|
||||
|
||||
def on_created(self, event):
|
||||
"""文件创建事件"""
|
||||
if event.is_directory:
|
||||
return
|
||||
if not self._is_supported_file(event.src_path):
|
||||
return
|
||||
self._debounce_change(event.src_path, ChangeType.ADDED)
|
||||
|
||||
def on_modified(self, event):
|
||||
"""文件修改事件"""
|
||||
if event.is_directory:
|
||||
return
|
||||
if not self._is_supported_file(event.src_path):
|
||||
return
|
||||
self._debounce_change(event.src_path, ChangeType.MODIFIED)
|
||||
|
||||
def on_deleted(self, event):
|
||||
"""文件删除事件"""
|
||||
if event.is_directory:
|
||||
return
|
||||
if not self._is_supported_file(event.src_path):
|
||||
return
|
||||
self._debounce_change(event.src_path, ChangeType.DELETED)
|
||||
|
||||
def on_moved(self, event):
|
||||
"""文件移动事件"""
|
||||
if event.is_directory:
|
||||
return
|
||||
# 移动视为删除旧文件 + 创建新文件
|
||||
if self._is_supported_file(event.src_path):
|
||||
self._debounce_change(event.src_path, ChangeType.DELETED)
|
||||
if self._is_supported_file(event.dest_path):
|
||||
self._debounce_change(event.dest_path, ChangeType.ADDED)
|
||||
|
||||
|
||||
class KnowledgeSyncService:
|
||||
"""知识库同步服务"""
|
||||
|
||||
def __init__(self, documents_path: str = None):
|
||||
"""
|
||||
初始化同步服务
|
||||
|
||||
Args:
|
||||
documents_path: 文档目录路径,默认为 ./documents
|
||||
"""
|
||||
self.documents_path = documents_path or os.path.join(
|
||||
os.path.dirname(os.path.abspath(__file__)), "documents"
|
||||
)
|
||||
self.db = SyncDatabase()
|
||||
|
||||
self._observer = None
|
||||
self._running = False
|
||||
self.on_change_callback: Optional[Callable] = None
|
||||
self.on_sync_callback: Optional[Callable] = None
|
||||
|
||||
|
||||
|
||||
@staticmethod
|
||||
def calculate_file_hash(file_path: str) -> str:
|
||||
"""计算文件哈希"""
|
||||
hasher = hashlib.md5()
|
||||
try:
|
||||
with open(file_path, 'rb') as f:
|
||||
for chunk in iter(lambda: f.read(8192), b''):
|
||||
hasher.update(chunk)
|
||||
return hasher.hexdigest()
|
||||
except Exception as e:
|
||||
logger.error(f"计算文件哈希失败: {file_path}, 错误: {e}")
|
||||
return ""
|
||||
|
||||
def scan_documents(self) -> Dict[str, Dict]:
|
||||
"""扫描文档目录,返回所有文档信息"""
|
||||
documents = {}
|
||||
# v5 统一解析支持的所有格式
|
||||
supported_extensions = {'.pdf', '.docx', '.doc', '.xlsx', '.xls', '.pptx', '.txt', '.png', '.jpg', '.jpeg', '.bmp', '.tiff'}
|
||||
|
||||
for root, dirs, files in os.walk(self.documents_path):
|
||||
for filename in files:
|
||||
ext = os.path.splitext(filename)[1].lower()
|
||||
if ext not in supported_extensions:
|
||||
continue
|
||||
|
||||
file_path = os.path.join(root, filename)
|
||||
rel_path = os.path.relpath(file_path, self.documents_path).replace(chr(92), "/")
|
||||
|
||||
try:
|
||||
file_stat = os.stat(file_path)
|
||||
documents[rel_path] = {
|
||||
"document_id": rel_path,
|
||||
"document_name": filename,
|
||||
"file_path": file_path,
|
||||
"file_size": file_stat.st_size,
|
||||
"last_modified": datetime.fromtimestamp(file_stat.st_mtime),
|
||||
"content_hash": self.calculate_file_hash(file_path)
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"扫描文档失败: {rel_path}, 错误: {e}")
|
||||
|
||||
return documents
|
||||
|
||||
def detect_changes(self) -> List[DocumentChange]:
|
||||
"""检测文档变更"""
|
||||
changes = []
|
||||
current_docs = self.scan_documents()
|
||||
stored_docs = self.db.get_all_document_hashes()
|
||||
|
||||
current_ids = set(current_docs.keys())
|
||||
stored_ids = set(stored_docs.keys())
|
||||
|
||||
# 新增的文档
|
||||
for doc_id in current_ids - stored_ids:
|
||||
doc = current_docs[doc_id]
|
||||
changes.append(DocumentChange(
|
||||
document_id=doc_id,
|
||||
document_name=doc["document_name"],
|
||||
change_type=ChangeType.ADDED,
|
||||
old_hash=None,
|
||||
new_hash=doc["content_hash"],
|
||||
change_time=datetime.now()
|
||||
))
|
||||
|
||||
# 删除的文档
|
||||
for doc_id in stored_ids - current_ids:
|
||||
doc = stored_docs[doc_id]
|
||||
changes.append(DocumentChange(
|
||||
document_id=doc_id,
|
||||
document_name=doc["document_name"],
|
||||
change_type=ChangeType.DELETED,
|
||||
old_hash=doc["content_hash"],
|
||||
new_hash=None,
|
||||
change_time=datetime.now()
|
||||
))
|
||||
|
||||
# 修改的文档
|
||||
for doc_id in current_ids & stored_ids:
|
||||
current_doc = current_docs[doc_id]
|
||||
stored_doc = stored_docs[doc_id]
|
||||
|
||||
if current_doc["content_hash"] != stored_doc["content_hash"]:
|
||||
changes.append(DocumentChange(
|
||||
document_id=doc_id,
|
||||
document_name=current_doc["document_name"],
|
||||
change_type=ChangeType.MODIFIED,
|
||||
old_hash=stored_doc["content_hash"],
|
||||
new_hash=current_doc["content_hash"],
|
||||
change_time=datetime.now()
|
||||
))
|
||||
|
||||
return changes
|
||||
|
||||
def process_change(self, change: DocumentChange) -> bool:
|
||||
"""处理单个变更"""
|
||||
try:
|
||||
file_path = os.path.join(self.documents_path, change.document_id)
|
||||
|
||||
# 从 document_id 中解析目标向量库
|
||||
# document_id 格式: "public/filename.pdf" 或 "finance/filename.pdf"
|
||||
kb_name = self._get_kb_name_from_path(change.document_id)
|
||||
|
||||
# 导入知识库管理器
|
||||
from knowledge.manager import get_kb_manager
|
||||
kb_manager = get_kb_manager()
|
||||
|
||||
if change.change_type == ChangeType.ADDED:
|
||||
# 新增文档 - 使用多向量库方法
|
||||
chunks_added = kb_manager.add_file_to_kb(
|
||||
kb_name=kb_name,
|
||||
filepath=file_path,
|
||||
extra_metadata={
|
||||
'status': 'active',
|
||||
'version': 'v1',
|
||||
'change_time': datetime.now().isoformat()
|
||||
}
|
||||
)
|
||||
# 更新哈希记录
|
||||
self.db.set_document_hash(
|
||||
change.document_id,
|
||||
change.document_name,
|
||||
change.new_hash,
|
||||
os.path.getsize(file_path) if os.path.exists(file_path) else 0,
|
||||
datetime.now()
|
||||
)
|
||||
|
||||
# 创建版本记录
|
||||
try:
|
||||
from knowledge.document_versions import get_version_query
|
||||
version_query = get_version_query()
|
||||
version_query.create_version_record(
|
||||
collection=kb_name,
|
||||
document_id=change.document_name,
|
||||
version="v1",
|
||||
status="active",
|
||||
change_summary="新增文档",
|
||||
created_by="sync_service",
|
||||
chunk_count=chunks_added
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"创建版本记录失败: {e}")
|
||||
|
||||
logger.info(f"已添加文档到 {kb_name}: {change.document_id}, 片段数: {chunks_added}")
|
||||
|
||||
elif change.change_type == ChangeType.MODIFIED:
|
||||
# 修改文档:使用版本管理策略
|
||||
# 1. 获取当前版本号
|
||||
old_version = self._get_current_version(kb_name, change.document_name)
|
||||
|
||||
# 2. 标记旧版本为 superseded(如果存在)
|
||||
if old_version:
|
||||
try:
|
||||
kb_manager.mark_document_as_superseded(
|
||||
kb_name,
|
||||
change.document_name,
|
||||
reason="文档更新"
|
||||
)
|
||||
logger.info(f"标记旧版本为 superseded: {change.document_name} {old_version}")
|
||||
except Exception as e:
|
||||
logger.warning(f"标记旧版本失败: {e}")
|
||||
|
||||
# 3. 生成新版本号
|
||||
new_version = self._generate_version_id(kb_name, change.document_name)
|
||||
|
||||
# 4. 添加新版本
|
||||
chunks_added = kb_manager.add_file_to_kb(
|
||||
kb_name=kb_name,
|
||||
filepath=file_path,
|
||||
extra_metadata={
|
||||
'status': 'active',
|
||||
'version': new_version,
|
||||
'previous_version': old_version or '',
|
||||
'change_time': datetime.now().isoformat()
|
||||
}
|
||||
)
|
||||
|
||||
# 5. 更新哈希记录
|
||||
self.db.set_document_hash(
|
||||
change.document_id,
|
||||
change.document_name,
|
||||
change.new_hash,
|
||||
os.path.getsize(file_path) if os.path.exists(file_path) else 0,
|
||||
datetime.now()
|
||||
)
|
||||
|
||||
# 6. 记录版本变更
|
||||
if old_version:
|
||||
self._record_version_change(
|
||||
kb_name,
|
||||
change.document_name,
|
||||
old_version,
|
||||
new_version,
|
||||
"文档更新"
|
||||
)
|
||||
|
||||
logger.info(f"已更新文档: {change.document_id}, 版本: {old_version} → {new_version}, 添加 {chunks_added} 片段")
|
||||
|
||||
elif change.change_type == ChangeType.DELETED:
|
||||
# 删除文档
|
||||
deleted = kb_manager.delete_document(kb_name, change.document_name)
|
||||
# 删除哈希记录
|
||||
self.db.delete_document_hash(change.document_id)
|
||||
logger.info(f"已删除文档: {change.document_id}, 删除 {deleted} 片段")
|
||||
|
||||
# ==================== 缓存失效 ====================
|
||||
# 文档变更后递增知识库版本号,使旧缓存自动失效
|
||||
if CACHE_AVAILABLE:
|
||||
try:
|
||||
cache = get_cache_manager()
|
||||
cache.increment_kb_version(kb_name)
|
||||
logger.debug(f"已递增知识库版本号: {kb_name}")
|
||||
except Exception as e:
|
||||
logger.warning(f"递增缓存版本号失败: {e}")
|
||||
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"处理变更失败: {change.document_id}, 错误: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
return False
|
||||
|
||||
def _get_kb_name_from_path(self, document_id: str) -> str:
|
||||
"""
|
||||
从文档ID中解析目标向量库名称
|
||||
|
||||
Args:
|
||||
document_id: 文档ID,格式如 "public_kb/filename.pdf" 或 "dept_hr/filename.pdf"
|
||||
|
||||
Returns:
|
||||
向量库名称(目录名 = 向量库名)
|
||||
"""
|
||||
# 统一路径分隔符(兼容 Windows 和 Linux)
|
||||
normalized = document_id.replace('\\', '/')
|
||||
# 获取第一级目录名(即向量库名)
|
||||
parts = normalized.split('/')
|
||||
if len(parts) > 1:
|
||||
return parts[0] # 目录名即向量库名
|
||||
else:
|
||||
return 'public_kb' # 默认公开库
|
||||
|
||||
def _get_current_version(self, kb_name: str, filename: str) -> str:
|
||||
"""
|
||||
获取文档当前版本号
|
||||
|
||||
Args:
|
||||
kb_name: 知识库名称
|
||||
filename: 文件名
|
||||
|
||||
Returns:
|
||||
当前版本号,如 "v1", "v2",不存在则返回 None
|
||||
"""
|
||||
try:
|
||||
from knowledge.document_versions import get_version_query
|
||||
version_query = get_version_query()
|
||||
active_version = version_query.get_active_version(kb_name, filename)
|
||||
return active_version.version if active_version else None
|
||||
except Exception as e:
|
||||
logger.warning(f"获取当前版本失败: {e}")
|
||||
return None
|
||||
|
||||
def _generate_version_id(self, kb_name: str, filename: str) -> str:
|
||||
"""
|
||||
生成新版本号
|
||||
|
||||
Args:
|
||||
kb_name: 知识库名称
|
||||
filename: 文件名
|
||||
|
||||
Returns:
|
||||
新版本号,如 "v1", "v2", "v3"
|
||||
"""
|
||||
current_version = self._get_current_version(kb_name, filename)
|
||||
if not current_version:
|
||||
return "v1"
|
||||
|
||||
# 从 "v1" 提取数字并递增
|
||||
try:
|
||||
version_num = int(current_version.replace('v', ''))
|
||||
return f"v{version_num + 1}"
|
||||
except (ValueError, AttributeError):
|
||||
return "v1"
|
||||
|
||||
def _record_version_change(
|
||||
self,
|
||||
kb_name: str,
|
||||
filename: str,
|
||||
old_version: str,
|
||||
new_version: str,
|
||||
reason: str
|
||||
):
|
||||
"""
|
||||
记录版本变更到数据库
|
||||
|
||||
Args:
|
||||
kb_name: 知识库名称
|
||||
filename: 文件名
|
||||
old_version: 旧版本号
|
||||
new_version: 新版本号
|
||||
reason: 变更原因
|
||||
"""
|
||||
try:
|
||||
from knowledge.document_versions import get_version_query
|
||||
version_query = get_version_query()
|
||||
version_query.log_version_change(
|
||||
collection=kb_name,
|
||||
document_id=filename,
|
||||
change_type="update",
|
||||
old_version=old_version,
|
||||
new_version=new_version,
|
||||
old_status="active",
|
||||
new_status="active",
|
||||
reason=reason,
|
||||
changed_by="sync_service"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"记录版本变更失败: {e}")
|
||||
|
||||
def sync_now(self) -> SyncResult:
|
||||
"""立即执行同步"""
|
||||
logger.info("开始同步...")
|
||||
|
||||
result = SyncResult(
|
||||
status=SyncStatus.RUNNING,
|
||||
start_time=datetime.now(),
|
||||
end_time=None,
|
||||
documents_processed=0,
|
||||
documents_added=0,
|
||||
documents_modified=0,
|
||||
documents_deleted=0,
|
||||
errors=[]
|
||||
)
|
||||
|
||||
try:
|
||||
# 检测变更
|
||||
changes = self.detect_changes()
|
||||
|
||||
# 处理变更
|
||||
for change in changes:
|
||||
success = self.process_change(change)
|
||||
result.documents_processed += 1
|
||||
|
||||
if success:
|
||||
if change.change_type == ChangeType.ADDED:
|
||||
result.documents_added += 1
|
||||
elif change.change_type == ChangeType.MODIFIED:
|
||||
result.documents_modified += 1
|
||||
elif change.change_type == ChangeType.DELETED:
|
||||
result.documents_deleted += 1
|
||||
else:
|
||||
result.errors.append(f"处理失败: {change.document_id}")
|
||||
|
||||
# 记录变更
|
||||
self.db.log_change(change)
|
||||
|
||||
result.status = SyncStatus.COMPLETED
|
||||
|
||||
except Exception as e:
|
||||
result.status = SyncStatus.FAILED
|
||||
result.errors.append(str(e))
|
||||
logger.error(f"同步失败: {e}")
|
||||
|
||||
result.end_time = datetime.now()
|
||||
|
||||
# 记录同步状态
|
||||
self.db.log_sync_status(result)
|
||||
|
||||
# 触发回调
|
||||
if self.on_sync_callback:
|
||||
self.on_sync_callback(result)
|
||||
|
||||
logger.info(f"同步完成: 处理 {result.documents_processed} 个文档, "
|
||||
f"新增 {result.documents_added}, "
|
||||
f"修改 {result.documents_modified}, "
|
||||
f"删除 {result.documents_deleted}")
|
||||
|
||||
return result
|
||||
|
||||
def start(self):
|
||||
"""启动文件监控"""
|
||||
if not HAS_WATCHDOG:
|
||||
logger.error("watchdog 未安装,无法启动文件监控")
|
||||
return False
|
||||
|
||||
if self._running:
|
||||
logger.warning("文件监控已在运行")
|
||||
return True
|
||||
|
||||
# 首次同步
|
||||
logger.info("执行首次同步...")
|
||||
self.sync_now()
|
||||
|
||||
# 启动监控
|
||||
event_handler = FileChangeHandler(self)
|
||||
self._observer = Observer()
|
||||
self._observer.schedule(event_handler, self.documents_path, recursive=True)
|
||||
self._observer.start()
|
||||
|
||||
self._running = True
|
||||
logger.info(f"文件监控已启动,监控目录: {self.documents_path}")
|
||||
return True
|
||||
|
||||
def stop(self):
|
||||
"""停止文件监控"""
|
||||
if self._observer:
|
||||
self._observer.stop()
|
||||
self._observer.join()
|
||||
self._observer = None
|
||||
|
||||
self._running = False
|
||||
logger.info("文件监控已停止")
|
||||
|
||||
def is_running(self) -> bool:
|
||||
"""检查监控是否在运行"""
|
||||
return self._running
|
||||
|
||||
|
||||
# 便捷函数
|
||||
def create_sync_service(documents_path: str = None) -> KnowledgeSyncService:
|
||||
"""创建同步服务实例"""
|
||||
return KnowledgeSyncService(documents_path)
|
||||
|
||||
|
||||
# 测试代码
|
||||
if __name__ == "__main__":
|
||||
print("=" * 60)
|
||||
print("知识库同步服务测试")
|
||||
print("=" * 60)
|
||||
|
||||
# 创建服务
|
||||
sync_service = KnowledgeSyncService()
|
||||
|
||||
# 测试扫描文档
|
||||
print("\n[1] 扫描文档...")
|
||||
docs = sync_service.scan_documents()
|
||||
print(f"找到 {len(docs)} 个文档")
|
||||
for doc_id, doc in list(docs.items())[:5]:
|
||||
print(f" - {doc['document_name']}: {doc['content_hash'][:8]}...")
|
||||
|
||||
# 测试变更检测
|
||||
print("\n[2] 检测变更...")
|
||||
changes = sync_service.detect_changes()
|
||||
print(f"检测到 {len(changes)} 个变更")
|
||||
for change in changes[:5]:
|
||||
print(f" - {change.document_name}: {change.change_type.value}")
|
||||
|
||||
# 测试同步
|
||||
print("\n[3] 执行同步...")
|
||||
result = sync_service.sync_now()
|
||||
print(f"同步状态: {result.status.value}")
|
||||
print(f"处理文档: {result.documents_processed}")
|
||||
print(f"新增: {result.documents_added}, 修改: {result.documents_modified}, 删除: {result.documents_deleted}")
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print("测试完成")
|
||||
Reference in New Issue
Block a user