Underground-Digital/Workflow-Engine
0
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 