diff --git a/apps/models_provider/constants/model_provider_constants.py b/apps/models_provider/constants/model_provider_constants.py index 1117204a1a0..a48d93a7dc7 100644 --- a/apps/models_provider/constants/model_provider_constants.py +++ b/apps/models_provider/constants/model_provider_constants.py @@ -1,52 +1,71 @@ # coding=utf-8 -from enum import Enum +import importlib +import threading +from typing import Iterator -from models_provider.impl.aliyun_bai_lian_model_provider.aliyun_bai_lian_model_provider import ( - AliyunBaiLianModelProvider, -) -from models_provider.impl.anthropic_model_provider.anthropic_model_provider import AnthropicModelProvider -from models_provider.impl.aws_bedrock_model_provider.aws_bedrock_model_provider import BedrockModelProvider -from models_provider.impl.azure_model_provider.azure_model_provider import AzureModelProvider -from models_provider.impl.deepseek_model_provider.deepseek_model_provider import DeepSeekModelProvider -from models_provider.impl.docker_ai_model_provider.docker_ai_model_provider import DockerModelProvider -from models_provider.impl.gemini_model_provider.gemini_model_provider import GeminiModelProvider -from models_provider.impl.kimi_model_provider.kimi_model_provider import KimiModelProvider -from models_provider.impl.local_model_provider.local_model_provider import LocalModelProvider -from models_provider.impl.minimax_model_provider.minimax_model_provider import MiniMaxModelProvider -from models_provider.impl.ollama_model_provider.ollama_model_provider import OllamaModelProvider -from models_provider.impl.openai_model_provider.openai_model_provider import OpenAIModelProvider -from models_provider.impl.regolo_model_provider.regolo_model_provider import RegoloModelProvider -from models_provider.impl.siliconCloud_model_provider.siliconCloud_model_provider import SiliconCloudModelProvider -from models_provider.impl.tencent_model_provider.tencent_model_provider import TencentModelProvider -from models_provider.impl.vllm_model_provider.vllm_model_provider import VllmModelProvider -from models_provider.impl.volcanic_engine_model_provider.volcanic_engine_model_provider import ( - VolcanicEngineModelProvider, +# 供应商注册表:模型供应商不再用 Enum 硬编码实例,而是注册为"模块路径 + 类名"的 +# 惰性工厂。__getitem__ 首次访问某供应商时才 importlib 加载对应模块并缓存实例, +# 因此导入 models_provider 时不会连带加载 21 个供应商的重依赖(openai/bedrock/ +# gemini/azure 等),只在真正使用某个供应商时按需加载,从而加快启动、降低内存。 + + +class _ProviderRegistry: + def __init__(self): + self._factories: dict[str, tuple[str, str]] = {} + self._instances: dict[str, object] = {} + self._lock = threading.Lock() + + def register(self, name: str, provider_dir: str, class_name: str) -> "_ProviderRegistry": + self._factories[name] = (f"models_provider.impl.{provider_dir}.{provider_dir}", class_name) + return self + + def __getitem__(self, name: str): + instance = self._instances.get(name) + if instance is None: + with self._lock: + instance = self._instances.get(name) + if instance is None: + module_path, class_name = self._factories[name] + module = importlib.import_module(module_path) + instance = getattr(module, class_name)() + self._instances[name] = instance + return instance + + def __iter__(self) -> Iterator[str]: + return iter(self._factories) + + def __contains__(self, name: object) -> bool: + return name in self._factories + + def __len__(self) -> int: + return len(self._factories) + + @property + def __members__(self) -> dict: + return self._factories + + +ModelProvideConstants = ( + _ProviderRegistry() + .register("model_azure_provider", "azure_model_provider", "AzureModelProvider") + .register("model_qianfan_provider", "qianfan_model_provider", "QianfanModelProvider") + .register("model_ollama_provider", "ollama_model_provider", "OllamaModelProvider") + .register("model_openai_provider", "openai_model_provider", "OpenAIModelProvider") + .register("model_docker_ai_provider", "docker_ai_model_provider", "DockerModelProvider") + .register("model_kimi_provider", "kimi_model_provider", "KimiModelProvider") + .register("model_zhipu_provider", "zhipu_model_provider", "ZhiPuModelProvider") + .register("model_xf_provider", "xf_model_provider", "XunFeiModelProvider") + .register("model_deepseek_provider", "deepseek_model_provider", "DeepSeekModelProvider") + .register("model_gemini_provider", "gemini_model_provider", "GeminiModelProvider") + .register("model_volcanic_engine_provider", "volcanic_engine_model_provider", "VolcanicEngineModelProvider") + .register("model_tencent_provider", "tencent_model_provider", "TencentModelProvider") + .register("model_aws_bedrock_provider", "aws_bedrock_model_provider", "BedrockModelProvider") + .register("model_local_provider", "local_model_provider", "LocalModelProvider") + .register("model_xinference_provider", "xinference_model_provider", "XinferenceModelProvider") + .register("model_vllm_provider", "vllm_model_provider", "VllmModelProvider") + .register("aliyun_bai_lian_model_provider", "aliyun_bai_lian_model_provider", "AliyunBaiLianModelProvider") + .register("model_anthropic_provider", "anthropic_model_provider", "AnthropicModelProvider") + .register("model_siliconCloud_provider", "siliconCloud_model_provider", "SiliconCloudModelProvider") + .register("model_regolo_provider", "regolo_model_provider", "RegoloModelProvider") + .register("model_minimax_provider", "minimax_model_provider", "MiniMaxModelProvider") ) -from models_provider.impl.qianfan_model_provider.qianfan_model_provider import QianfanModelProvider -from models_provider.impl.xf_model_provider.xf_model_provider import XunFeiModelProvider -from models_provider.impl.xinference_model_provider.xinference_model_provider import XinferenceModelProvider -from models_provider.impl.zhipu_model_provider.zhipu_model_provider import ZhiPuModelProvider - - -class ModelProvideConstants(Enum): - model_azure_provider = AzureModelProvider() - model_qianfan_provider = QianfanModelProvider() - model_ollama_provider = OllamaModelProvider() - model_openai_provider = OpenAIModelProvider() - model_docker_ai_provider = DockerModelProvider() - model_kimi_provider = KimiModelProvider() - model_zhipu_provider = ZhiPuModelProvider() - model_xf_provider = XunFeiModelProvider() - model_deepseek_provider = DeepSeekModelProvider() - model_gemini_provider = GeminiModelProvider() - model_volcanic_engine_provider = VolcanicEngineModelProvider() - model_tencent_provider = TencentModelProvider() - model_aws_bedrock_provider = BedrockModelProvider() - model_local_provider = LocalModelProvider() - model_xinference_provider = XinferenceModelProvider() - model_vllm_provider = VllmModelProvider() - aliyun_bai_lian_model_provider = AliyunBaiLianModelProvider() - model_anthropic_provider = AnthropicModelProvider() - model_siliconCloud_provider = SiliconCloudModelProvider() - model_regolo_provider = RegoloModelProvider() - model_minimax_provider = MiniMaxModelProvider() diff --git a/apps/models_provider/impl/base_chat_open_ai.py b/apps/models_provider/impl/base_chat_open_ai.py index bc216c824a8..f35641257f8 100644 --- a/apps/models_provider/impl/base_chat_open_ai.py +++ b/apps/models_provider/impl/base_chat_open_ai.py @@ -32,6 +32,10 @@ def custom_get_token_ids(text: str): return tokenizer.encode(text) +# 复用固定线程池,避免每次 token 计数都新建/销毁线程;线程只在首次调用时创建 +_token_count_executor = ThreadPoolExecutor(max_workers=4, thread_name_prefix="maxkb-token-count") + + def _convert_delta_to_message_chunk( _dict: Mapping[str, Any], default_class: type[BaseMessageChunk] ) -> BaseMessageChunk: @@ -100,18 +104,17 @@ def get_num_tokens_from_messages( timeout: Optional[float] = 0.5, ) -> int: if self.usage_metadata is None or self.usage_metadata == {}: - with ThreadPoolExecutor(max_workers=1) as executor: - future = executor.submit(super().get_num_tokens_from_messages, messages, tools) - try: - response = future.result(timeout=timeout) - maxkb_logger.info("请求成功(未超时)") - return response - except Exception as e: - if isinstance(e, ReadTimeout): - raise # 继续抛出 - else: - tokenizer = TokenizerManage.get_tokenizer() - return sum([len(tokenizer.encode(get_buffer_string([m]))) for m in messages]) + future = _token_count_executor.submit(super().get_num_tokens_from_messages, messages, tools) + try: + response = future.result(timeout=timeout) + maxkb_logger.info("请求成功(未超时)") + return response + except Exception as e: + if isinstance(e, ReadTimeout): + raise # 继续抛出 + else: + tokenizer = TokenizerManage.get_tokenizer() + return sum([len(tokenizer.encode(get_buffer_string([m]))) for m in messages]) return self.usage_metadata.get("input_tokens", self.usage_metadata.get("prompt_tokens", 0)) diff --git a/apps/models_provider/serializers/model_serializer.py b/apps/models_provider/serializers/model_serializer.py index d528f483c3f..7b3428a558d 100644 --- a/apps/models_provider/serializers/model_serializer.py +++ b/apps/models_provider/serializers/model_serializer.py @@ -72,9 +72,7 @@ class ModelPullManage: @staticmethod def pull(model: Model, credential: Dict): try: - response = ModelProvideConstants[model.provider].value.down_model( - model.model_type, model.model_name, credential - ) + response = ModelProvideConstants[model.provider].down_model(model.model_type, model.model_name, credential) down_model_chunk = {} last_update_time = time.time() @@ -117,7 +115,7 @@ def model_to_dict(model: Model): "status": model.status, "meta": model.meta, "credential": ModelProvideConstants[model.provider] - .value.get_model_credential(model.model_type, model.model_name) + .get_model_credential(model.model_type, model.model_name) .encryption_dict(credential), "workspace_id": model.workspace_id, "nick_name": model.user.nick_name if model.user else "", @@ -272,8 +270,8 @@ def is_valid(self, model=None, raise_exception=False): model_type = self.data.get("model_type") model_name = self.data.get("model_name") credential = self.data.get("credential") - provider_handler = ModelProvideConstants[provider].value - model_credential = ModelProvideConstants[provider].value.get_model_credential(model_type, model_name) + provider_handler = ModelProvideConstants[provider] + model_credential = ModelProvideConstants[provider].get_model_credential(model_type, model_name) source_model_credential = json.loads(rsa_long_decrypt(model.credential)) source_encryption_model_credential = model_credential.encryption_dict(source_model_credential) if credential is not None: @@ -303,7 +301,7 @@ def is_valid(self, *, raise_exception=False): 500, _("base model【{model_name}】already exists").format(model_name=self.data.get("name")) ) default_params = {item["field"]: item["default_value"] for item in self.data.get("model_params_form")} - ModelProvideConstants[self.data.get("provider")].value.is_valid_credential( + ModelProvideConstants[self.data.get("provider")].is_valid_credential( self.data.get("model_type"), self.data.get("model_name"), self.data.get("credential"), diff --git a/apps/models_provider/tools.py b/apps/models_provider/tools.py index 4ec31303a03..58a6df7572d 100644 --- a/apps/models_provider/tools.py +++ b/apps/models_provider/tools.py @@ -59,7 +59,7 @@ def get_provider(provider): @param provider: 供应商字符串 @return: 供应商实例 """ - return ModelProvideConstants[provider].value + return ModelProvideConstants[provider] def get_model_list(provider, model_type): diff --git a/apps/models_provider/views/provide.py b/apps/models_provider/views/provide.py index 7e7ce5aadcb..2f41a54755d 100644 --- a/apps/models_provider/views/provide.py +++ b/apps/models_provider/views/provide.py @@ -33,19 +33,16 @@ def get(self, request: Request): len( [ item - for item in ModelProvideConstants[key].value.get_model_type_list() + for item in ModelProvideConstants[key].get_model_type_list() if item["value"] == model_type ] ) > 0 ): - providers.append(ModelProvideConstants[key].value.get_model_provide_info().to_dict()) + providers.append(ModelProvideConstants[key].get_model_provide_info().to_dict()) return result.success(providers) return result.success( - [ - ModelProvideConstants[key].value.get_model_provide_info().to_dict() - for key in ModelProvideConstants.__members__ - ] + [ModelProvideConstants[key].get_model_provide_info().to_dict() for key in ModelProvideConstants.__members__] ) class ModelTypeList(APIView): @@ -62,7 +59,7 @@ class ModelTypeList(APIView): ) def get(self, request: Request): provider = request.query_params.get("provider") - return result.success(ModelProvideConstants[provider].value.get_model_type_list()) + return result.success(ModelProvideConstants[provider].get_model_type_list()) class ModelList(APIView): authentication_classes = [TokenAuth] @@ -80,7 +77,7 @@ def get(self, request: Request): provider = request.query_params.get("provider") model_type = request.query_params.get("model_type") - return result.success(ModelProvideConstants[provider].value.get_model_list(model_type)) + return result.success(ModelProvideConstants[provider].get_model_list(model_type)) class ModelParamsForm(APIView): authentication_classes = [TokenAuth] @@ -118,5 +115,5 @@ def get(self, request: Request): model_type = request.query_params.get("model_type") model_name = request.query_params.get("model_name") return result.success( - ModelProvideConstants[provider].value.get_model_credential(model_type, model_name).to_form_list() + ModelProvideConstants[provider].get_model_credential(model_type, model_name).to_form_list() )