添加对 Azure OpenAI 的支持,新增适配器并更新界面选项

This commit is contained in:
桑榆肖物
2025-02-09 22:50:52 +08:00
parent 1ce0edceac
commit 8afa1083e0
3 changed files with 76 additions and 4 deletions
+30 -1
View File
@@ -4,7 +4,7 @@ import logging
import requests import requests
import traceback import traceback
from typing import List from typing import List
from langchain_openai import OpenAIEmbeddings from langchain_openai import OpenAIEmbeddings, AzureOpenAIEmbeddings
def ensure_openai_base_url_has_v1(url: str) -> str: def ensure_openai_base_url_has_v1(url: str) -> str:
""" """
@@ -46,6 +46,33 @@ class OpenAIEmbeddingAdapter(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 AzureOpenAIEmbeddingAdapter(BaseEmbeddingAdapter):
"""
基于 AzureOpenAIEmbeddings(或兼容接口)的适配器
"""
def __init__(self, api_key: str, base_url: str, model_name: str):
import re
match = re.match(r'https://(.+?)/openai/deployments/(.+?)/embeddings\?api-version=(.+)', base_url)
if match:
self.azure_endpoint = f"https://{match.group(1)}"
self.azure_deployment = match.group(2)
self.api_version = match.group(3)
else:
raise ValueError("Invalid Azure OpenAI base_url format")
self._embedding = AzureOpenAIEmbeddings(
azure_endpoint=self.azure_endpoint,
azure_deployment=self.azure_deployment,
openai_api_key=api_key,
api_version=self.api_version,
)
def embed_documents(self, texts: List[str]) -> List[List[float]]:
return self._embedding.embed_documents(texts)
def embed_query(self, query: str) -> List[float]:
return self._embedding.embed_query(query)
class OllamaEmbeddingAdapter(BaseEmbeddingAdapter): class OllamaEmbeddingAdapter(BaseEmbeddingAdapter):
""" """
其接口路径为 /api/embeddings 其接口路径为 /api/embeddings
@@ -112,6 +139,8 @@ def create_embedding_adapter(
""" """
if interface_format.lower() == "openai": if interface_format.lower() == "openai":
return OpenAIEmbeddingAdapter(api_key, base_url, model_name) return OpenAIEmbeddingAdapter(api_key, base_url, model_name)
elif interface_format.lower() == "azure openai":
return AzureOpenAIEmbeddingAdapter(api_key, base_url, model_name)
elif interface_format.lower() == "ollama": elif interface_format.lower() == "ollama":
return OllamaEmbeddingAdapter(model_name, base_url) return OllamaEmbeddingAdapter(model_name, base_url)
elif interface_format.lower() == "ml studio": elif interface_format.lower() == "ml studio":
+40 -1
View File
@@ -2,7 +2,7 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
import logging import logging
from typing import Optional from typing import Optional
from langchain_openai import ChatOpenAI from langchain_openai import ChatOpenAI, AzureChatOpenAI
def ensure_openai_base_url_has_v1(url: str) -> str: def ensure_openai_base_url_has_v1(url: str) -> str:
import re import re
@@ -77,6 +77,43 @@ class OpenAIAdapter(BaseLLMAdapter):
return "" return ""
return response.content return response.content
class AzureOpenAIAdapter(BaseLLMAdapter):
"""
适配 Azure OpenAI 接口(使用 langchain.ChatOpenAI
"""
def __init__(self, api_key: str, base_url: str, model_name: str, max_tokens: int, temperature: float = 0.7, timeout: Optional[int] = 600):
import re
match = re.match(r'https://(.+?)/openai/deployments/(.+?)/chat/completions\?api-version=(.+)', base_url)
if match:
self.azure_endpoint = f"https://{match.group(1)}"
self.azure_deployment = match.group(2)
self.api_version = match.group(3)
else:
raise ValueError("Invalid Azure OpenAI base_url format")
self.api_key = api_key
self.model_name = self.azure_deployment
self.max_tokens = max_tokens
self.temperature = temperature
self.timeout = timeout
self._client = AzureChatOpenAI(
azure_endpoint=self.azure_endpoint,
azure_deployment=self.azure_deployment,
api_version=self.api_version,
api_key=self.api_key,
max_tokens=self.max_tokens,
temperature=self.temperature,
timeout=self.timeout
)
def invoke(self, prompt: str) -> str:
response = self._client.invoke(prompt)
if not response:
logging.warning("No response from AzureOpenAIAdapter.")
return ""
return response.content
class OllamaAdapter(BaseLLMAdapter): class OllamaAdapter(BaseLLMAdapter):
""" """
Ollama 同样有一个 OpenAI-like /v1/chat 接口,可直接使用 ChatOpenAI。 Ollama 同样有一个 OpenAI-like /v1/chat 接口,可直接使用 ChatOpenAI。
@@ -147,6 +184,8 @@ def create_llm_adapter(
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 interface_format.lower() == "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":
return AzureOpenAIAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout)
elif interface_format.lower() == "ollama": elif interface_format.lower() == "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 interface_format.lower() == "ml studio":
+6 -2
View File
@@ -290,6 +290,8 @@ class NovelGeneratorGUI:
self.base_url_var.set("http://localhost:1234/v1") self.base_url_var.set("http://localhost:1234/v1")
elif new_value == "OpenAI": elif new_value == "OpenAI":
self.base_url_var.set("https://api.openai.com/v1") self.base_url_var.set("https://api.openai.com/v1")
elif new_value == "Azure OpenAI":
self.base_url_var.set("https://[az].openai.azure.com/openai/deployments/[model]/chat/completions?api-version=2024-08-01-preview")
elif new_value == "DeepSeek": elif new_value == "DeepSeek":
self.base_url_var.set("https://api.deepseek.com/v1") self.base_url_var.set("https://api.deepseek.com/v1")
@@ -332,7 +334,7 @@ class NovelGeneratorGUI:
column=0, column=0,
font=("Microsoft YaHei", 12) font=("Microsoft YaHei", 12)
) )
interface_options = ["DeepSeek", "OpenAI", "Ollama", "ML Studio"] interface_options = ["DeepSeek", "OpenAI", "Azure OpenAI", "Ollama", "ML Studio"]
interface_dropdown = ctk.CTkOptionMenu( interface_dropdown = ctk.CTkOptionMenu(
self.ai_config_tab, self.ai_config_tab,
values=interface_options, values=interface_options,
@@ -452,6 +454,8 @@ class NovelGeneratorGUI:
self.embedding_url_var.set("http://localhost:1234/v1") self.embedding_url_var.set("http://localhost:1234/v1")
elif new_value == "OpenAI": elif new_value == "OpenAI":
self.embedding_url_var.set("https://api.openai.com/v1") self.embedding_url_var.set("https://api.openai.com/v1")
elif new_value == "Azure OpenAI":
self.embedding_url_var.set("https://[az].openai.azure.com/openai/deployments/[model]/embeddings?api-version=2023-05-15")
elif new_value == "DeepSeek": elif new_value == "DeepSeek":
self.embedding_url_var.set("https://api.deepseek.com/v1") self.embedding_url_var.set("https://api.deepseek.com/v1")
@@ -482,7 +486,7 @@ class NovelGeneratorGUI:
column=0, column=0,
font=("Microsoft YaHei", 12) font=("Microsoft YaHei", 12)
) )
emb_interface_options = ["DeepSeek", "OpenAI", "Ollama", "ML Studio"] emb_interface_options = ["DeepSeek", "OpenAI", "Azure OpenAI", "Ollama", "ML Studio"]
emb_interface_dropdown = ctk.CTkOptionMenu( emb_interface_dropdown = ctk.CTkOptionMenu(
self.embeddings_config_tab, self.embeddings_config_tab,
values=emb_interface_options, values=emb_interface_options,