219 lines
6.0 KiB
Python
219 lines
6.0 KiB
Python
"""
|
|
LLM Provider 抽象层
|
|
|
|
支持多种 LLM 提供商的统一接口
|
|
"""
|
|
|
|
from typing import Any, Dict, List, Optional, Protocol, Type, runtime_checkable
|
|
from abc import ABC, abstractmethod
|
|
import logging
|
|
|
|
from langchain_core.language_models import BaseChatModel
|
|
|
|
from config import Config
|
|
from .registry import ProviderRegistry, register_provider
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
@runtime_checkable
|
|
class LLMProvider(Protocol):
|
|
"""
|
|
LLM Provider 协议
|
|
|
|
定义所有 LLM 提供商必须实现的接口
|
|
"""
|
|
|
|
@property
|
|
def name(self) -> str:
|
|
"""Provider 名称"""
|
|
...
|
|
|
|
@property
|
|
def supported_models(self) -> List[str]:
|
|
"""支持的模型列表"""
|
|
...
|
|
|
|
def create_model(self, config: Dict[str, Any]) -> BaseChatModel:
|
|
"""
|
|
创建 LLM 模型实例
|
|
|
|
Args:
|
|
config: 模型配置
|
|
|
|
Returns:
|
|
BaseChatModel 实例
|
|
"""
|
|
...
|
|
|
|
|
|
class BaseLLMProvider(ABC):
|
|
"""LLM Provider 基类"""
|
|
|
|
def __init__(self, name: str, supported_models: List[str]):
|
|
self._name = name
|
|
self._supported_models = supported_models
|
|
|
|
@property
|
|
def name(self) -> str:
|
|
return self._name
|
|
|
|
@property
|
|
def supported_models(self) -> List[str]:
|
|
return self._supported_models
|
|
|
|
@abstractmethod
|
|
def create_model(self, config: Dict[str, Any]) -> BaseChatModel:
|
|
"""子类实现具体的模型创建逻辑"""
|
|
pass
|
|
|
|
|
|
class OpenAICompatibleProvider(BaseLLMProvider):
|
|
"""
|
|
OpenAI 兼容的 Provider
|
|
|
|
支持所有兼容 OpenAI API 的服务:
|
|
- OpenAI
|
|
- Azure OpenAI
|
|
- 本地部署的兼容服务
|
|
- 国产大模型(通义千问、文心一言等)
|
|
"""
|
|
|
|
def __init__(self, name: str = "openai_compatible", supported_models: Optional[List[str]] = None):
|
|
super().__init__(name, supported_models or [])
|
|
|
|
def create_model(self, config: Dict[str, Any]) -> BaseChatModel:
|
|
from langchain_openai import ChatOpenAI
|
|
|
|
return ChatOpenAI(
|
|
model=config.get("model_name") or config.get("MODEL_NAME"),
|
|
openai_api_key=config.get("api_key") or config.get("OPENAI_API_KEY"),
|
|
openai_api_base=config.get("base_url") or config.get("URL"),
|
|
temperature=config.get("temperature", 0.7),
|
|
max_tokens=config.get("max_tokens"),
|
|
timeout=config.get("timeout", 30),
|
|
max_retries=config.get("max_retries", 3),
|
|
)
|
|
|
|
|
|
class AzureOpenAIProvider(BaseLLMProvider):
|
|
"""Azure OpenAI Provider"""
|
|
|
|
def __init__(self):
|
|
super().__init__("azure", ["gpt-4", "gpt-35-turbo"])
|
|
|
|
def create_model(self, config: Dict[str, Any]) -> BaseChatModel:
|
|
from langchain_openai import AzureChatOpenAI
|
|
|
|
return AzureChatOpenAI(
|
|
azure_deployment=config.get("deployment_name"),
|
|
openai_api_version=config.get("api_version", "2024-02-15-preview"),
|
|
azure_endpoint=config.get("endpoint"),
|
|
openai_api_key=config.get("api_key"),
|
|
temperature=config.get("temperature", 0.7),
|
|
)
|
|
|
|
|
|
class LLMFactory:
|
|
"""
|
|
LLM 工厂类
|
|
|
|
统一管理 LLM 实例的创建,支持:
|
|
- 多种 Provider
|
|
- 配置驱动
|
|
- 单例缓存
|
|
"""
|
|
|
|
_instances: Dict[str, BaseChatModel] = {}
|
|
_default_provider: str = "openai_compatible"
|
|
|
|
@classmethod
|
|
def register_provider(cls, name: str, provider: LLMProvider) -> None:
|
|
"""注册 Provider"""
|
|
ProviderRegistry._entries[name] = type(
|
|
"RegistryEntry",
|
|
(),
|
|
{"instance": provider, "metadata": {}}
|
|
)()
|
|
logger.info(f"Registered LLM provider: {name}")
|
|
|
|
@classmethod
|
|
def get_provider(cls, name: str) -> Optional[LLMProvider]:
|
|
"""获取 Provider"""
|
|
entry = ProviderRegistry.get_entry(name)
|
|
if entry:
|
|
return entry.instance
|
|
return None
|
|
|
|
@classmethod
|
|
def create(
|
|
cls,
|
|
provider: Optional[str] = None,
|
|
config: Optional[Dict[str, Any]] = None,
|
|
model_section: Optional[str] = None,
|
|
use_cache: bool = True,
|
|
) -> BaseChatModel:
|
|
"""
|
|
创建 LLM 实例
|
|
|
|
Args:
|
|
provider: Provider 名称,默认使用 openai_compatible
|
|
config: 模型配置
|
|
model_section: 配置文件中的模型段名
|
|
use_cache: 是否使用缓存
|
|
|
|
Returns:
|
|
BaseChatModel 实例
|
|
"""
|
|
provider_name = provider or cls._default_provider
|
|
|
|
if model_section:
|
|
config = Config.get_section(model_section)
|
|
cache_key = f"{provider_name}:{model_section}"
|
|
else:
|
|
config = config or {}
|
|
cache_key = f"{provider_name}:{hash(frozenset(config.items()))}"
|
|
|
|
if use_cache and cache_key in cls._instances:
|
|
return cls._instances[cache_key]
|
|
|
|
provider_instance = cls.get_provider(provider_name)
|
|
|
|
if provider_instance is None:
|
|
provider_instance = OpenAICompatibleProvider()
|
|
|
|
model = provider_instance.create_model(config)
|
|
|
|
if use_cache:
|
|
cls._instances[cache_key] = model
|
|
|
|
return model
|
|
|
|
@classmethod
|
|
def clear_cache(cls) -> None:
|
|
"""清空缓存"""
|
|
cls._instances.clear()
|
|
|
|
@classmethod
|
|
def list_providers(cls) -> List[str]:
|
|
"""列出所有注册的 Provider"""
|
|
return ProviderRegistry.list_names()
|
|
|
|
|
|
OpenAICompatibleProvider()
|
|
ProviderRegistry._entries["openai_compatible"] = type(
|
|
"RegistryEntry",
|
|
(),
|
|
{"instance": OpenAICompatibleProvider(), "metadata": {}}
|
|
)()
|
|
|
|
|
|
def create_chat_model(model_section: Optional[str] = None) -> BaseChatModel:
|
|
"""
|
|
创建聊天模型(兼容现有代码)
|
|
|
|
这是现有 llm_factory.create_chat_model 的替代实现,
|
|
使用新的 Provider 架构但保持接口兼容
|
|
"""
|
|
return LLMFactory.create(model_section=model_section)
|