fixed #98
This commit is contained in:
+40
-4
@@ -133,6 +133,38 @@ class MLStudioEmbeddingAdapter(BaseEmbeddingAdapter):
|
|||||||
def embed_query(self, query: str) -> List[float]:
|
def embed_query(self, query: str) -> List[float]:
|
||||||
return self._embedding.embed_query(query)
|
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(
|
def create_embedding_adapter(
|
||||||
interface_format: str,
|
interface_format: str,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
@@ -142,13 +174,17 @@ def create_embedding_adapter(
|
|||||||
"""
|
"""
|
||||||
工厂函数:根据 interface_format 返回不同的 embedding 适配器实例
|
工厂函数:根据 interface_format 返回不同的 embedding 适配器实例
|
||||||
"""
|
"""
|
||||||
if interface_format.lower() == "openai":
|
fmt = interface_format.strip().lower()
|
||||||
|
if fmt == "openai":
|
||||||
return OpenAIEmbeddingAdapter(api_key, base_url, model_name)
|
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)
|
return AzureOpenAIEmbeddingAdapter(api_key, base_url, model_name)
|
||||||
elif interface_format.lower() == "ollama":
|
elif fmt == "ollama":
|
||||||
return OllamaEmbeddingAdapter(model_name, base_url)
|
return OllamaEmbeddingAdapter(model_name, base_url)
|
||||||
elif interface_format.lower() == "ml studio":
|
elif fmt == "ml studio":
|
||||||
return MLStudioEmbeddingAdapter(api_key, base_url, model_name)
|
return MLStudioEmbeddingAdapter(api_key, base_url, model_name)
|
||||||
|
elif fmt == "gemini":
|
||||||
|
# base_url 对 Gemini 暂无用处,可忽略
|
||||||
|
return GeminiEmbeddingAdapter(api_key, model_name)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unknown embedding interface_format: {interface_format}")
|
raise ValueError(f"Unknown embedding interface_format: {interface_format}")
|
||||||
|
|||||||
+10
-9
@@ -18,7 +18,7 @@ def ensure_openai_base_url_has_v1(url: str) -> str:
|
|||||||
|
|
||||||
class BaseLLMAdapter:
|
class BaseLLMAdapter:
|
||||||
"""
|
"""
|
||||||
统一的 LLM 接口基类,为不同后端(OpenAI、Ollama、ML Studio 等)提供一致的方法签名。
|
统一的 LLM 接口基类,为不同后端(OpenAI、Ollama、ML Studio、Gemini等)提供一致的方法签名。
|
||||||
"""
|
"""
|
||||||
def invoke(self, prompt: str) -> str:
|
def invoke(self, prompt: str) -> str:
|
||||||
raise NotImplementedError("Subclasses must implement .invoke(prompt) method.")
|
raise NotImplementedError("Subclasses must implement .invoke(prompt) method.")
|
||||||
@@ -81,7 +81,7 @@ class OpenAIAdapter(BaseLLMAdapter):
|
|||||||
|
|
||||||
class GeminiAdapter(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):
|
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
|
self.api_key = api_key
|
||||||
@@ -151,7 +151,6 @@ class AzureOpenAIAdapter(BaseLLMAdapter):
|
|||||||
class OllamaAdapter(BaseLLMAdapter):
|
class OllamaAdapter(BaseLLMAdapter):
|
||||||
"""
|
"""
|
||||||
Ollama 同样有一个 OpenAI-like /v1/chat 接口,可直接使用 ChatOpenAI。
|
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):
|
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)
|
self.base_url = ensure_openai_base_url_has_v1(base_url)
|
||||||
@@ -214,17 +213,19 @@ def create_llm_adapter(
|
|||||||
"""
|
"""
|
||||||
工厂函数:根据 interface_format 返回不同的适配器实例。
|
工厂函数:根据 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)
|
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)
|
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)
|
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)
|
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)
|
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)
|
return GeminiAdapter(api_key, model_name, max_tokens, temperature, timeout)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unknown interface_format: {interface_format}")
|
raise ValueError(f"Unknown interface_format: {interface_format}")
|
||||||
|
|||||||
Reference in New Issue
Block a user