Files
AI_NovelGenerator/llm_adapters.py
T
2025-02-06 19:37:37 +08:00

150 lines
5.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# 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,
max_tokens: int
) -> BaseLLMAdapter:
"""
工厂函数:根据 interface_format 返回不同的适配器实例。
"""
if interface_format.lower() == "deepseek":
return DeepSeekAdapter(api_key, base_url, model_name, max_tokens, temperature)
elif interface_format.lower() == "openai":
return OpenAIAdapter(api_key, base_url, model_name, max_tokens, temperature)
elif interface_format.lower() == "ollama":
return OllamaAdapter(api_key, base_url, model_name, max_tokens, temperature)
elif interface_format.lower() == "ml studio":
return MLStudioAdapter(api_key, base_url, model_name, max_tokens, temperature)
else:
raise ValueError(f"Unknown interface_format: {interface_format}")