refactor(kb_utils): 重构获取嵌入配置的逻辑以提高可维护性

优先使用新版配置格式并添加断言检查
移除旧版兼容性代码,简化逻辑
添加异常处理和日志记录
This commit is contained in:
Wenjie Zhang 2025-12-18 11:29:17 +08:00
parent 98fdce0e79
commit cd5c6947e2

View File

@ -1,6 +1,7 @@
import hashlib
import os
import time
import traceback
from pathlib import Path
import aiofiles
@ -290,39 +291,33 @@ def get_embedding_config(embed_info: dict) -> dict:
dict: 标准化的嵌入配置
"""
try:
# 检查 embed_info 是否有效
# 使用最新配置
assert isinstance(embed_info, dict), f"embed_info must be a dict, got {type(embed_info)}"
assert "model_id" in embed_info, f"embed_info must contain 'model_id', got {embed_info}"
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
except AssertionError as e:
logger.error(f"AssertionError in get_embedding_config: {e}, embed_info={embed_info}")
# 兼容性检查:旧版配置字段
try:
# 1. 检查 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 对象(在某些情况下可能直接传入对象)
# 2. 检查是否是 EmbedModelInfo 对象(在某些情况下可能直接传入对象)
if hasattr(embed_info, "name") and isinstance(embed_info, EmbedModelInfo):
logger.debug(f"Using EmbedModelInfo object: {embed_info.name}")
logger.debug(f"Using EmbedModelInfo object: {embed_info.name}, {traceback.format_exc()}")
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
raise ValueError(f"Unsupported embed_info format: {embed_info}")
except Exception as e:
logger.error(f"Error in get_embedding_config: {e}, embed_info={embed_info}")