2025-07-02 02:38:36 +08:00
|
|
|
|
import os
|
|
|
|
|
|
import traceback
|
2025-04-02 13:00:25 +08:00
|
|
|
|
|
2025-12-17 22:27:46 +08:00
|
|
|
|
from langchain.chat_models import BaseChatModel, init_chat_model
|
2025-07-02 02:38:36 +08:00
|
|
|
|
from pydantic import SecretStr
|
2025-03-25 05:40:07 +08:00
|
|
|
|
|
2025-09-01 22:37:03 +08:00
|
|
|
|
from src import config
|
|
|
|
|
|
from src.utils import get_docker_safe_url
|
2025-12-17 22:27:46 +08:00
|
|
|
|
from src.utils.logging_config import logger
|
2025-03-29 17:33:09 +08:00
|
|
|
|
|
|
|
|
|
|
|
2025-04-02 00:00:04 +08:00
|
|
|
|
def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel:
|
2025-07-02 02:38:36 +08:00
|
|
|
|
"""
|
|
|
|
|
|
Load a chat model from a fully specified name.
|
2025-03-29 17:33:09 +08:00
|
|
|
|
"""
|
|
|
|
|
|
provider, model = fully_specified_name.split("/", maxsplit=1)
|
2025-04-02 00:00:04 +08:00
|
|
|
|
|
2025-10-22 11:51:32 +08:00
|
|
|
|
assert provider != "custom", "[弃用] 自定义模型已移除,请在 src/config/static/models.py 中配置"
|
2025-07-02 02:38:36 +08:00
|
|
|
|
|
2025-10-22 11:51:32 +08:00
|
|
|
|
model_info = config.model_names.get(provider)
|
|
|
|
|
|
if not model_info:
|
|
|
|
|
|
raise ValueError(f"Unknown model provider: {provider}")
|
|
|
|
|
|
|
|
|
|
|
|
env_var = model_info.env
|
2025-10-10 14:59:12 +08:00
|
|
|
|
|
2025-11-12 19:53:56 +08:00
|
|
|
|
api_key = os.getenv(env_var) or env_var
|
2025-10-10 14:59:12 +08:00
|
|
|
|
|
2025-10-22 11:51:32 +08:00
|
|
|
|
base_url = get_docker_safe_url(model_info.base_url)
|
2025-07-02 02:38:36 +08:00
|
|
|
|
|
2025-12-17 22:27:46 +08:00
|
|
|
|
if provider in ["openai", "deepseek"]:
|
|
|
|
|
|
model_spec = f"{provider}:{model}"
|
|
|
|
|
|
logger.debug(f"[offical] Loading model {model_spec} with kwargs {kwargs}")
|
|
|
|
|
|
return init_chat_model(model_spec, **kwargs)
|
|
|
|
|
|
|
|
|
|
|
|
elif provider in ["dashscope"]:
|
2025-07-02 02:38:36 +08:00
|
|
|
|
from langchain_deepseek import ChatDeepSeek
|
2025-09-01 22:37:03 +08:00
|
|
|
|
|
2025-07-02 02:38:36 +08:00
|
|
|
|
return ChatDeepSeek(
|
|
|
|
|
|
model=model,
|
|
|
|
|
|
api_key=SecretStr(api_key),
|
|
|
|
|
|
base_url=base_url,
|
|
|
|
|
|
api_base=base_url,
|
2025-10-06 21:07:33 +08:00
|
|
|
|
stream_usage=True,
|
2025-07-02 02:38:36 +08:00
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
else:
|
|
|
|
|
|
try: # 其他模型,默认使用OpenAIBase, like openai, zhipuai
|
|
|
|
|
|
from langchain_openai import ChatOpenAI
|
2025-09-01 22:37:03 +08:00
|
|
|
|
|
2025-07-02 02:38:36 +08:00
|
|
|
|
return ChatOpenAI(
|
|
|
|
|
|
model=model,
|
|
|
|
|
|
api_key=SecretStr(api_key),
|
|
|
|
|
|
base_url=base_url,
|
2025-10-06 21:07:33 +08:00
|
|
|
|
stream_usage=True,
|
2025-07-02 02:38:36 +08:00
|
|
|
|
)
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
raise ValueError(f"Model provider {provider} load failed, {e} \n {traceback.format_exc()}")
|