From b8e83b71a48ea14415138b7dae76c50a46089de4 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Mon, 15 Dec 2025 13:35:24 +0800 Subject: [PATCH] =?UTF-8?q?fix(embedding):=20=E4=BF=AE=E5=A4=8D=E5=B5=8C?= =?UTF-8?q?=E5=85=A5=E6=A8=A1=E5=9E=8B=E9=85=8D=E7=BD=AE=E4=B8=AD=E7=9A=84?= =?UTF-8?q?=E5=85=BC=E5=AE=B9=E6=80=A7=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 修复 get_embedding_config 中对 model/name 字段的兼容性处理 - 改进错误处理和日志记录 - 支持多种配置格式以保持向后兼容性 - 修复 lightrag.py 中 embedding 函数的模型名称获取逻辑 --- src/knowledge/implementations/lightrag.py | 12 ++++- src/knowledge/utils/kb_utils.py | 64 +++++++++++++++-------- 2 files changed, 53 insertions(+), 23 deletions(-) diff --git a/src/knowledge/implementations/lightrag.py b/src/knowledge/implementations/lightrag.py index 08feed01..cdddd5f2 100644 --- a/src/knowledge/implementations/lightrag.py +++ b/src/knowledge/implementations/lightrag.py @@ -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]: diff --git a/src/knowledge/utils/kb_utils.py b/src/knowledge/utils/kb_utils.py index ac9411ed..d955fffd 100644 --- a/src/knowledge/utils/kb_utils.py +++ b/src/knowledge/utils/kb_utils.py @@ -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: