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 225ccce69c
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): def _get_embedding_func(self, embed_info: dict):
"""获取 embedding 函数""" """获取 embedding 函数"""
config_dict = get_embedding_config(embed_info) 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"): if config_dict.get("model_id") and config_dict["model_id"].startswith("ollama"):
from lightrag.llm.ollama import ollama_embed 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( return EmbeddingFunc(
embedding_dim=config_dict["dimension"], embedding_dim=config_dict["dimension"],
max_token_size=8192, max_token_size=8192,
func=lambda texts: openai_embed( func=lambda texts: openai_embed(
texts=texts, texts=texts,
model=config_dict["model"], model=model_name,
api_key=config_dict["api_key"], api_key=config_dict["api_key"],
base_url=config_dict["base_url"].replace("/embeddings", ""), 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]: 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: Returns:
dict: 标准化的嵌入配置 dict: 标准化的嵌入配置
""" """
config_dict = {}
try: try:
if embed_info: # 检查 embed_info 是否有效
# 优先检查是否有 model_id 字段 if not embed_info or ("model" not in embed_info and "name" not in embed_info):
if "model_id" in embed_info: logger.error(f"Invalid embed_info: {embed_info}, using default embedding model config")
return config.embed_model_names[embed_info["model_id"]].model_dump() raise ValueError("Invalid embed_info: must be a non-empty dictionary")
elif hasattr(embed_info, "name") and isinstance(embed_info, EmbedModelInfo):
return embed_info.model_dump() # 优先检查是否有 model_id 字段
else: if "model_id" in embed_info and embed_info["model_id"]:
# 字典形式(保持向后兼容) logger.warning(f"Using model_id: {embed_info['model_id']}")
config_dict["model"] = embed_info["name"] config_dict = config.embed_model_names[embed_info["model_id"]].model_dump()
config_dict["api_key"] = os.getenv(embed_info["api_key"]) or embed_info["api_key"] config_dict["api_key"] = os.getenv(config_dict["api_key"]) or config_dict["api_key"]
config_dict["base_url"] = embed_info["base_url"] return config_dict
config_dict["dimension"] = embed_info.get("dimension", 1024)
else: # 检查是否是 EmbedModelInfo 对象(在某些情况下可能直接传入对象)
return config.embed_model_names[config.embed_model].model_dump() 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: except Exception as e:
logger.error(f"Error in get_embedding_config: {e}, {embed_info}") logger.error(f"Error in get_embedding_config: {e}, embed_info={embed_info}")
raise ValueError(f"Error in get_embedding_config: {e}") # 返回默认配置作为fallback
logger.warning("Falling back to default embedding model config")
logger.debug(f"Embedding config: {config_dict}") try:
return config_dict 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: def is_minio_url(file_path: str) -> bool: