ForcePilot/src/agents/common/models.py

52 lines
1.5 KiB
Python
Raw Normal View History

import os
import traceback
2025-04-02 13:00:25 +08:00
from langchain.chat_models import BaseChatModel
from pydantic import SecretStr
2025-03-25 05:40:07 +08:00
from src import config
from src.utils import get_docker_safe_url
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:
"""
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-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
api_key = os.getenv(env_var, env_var)
2025-10-22 11:51:32 +08:00
base_url = get_docker_safe_url(model_info.base_url)
if provider in ["deepseek", "dashscope"]:
from langchain_deepseek import ChatDeepSeek
return ChatDeepSeek(
model=model,
api_key=SecretStr(api_key),
base_url=base_url,
api_base=base_url,
stream_usage=True,
)
else:
try: # 其他模型默认使用OpenAIBase, like openai, zhipuai
from langchain_openai import ChatOpenAI
return ChatOpenAI(
model=model,
api_key=SecretStr(api_key),
base_url=base_url,
stream_usage=True,
)
except Exception as e:
raise ValueError(f"Model provider {provider} load failed, {e} \n {traceback.format_exc()}")