Merge pull request #107 from huanshang141/OpenAI-force-link

对支持openai格式的接口,base url带#号时强制使用原base url
This commit is contained in:
IdleCloud
2025-02-15 12:00:37 +08:00
committed by GitHub
+14 -5
View File
@@ -9,11 +9,20 @@ from azure.ai.inference import ChatCompletionsClient
from azure.core.credentials import AzureKeyCredential from azure.core.credentials import AzureKeyCredential
from azure.ai.inference.models import SystemMessage, UserMessage from azure.ai.inference.models import SystemMessage, UserMessage
def ensure_openai_base_url_has_v1(url: str) -> str: def check_base_url(url: str) -> str:
"""
处理base_url的规则:
1. 如果url以#结尾,则移除#并直接使用用户提供的url
2. 否则检查是否需要添加/v1后缀
"""
import re import re
url = url.strip() url = url.strip()
if not url: if not url:
return url return url
if url.endswith('#'):
return url.rstrip('#')
if not re.search(r'/v\d+$', url): if not re.search(r'/v\d+$', url):
if '/v1' not in url: if '/v1' not in url:
url = url.rstrip('/') + '/v1' url = url.rstrip('/') + '/v1'
@@ -31,7 +40,7 @@ class DeepSeekAdapter(BaseLLMAdapter):
适配官方/OpenAI兼容接口(使用 langchain.ChatOpenAI 适配官方/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): 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 = check_base_url(base_url)
self.api_key = api_key self.api_key = api_key
self.model_name = model_name self.model_name = model_name
self.max_tokens = max_tokens self.max_tokens = max_tokens
@@ -59,7 +68,7 @@ class OpenAIAdapter(BaseLLMAdapter):
适配官方/OpenAI兼容接口(使用 langchain.ChatOpenAI 适配官方/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): 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 = check_base_url(base_url)
self.api_key = api_key self.api_key = api_key
self.model_name = model_name self.model_name = model_name
self.max_tokens = max_tokens self.max_tokens = max_tokens
@@ -156,7 +165,7 @@ class OllamaAdapter(BaseLLMAdapter):
Ollama 同样有一个 OpenAI-like /v1/chat 接口,可直接使用 ChatOpenAI。 Ollama 同样有一个 OpenAI-like /v1/chat 接口,可直接使用 ChatOpenAI。
""" """
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 = check_base_url(base_url)
self.api_key = api_key self.api_key = api_key
self.model_name = model_name self.model_name = model_name
self.max_tokens = max_tokens self.max_tokens = max_tokens
@@ -181,7 +190,7 @@ class OllamaAdapter(BaseLLMAdapter):
class MLStudioAdapter(BaseLLMAdapter): class MLStudioAdapter(BaseLLMAdapter):
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 = check_base_url(base_url)
self.api_key = api_key self.api_key = api_key
self.model_name = model_name self.model_name = model_name
self.max_tokens = max_tokens self.max_tokens = max_tokens