2024-07-16 18:14:27 +08:00
|
|
|
import os
|
|
|
|
|
import json
|
|
|
|
|
import time
|
2024-07-28 16:16:52 +08:00
|
|
|
from src.utils import hashstr, setup_logger, is_text_pdf
|
|
|
|
|
from src.models.embedding import get_embedding_model
|
2024-07-14 23:59:52 +08:00
|
|
|
|
2024-07-16 18:14:27 +08:00
|
|
|
logger = setup_logger("DataBaseManager")
|
|
|
|
|
|
|
|
|
|
|
2024-07-14 23:59:52 +08:00
|
|
|
class DataBaseManager:
|
|
|
|
|
|
|
|
|
|
def __init__(self, config=None) -> None:
|
|
|
|
|
self.config = config
|
2024-07-28 16:16:52 +08:00
|
|
|
self.database_path = os.path.join(config.save_dir, "data", "database.json")
|
2024-07-22 00:00:54 +08:00
|
|
|
self.embed_model = get_embedding_model(config)
|
2024-07-16 18:14:27 +08:00
|
|
|
|
2024-07-31 20:22:05 +08:00
|
|
|
if self.config.enable_knowledge_base:
|
2024-08-07 17:24:46 +08:00
|
|
|
from src.core.knowledgebase import KnowledgeBase
|
2024-07-31 20:22:05 +08:00
|
|
|
self.knowledge_base = KnowledgeBase(config, self.embed_model)
|
2024-09-11 01:08:13 +08:00
|
|
|
if self.config.enable_knowledge_graph:
|
|
|
|
|
from src.core.graphbase import GraphDatabase
|
|
|
|
|
self.graph_base = GraphDatabase(self.config, self.embed_model)
|
|
|
|
|
self.graph_base.start()
|
|
|
|
|
else:
|
|
|
|
|
self.graph_base = None
|
2024-07-16 18:14:27 +08:00
|
|
|
|
2024-07-31 20:22:05 +08:00
|
|
|
self.data = {"databases": [], "graph": {}}
|
|
|
|
|
|
2024-07-16 18:14:27 +08:00
|
|
|
self._load_databases()
|
2024-07-28 16:16:52 +08:00
|
|
|
self._update_database()
|
2024-07-16 18:14:27 +08:00
|
|
|
|
|
|
|
|
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 = {
|
|
|
|
|
"databases": [DataBaseLite(**db) for db in data["databases"]],
|
|
|
|
|
"graph": data["graph"]
|
|
|
|
|
}
|
|
|
|
|
|
2024-08-25 12:34:35 +08:00
|
|
|
# 检查所有文件,如果出现状态是 processing 的,那么设置为 failed
|
|
|
|
|
for db in self.data["databases"]:
|
|
|
|
|
for file in db.files:
|
|
|
|
|
if file["status"] == "processing" or file["status"] == "waiting":
|
|
|
|
|
file["status"] = "failed"
|
|
|
|
|
|
2024-07-16 18:14:27 +08:00
|
|
|
def _save_databases(self):
|
|
|
|
|
"""将数据库的信息保存到本地的文件里面"""
|
2024-07-28 16:16:52 +08:00
|
|
|
self._update_database()
|
2024-07-16 18:14:27 +08:00
|
|
|
with open(self.database_path, "w+") as f:
|
|
|
|
|
json.dump({
|
|
|
|
|
"databases": [db.to_dict() for db in self.data["databases"]],
|
|
|
|
|
"graph": self.data["graph"]
|
|
|
|
|
}, f, ensure_ascii=False, indent=4)
|
2024-07-14 23:59:52 +08:00
|
|
|
|
2024-07-28 16:16:52 +08:00
|
|
|
def _update_database(self):
|
|
|
|
|
self.id2db = {db.db_id: db for db in self.data["databases"]}
|
|
|
|
|
self.name2db = {db.name: db for db in self.data["databases"]}
|
|
|
|
|
self.metaname2db = {db.metaname: db for db in self.data["databases"]}
|
2024-07-14 23:59:52 +08:00
|
|
|
|
2024-07-28 16:16:52 +08:00
|
|
|
def get_databases(self):
|
|
|
|
|
self._update_database()
|
2024-07-16 18:14:27 +08:00
|
|
|
knowledge_base_collections = self.knowledge_base.get_collection_names()
|
|
|
|
|
if len(self.data["databases"]) != len(knowledge_base_collections):
|
|
|
|
|
logger.warning(f"Database number not match, {knowledge_base_collections}")
|
2024-07-14 23:59:52 +08:00
|
|
|
|
2024-07-16 18:14:27 +08:00
|
|
|
for db in self.data["databases"]:
|
|
|
|
|
db.update(self.knowledge_base.get_collection_info(db.metaname))
|
|
|
|
|
|
|
|
|
|
return {"databases": [db.to_dict() for db in self.data["databases"]]}
|
2024-07-14 23:59:52 +08:00
|
|
|
|
2024-07-16 18:14:27 +08:00
|
|
|
def get_graph(self):
|
2024-09-06 12:54:17 +08:00
|
|
|
if self.config.enable_knowledge_graph:
|
2024-07-16 18:14:27 +08:00
|
|
|
self.data["graph"].update(self.graph_base.get_database_info("neo4j"))
|
|
|
|
|
return {"graph": self.data["graph"]}
|
|
|
|
|
else:
|
2024-07-31 20:22:05 +08:00
|
|
|
return {"message": "Graph base not enabled", "graph": {}}
|
2024-07-16 18:14:27 +08:00
|
|
|
|
2024-08-25 20:29:24 +08:00
|
|
|
def create_database(self, database_name, description, db_type, dimension):
|
2024-09-06 12:54:17 +08:00
|
|
|
from src.config import EMBED_MODEL_INFO
|
|
|
|
|
dimension = dimension or EMBED_MODEL_INFO[self.config.embed_model]["dimension"]
|
|
|
|
|
|
2024-08-25 20:29:24 +08:00
|
|
|
new_database = DataBaseLite(database_name,
|
|
|
|
|
description,
|
|
|
|
|
db_type,
|
|
|
|
|
embed_model=self.config.embed_model,
|
|
|
|
|
dimension=dimension)
|
2024-07-16 18:14:27 +08:00
|
|
|
|
2024-08-25 20:29:24 +08:00
|
|
|
self.knowledge_base.add_collection(new_database.metaname, dimension)
|
2024-07-16 18:14:27 +08:00
|
|
|
self.data["databases"].append(new_database)
|
|
|
|
|
self._save_databases()
|
2024-07-14 23:59:52 +08:00
|
|
|
return self.get_databases()
|
|
|
|
|
|
2024-09-25 13:46:23 +08:00
|
|
|
def add_files(self, db_id, files, params=None):
|
2024-07-16 18:14:27 +08:00
|
|
|
db = self.get_kb_by_id(db_id)
|
2024-07-22 00:00:54 +08:00
|
|
|
|
|
|
|
|
if db.embed_model != self.config.embed_model:
|
|
|
|
|
logger.error(f"Embed model not match, {db.embed_model} != {self.config.embed_model}")
|
2024-09-06 12:54:17 +08:00
|
|
|
return {"message": f"Embed model not match, cur: {self.config.embed_model}", "status": "failed"}
|
2024-07-22 00:00:54 +08:00
|
|
|
|
2024-09-25 13:46:23 +08:00
|
|
|
# Preprocessing the files to the queue
|
2024-07-16 18:14:27 +08:00
|
|
|
new_files = []
|
|
|
|
|
for file in files:
|
2024-08-25 12:34:35 +08:00
|
|
|
new_file = {
|
2024-07-16 18:14:27 +08:00
|
|
|
"file_id": "file_" + hashstr(file + str(time.time())),
|
|
|
|
|
"filename": os.path.basename(file),
|
|
|
|
|
"path": file,
|
|
|
|
|
"type": file.split(".")[-1],
|
|
|
|
|
"status": "waiting",
|
|
|
|
|
"created_at": time.time()
|
2024-08-25 12:34:35 +08:00
|
|
|
}
|
|
|
|
|
db.files.append(new_file)
|
|
|
|
|
new_files.append(new_file)
|
2024-07-16 18:14:27 +08:00
|
|
|
|
2024-09-26 22:45:02 +08:00
|
|
|
from src.core.indexing import chunk
|
2024-08-25 12:34:35 +08:00
|
|
|
for new_file in new_files:
|
|
|
|
|
file_id = new_file["file_id"]
|
2024-09-25 13:46:23 +08:00
|
|
|
idx = self.get_idx_by_fileid(db, file_id)
|
2024-07-16 18:14:27 +08:00
|
|
|
db.files[idx]["status"] = "processing"
|
|
|
|
|
|
|
|
|
|
try:
|
2024-09-26 22:45:02 +08:00
|
|
|
nodes = chunk(new_file["path"], params=params)
|
2024-07-16 18:14:27 +08:00
|
|
|
self.knowledge_base.add_documents(
|
2024-09-26 22:45:02 +08:00
|
|
|
docs=[node.text for node in nodes],
|
2024-07-16 18:14:27 +08:00
|
|
|
collection_name=db.metaname,
|
2024-08-25 12:34:35 +08:00
|
|
|
file_id=file_id)
|
|
|
|
|
|
2024-09-25 13:46:23 +08:00
|
|
|
idx = self.get_idx_by_fileid(db, file_id)
|
2024-07-16 18:14:27 +08:00
|
|
|
db.files[idx]["status"] = "done"
|
2024-08-25 12:34:35 +08:00
|
|
|
|
2024-07-16 18:14:27 +08:00
|
|
|
except Exception as e:
|
|
|
|
|
logger.error(f"Failed to add documents to collection {db.metaname}, {e}")
|
2024-09-25 13:46:23 +08:00
|
|
|
idx = self.get_idx_by_fileid(db, file_id)
|
2024-07-16 18:14:27 +08:00
|
|
|
db.files[idx]["status"] = "failed"
|
|
|
|
|
|
|
|
|
|
self._save_databases()
|
|
|
|
|
|
2024-07-22 00:00:54 +08:00
|
|
|
return {"message": "全部解析完成", "status": "success"}
|
|
|
|
|
|
2024-07-16 18:14:27 +08:00
|
|
|
def get_database_info(self, db_id):
|
|
|
|
|
db = self.get_kb_by_id(db_id)
|
|
|
|
|
if db is None:
|
|
|
|
|
return None
|
|
|
|
|
else:
|
|
|
|
|
db.update(self.knowledge_base.get_collection_info(db.metaname))
|
|
|
|
|
return db.to_dict()
|
|
|
|
|
|
2024-09-25 13:46:23 +08:00
|
|
|
def read_text(self, file, params=None):
|
2024-07-24 18:32:14 +08:00
|
|
|
support_format = [".pdf", ".txt", ".md"]
|
2024-07-16 18:14:27 +08:00
|
|
|
assert os.path.exists(file), "File not found"
|
|
|
|
|
logger.info(f"Try to read file {file}")
|
2024-07-17 18:52:20 +08:00
|
|
|
|
|
|
|
|
if not os.path.isfile(file):
|
|
|
|
|
logger.error(f"Directory not supported now!")
|
|
|
|
|
raise NotImplementedError("Directory not supported now!")
|
|
|
|
|
|
|
|
|
|
if file.endswith(".pdf"):
|
|
|
|
|
if is_text_pdf(file):
|
2024-09-25 13:46:23 +08:00
|
|
|
from src.core.filereader import pdfreader
|
2024-07-16 18:14:27 +08:00
|
|
|
return pdfreader(file)
|
|
|
|
|
else:
|
2024-08-07 17:24:46 +08:00
|
|
|
from src.plugins import pdf2txt
|
2024-07-17 18:52:20 +08:00
|
|
|
return pdf2txt(file, return_text=True)
|
|
|
|
|
|
|
|
|
|
elif file.endswith(".txt") or file.endswith(".md"):
|
2024-09-25 13:46:23 +08:00
|
|
|
from src.core.filereader import plainreader
|
2024-07-17 18:52:20 +08:00
|
|
|
return plainreader(file)
|
|
|
|
|
|
2024-07-16 18:14:27 +08:00
|
|
|
else:
|
2024-07-17 18:52:20 +08:00
|
|
|
logger.error(f"File format not supported, only support {support_format}")
|
|
|
|
|
raise Exception(f"File format not supported, only support {support_format}")
|
|
|
|
|
|
2024-07-16 18:14:27 +08:00
|
|
|
def delete_file(self, db_id, file_id):
|
|
|
|
|
db = self.get_kb_by_id(db_id)
|
2024-09-25 13:46:23 +08:00
|
|
|
file_idx_to_delete = self.get_idx_by_fileid(db, file_id)
|
2024-07-16 18:14:27 +08:00
|
|
|
|
|
|
|
|
self.knowledge_base.client.delete(
|
|
|
|
|
collection_name=db.metaname,
|
|
|
|
|
filter=f"file_id == '{file_id}'"),
|
|
|
|
|
|
|
|
|
|
del db.files[file_idx_to_delete]
|
|
|
|
|
self._save_databases()
|
|
|
|
|
|
|
|
|
|
def get_file_info(self, db_id, file_id):
|
|
|
|
|
db = self.get_kb_by_id(db_id)
|
|
|
|
|
if db is None:
|
|
|
|
|
return {"message": "database not found"}, 404
|
|
|
|
|
lines = self.knowledge_base.client.query(
|
|
|
|
|
collection_name=db.metaname,
|
|
|
|
|
filter=f"file_id == '{file_id}'",
|
|
|
|
|
output_fields=["id", "text", "file_id", "hash"]
|
|
|
|
|
)
|
|
|
|
|
return {"lines": lines}
|
|
|
|
|
|
2024-09-25 13:46:23 +08:00
|
|
|
def chunking(self, text, params=None):
|
|
|
|
|
chunk_method = params.get("chunk_method", "fixed")
|
|
|
|
|
chunk_size = params.get("chunk_size", 500)
|
|
|
|
|
|
2024-07-16 18:14:27 +08:00
|
|
|
"""将文本切分成固定大小的块"""
|
|
|
|
|
chunks = []
|
|
|
|
|
for i in range(0, len(text), chunk_size):
|
|
|
|
|
chunks.append(text[i:i + chunk_size])
|
|
|
|
|
return chunks
|
2024-07-14 23:59:52 +08:00
|
|
|
|
2024-07-17 18:52:20 +08:00
|
|
|
def delete_database(self, db_id):
|
|
|
|
|
db = self.get_kb_by_id(db_id)
|
|
|
|
|
if db is None:
|
|
|
|
|
return {"message": "database not found"}, 404
|
|
|
|
|
|
|
|
|
|
self.knowledge_base.client.drop_collection(db.metaname)
|
|
|
|
|
self.data["databases"] = [d for d in self.data["databases"] if d.db_id != db_id]
|
|
|
|
|
self._save_databases()
|
|
|
|
|
return {"message": "删除成功"}
|
|
|
|
|
|
2024-07-16 18:14:27 +08:00
|
|
|
def get_kb_by_id(self, db_id):
|
|
|
|
|
for db in self.data["databases"]:
|
|
|
|
|
if db.db_id == db_id:
|
2024-07-28 16:16:52 +08:00
|
|
|
return db
|
2024-08-30 11:54:21 +08:00
|
|
|
return None
|
|
|
|
|
|
2024-09-25 13:46:23 +08:00
|
|
|
def get_idx_by_fileid(self, db, file_id):
|
|
|
|
|
for idx, f in enumerate(db.files):
|
|
|
|
|
if f["file_id"] == file_id:
|
|
|
|
|
return idx
|
|
|
|
|
|
2024-08-30 11:54:21 +08:00
|
|
|
|
|
|
|
|
class DataBaseLite:
|
|
|
|
|
def __init__(self, name, description, db_type, dimension=None, **kwargs) -> None:
|
|
|
|
|
self.name = name
|
|
|
|
|
self.description = description
|
|
|
|
|
self.db_type = db_type
|
|
|
|
|
self.dimension = dimension
|
|
|
|
|
self.db_id = kwargs.get("db_id", hashstr(name))
|
|
|
|
|
self.metaname = kwargs.get("metaname", f"{db_type[:1]}{hashstr(name)}")
|
|
|
|
|
self.metadata = kwargs.get("metaname", {})
|
|
|
|
|
self.files = kwargs.get("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_type": self.db_type,
|
|
|
|
|
"db_id": self.db_id,
|
|
|
|
|
"embed_model": self.embed_model,
|
|
|
|
|
"metaname": self.metaname,
|
|
|
|
|
"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()
|