init: RAG 知识库服务初始提交

- 后端 API(Flask + Gunicorn)
- RAG 引擎(混合检索 + 云端 Reranker + 引用溯源)
- 文档解析(MinerU + 多格式支持)
- Docker 生产部署配置
- 排除前端项目、敏感配置、模型文件
This commit is contained in:
lacerate551
2026-06-04 17:35:27 +08:00
commit 100d1a06eb
158 changed files with 64534 additions and 0 deletions

40
knowledge/__init__.py Normal file
View 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
View 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
View 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: 文档 IDsource 字段),用于过滤特定文档的切片
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
View 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
View 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
View 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

View 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
View 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
View 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
View 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
View 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
View 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.12020-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
View 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
View 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.5BM25 权重 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
View 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("测试完成")