Files
more_dots/core/providers.py
T

219 lines
6.0 KiB
Python
Raw Normal View History

2026-03-11 23:40:39 +08:00
"""
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)