Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
model_provider_factory.py280 linesDownload Raw Back to model_providers
1import logging2import os3from collections.abc import Sequence4from typing import Optional5 6from pydantic import BaseModel, ConfigDict7 8from core.helper.module_import_helper import load_single_subclass_from_source9from core.helper.position_helper import get_provider_position_map, sort_to_dict_by_position_map10from core.model_runtime.entities.model_entities import ModelType11from core.model_runtime.entities.provider_entities import ProviderConfig, ProviderEntity, SimpleProviderEntity12from core.model_runtime.model_providers.__base.model_provider import ModelProvider13from core.model_runtime.schema_validators.model_credential_schema_validator import ModelCredentialSchemaValidator14from core.model_runtime.schema_validators.provider_credential_schema_validator import ProviderCredentialSchemaValidator15 16logger = logging.getLogger(__name__)17 18 19class ModelProviderExtension(BaseModel):20    model_config = ConfigDict(arbitrary_types_allowed=True)21 22    provider_instance: ModelProvider23    name: str24    position: Optional[int] = None25 26 27class ModelProviderFactory:28    model_provider_extensions: Optional[dict[str, ModelProviderExtension]] = None29 30    def __init__(self) -> None:31        # for cache in memory32        self.get_providers()33 34    def get_providers(self) -> Sequence[ProviderEntity]:35        """36        Get all providers37        :return: list of providers38        """39        # scan all providers40        model_provider_extensions = self._get_model_provider_map()41 42        # traverse all model_provider_extensions43        providers = []44        for model_provider_extension in model_provider_extensions.values():45            # get model_provider instance46            model_provider_instance = model_provider_extension.provider_instance47 48            # get provider schema49            provider_schema = model_provider_instance.get_provider_schema()50 51            for model_type in provider_schema.supported_model_types:52                # get predefined models for given model type53                models = model_provider_instance.models(model_type)54                if models:55                    provider_schema.models.extend(models)56 57            providers.append(provider_schema)58 59        # return providers60        return providers61 62    def provider_credentials_validate(self, *, provider: str, credentials: dict) -> dict:63        """64        Validate provider credentials65 66        :param provider: provider name67        :param credentials: provider credentials, credentials form defined in `provider_credential_schema`.68        :return:69        """70        # get the provider instance71        model_provider_instance = self.get_provider_instance(provider)72 73        # get provider schema74        provider_schema = model_provider_instance.get_provider_schema()75 76        # get provider_credential_schema and validate credentials according to the rules77        provider_credential_schema = provider_schema.provider_credential_schema78 79        if not provider_credential_schema:80            raise ValueError(f"Provider {provider} does not have provider_credential_schema")81 82        # validate provider credential schema83        validator = ProviderCredentialSchemaValidator(provider_credential_schema)84        filtered_credentials = validator.validate_and_filter(credentials)85 86        # validate the credentials, raise exception if validation failed87        model_provider_instance.validate_provider_credentials(filtered_credentials)88 89        return filtered_credentials90 91    def model_credentials_validate(92        self, *, provider: str, model_type: ModelType, model: str, credentials: dict93    ) -> dict:94        """95        Validate model credentials96 97        :param provider: provider name98        :param model_type: model type99        :param model: model name100        :param credentials: model credentials, credentials form defined in `model_credential_schema`.101        :return:102        """103        # get the provider instance104        model_provider_instance = self.get_provider_instance(provider)105 106        # get provider schema107        provider_schema = model_provider_instance.get_provider_schema()108 109        # get model_credential_schema and validate credentials according to the rules110        model_credential_schema = provider_schema.model_credential_schema111 112        if not model_credential_schema:113            raise ValueError(f"Provider {provider} does not have model_credential_schema")114 115        # validate model credential schema116        validator = ModelCredentialSchemaValidator(model_type, model_credential_schema)117        filtered_credentials = validator.validate_and_filter(credentials)118 119        # get model instance of the model type120        model_instance = model_provider_instance.get_model_instance(model_type)121 122        # call validate_credentials method of model type to validate credentials, raise exception if validation failed123        model_instance.validate_credentials(model, filtered_credentials)124 125        return filtered_credentials126 127    def get_models(128        self,129        *,130        provider: Optional[str] = None,131        model_type: Optional[ModelType] = None,132        provider_configs: Optional[list[ProviderConfig]] = None,133    ) -> list[SimpleProviderEntity]:134        """135        Get all models for given model type136 137        :param provider: provider name138        :param model_type: model type139        :param provider_configs: list of provider configs140        :return: list of models141        """142        provider_configs = provider_configs or []143 144        # scan all providers145        model_provider_extensions = self._get_model_provider_map()146 147        # convert provider_configs to dict148        provider_credentials_dict = {}149        for provider_config in provider_configs:150            provider_credentials_dict[provider_config.provider] = provider_config.credentials151 152        # traverse all model_provider_extensions153        providers = []154        for name, model_provider_extension in model_provider_extensions.items():155            # filter by provider if provider is present156            if provider and name != provider:157                continue158 159            # get model_provider instance160            model_provider_instance = model_provider_extension.provider_instance161 162            # get provider schema163            provider_schema = model_provider_instance.get_provider_schema()164 165            model_types = provider_schema.supported_model_types166            if model_type:167                if model_type not in model_types:168                    continue169 170                model_types = [model_type]171 172            all_model_type_models = []173            for model_type in model_types:174                # get predefined models for given model type175                models = model_provider_instance.models(176                    model_type=model_type,177                )178 179                all_model_type_models.extend(models)180 181            simple_provider_schema = provider_schema.to_simple_provider()182            simple_provider_schema.models.extend(all_model_type_models)183 184            providers.append(simple_provider_schema)185 186        return providers187 188    def get_provider_instance(self, provider: str) -> ModelProvider:189        """190        Get provider instance by provider name191        :param provider: provider name192        :return: provider instance193        """194        # scan all providers195        model_provider_extensions = self._get_model_provider_map()196 197        # get the provider extension198        model_provider_extension = model_provider_extensions.get(provider)199        if not model_provider_extension:200            raise Exception(f"Invalid provider: {provider}")201 202        # get the provider instance203        model_provider_instance = model_provider_extension.provider_instance204 205        return model_provider_instance206 207    def _get_model_provider_map(self) -> dict[str, ModelProviderExtension]:208        """209        Retrieves the model provider map.210 211        This method retrieves the model provider map, which is a dictionary containing the model provider names as keys212        and instances of `ModelProviderExtension` as values. The model provider map is used to store information about213        available model providers.214 215        Returns:216            A dictionary containing the model provider map.217 218        Raises:219            None.220        """221        if self.model_provider_extensions:222            return self.model_provider_extensions223 224        # get the path of current classes225        current_path = os.path.abspath(__file__)226        model_providers_path = os.path.dirname(current_path)227 228        # get all folders path under model_providers_path that do not start with __229        model_provider_dir_paths = [230            os.path.join(model_providers_path, model_provider_dir)231            for model_provider_dir in os.listdir(model_providers_path)232            if not model_provider_dir.startswith("__")233            and os.path.isdir(os.path.join(model_providers_path, model_provider_dir))234        ]235 236        # get _position.yaml file path237        position_map = get_provider_position_map(model_providers_path)238 239        # traverse all model_provider_dir_paths240        model_providers: list[ModelProviderExtension] = []241        for model_provider_dir_path in model_provider_dir_paths:242            # get model_provider dir name243            model_provider_name = os.path.basename(model_provider_dir_path)244 245            file_names = os.listdir(model_provider_dir_path)246 247            if (model_provider_name + ".py") not in file_names:248                logger.warning(f"Missing {model_provider_name}.py file in {model_provider_dir_path}, Skip.")249                continue250 251            # Dynamic loading {model_provider_name}.py file and find the subclass of ModelProvider252            py_path = os.path.join(model_provider_dir_path, model_provider_name + ".py")253            model_provider_class = load_single_subclass_from_source(254                module_name=f"core.model_runtime.model_providers.{model_provider_name}.{model_provider_name}",255                script_path=py_path,256                parent_type=ModelProvider,257            )258 259            if not model_provider_class:260                logger.warning(f"Missing Model Provider Class that extends ModelProvider in {py_path}, Skip.")261                continue262 263            if f"{model_provider_name}.yaml" not in file_names:264                logger.warning(f"Missing {model_provider_name}.yaml file in {model_provider_dir_path}, Skip.")265                continue266 267            model_providers.append(268                ModelProviderExtension(269                    name=model_provider_name,270                    provider_instance=model_provider_class(),271                    position=position_map.get(model_provider_name),272                )273            )274 275        sorted_extensions = sort_to_dict_by_position_map(position_map, model_providers, lambda x: x.name)276 277        self.model_provider_extensions = sorted_extensions278 279        return sorted_extensions280