中途保存

This commit is contained in:
YILING0013
2025-02-06 18:35:28 +08:00
parent bf55e74ac9
commit a2b86a90dc
5 changed files with 523 additions and 331 deletions
+148
View File
@@ -0,0 +1,148 @@
# llm_adapters.py
# -*- coding: utf-8 -*-
import logging
from typing import Optional
from langchain_openai import ChatOpenAI
def ensure_openai_base_url_has_v1(url: str) -> str:
"""
若用户输入的 url 不包含 '/v1',则在末尾追加 '/v1'
"""
import re
url = url.strip()
if not url:
return url
if not re.search(r'/v\d+$', url):
if '/v1' not in url:
url = url.rstrip('/') + '/v1'
return url
class BaseLLMAdapter:
"""
统一的 LLM 接口基类,为不同后端(OpenAI、Ollama、ML Studio 等)提供一致的方法签名。
"""
def invoke(self, prompt: str) -> str:
raise NotImplementedError("Subclasses must implement .invoke(prompt) method.")
class DeepSeekAdapter(BaseLLMAdapter):
"""
适配官方/OpenAI兼容接口(使用 langchain.ChatOpenAI
"""
def __init__(self, api_key: str, base_url: str, model_name: str, max_tokens:int, temperature: float = 0.7):
self.base_url = ensure_openai_base_url_has_v1(base_url)
self.api_key = api_key
self.model_name = model_name
self.max_tokens = max_tokens
self.temperature = temperature
self._client = ChatOpenAI(
model=self.model_name,
api_key=self.api_key,
base_url=self.base_url,
max_tokens=self.max_tokens,
temperature=self.temperature
)
def invoke(self, prompt: str) -> str:
response = self._client.invoke(prompt)
if not response:
logging.warning("No response from DeepSeekAdapter.")
return ""
return response.content
class OpenAIAdapter(BaseLLMAdapter):
"""
适配官方/OpenAI兼容接口(使用 langchain.ChatOpenAI
"""
def __init__(self, api_key: str, base_url: str, model_name: str, max_tokens:int, temperature: float = 0.7):
self.base_url = ensure_openai_base_url_has_v1(base_url)
self.api_key = api_key
self.model_name = model_name
self.max_tokens = max_tokens
self.temperature = temperature
self._client = ChatOpenAI(
model=self.model_name,
api_key=self.api_key,
base_url=self.base_url,
max_tokens=self.max_tokens,
temperature=self.temperature
)
def invoke(self, prompt: str) -> str:
response = self._client.invoke(prompt)
if not response:
logging.warning("No response from OpenAIAdapter.")
return ""
return response.content
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):
self.base_url = ensure_openai_base_url_has_v1(base_url)
self.api_key = api_key
self.model_name = model_name
self.max_tokens = max_tokens
self.temperature = temperature
self._client = ChatOpenAI(
model=self.model_name,
api_key=self.api_key,
base_url=self.base_url,
max_tokens=self.max_tokens,
temperature=self.temperature
)
def invoke(self, prompt: str) -> str:
response = self._client.invoke(prompt)
if not response:
logging.warning("No response from OllamaAdapter.")
return ""
return response.content
class MLStudioAdapter(BaseLLMAdapter):
def __init__(self, api_key: str, base_url: str, model_name: str, max_tokens:int, temperature: float = 0.7):
self.base_url = ensure_openai_base_url_has_v1(base_url)
self.api_key = api_key
self.model_name = model_name
self.max_tokens = max_tokens
self.temperature = temperature
self._client = ChatOpenAI(
model=self.model_name,
api_key=self.api_key,
base_url=self.base_url,
max_tokens=self.max_tokens,
temperature=self.temperature
)
def invoke(self, prompt: str) -> str:
response = self._client.invoke(prompt)
if not response:
logging.warning("No response from MLStudioAdapter.")
return ""
return response.content
def create_llm_adapter(
interface_format: str,
base_url: str,
model_name: str,
api_key: str,
temperature: float
) -> BaseLLMAdapter:
"""
工厂函数:根据 interface_format 返回不同的适配器实例。
"""
if interface_format.lower() == "deepseek":
return DeepSeekAdapter(api_key, base_url, model_name, temperature)
elif interface_format.lower() == "openai":
return OpenAIAdapter(api_key, base_url, model_name, temperature)
elif interface_format.lower() == "ollama":
return OllamaAdapter(api_key, base_url, model_name, temperature)
elif interface_format.lower() == "ml studio":
return MLStudioAdapter(api_key, base_url, model_name, temperature)
else:
raise ValueError(f"Unknown interface_format: {interface_format}")