""" 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)