fix: 处理 embed_info 可能为字典或 EmbedModelInfo 对象的情况,确保获取模型名称和其他属性的兼容性

This commit is contained in:
Wenjie Zhang 2025-10-23 22:57:01 +08:00
parent 59727d6d7e
commit f6547fef13
7 changed files with 404 additions and 19 deletions

View File

@ -7,7 +7,7 @@
## Bugs ## Bugs
- - [x] 修复本地知识库的 metadata 和 向量数据库中不一致的情况。
## Next ## Next

View File

@ -1,5 +1,7 @@
import json import json
import os import os
import tempfile
import shutil
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Any from typing import Any
@ -524,6 +526,7 @@ class KnowledgeBase(ABC):
def _load_metadata(self): def _load_metadata(self):
"""加载元数据""" """加载元数据"""
meta_file = os.path.join(self.work_dir, f"metadata_{self.kb_type}.json") meta_file = os.path.join(self.work_dir, f"metadata_{self.kb_type}.json")
if os.path.exists(meta_file): if os.path.exists(meta_file):
try: try:
with open(meta_file, encoding="utf-8") as f: with open(meta_file, encoding="utf-8") as f:
@ -533,19 +536,74 @@ class KnowledgeBase(ABC):
logger.info(f"Loaded {self.kb_type} metadata for {len(self.databases_meta)} databases") logger.info(f"Loaded {self.kb_type} metadata for {len(self.databases_meta)} databases")
except Exception as e: except Exception as e:
logger.error(f"Failed to load {self.kb_type} metadata: {e}") logger.error(f"Failed to load {self.kb_type} metadata: {e}")
# 尝试从备份恢复
backup_file = f"{meta_file}.backup"
if os.path.exists(backup_file):
try:
with open(backup_file, encoding="utf-8") as f:
data = json.load(f)
self.databases_meta = data.get("databases", {})
self.files_meta = data.get("files", {})
logger.info(f"Loaded {self.kb_type} metadata from backup")
# 恢复备份文件
shutil.copy2(backup_file, meta_file)
return
except Exception as backup_e:
logger.error(f"Failed to load backup: {backup_e}")
# 如果加载失败,初始化为空状态
logger.warning(f"Initializing empty {self.kb_type} metadata")
self.databases_meta = {}
self.files_meta = {}
def _serialize_metadata(self, obj):
"""递归序列化元数据中的 Pydantic 模型"""
if hasattr(obj, 'dict'):
return obj.dict()
elif isinstance(obj, dict):
return {k: self._serialize_metadata(v) for k, v in obj.items()}
elif isinstance(obj, list):
return [self._serialize_metadata(item) for item in obj]
else:
return obj
def _save_metadata(self): def _save_metadata(self):
"""保存元数据""" """保存元数据"""
self._normalize_metadata_state() self._normalize_metadata_state()
meta_file = os.path.join(self.work_dir, f"metadata_{self.kb_type}.json") meta_file = os.path.join(self.work_dir, f"metadata_{self.kb_type}.json")
backup_file = f"{meta_file}.backup"
try: try:
# 创建简单备份
if os.path.exists(meta_file):
shutil.copy2(meta_file, backup_file)
# 准备数据并序列化 Pydantic 模型
data = { data = {
"databases": self.databases_meta, "databases": self._serialize_metadata(self.databases_meta),
"files": self.files_meta, "files": self._serialize_metadata(self.files_meta),
"kb_type": self.kb_type, "kb_type": self.kb_type,
"updated_at": utc_isoformat(), "updated_at": utc_isoformat(),
} }
with open(meta_file, "w", encoding="utf-8") as f:
json.dump(data, f, ensure_ascii=False, indent=2) # 原子性写入(使用临时文件)
with tempfile.NamedTemporaryFile(
mode='w', dir=os.path.dirname(meta_file),
prefix='.tmp_', suffix='.json', delete=False
) as tmp_file:
json.dump(data, tmp_file, ensure_ascii=False, indent=2)
temp_path = tmp_file.name
os.replace(temp_path, meta_file)
logger.debug(f"Saved {self.kb_type} metadata")
except Exception as e: except Exception as e:
logger.error(f"Failed to save {self.kb_type} metadata: {e}") logger.error(f"Failed to save {self.kb_type} metadata: {e}")
# 尝试恢复备份
if os.path.exists(backup_file):
try:
shutil.copy2(backup_file, meta_file)
logger.info("Restored metadata from backup")
except Exception as restore_e:
logger.error(f"Failed to restore backup: {restore_e}")
raise e

View File

@ -299,7 +299,7 @@ class GraphDatabase:
logger.info(f"Adding entity to {kgdb_name}") logger.info(f"Adding entity to {kgdb_name}")
session.execute_write(_create_graph, triples) session.execute_write(_create_graph, triples)
logger.info(f"Creating vector index for {kgdb_name} with {config.embed_model}") logger.info(f"Creating vector index for {kgdb_name} with {config.embed_model}")
session.execute_write(_create_vector_index, cur_embed_info["dimension"]) session.execute_write(_create_vector_index, getattr(cur_embed_info, 'dimension', 1024))
# 收集所有需要处理的实体名称,去重 # 收集所有需要处理的实体名称,去重
all_entities = [] all_entities = []

View File

@ -72,7 +72,12 @@ class ChromaKB(KnowledgeBase):
logger.info(f"Retrieved existing collection: {collection_name}") logger.info(f"Retrieved existing collection: {collection_name}")
# 检查现有集合的配置是否匹配当前的 embed_info # 检查现有集合的配置是否匹配当前的 embed_info
expected_model = embed_info.get("name") if embed_info else "default" expected_model = getattr(embed_info, 'name', None) if embed_info else None
if expected_model is None and hasattr(embed_info, 'get'):
expected_model = embed_info.get('name')
elif embed_info and isinstance(embed_info, dict):
expected_model = embed_info.get('name')
expected_model = expected_model or "default"
collection_metadata = collection.metadata or {} collection_metadata = collection.metadata or {}
current_model = collection_metadata.get("embedding_model", "unknown") current_model = collection_metadata.get("embedding_model", "unknown")
@ -88,11 +93,18 @@ class ChromaKB(KnowledgeBase):
except Exception: except Exception:
# 创建新集合 # 创建新集合
logger.info(f"Creating new collection with embedding model: {embed_info.get('name', 'default')}") model_name = getattr(embed_info, 'name', None) if embed_info else None
if model_name is None and hasattr(embed_info, 'get'):
model_name = embed_info.get('name')
elif embed_info and isinstance(embed_info, dict):
model_name = embed_info.get('name')
model_name = model_name or 'default'
logger.info(f"Creating new collection with embedding model: {model_name}")
collection_metadata = { collection_metadata = {
"db_id": db_id, "db_id": db_id,
"created_at": utc_isoformat(), "created_at": utc_isoformat(),
"embedding_model": embed_info.get("name") if embed_info else "default", "embedding_model": model_name,
} }
collection = self.chroma_client.create_collection( collection = self.chroma_client.create_collection(
name=collection_name, embedding_function=embedding_function, metadata=collection_metadata name=collection_name, embedding_function=embedding_function, metadata=collection_metadata

View File

@ -103,7 +103,7 @@ class MilvusKB(KnowledgeBase):
# 检查嵌入模型是否匹配 # 检查嵌入模型是否匹配
description = collection.description description = collection.description
expected_model = embed_info.get("name") if embed_info else "default" expected_model = getattr(embed_info, 'name', 'default') if embed_info else "default"
if expected_model not in description: if expected_model not in description:
logger.warning(f"Collection {collection_name} model mismatch, recreating...") logger.warning(f"Collection {collection_name} model mismatch, recreating...")
@ -116,8 +116,8 @@ class MilvusKB(KnowledgeBase):
except Exception: except Exception:
# 创建新集合 # 创建新集合
embedding_dim = embed_info.get("dimension", 1024) if embed_info else 1024 embedding_dim = getattr(embed_info, 'dimension', 1024) if embed_info else 1024
model_name = embed_info.get("name", "default") if embed_info else "default" model_name = getattr(embed_info, 'name', 'default') if embed_info else "default"
# 定义集合Schema # 定义集合Schema
fields = [ fields = [

View File

@ -1,6 +1,8 @@
import asyncio import asyncio
import json import json
import os import os
import shutil
import tempfile
from src.knowledge.base import KBNotFoundError, KnowledgeBase from src.knowledge.base import KBNotFoundError, KnowledgeBase
from src.knowledge.factory import KnowledgeBaseFactory from src.knowledge.factory import KnowledgeBaseFactory
@ -43,9 +45,27 @@ class KnowledgeBaseManager:
logger.info("KnowledgeBaseManager initialized") logger.info("KnowledgeBaseManager initialized")
# 在后台运行数据一致性检测(不阻塞初始化)
try:
# 尝试获取当前事件循环,如果没有则创建新的
try:
loop = asyncio.get_event_loop()
if loop.is_running():
# 如果已经在事件循环中,创建任务
asyncio.create_task(self.detect_data_inconsistencies())
else:
# 如果事件循环未运行,直接运行
loop.run_until_complete(self.detect_data_inconsistencies())
except RuntimeError:
# 没有事件循环,创建一个来运行检测
asyncio.run(self.detect_data_inconsistencies())
except Exception as e:
logger.warning(f"初始化时运行数据一致性检测失败: {e}")
def _load_global_metadata(self): def _load_global_metadata(self):
"""加载全局元数据""" """加载全局元数据"""
meta_file = os.path.join(self.work_dir, "global_metadata.json") meta_file = os.path.join(self.work_dir, "global_metadata.json")
if os.path.exists(meta_file): if os.path.exists(meta_file):
try: try:
with open(meta_file, encoding="utf-8") as f: with open(meta_file, encoding="utf-8") as f:
@ -54,13 +74,63 @@ class KnowledgeBaseManager:
logger.info(f"Loaded global metadata for {len(self.global_databases_meta)} databases") logger.info(f"Loaded global metadata for {len(self.global_databases_meta)} databases")
except Exception as e: except Exception as e:
logger.error(f"Failed to load global metadata: {e}") logger.error(f"Failed to load global metadata: {e}")
# 尝试从备份恢复
backup_file = f"{meta_file}.backup"
if os.path.exists(backup_file):
try:
with open(backup_file, encoding="utf-8") as f:
data = json.load(f)
self.global_databases_meta = data.get("databases", {})
logger.info("Loaded global metadata from backup")
# 恢复备份文件
shutil.copy2(backup_file, meta_file)
return
except Exception as backup_e:
logger.error(f"Failed to load backup: {backup_e}")
# 如果加载失败,初始化为空状态
logger.warning("Initializing empty global metadata")
self.global_databases_meta = {}
def _save_global_metadata(self): def _save_global_metadata(self):
"""保存全局元数据""" """保存全局元数据"""
self._normalize_global_metadata()
meta_file = os.path.join(self.work_dir, "global_metadata.json") meta_file = os.path.join(self.work_dir, "global_metadata.json")
data = {"databases": self.global_databases_meta, "updated_at": utc_isoformat(), "version": "2.0"} backup_file = f"{meta_file}.backup"
with open(meta_file, "w", encoding="utf-8") as f:
json.dump(data, f, ensure_ascii=False, indent=2) try:
# 创建简单备份
if os.path.exists(meta_file):
shutil.copy2(meta_file, backup_file)
# 准备数据
data = {
"databases": self.global_databases_meta,
"updated_at": utc_isoformat(),
"version": "2.0"
}
# 原子性写入(使用临时文件)
with tempfile.NamedTemporaryFile(
mode='w', dir=os.path.dirname(meta_file),
prefix='.tmp_', suffix='.json', delete=False
) as tmp_file:
json.dump(data, tmp_file, ensure_ascii=False, indent=2)
temp_path = tmp_file.name
os.replace(temp_path, meta_file)
logger.debug("Saved global metadata")
except Exception as e:
logger.error(f"Failed to save global metadata: {e}")
# 尝试恢复备份
if os.path.exists(backup_file):
try:
shutil.copy2(backup_file, meta_file)
logger.info("Restored global metadata from backup")
except Exception as restore_e:
logger.error(f"Failed to restore backup: {restore_e}")
raise e
def _normalize_global_metadata(self) -> None: def _normalize_global_metadata(self) -> None:
"""Normalize stored timestamps within the global metadata cache.""" """Normalize stored timestamps within the global metadata cache."""
@ -434,3 +504,239 @@ class KnowledgeBaseManager:
lightrag_databases.append(db) lightrag_databases.append(db)
return lightrag_databases return lightrag_databases
# =============================================================================
# 数据一致性检测方法
# =============================================================================
async def detect_data_inconsistencies(self) -> dict:
"""
检测向量数据库中存在但在 metadata 中缺失的数据
Returns:
包含不一致信息的字典按知识库类型分组
"""
inconsistencies = {
"chroma": {"missing_collections": [], "missing_files": []},
"milvus": {"missing_collections": [], "missing_files": []},
"total_missing_collections": 0,
"total_missing_files": 0
}
logger.info("开始检测向量数据库与元数据的一致性...")
# 检测 ChromaDB 数据不一致
if "chroma" in self.kb_instances:
try:
chroma_inconsistencies = await self._detect_chroma_inconsistencies()
inconsistencies["chroma"] = chroma_inconsistencies
inconsistencies["total_missing_collections"] += len(chroma_inconsistencies["missing_collections"])
inconsistencies["total_missing_files"] += len(chroma_inconsistencies["missing_files"])
except Exception as e:
logger.error(f"检测 ChromaDB 数据不一致时出错: {e}")
# 检测 Milvus 数据不一致
if "milvus" in self.kb_instances:
try:
milvus_inconsistencies = await self._detect_milvus_inconsistencies()
inconsistencies["milvus"] = milvus_inconsistencies
inconsistencies["total_missing_collections"] += len(milvus_inconsistencies["missing_collections"])
inconsistencies["total_missing_files"] += len(milvus_inconsistencies["missing_files"])
except Exception as e:
logger.error(f"检测 Milvus 数据不一致时出错: {e}")
# 输出检测结果到日志
self._log_inconsistencies(inconsistencies)
return inconsistencies
async def _detect_chroma_inconsistencies(self) -> dict:
"""检测 ChromaDB 中的数据不一致"""
inconsistencies = {"missing_collections": [], "missing_files": []}
chroma_kb = self.kb_instances["chroma"]
# 获取 ChromaDB 中所有实际的集合
try:
actual_collections = chroma_kb.chroma_client.list_collections()
actual_collection_names = {col.name for col in actual_collections}
# 获取 metadata 中记录的数据库ID
metadata_collection_names = set()
for db_id, db_meta in chroma_kb.databases_meta.items():
metadata_collection_names.add(db_id)
# 找出存在于 ChromaDB 但不在 metadata 中的集合
missing_collections = actual_collection_names - metadata_collection_names
for collection_name in missing_collections:
# 跳过一些系统集合
if not collection_name.startswith("kb_"):
continue
collection_info = {
"collection_name": collection_name,
"detected_at": utc_isoformat()
}
# 尝试获取集合的基本信息
try:
collection = chroma_kb.chroma_client.get_collection(name=collection_name)
collection_info["count"] = collection.count()
collection_info["metadata"] = collection.metadata
except Exception as e:
logger.warning(f"无法获取集合 {collection_name} 的详细信息: {e}")
collection_info["count"] = "unknown"
inconsistencies["missing_collections"].append(collection_info)
logger.warning(f"发现 ChromaDB 中存在但 metadata 中缺失的集合: {collection_name} (文档数: {collection_info['count']})")
# 检查文件级别的不一致(针对已知的数据库)
for db_id in metadata_collection_names:
try:
collection = chroma_kb.chroma_client.get_collection(name=db_id)
actual_count = collection.count()
# 获取 metadata 中记录的文件数量
metadata_files_count = sum(1 for file_info in chroma_kb.files_meta.values()
if file_info.get("database_id") == db_id)
# 如果向量数据库中有数据但 metadata 中没有文件记录,可能存在文件缺失
if actual_count > 0 and metadata_files_count == 0:
inconsistencies["missing_files"].append({
"database_id": db_id,
"vector_count": actual_count,
"metadata_files_count": metadata_files_count,
"detected_at": utc_isoformat()
})
logger.warning(f"发现数据库 {db_id} 在 ChromaDB 中有 {actual_count} 条向量数据,但 metadata 中没有文件记录")
except Exception as e:
logger.debug(f"检查数据库 {db_id} 的文件一致性时出错: {e}")
except Exception as e:
logger.error(f"检测 ChromaDB 数据不一致时出错: {e}")
return inconsistencies
async def _detect_milvus_inconsistencies(self) -> dict:
"""检测 Milvus 中的数据不一致"""
inconsistencies = {"missing_collections": [], "missing_files": []}
milvus_kb = self.kb_instances["milvus"]
try:
from pymilvus import utility
# 获取 Milvus 中所有实际的集合
actual_collection_names = set(utility.list_collections(using=milvus_kb.connection_alias))
# 获取 metadata 中记录的数据库ID
metadata_collection_names = set(milvus_kb.databases_meta.keys())
# 找出存在于 Milvus 但不在 metadata 中的集合
missing_collections = actual_collection_names - metadata_collection_names
for collection_name in missing_collections:
# 跳过一些系统集合
if not collection_name.startswith("kb_"):
continue
collection_info = {
"collection_name": collection_name,
"detected_at": utc_isoformat()
}
# 尝试获取集合的基本信息
try:
from pymilvus import Collection
collection = Collection(name=collection_name, using=milvus_kb.connection_alias)
collection_info["count"] = collection.num_entities
collection_info["description"] = collection.description
except Exception as e:
logger.warning(f"无法获取集合 {collection_name} 的详细信息: {e}")
collection_info["count"] = "unknown"
inconsistencies["missing_collections"].append(collection_info)
logger.warning(f"发现 Milvus 中存在但 metadata 中缺失的集合: {collection_name} (实体数: {collection_info['count']})")
# 检查文件级别的不一致(针对已知的数据库)
for db_id in metadata_collection_names:
try:
if utility.has_collection(db_id, using=milvus_kb.connection_alias):
from pymilvus import Collection
collection = Collection(name=db_id, using=milvus_kb.connection_alias)
actual_count = collection.num_entities
# 获取 metadata 中记录的文件数量
metadata_files_count = sum(1 for file_info in milvus_kb.files_meta.values()
if file_info.get("database_id") == db_id)
# 如果向量数据库中有数据但 metadata 中没有文件记录,可能存在文件缺失
if actual_count > 0 and metadata_files_count == 0:
inconsistencies["missing_files"].append({
"database_id": db_id,
"vector_count": actual_count,
"metadata_files_count": metadata_files_count,
"detected_at": utc_isoformat()
})
logger.warning(f"发现数据库 {db_id} 在 Milvus 中有 {actual_count} 条向量数据,但 metadata 中没有文件记录")
except Exception as e:
logger.debug(f"检查数据库 {db_id} 的文件一致性时出错: {e}")
except Exception as e:
logger.error(f"检测 Milvus 数据不一致时出错: {e}")
return inconsistencies
def _log_inconsistencies(self, inconsistencies: dict) -> None:
"""将不一致检测结果输出到日志"""
total_missing_collections = inconsistencies["total_missing_collections"]
total_missing_files = inconsistencies["total_missing_files"]
if total_missing_collections == 0 and total_missing_files == 0:
logger.info("数据一致性检测完成,未发现不一致情况")
return
logger.warning("=" * 80)
logger.warning("数据一致性检测完成,发现以下不一致情况:")
logger.warning("=" * 80)
# ChromaDB 不一致情况
chroma_missing = inconsistencies["chroma"]["missing_collections"]
chroma_files_missing = inconsistencies["chroma"]["missing_files"]
if chroma_missing or chroma_files_missing:
logger.warning(f"ChromaDB 不一致情况:")
logger.warning(f" 缺失集合数量: {len(chroma_missing)}")
for collection_info in chroma_missing:
logger.warning(f" - 集合: {collection_info['collection_name']}, 向量数: {collection_info['count']}")
logger.warning(f" 缺失文件记录数量: {len(chroma_files_missing)}")
for file_info in chroma_files_missing:
logger.warning(f" - 数据库: {file_info['database_id']}, 向量数: {file_info['vector_count']}, 元数据文件数: {file_info['metadata_files_count']}")
# Milvus 不一致情况
milvus_missing = inconsistencies["milvus"]["missing_collections"]
milvus_files_missing = inconsistencies["milvus"]["missing_files"]
if milvus_missing or milvus_files_missing:
logger.warning(f"Milvus 不一致情况:")
logger.warning(f" 缺失集合数量: {len(milvus_missing)}")
for collection_info in milvus_missing:
logger.warning(f" - 集合: {collection_info['collection_name']}, 实体数: {collection_info['count']}")
logger.warning(f" 缺失文件记录数量: {len(milvus_files_missing)}")
for file_info in milvus_files_missing:
logger.warning(f" - 数据库: {file_info['database_id']}, 向量数: {file_info['vector_count']}, 元数据文件数: {file_info['metadata_files_count']}")
logger.warning("=" * 80)
logger.warning(f"总计:缺失集合 {total_missing_collections} 个,缺失文件记录 {total_missing_files}")
logger.warning("建议:检查这些不一致的数据,必要时进行数据清理或元数据修复")
logger.warning("=" * 80)
async def manual_consistency_check(self) -> dict:
"""
手动触发数据一致性检测
Returns:
检测结果字典
"""
logger.info("手动触发数据一致性检测...")
return await self.detect_data_inconsistencies()

View File

@ -210,6 +210,15 @@ def get_embedding_config(embed_info: dict) -> dict:
try: try:
if embed_info: if embed_info:
# 处理 embed_info 可能是字典或 EmbedModelInfo 对象的情况
if hasattr(embed_info, 'name'):
# EmbedModelInfo 对象
config_dict["model"] = embed_info.name
config_dict["api_key"] = os.getenv(embed_info.api_key, embed_info.api_key)
config_dict["base_url"] = embed_info.base_url
config_dict["dimension"] = embed_info.dimension
else:
# 字典形式
config_dict["model"] = embed_info["name"] config_dict["model"] = embed_info["name"]
config_dict["api_key"] = os.getenv(embed_info["api_key"], embed_info["api_key"]) config_dict["api_key"] = os.getenv(embed_info["api_key"], embed_info["api_key"])
config_dict["base_url"] = embed_info["base_url"] config_dict["base_url"] = embed_info["base_url"]