使用数据库替代现有的 json 形式
This commit is contained in:
parent
5fd20b9bbd
commit
642b8411a7
234
src/core/kb_db_manager.py
Normal file
234
src/core/kb_db_manager.py
Normal file
@ -0,0 +1,234 @@
|
||||
import os
|
||||
import pathlib
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker, joinedload
|
||||
from contextlib import contextmanager
|
||||
from sqlalchemy.orm.attributes import instance_state
|
||||
|
||||
from src import config
|
||||
from src.models.kb_models import Base, KnowledgeDatabase, KnowledgeFile, KnowledgeNode
|
||||
from src.utils import logger
|
||||
|
||||
class KBDBManager:
|
||||
"""知识库数据库管理器"""
|
||||
|
||||
def __init__(self):
|
||||
self.db_path = os.path.join(config.save_dir, "data", "knowledge.db")
|
||||
self.ensure_db_dir()
|
||||
|
||||
# 创建SQLAlchemy引擎
|
||||
self.engine = create_engine(f"sqlite:///{self.db_path}")
|
||||
|
||||
# 创建会话工厂
|
||||
self.Session = sessionmaker(bind=self.engine)
|
||||
|
||||
# 确保表存在
|
||||
self.create_tables()
|
||||
|
||||
def ensure_db_dir(self):
|
||||
"""确保数据库目录存在"""
|
||||
db_dir = os.path.dirname(self.db_path)
|
||||
pathlib.Path(db_dir).mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def create_tables(self):
|
||||
"""创建数据库表"""
|
||||
Base.metadata.create_all(self.engine)
|
||||
|
||||
@contextmanager
|
||||
def get_session(self):
|
||||
"""获取数据库会话的上下文管理器"""
|
||||
session = self.Session()
|
||||
try:
|
||||
yield session
|
||||
session.commit()
|
||||
except Exception as e:
|
||||
session.rollback()
|
||||
logger.error(f"数据库操作失败: {e}")
|
||||
raise
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
def _detach_safely(self, obj):
|
||||
"""安全地分离对象,确保属性已加载"""
|
||||
if obj is None:
|
||||
return None
|
||||
|
||||
# 确保主键已加载
|
||||
if hasattr(obj, 'id'):
|
||||
_ = obj.id
|
||||
if hasattr(obj, 'db_id'):
|
||||
_ = obj.db_id
|
||||
|
||||
# 根据需要添加其他必须预加载的属性
|
||||
|
||||
return obj
|
||||
|
||||
# 知识库操作方法
|
||||
def get_all_databases(self):
|
||||
"""获取所有知识库"""
|
||||
with self.get_session() as session:
|
||||
# 使用eager loading加载关联的files
|
||||
databases = session.query(KnowledgeDatabase).options(
|
||||
joinedload(KnowledgeDatabase.files)
|
||||
).all()
|
||||
|
||||
# 转换为字典并返回,避免后续延迟加载
|
||||
return [self._to_dict_safely(db) for db in databases]
|
||||
|
||||
def get_database_by_id(self, db_id):
|
||||
"""根据ID获取知识库"""
|
||||
with self.get_session() as session:
|
||||
# 使用eager loading加载关联的files
|
||||
db = session.query(KnowledgeDatabase).options(
|
||||
joinedload(KnowledgeDatabase.files).joinedload(KnowledgeFile.nodes)
|
||||
).filter_by(db_id=db_id).first()
|
||||
|
||||
# 转换为字典并返回,避免后续延迟加载
|
||||
return self._to_dict_safely(db) if db else None
|
||||
|
||||
def _to_dict_safely(self, obj):
|
||||
"""安全地将对象转换为字典,避免延迟加载问题"""
|
||||
if hasattr(obj, 'to_dict'):
|
||||
return obj.to_dict()
|
||||
return obj
|
||||
|
||||
def create_database(self, db_id, name, description, embed_model=None, dimension=None, metadata=None):
|
||||
"""创建知识库"""
|
||||
with self.get_session() as session:
|
||||
db = KnowledgeDatabase(
|
||||
db_id=db_id,
|
||||
name=name,
|
||||
description=description,
|
||||
embed_model=embed_model,
|
||||
dimension=dimension,
|
||||
meta_info=metadata or {} # 存储到meta_info字段
|
||||
)
|
||||
session.add(db)
|
||||
session.flush() # 立即写入数据库,获取ID
|
||||
|
||||
# 手动将必要的数据加载到内存中
|
||||
db_dict = {
|
||||
"db_id": db_id,
|
||||
"name": name,
|
||||
"description": description,
|
||||
"embed_model": embed_model,
|
||||
"dimension": dimension,
|
||||
"metadata": metadata or {}, # 返回时使用metadata键
|
||||
"files": {}
|
||||
}
|
||||
return db_dict
|
||||
|
||||
def delete_database(self, db_id):
|
||||
"""删除知识库"""
|
||||
with self.get_session() as session:
|
||||
db = session.query(KnowledgeDatabase).filter_by(db_id=db_id).first()
|
||||
if db:
|
||||
session.delete(db)
|
||||
return True
|
||||
return False
|
||||
|
||||
# 文件操作方法
|
||||
def add_file(self, db_id, file_id, filename, path, file_type, status="waiting"):
|
||||
"""添加文件"""
|
||||
with self.get_session() as session:
|
||||
file = KnowledgeFile(
|
||||
file_id=file_id,
|
||||
database_id=db_id,
|
||||
filename=filename,
|
||||
path=path,
|
||||
file_type=file_type,
|
||||
status=status
|
||||
)
|
||||
session.add(file)
|
||||
session.flush()
|
||||
|
||||
# 返回字典而非对象,避免会话关闭后的延迟加载问题
|
||||
return {
|
||||
"file_id": file_id,
|
||||
"filename": filename,
|
||||
"path": path,
|
||||
"type": file_type,
|
||||
"status": status,
|
||||
"created_at": file.created_at.timestamp() if file.created_at else None,
|
||||
"nodes": []
|
||||
}
|
||||
|
||||
def update_file_status(self, file_id, status):
|
||||
"""更新文件状态"""
|
||||
with self.get_session() as session:
|
||||
file = session.query(KnowledgeFile).filter_by(file_id=file_id).first()
|
||||
if file:
|
||||
file.status = status
|
||||
return True
|
||||
return False
|
||||
|
||||
def delete_file(self, file_id):
|
||||
"""删除文件"""
|
||||
with self.get_session() as session:
|
||||
file = session.query(KnowledgeFile).filter_by(file_id=file_id).first()
|
||||
if file:
|
||||
session.delete(file)
|
||||
return True
|
||||
return False
|
||||
|
||||
def get_files_by_database(self, db_id):
|
||||
"""获取知识库下的所有文件"""
|
||||
with self.get_session() as session:
|
||||
files = session.query(KnowledgeFile).options(
|
||||
joinedload(KnowledgeFile.nodes)
|
||||
).filter_by(database_id=db_id).all()
|
||||
return [self._to_dict_safely(file) for file in files]
|
||||
|
||||
def get_file_by_id(self, file_id):
|
||||
"""根据ID获取文件"""
|
||||
with self.get_session() as session:
|
||||
file = session.query(KnowledgeFile).options(
|
||||
joinedload(KnowledgeFile.nodes)
|
||||
).filter_by(file_id=file_id).first()
|
||||
return self._to_dict_safely(file) if file else None
|
||||
|
||||
# 知识块操作方法
|
||||
def add_node(self, file_id, text, hash_value=None, start_char_idx=None, end_char_idx=None, metadata=None):
|
||||
"""添加知识块"""
|
||||
with self.get_session() as session:
|
||||
node = KnowledgeNode(
|
||||
file_id=file_id,
|
||||
text=text,
|
||||
hash=hash_value,
|
||||
start_char_idx=start_char_idx,
|
||||
end_char_idx=end_char_idx,
|
||||
meta_info=metadata or {}
|
||||
)
|
||||
session.add(node)
|
||||
session.flush()
|
||||
|
||||
# 返回字典而非对象,避免会话关闭后的延迟加载问题
|
||||
return {
|
||||
"id": node.id,
|
||||
"file_id": file_id,
|
||||
"text": text,
|
||||
"hash": hash_value,
|
||||
"start_char_idx": start_char_idx,
|
||||
"end_char_idx": end_char_idx,
|
||||
"metadata": metadata or {}
|
||||
}
|
||||
|
||||
def get_nodes_by_file(self, file_id):
|
||||
"""获取文件下的所有知识块"""
|
||||
with self.get_session() as session:
|
||||
nodes = session.query(KnowledgeNode).filter_by(file_id=file_id).all()
|
||||
return [self._to_dict_safely(node) for node in nodes]
|
||||
|
||||
def get_nodes_by_filter(self, file_id=None, search_text=None, limit=100):
|
||||
"""根据条件筛选知识块"""
|
||||
with self.get_session() as session:
|
||||
query = session.query(KnowledgeNode)
|
||||
if file_id:
|
||||
query = query.filter_by(file_id=file_id)
|
||||
if search_text:
|
||||
query = query.filter(KnowledgeNode.text.like(f"%{search_text}%"))
|
||||
nodes = query.limit(limit).all()
|
||||
return [self._to_dict_safely(node) for node in nodes]
|
||||
|
||||
# 创建全局知识库数据库管理器实例
|
||||
kb_db_manager = KBDBManager()
|
||||
@ -9,24 +9,42 @@ from pymilvus import MilvusClient, MilvusException
|
||||
from src import config
|
||||
from src.utils import logger, hashstr
|
||||
from src.core.indexing import chunk, read_text
|
||||
|
||||
from src.core.kb_db_manager import kb_db_manager
|
||||
|
||||
|
||||
class KnowledgeBase:
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.data = []
|
||||
self.client = None
|
||||
self.work_dir = os.path.join(config.save_dir, "data")
|
||||
self.database_path = os.path.join(self.work_dir, "database.json")
|
||||
|
||||
# 数据库管理器
|
||||
self.db_manager = kb_db_manager
|
||||
|
||||
# Configuration
|
||||
self.default_distance_threshold = 0.5
|
||||
self.default_rerank_threshold = 0.1
|
||||
self.default_max_query_count = 20
|
||||
|
||||
# 检查是否需要从JSON文件迁移到SQLite
|
||||
self._check_migration()
|
||||
|
||||
self._load_models()
|
||||
self._load_databases()
|
||||
|
||||
def _check_migration(self):
|
||||
"""检查是否需要从JSON文件迁移到SQLite"""
|
||||
json_path = os.path.join(self.work_dir, "database.json")
|
||||
if os.path.exists(json_path):
|
||||
logger.info("检测到旧的JSON格式知识库数据,准备迁移到SQLite...")
|
||||
try:
|
||||
from src.core.migrate_kb_to_sqlite import migrate_json_to_sqlite
|
||||
result = migrate_json_to_sqlite()
|
||||
if result:
|
||||
logger.info("知识库数据已成功迁移到SQLite")
|
||||
else:
|
||||
logger.warning("知识库数据迁移失败或无需迁移")
|
||||
except Exception as e:
|
||||
logger.error(f"迁移过程中出错: {e}")
|
||||
|
||||
def _load_models(self):
|
||||
"""所有需要重启的模型"""
|
||||
@ -43,44 +61,27 @@ class KnowledgeBase:
|
||||
if not self.connect_to_milvus():
|
||||
raise ConnectionError("Failed to connect to Milvus")
|
||||
|
||||
def _load_databases(self):
|
||||
"""将数据库的信息保存到本地的文件里面"""
|
||||
if not os.path.exists(self.database_path):
|
||||
return
|
||||
|
||||
with open(self.database_path, "r") as f:
|
||||
data = json.load(f)
|
||||
self.data = [DataBaseLite(**db) for db in data["databases"]]
|
||||
|
||||
self._update_database()
|
||||
|
||||
def _save_databases(self):
|
||||
"""将数据库的信息保存到本地的文件里面"""
|
||||
self._update_database()
|
||||
os.makedirs(os.path.dirname(self.database_path), exist_ok=True)
|
||||
with open(self.database_path, "w") as f:
|
||||
json.dump({
|
||||
"databases": [db.to_dict() for db in self.data],
|
||||
}, f, ensure_ascii=False, indent=4)
|
||||
|
||||
def _update_database(self):
|
||||
self.id2db = {db.db_id: db for db in self.data}
|
||||
self.name2db = {db.name: db for db in self.data}
|
||||
|
||||
def create_database(self, database_name, description, dimension=None):
|
||||
"""创建一个数据库"""
|
||||
dimension = dimension or self.embed_model.get_dimension()
|
||||
db = DataBaseLite(database_name,
|
||||
description,
|
||||
embed_model=self.embed_model.embed_model_fullname,
|
||||
dimension=dimension)
|
||||
db_id = f"kb_{hashstr(database_name, with_salt=True)}"
|
||||
|
||||
# 创建数据库记录
|
||||
db_dict = self.db_manager.create_database(
|
||||
db_id=db_id,
|
||||
name=database_name,
|
||||
description=description,
|
||||
embed_model=self.embed_model.embed_model_fullname,
|
||||
dimension=dimension
|
||||
)
|
||||
|
||||
# 创建数据库对应的文件夹
|
||||
self._ensure_db_folders(db.db_id)
|
||||
self._ensure_db_folders(db_id)
|
||||
|
||||
self.add_collection(db.db_id, dimension)
|
||||
self.data.append(db)
|
||||
self._save_databases()
|
||||
# 在Milvus中创建集合
|
||||
self.add_collection(db_id, dimension)
|
||||
|
||||
return db_dict
|
||||
|
||||
def _ensure_db_folders(self, db_id):
|
||||
"""确保数据库文件夹存在"""
|
||||
@ -98,33 +99,67 @@ class KnowledgeBase:
|
||||
def get_databases(self):
|
||||
assert config.enable_knowledge_base, "知识库未启用"
|
||||
|
||||
for db in self.data:
|
||||
db.update(self.get_collection_info(db.db_id))
|
||||
processing_files = [f for fid, f in db.files.items() if f["status"] in ["processing", "waiting"]]
|
||||
if processing_files:
|
||||
logger.info(f"数据库 {db.name} 有 {len(processing_files)} 个文件正在处理中")
|
||||
# 从数据库获取所有知识库
|
||||
databases = self.db_manager.get_all_databases()
|
||||
|
||||
self._save_databases()
|
||||
return {"databases": [db.to_dict() for db in self.data]}
|
||||
# 检查和更新Milvus信息
|
||||
databases_with_milvus = []
|
||||
for db in databases:
|
||||
db_copy = db.copy() # 创建字典的副本以避免修改原始数据
|
||||
# 更新Milvus集合信息
|
||||
try:
|
||||
milvus_info = self.get_collection_info(db["db_id"])
|
||||
db_copy["metadata"] = milvus_info
|
||||
logger.debug(f"获取知识库 {db['name']} (ID: {db['db_id']}) 的Milvus信息成功: {milvus_info}")
|
||||
except Exception as e:
|
||||
logger.warning(f"获取知识库 {db['name']} (ID: {db['db_id']}) 的Milvus信息失败: {e}")
|
||||
# 添加一个默认的Milvus状态
|
||||
db_copy.update({
|
||||
"row_count": 0,
|
||||
"status": "未连接",
|
||||
"error": str(e)
|
||||
})
|
||||
|
||||
# 检查处理中的文件
|
||||
processing_files = [f for f_id, f in db_copy.get("files", {}).items()
|
||||
if f["status"] in ["processing", "waiting"]]
|
||||
if processing_files:
|
||||
logger.info(f"数据库 {db['name']} 有 {len(processing_files)} 个文件正在处理中")
|
||||
|
||||
databases_with_milvus.append(db_copy)
|
||||
|
||||
return {"databases": databases_with_milvus}
|
||||
|
||||
def get_database_info(self, db_id):
|
||||
db = self.get_kb_by_id(db_id)
|
||||
if db is None:
|
||||
db_dict = self.db_manager.get_database_by_id(db_id)
|
||||
if db_dict is None:
|
||||
return None
|
||||
else:
|
||||
db.update(self.get_collection_info(db.db_id))
|
||||
return db.to_dict()
|
||||
db_copy = db_dict.copy()
|
||||
try:
|
||||
milvus_info = self.get_collection_info(db_id)
|
||||
db_copy.update(milvus_info)
|
||||
except Exception as e:
|
||||
logger.warning(f"获取知识库 ID: {db_id} 的Milvus信息失败: {e}")
|
||||
# 添加一个默认的Milvus状态
|
||||
db_copy.update({
|
||||
"row_count": 0,
|
||||
"status": "未连接",
|
||||
"error": str(e)
|
||||
})
|
||||
return db_copy
|
||||
|
||||
def get_database_id(self):
|
||||
return [db.db_id for db in self.data]
|
||||
databases = self.db_manager.get_all_databases()
|
||||
return [db["db_id"] for db in databases]
|
||||
|
||||
def get_file_info(self, db_id, file_id):
|
||||
db = self.get_kb_by_id(db_id)
|
||||
db = self.db_manager.get_database_by_id(db_id)
|
||||
if db is None:
|
||||
raise Exception(f"database not found, {db_id}")
|
||||
|
||||
lines = self.client.query(
|
||||
collection_name=db.db_id,
|
||||
collection_name=db_id,
|
||||
filter=f"file_id == '{file_id}'",
|
||||
output_fields=None
|
||||
)
|
||||
@ -133,14 +168,13 @@ class KnowledgeBase:
|
||||
line.pop("vector")
|
||||
|
||||
lines.sort(key=lambda x: x.get("start_char_idx") or 0)
|
||||
# logger.debug(f"lines[0]: {lines[0]}")
|
||||
return {"lines": lines}
|
||||
|
||||
def get_kb_by_id(self, db_id):
|
||||
if not config.enable_knowledge_base:
|
||||
return None
|
||||
|
||||
return next((db for db in self.data if db.db_id == db_id), None)
|
||||
return self.db_manager.get_database_by_id(db_id)
|
||||
|
||||
def file_to_chunk(self, files, params=None):
|
||||
"""将文件转换为分块
|
||||
@ -183,93 +217,95 @@ class KnowledgeBase:
|
||||
"""添加分块"""
|
||||
db = self.get_kb_by_id(db_id)
|
||||
|
||||
if db.embed_model != config.embed_model:
|
||||
logger.error(f"Embed model not match, {db.embed_model} != {config.embed_model}")
|
||||
return {"message": f"Embed model not match, cur: {config.embed_model}, req: {db.embed_model}", "status": "failed"}
|
||||
if db["embed_model"] != self.embed_model.embed_model_fullname:
|
||||
logger.error(f"Embed model not match, {db['embed_model']} != {self.embed_model.embed_model_fullname}")
|
||||
return {"message": f"Embed model not match, cur: {self.embed_model.embed_model_fullname}, req: {db['embed_model']}", "status": "failed"}
|
||||
|
||||
db.files.update(file_chunks)
|
||||
self._save_databases()
|
||||
|
||||
for file_id, chunk in file_chunks.items():
|
||||
db.files[file_id]["status"] = "processing"
|
||||
self._save_databases()
|
||||
for file_id, chunk_info in file_chunks.items():
|
||||
# 在数据库中创建文件记录
|
||||
self.db_manager.add_file(
|
||||
db_id=db_id,
|
||||
file_id=file_id,
|
||||
filename=chunk_info["filename"],
|
||||
path=chunk_info["path"],
|
||||
file_type=chunk_info["type"],
|
||||
status="processing"
|
||||
)
|
||||
|
||||
try:
|
||||
self.add_documents(
|
||||
file_id=file_id,
|
||||
collection_name=db.db_id,
|
||||
docs=[node["text"] for node in chunk["nodes"]],
|
||||
chunk_infos=chunk["nodes"])
|
||||
collection_name=db_id,
|
||||
docs=[node["text"] for node in chunk_info["nodes"]],
|
||||
chunk_infos=chunk_info["nodes"])
|
||||
|
||||
db.files[file_id]["status"] = "done"
|
||||
# 更新文件状态为完成
|
||||
self.db_manager.update_file_status(file_id, "done")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to add documents to collection {db.db_id}, {e}, {traceback.format_exc()}")
|
||||
db.files[file_id]["status"] = "failed"
|
||||
|
||||
self._save_databases()
|
||||
logger.error(f"Failed to add documents to collection {db_id}, {e}, {traceback.format_exc()}")
|
||||
# 更新文件状态为失败
|
||||
self.db_manager.update_file_status(file_id, "failed")
|
||||
|
||||
def add_files(self, db_id, files, params=None):
|
||||
db = self.get_kb_by_id(db_id)
|
||||
|
||||
if db.embed_model != config.embed_model:
|
||||
logger.error(f"Embed model not match, {db.embed_model} != {config.embed_model}")
|
||||
return {"message": f"Embed model not match, cur: {config.embed_model}, req: {db.embed_model}", "status": "failed"}
|
||||
if db["embed_model"] != self.embed_model.embed_model_fullname:
|
||||
logger.error(f"Embed model not match, {db['embed_model']} != {self.embed_model.embed_model_fullname}")
|
||||
return {"message": f"Embed model not match, cur: {self.embed_model.embed_model_fullname}, req: {db['embed_model']}", "status": "failed"}
|
||||
|
||||
# Preprocessing the files to the queue
|
||||
new_files = self.file_to_chunk(files, params=params)
|
||||
db.files.update(new_files) # 更新数据库状态
|
||||
|
||||
# 先保存一次数据库状态,确保waiting状态被记录
|
||||
self._save_databases()
|
||||
|
||||
for file_id, new_file in new_files.items():
|
||||
db.files[file_id]["status"] = "processing"
|
||||
# 更新处理状态
|
||||
self._save_databases()
|
||||
# 在数据库中创建文件记录
|
||||
self.db_manager.add_file(
|
||||
db_id=db_id,
|
||||
file_id=file_id,
|
||||
filename=new_file["filename"],
|
||||
path=new_file["path"],
|
||||
file_type=new_file["type"],
|
||||
status="processing"
|
||||
)
|
||||
|
||||
try:
|
||||
self.add_documents(
|
||||
file_id=file_id,
|
||||
collection_name=db.db_id,
|
||||
collection_name=db_id,
|
||||
docs=[node["text"] for node in new_file["nodes"]],
|
||||
chunk_infos=new_file["nodes"])
|
||||
|
||||
db.files[file_id]["status"] = "done"
|
||||
# 更新文件状态为完成
|
||||
self.db_manager.update_file_status(file_id, "done")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to add documents to collection {db.db_id}, {e}, {traceback.format_exc()}")
|
||||
db.files[file_id]["status"] = "failed"
|
||||
|
||||
# 每个文件处理完成后立即保存数据库状态
|
||||
self._save_databases()
|
||||
logger.error(f"Failed to add documents to collection {db_id}, {e}, {traceback.format_exc()}")
|
||||
# 更新文件状态为失败
|
||||
self.db_manager.update_file_status(file_id, "failed")
|
||||
|
||||
def delete_file(self, db_id, file_id):
|
||||
db = self.get_kb_by_id(db_id)
|
||||
if db is None:
|
||||
raise Exception(f"database not found, {db_id}")
|
||||
# 从Milvus中删除文件的向量
|
||||
self.client.delete(collection_name=db_id, filter=f"file_id == '{file_id}'")
|
||||
|
||||
self.client.delete(collection_name=db.db_id, filter=f"file_id == '{file_id}'")
|
||||
del db.files[file_id]
|
||||
self._save_databases()
|
||||
# 从SQLite中删除文件记录
|
||||
self.db_manager.delete_file(file_id)
|
||||
|
||||
def delete_database(self, db_id):
|
||||
db = self.get_kb_by_id(db_id)
|
||||
if db is None:
|
||||
raise Exception(f"database not found, {db_id}")
|
||||
# 从Milvus中删除集合
|
||||
self.client.drop_collection(collection_name=db_id)
|
||||
|
||||
# 从SQLite中删除数据库记录
|
||||
self.db_manager.delete_database(db_id)
|
||||
|
||||
self.client.drop_collection(collection_name=db.db_id)
|
||||
self.data.remove(db)
|
||||
# 删除数据库对应的文件夹
|
||||
db_folder = os.path.join(self.work_dir, db.db_id)
|
||||
db_folder = os.path.join(self.work_dir, db_id)
|
||||
if os.path.exists(db_folder):
|
||||
shutil.rmtree(db_folder)
|
||||
self._save_databases()
|
||||
|
||||
return {"message": "删除成功"}
|
||||
|
||||
def restart(self):
|
||||
self._load_models()
|
||||
self._load_databases()
|
||||
|
||||
###################################
|
||||
#* Below is the code for retriever #
|
||||
@ -283,8 +319,12 @@ class KnowledgeBase:
|
||||
max_query_count = kwargs.get("max_query_count", self.default_max_query_count)
|
||||
|
||||
all_db_result = self.search(query, db_id, limit=max_query_count)
|
||||
|
||||
# 获取文件信息并添加到结果中
|
||||
for res in all_db_result:
|
||||
res["file"] = db.files[res["entity"]["file_id"]]
|
||||
file = self.db_manager.get_file_by_id(res["entity"]["file_id"])
|
||||
if file:
|
||||
res["file"] = file
|
||||
|
||||
db_result = [r for r in all_db_result if r["distance"] > distance_threshold]
|
||||
|
||||
@ -318,7 +358,6 @@ class KnowledgeBase:
|
||||
|
||||
return retriever
|
||||
|
||||
|
||||
################################
|
||||
#* Below is the code for milvus #
|
||||
################################
|
||||
@ -351,10 +390,20 @@ class KnowledgeBase:
|
||||
return collections
|
||||
|
||||
def get_collection_info(self, collection_name):
|
||||
collection = self.client.describe_collection(collection_name)
|
||||
collection.update(self.client.get_collection_stats(collection_name))
|
||||
# collection["id"] = hashstr(collection_name)
|
||||
return collection
|
||||
"""获取Milvus集合信息,处理可能的错误"""
|
||||
try:
|
||||
collection = self.client.describe_collection(collection_name)
|
||||
collection.update(self.client.get_collection_stats(collection_name))
|
||||
return collection
|
||||
except MilvusException as e:
|
||||
logger.warning(f"获取集合 {collection_name} 信息失败: {e}")
|
||||
# 返回一个带有错误信息的基本结构
|
||||
return {
|
||||
"name": collection_name,
|
||||
"row_count": 0,
|
||||
"status": "错误",
|
||||
"error_message": str(e)
|
||||
}
|
||||
|
||||
def add_collection(self, collection_name, dimension=None):
|
||||
if self.client.has_collection(collection_name=collection_name):
|
||||
@ -415,47 +464,4 @@ class KnowledgeBase:
|
||||
|
||||
def search_by_id(self, collection_name, id, output_fields=["id", "text"]):
|
||||
res = self.client.get(collection_name, id, output_fields=output_fields)
|
||||
return res
|
||||
|
||||
|
||||
class DataBaseLite:
|
||||
def __init__(self, name, description, dimension=None, **kwargs) -> None:
|
||||
self.name = name
|
||||
self.description = description
|
||||
self.dimension = dimension
|
||||
self.metadata = kwargs.get("metadata", {})
|
||||
# logger.debug(f"DataBaseLite init: {self.metadata}")
|
||||
self.db_id = self.metadata.get("collection_name", kwargs.get("db_id")) # metaname 的历史遗留问题
|
||||
self.db_id = self.db_id or f"kb_{hashstr(name, with_salt=True)}"
|
||||
self.files = kwargs.get("files", [])
|
||||
|
||||
if isinstance(self.files, list):
|
||||
self.files = {f["file_id"]: f for f in self.files}
|
||||
|
||||
self.embed_model = kwargs.get("embed_model", None)
|
||||
|
||||
def id2file(self, file_id):
|
||||
for f in self.files:
|
||||
if f["file_id"] == file_id:
|
||||
return f
|
||||
return None
|
||||
|
||||
def update(self, metadata):
|
||||
self.metadata = metadata
|
||||
|
||||
def to_dict(self):
|
||||
return {
|
||||
"name": self.name,
|
||||
"description": self.description,
|
||||
"db_id": self.db_id,
|
||||
"embed_model": self.embed_model,
|
||||
"metadata": self.metadata,
|
||||
"files": self.files,
|
||||
"dimension": self.dimension
|
||||
}
|
||||
|
||||
def to_json(self):
|
||||
return json.dumps(self.to_dict(), ensure_ascii=False)
|
||||
|
||||
def __str__(self):
|
||||
return self.to_json()
|
||||
return res
|
||||
107
src/core/migrate_kb_to_sqlite.py
Normal file
107
src/core/migrate_kb_to_sqlite.py
Normal file
@ -0,0 +1,107 @@
|
||||
import os
|
||||
import json
|
||||
import time
|
||||
from pathlib import Path
|
||||
import traceback
|
||||
|
||||
from src import config
|
||||
from src.utils import logger
|
||||
from src.core.kb_db_manager import kb_db_manager
|
||||
|
||||
def migrate_json_to_sqlite():
|
||||
"""将JSON文件数据迁移到SQLite数据库"""
|
||||
# 原始JSON文件路径
|
||||
json_path = os.path.join(config.save_dir, "data", "database.json")
|
||||
|
||||
if not os.path.exists(json_path):
|
||||
logger.info(f"未找到原始JSON文件: {json_path},无需迁移")
|
||||
return False
|
||||
|
||||
try:
|
||||
# 读取JSON文件
|
||||
with open(json_path, "r", encoding='utf-8') as f:
|
||||
data = json.load(f)
|
||||
|
||||
if not data or "databases" not in data or not data["databases"]:
|
||||
logger.info("JSON文件中没有数据库信息,无需迁移")
|
||||
return False
|
||||
|
||||
# 开始迁移
|
||||
logger.info(f"开始迁移知识库数据,共 {len(data['databases'])} 个数据库")
|
||||
|
||||
# 遍历所有数据库
|
||||
for db_info in data["databases"]:
|
||||
db_id = db_info["db_id"]
|
||||
name = db_info["name"]
|
||||
description = db_info["description"]
|
||||
embed_model = db_info.get("embed_model")
|
||||
dimension = db_info.get("dimension")
|
||||
metadata = db_info.get("metadata", {})
|
||||
|
||||
logger.info(f"处理数据库: {name} (ID: {db_id}), metadata类型: {type(metadata)}")
|
||||
|
||||
# 检查数据库是否已存在
|
||||
existing_db = kb_db_manager.get_database_by_id(db_id)
|
||||
if existing_db:
|
||||
logger.info(f"数据库 {name} (ID: {db_id}) 已存在,跳过创建")
|
||||
continue
|
||||
|
||||
# 创建数据库
|
||||
db = kb_db_manager.create_database(
|
||||
db_id=db_id,
|
||||
name=name,
|
||||
description=description,
|
||||
embed_model=embed_model,
|
||||
dimension=dimension,
|
||||
metadata=metadata # 这里传入metadata,在kb_db_manager中会被正确存储为meta_info
|
||||
)
|
||||
|
||||
# 处理文件
|
||||
files = db_info.get("files", {})
|
||||
if isinstance(files, list):
|
||||
files = {f["file_id"]: f for f in files}
|
||||
|
||||
for file_id, file_info in files.items():
|
||||
# 添加文件
|
||||
kb_db_manager.add_file(
|
||||
db_id=db_id,
|
||||
file_id=file_id,
|
||||
filename=file_info["filename"],
|
||||
path=file_info["path"],
|
||||
file_type=file_info["type"],
|
||||
status=file_info["status"]
|
||||
)
|
||||
|
||||
# 处理节点
|
||||
nodes = file_info.get("nodes", [])
|
||||
for node in nodes:
|
||||
node_metadata = node.get("metadata", {})
|
||||
if node_metadata is None:
|
||||
node_metadata = {}
|
||||
logger.debug(f"节点metadata类型: {type(node_metadata)}")
|
||||
|
||||
kb_db_manager.add_node(
|
||||
file_id=file_id,
|
||||
text=node["text"],
|
||||
hash_value=node.get("hash"),
|
||||
start_char_idx=node.get("start_char_idx"),
|
||||
end_char_idx=node.get("end_char_idx"),
|
||||
metadata=node_metadata # 在kb_db_manager中会被正确存储为meta_info
|
||||
)
|
||||
|
||||
logger.info(f"数据库 {name} (ID: {db_id}) 迁移完成,共 {len(files)} 个文件")
|
||||
|
||||
# 备份原始JSON文件
|
||||
backup_path = json_path + f".bak.{int(time.time())}"
|
||||
os.rename(json_path, backup_path)
|
||||
logger.info(f"迁移完成,原始JSON文件已备份为: {backup_path}")
|
||||
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"迁移过程中出错: {e}")
|
||||
logger.error(traceback.format_exc())
|
||||
return False
|
||||
|
||||
if __name__ == "__main__":
|
||||
migrate_json_to_sqlite()
|
||||
107
src/models/kb_models.py
Normal file
107
src/models/kb_models.py
Normal file
@ -0,0 +1,107 @@
|
||||
from sqlalchemy import Column, Integer, String, DateTime, JSON, Float, ForeignKey, Text
|
||||
from sqlalchemy.orm import relationship
|
||||
from sqlalchemy.ext.declarative import declarative_base
|
||||
from sqlalchemy.sql import func
|
||||
import time
|
||||
|
||||
Base = declarative_base()
|
||||
|
||||
class KnowledgeDatabase(Base):
|
||||
"""知识库模型"""
|
||||
__tablename__ = 'knowledge_databases'
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
db_id = Column(String, nullable=False, unique=True, index=True) # 数据库ID
|
||||
name = Column(String, nullable=False) # 数据库名称
|
||||
description = Column(Text, nullable=True) # 描述
|
||||
embed_model = Column(String, nullable=True) # 嵌入模型名称
|
||||
dimension = Column(Integer, nullable=True) # 向量维度
|
||||
meta_info = Column(JSON, nullable=True) # 元数据
|
||||
created_at = Column(DateTime, default=func.now()) # 创建时间
|
||||
|
||||
# 关系
|
||||
files = relationship("KnowledgeFile", back_populates="database", cascade="all, delete-orphan")
|
||||
|
||||
def to_dict(self):
|
||||
"""转换为字典格式,确保meta_info映射为metadata"""
|
||||
result = {
|
||||
"id": self.id,
|
||||
"db_id": self.db_id,
|
||||
"name": self.name,
|
||||
"description": self.description,
|
||||
"embed_model": self.embed_model,
|
||||
"dimension": self.dimension,
|
||||
"metadata": self.meta_info or {}, # 确保映射正确
|
||||
"created_at": self.created_at.isoformat() if self.created_at else None
|
||||
}
|
||||
|
||||
# 添加文件信息
|
||||
if self.files:
|
||||
result["files"] = {file.file_id: file.to_dict() for file in self.files}
|
||||
else:
|
||||
result["files"] = {}
|
||||
|
||||
return result
|
||||
|
||||
class KnowledgeFile(Base):
|
||||
"""知识库文件模型"""
|
||||
__tablename__ = 'knowledge_files'
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
file_id = Column(String, nullable=False, index=True) # 文件ID
|
||||
database_id = Column(String, ForeignKey('knowledge_databases.db_id'), nullable=False) # 所属数据库ID
|
||||
filename = Column(String, nullable=False) # 文件名
|
||||
path = Column(String, nullable=False) # 文件路径
|
||||
file_type = Column(String, nullable=False) # 文件类型
|
||||
status = Column(String, nullable=False) # 处理状态
|
||||
created_at = Column(DateTime, default=func.now()) # 创建时间
|
||||
|
||||
# 关系
|
||||
database = relationship("KnowledgeDatabase", back_populates="files")
|
||||
nodes = relationship("KnowledgeNode", back_populates="file", cascade="all, delete-orphan")
|
||||
|
||||
def to_dict(self):
|
||||
"""转换为字典格式"""
|
||||
result = {
|
||||
"file_id": self.file_id,
|
||||
"filename": self.filename,
|
||||
"path": self.path,
|
||||
"type": self.file_type,
|
||||
"status": self.status,
|
||||
"created_at": self.created_at.timestamp() if self.created_at else time.time()
|
||||
}
|
||||
|
||||
# 添加节点信息
|
||||
if self.nodes:
|
||||
result["nodes"] = [node.to_dict() for node in self.nodes]
|
||||
else:
|
||||
result["nodes"] = []
|
||||
|
||||
return result
|
||||
|
||||
class KnowledgeNode(Base):
|
||||
"""知识块模型"""
|
||||
__tablename__ = 'knowledge_nodes'
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
file_id = Column(String, ForeignKey('knowledge_files.file_id'), nullable=False) # 所属文件ID
|
||||
text = Column(Text, nullable=False) # 文本内容
|
||||
hash = Column(String, nullable=True) # 文本哈希值
|
||||
start_char_idx = Column(Integer, nullable=True) # 开始字符索引
|
||||
end_char_idx = Column(Integer, nullable=True) # 结束字符索引
|
||||
meta_info = Column(JSON, nullable=True) # 元数据
|
||||
|
||||
# 关系
|
||||
file = relationship("KnowledgeFile", back_populates="nodes")
|
||||
|
||||
def to_dict(self):
|
||||
"""转换为字典格式,确保meta_info映射为metadata"""
|
||||
return {
|
||||
"id": self.id,
|
||||
"file_id": self.file_id,
|
||||
"text": self.text,
|
||||
"hash": self.hash,
|
||||
"start_char_idx": self.start_char_idx,
|
||||
"end_char_idx": self.end_char_idx,
|
||||
"metadata": self.meta_info or {} # 确保映射正确
|
||||
}
|
||||
@ -7,7 +7,7 @@
|
||||
<div class="database-info">
|
||||
<a-tag color="blue" v-if="database.embed_model">{{ database.embed_model }}</a-tag>
|
||||
<a-tag color="green" v-if="database.dimension">{{ database.dimension }}</a-tag>
|
||||
<span class="row-count">{{ database.metadata?.row_count }} 行 · {{ database.files ? Object.keys(database.files).length : 0 }} 文件</span>
|
||||
<span class="row-count">{{ database.files ? Object.keys(database.files).length : 0 }} 文件 · {{ database.db_id }}</span>
|
||||
</div>
|
||||
</template>
|
||||
<template #actions>
|
||||
@ -564,7 +564,7 @@ const deleteFile = (fileId) => {
|
||||
content: '确定要删除该文件吗?',
|
||||
okText: '确认',
|
||||
cancelText: '取消',
|
||||
onOk: () => {
|
||||
onOk: () => {
|
||||
state.lock = true
|
||||
fetch('/api/data/document', {
|
||||
method: "DELETE",
|
||||
|
||||
@ -45,7 +45,7 @@
|
||||
<div class="icon"><ReadFilled /></div>
|
||||
<div class="info">
|
||||
<h3>{{ database.name }}</h3>
|
||||
<p><span>{{ database.metadata.row_count }} 行</span> · <span>{{ database.files ? Object.keys(database.files).length : 0 }} 文件</span></p>
|
||||
<p><span>{{ database.files ? Object.keys(database.files).length : 0 }} 文件</span></p>
|
||||
</div>
|
||||
</div>
|
||||
<p class="description">{{ database.description || '暂无描述' }}</p>
|
||||
|
||||
Loading…
Reference in New Issue
Block a user