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: