diff --git a/src/core/kb_db_manager.py b/src/core/kb_db_manager.py new file mode 100644 index 00000000..691ca19f --- /dev/null +++ b/src/core/kb_db_manager.py @@ -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() \ No newline at end of file diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index dd2ac706..2c5709ee 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -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() \ No newline at end of file + return res \ No newline at end of file diff --git a/src/core/migrate_kb_to_sqlite.py b/src/core/migrate_kb_to_sqlite.py new file mode 100644 index 00000000..55933164 --- /dev/null +++ b/src/core/migrate_kb_to_sqlite.py @@ -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() \ No newline at end of file diff --git a/src/models/kb_models.py b/src/models/kb_models.py new file mode 100644 index 00000000..9770051e --- /dev/null +++ b/src/models/kb_models.py @@ -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 {} # 确保映射正确 + } \ No newline at end of file diff --git a/web/src/views/DataBaseInfoView.vue b/web/src/views/DataBaseInfoView.vue index 2a4bcfd6..119054dc 100644 --- a/web/src/views/DataBaseInfoView.vue +++ b/web/src/views/DataBaseInfoView.vue @@ -7,7 +7,7 @@
{{ database.embed_model }} {{ database.dimension }} - {{ database.metadata?.row_count }} 行 · {{ database.files ? Object.keys(database.files).length : 0 }} 文件 + {{ database.files ? Object.keys(database.files).length : 0 }} 文件 · {{ database.db_id }}