diff --git a/embedding_adapters.py b/embedding_adapters.py index 8e9cf6f..db4738b 100644 --- a/embedding_adapters.py +++ b/embedding_adapters.py @@ -45,7 +45,7 @@ class OpenAIEmbeddingAdapter(BaseEmbeddingAdapter): def embed_query(self, query: str) -> List[float]: return self._embedding.embed_query(query) - + class AzureOpenAIEmbeddingAdapter(BaseEmbeddingAdapter): """ 基于 AzureOpenAIEmbeddings(或兼容接口)的适配器 @@ -133,6 +133,38 @@ class MLStudioEmbeddingAdapter(BaseEmbeddingAdapter): def embed_query(self, query: str) -> List[float]: return self._embedding.embed_query(query) +class GeminiEmbeddingAdapter(BaseEmbeddingAdapter): + """ + 基于 Google Generative AI (Gemini)接口的 Embedding 适配器 + """ + def __init__(self, api_key: str, model_name: str): + from google import genai + # 全局配置,也可根据需要改成 Client(...) 初始化方式 + genai.configure(api_key=api_key) + self.model_name = model_name + + def embed_documents(self, texts: List[str]) -> List[List[float]]: + from google import genai + embeddings = [] + for text in texts: + try: + result = genai.embed_content(model=self.model_name, content=text) + # 返回结构中包含 'embedding' 字段 + embeddings.append(result.get('embedding', [])) + except Exception as e: + logging.error(f"Gemini embed_content error: {e}") + embeddings.append([]) + return embeddings + + def embed_query(self, query: str) -> List[float]: + from google import genai + try: + result = genai.embed_content(model=self.model_name, content=query) + return result.get('embedding', []) + except Exception as e: + logging.error(f"Gemini embed_content error: {e}") + return [] + def create_embedding_adapter( interface_format: str, api_key: str, @@ -142,13 +174,17 @@ def create_embedding_adapter( """ 工厂函数:根据 interface_format 返回不同的 embedding 适配器实例 """ - if interface_format.lower() == "openai": + fmt = interface_format.strip().lower() + if fmt == "openai": return OpenAIEmbeddingAdapter(api_key, base_url, model_name) - elif interface_format.lower() == "azure openai": + elif fmt == "azure openai": return AzureOpenAIEmbeddingAdapter(api_key, base_url, model_name) - elif interface_format.lower() == "ollama": + elif fmt == "ollama": return OllamaEmbeddingAdapter(model_name, base_url) - elif interface_format.lower() == "ml studio": + elif fmt == "ml studio": return MLStudioEmbeddingAdapter(api_key, base_url, model_name) + elif fmt == "gemini": + # base_url 对 Gemini 暂无用处,可忽略 + return GeminiEmbeddingAdapter(api_key, model_name) else: raise ValueError(f"Unknown embedding interface_format: {interface_format}") diff --git a/llm_adapters.py b/llm_adapters.py index 525b117..6437509 100644 --- a/llm_adapters.py +++ b/llm_adapters.py @@ -18,7 +18,7 @@ def ensure_openai_base_url_has_v1(url: str) -> str: class BaseLLMAdapter: """ - 统一的 LLM 接口基类,为不同后端(OpenAI、Ollama、ML Studio 等)提供一致的方法签名。 + 统一的 LLM 接口基类,为不同后端(OpenAI、Ollama、ML Studio、Gemini等)提供一致的方法签名。 """ def invoke(self, prompt: str) -> str: raise NotImplementedError("Subclasses must implement .invoke(prompt) method.") @@ -81,7 +81,7 @@ class OpenAIAdapter(BaseLLMAdapter): class GeminiAdapter(BaseLLMAdapter): """ - 适配 Google Gemini 接口 + 适配 Google Gemini (Google Generative AI) 接口 """ def __init__(self, api_key: str, model_name: str, max_tokens: int, temperature: float = 0.7, timeout: Optional[int] = 600): self.api_key = api_key @@ -151,7 +151,6 @@ class AzureOpenAIAdapter(BaseLLMAdapter): class OllamaAdapter(BaseLLMAdapter): """ Ollama 同样有一个 OpenAI-like /v1/chat 接口,可直接使用 ChatOpenAI。 - 但是通常 Ollama 默认本地服务在 http://localhost:11434,如果符合OpenAI风格即可直接传参。 """ def __init__(self, api_key: str, base_url: str, model_name: str, max_tokens: int, temperature: float = 0.7, timeout: Optional[int] = 600): self.base_url = ensure_openai_base_url_has_v1(base_url) @@ -214,17 +213,19 @@ def create_llm_adapter( """ 工厂函数:根据 interface_format 返回不同的适配器实例。 """ - if interface_format.lower() == "deepseek": + fmt = interface_format.strip().lower() + if fmt == "deepseek": return DeepSeekAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout) - elif interface_format.lower() == "openai": + elif fmt == "openai": return OpenAIAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout) - elif interface_format.lower() == "azure openai": + elif fmt == "azure openai": return AzureOpenAIAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout) - elif interface_format.lower() == "ollama": + elif fmt == "ollama": return OllamaAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout) - elif interface_format.lower() == "ml studio": + elif fmt == "ml studio": return MLStudioAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout) - elif interface_format.lower() == "gemini": + elif fmt == "gemini": + # base_url 对 Gemini 暂无用处,可忽略 return GeminiAdapter(api_key, model_name, max_tokens, temperature, timeout) else: raise ValueError(f"Unknown interface_format: {interface_format}")