x
This commit is contained in:
@@ -0,0 +1,218 @@
|
||||
"""
|
||||
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)
|
||||
Reference in New Issue
Block a user