添加对 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
+40 -1
View File
@@ -2,7 +2,7 @@
# -*- coding: utf-8 -*-
import logging
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:
import re
@@ -77,6 +77,43 @@ class OpenAIAdapter(BaseLLMAdapter):
return ""
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):
"""
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)
elif interface_format.lower() == "openai":
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":
return OllamaAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout)
elif interface_format.lower() == "ml studio":