fix(embedding): 修复嵌入模型配置中的兼容性问题

- 修复 get_embedding_config 中对 model/name 字段的兼容性处理
- 改进错误处理和日志记录
- 支持多种配置格式以保持向后兼容性
- 修复 lightrag.py 中 embedding 函数的模型名称获取逻辑
This commit is contained in:
Wenjie Zhang 2025-12-15 13:35:24 +08:00
parent 3d6a38a73d
commit b8e83b71a4
2 changed files with 53 additions and 23 deletions

View File

@ -200,6 +200,7 @@ class LightRagKB(KnowledgeBase):
def _get_embedding_func(self, embed_info: dict):
"""获取 embedding 函数"""
config_dict = get_embedding_config(embed_info)
logger.debug(f"Embedding config dict: {config_dict}")
if config_dict.get("model_id") and config_dict["model_id"].startswith("ollama"):
from lightrag.llm.ollama import ollama_embed
@ -219,15 +220,22 @@ class LightRagKB(KnowledgeBase):
),
)
# 尝试获取模型名称,支持多种键名以保持兼容性
if "name" in config_dict and config_dict["name"]:
model_name = config_dict["name"]
elif "model" in config_dict and config_dict["model"]:
model_name = config_dict["model"]
else:
raise ValueError(f"Neither 'name' nor 'model' found in config_dict or both are empty: {config_dict}")
return EmbeddingFunc(
embedding_dim=config_dict["dimension"],
max_token_size=8192,
func=lambda texts: openai_embed(
texts=texts,
model=config_dict["model"],
model=model_name,
api_key=config_dict["api_key"],
base_url=config_dict["base_url"].replace("/embeddings", ""),
),
)
)
async def add_content(self, db_id: str, items: list[str], params: dict | None = None) -> list[dict]:

View File

@ -289,30 +289,52 @@ def get_embedding_config(embed_info: dict) -> dict:
Returns:
dict: 标准化的嵌入配置
"""
config_dict = {}
try:
if embed_info:
# 优先检查是否有 model_id 字段
if "model_id" in embed_info:
return config.embed_model_names[embed_info["model_id"]].model_dump()
elif hasattr(embed_info, "name") and isinstance(embed_info, EmbedModelInfo):
return embed_info.model_dump()
else:
# 字典形式(保持向后兼容)
config_dict["model"] = embed_info["name"]
config_dict["api_key"] = os.getenv(embed_info["api_key"]) or embed_info["api_key"]
config_dict["base_url"] = embed_info["base_url"]
config_dict["dimension"] = embed_info.get("dimension", 1024)
else:
return config.embed_model_names[config.embed_model].model_dump()
# 检查 embed_info 是否有效
if not embed_info or ("model" not in embed_info and "name" not in embed_info):
logger.error(f"Invalid embed_info: {embed_info}, using default embedding model config")
raise ValueError("Invalid embed_info: must be a non-empty dictionary")
# 优先检查是否有 model_id 字段
if "model_id" in embed_info and embed_info["model_id"]:
logger.warning(f"Using model_id: {embed_info['model_id']}")
config_dict = config.embed_model_names[embed_info["model_id"]].model_dump()
config_dict["api_key"] = os.getenv(config_dict["api_key"]) or config_dict["api_key"]
return config_dict
# 检查是否是 EmbedModelInfo 对象(在某些情况下可能直接传入对象)
if hasattr(embed_info, "name") and isinstance(embed_info, EmbedModelInfo):
logger.debug(f"Using EmbedModelInfo object: {embed_info.name}")
config_dict = embed_info.model_dump()
config_dict["api_key"] = os.getenv(config_dict["api_key"]) or config_dict["api_key"]
return config_dict
# 字典形式(保持向后兼容)
# 检查必需字段是否存在
if not embed_info.get("name") or not embed_info.get("base_url"):
logger.warning(f"embed_info missing required 'name' or 'base_url' field: {embed_info}, using default")
raise ValueError("embed_info missing required 'name' or 'base_url' field")
config_dict = {
"model": embed_info["name"],
"api_key": os.getenv(embed_info["api_key"]) or embed_info["api_key"],
"base_url": embed_info["base_url"],
"dimension": embed_info.get("dimension", 1024)
}
logger.debug(f"Embedding config from dict: {config_dict}")
return config_dict
except Exception as e:
logger.error(f"Error in get_embedding_config: {e}, {embed_info}")
raise ValueError(f"Error in get_embedding_config: {e}")
logger.debug(f"Embedding config: {config_dict}")
return config_dict
logger.error(f"Error in get_embedding_config: {e}, embed_info={embed_info}")
# 返回默认配置作为fallback
logger.warning("Falling back to default embedding model config")
try:
config_dict = config.embed_model_names[config.embed_model].model_dump()
config_dict["api_key"] = os.getenv(config_dict["api_key"]) or config_dict["api_key"]
return config_dict
except Exception as fallback_error:
logger.error(f"Failed to get default embedding config: {fallback_error}")
raise ValueError(f"Failed to get embedding config and fallback failed: {e}")
def is_minio_url(file_path: str) -> bool: