65 lines
2.0 KiB
Python
65 lines
2.0 KiB
Python
"""LLM 模型枚举与实例获取。
|
||
|
||
Usage:
|
||
from app.ai.models import LLM
|
||
|
||
llm = LLM.DOUBAO_SEED_LITE.create(temperature=0)
|
||
llm = LLM.CLAUDE_OPUS.create(temperature=0)
|
||
"""
|
||
|
||
from enum import Enum
|
||
from typing import Callable
|
||
|
||
from langchain_anthropic import ChatAnthropic
|
||
from langchain_core.language_models import BaseChatModel
|
||
from langchain_openai import ChatOpenAI
|
||
|
||
from app.config import settings
|
||
|
||
ConfigGetter = Callable[[], str]
|
||
|
||
# 供应商连接配置 = (api_key函数, base_url函数)
|
||
_VOLCENGINE = (
|
||
lambda: settings.volcengine_api_key,
|
||
lambda: settings.volcengine_base_url,
|
||
)
|
||
_ANTHROPIC = (
|
||
lambda: settings.anthropic_api_key,
|
||
lambda: settings.anthropic_base_url,
|
||
)
|
||
|
||
|
||
class LLM(Enum):
|
||
"""所有可用模型,每个枚举值 = (模型名, 封装类, api_key函数, base_url函数)。"""
|
||
|
||
# 火山引擎(OpenAI 兼容)
|
||
DOUBAO_PRO_32K = ("doubao-1-5-pro-32k-250115", ChatOpenAI, *_VOLCENGINE)
|
||
DOUBAO_LITE_32K = ("doubao-1-5-lite-32k-250115", ChatOpenAI, *_VOLCENGINE)
|
||
DOUBAO_SEED_LITE = ("doubao-seed-2-0-lite-260215", ChatOpenAI, *_VOLCENGINE)
|
||
DOUBAO_SEED_PRO = ("doubao-seed-2-0-pro-260215", ChatOpenAI, *_VOLCENGINE)
|
||
DEEPSEEK_V4_FLASH = ("deepseek-v4-flash-260425", ChatOpenAI, *_VOLCENGINE)
|
||
|
||
# Claude(Anthropic 风格)
|
||
CLAUDE_OPUS = ("claude-opus-4-6", ChatAnthropic, *_ANTHROPIC)
|
||
|
||
def __init__(
|
||
self,
|
||
model_name: str,
|
||
model_class: type[BaseChatModel],
|
||
api_key_getter: ConfigGetter,
|
||
base_url_getter: ConfigGetter,
|
||
) -> None:
|
||
self.model_name = model_name
|
||
self._model_class = model_class
|
||
self._api_key_getter = api_key_getter
|
||
self._base_url_getter = base_url_getter
|
||
|
||
def create(self, **kwargs) -> BaseChatModel:
|
||
"""创建模型实例,其他参数透传给 LangChain 模型类。"""
|
||
return self._model_class(
|
||
model=self.model_name,
|
||
api_key=self._api_key_getter(),
|
||
base_url=self._base_url_getter(),
|
||
**kwargs,
|
||
)
|