Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
model_load_balancing_service.py576 linesDownload Raw Back to services
1import datetime2import json3import logging4from json import JSONDecodeError5from typing import Optional6 7from constants import HIDDEN_VALUE8from core.entities.provider_configuration import ProviderConfiguration9from core.helper import encrypter10from core.helper.model_provider_cache import ProviderCredentialsCache, ProviderCredentialsCacheType11from core.model_manager import LBModelManager12from core.model_runtime.entities.model_entities import ModelType13from core.model_runtime.entities.provider_entities import (14    ModelCredentialSchema,15    ProviderCredentialSchema,16)17from core.model_runtime.model_providers import model_provider_factory18from core.provider_manager import ProviderManager19from extensions.ext_database import db20from models.provider import LoadBalancingModelConfig21 22logger = logging.getLogger(__name__)23 24 25class ModelLoadBalancingService:26    def __init__(self) -> None:27        self.provider_manager = ProviderManager()28 29    def enable_model_load_balancing(self, tenant_id: str, provider: str, model: str, model_type: str) -> None:30        """31        enable model load balancing.32 33        :param tenant_id: workspace id34        :param provider: provider name35        :param model: model name36        :param model_type: model type37        :return:38        """39        # Get all provider configurations of the current workspace40        provider_configurations = self.provider_manager.get_configurations(tenant_id)41 42        # Get provider configuration43        provider_configuration = provider_configurations.get(provider)44        if not provider_configuration:45            raise ValueError(f"Provider {provider} does not exist.")46 47        # Enable model load balancing48        provider_configuration.enable_model_load_balancing(model=model, model_type=ModelType.value_of(model_type))49 50    def disable_model_load_balancing(self, tenant_id: str, provider: str, model: str, model_type: str) -> None:51        """52        disable model load balancing.53 54        :param tenant_id: workspace id55        :param provider: provider name56        :param model: model name57        :param model_type: model type58        :return:59        """60        # Get all provider configurations of the current workspace61        provider_configurations = self.provider_manager.get_configurations(tenant_id)62 63        # Get provider configuration64        provider_configuration = provider_configurations.get(provider)65        if not provider_configuration:66            raise ValueError(f"Provider {provider} does not exist.")67 68        # disable model load balancing69        provider_configuration.disable_model_load_balancing(model=model, model_type=ModelType.value_of(model_type))70 71    def get_load_balancing_configs(72        self, tenant_id: str, provider: str, model: str, model_type: str73    ) -> tuple[bool, list[dict]]:74        """75        Get load balancing configurations.76        :param tenant_id: workspace id77        :param provider: provider name78        :param model: model name79        :param model_type: model type80        :return:81        """82        # Get all provider configurations of the current workspace83        provider_configurations = self.provider_manager.get_configurations(tenant_id)84 85        # Get provider configuration86        provider_configuration = provider_configurations.get(provider)87        if not provider_configuration:88            raise ValueError(f"Provider {provider} does not exist.")89 90        # Convert model type to ModelType91        model_type = ModelType.value_of(model_type)92 93        # Get provider model setting94        provider_model_setting = provider_configuration.get_provider_model_setting(95            model_type=model_type,96            model=model,97        )98 99        is_load_balancing_enabled = False100        if provider_model_setting and provider_model_setting.load_balancing_enabled:101            is_load_balancing_enabled = True102 103        # Get load balancing configurations104        load_balancing_configs = (105            db.session.query(LoadBalancingModelConfig)106            .filter(107                LoadBalancingModelConfig.tenant_id == tenant_id,108                LoadBalancingModelConfig.provider_name == provider_configuration.provider.provider,109                LoadBalancingModelConfig.model_type == model_type.to_origin_model_type(),110                LoadBalancingModelConfig.model_name == model,111            )112            .order_by(LoadBalancingModelConfig.created_at)113            .all()114        )115 116        if provider_configuration.custom_configuration.provider:117            # check if the inherit configuration exists,118            # inherit is represented for the provider or model custom credentials119            inherit_config_exists = False120            for load_balancing_config in load_balancing_configs:121                if load_balancing_config.name == "__inherit__":122                    inherit_config_exists = True123                    break124 125            if not inherit_config_exists:126                # Initialize the inherit configuration127                inherit_config = self._init_inherit_config(tenant_id, provider, model, model_type)128 129                # prepend the inherit configuration130                load_balancing_configs.insert(0, inherit_config)131            else:132                # move the inherit configuration to the first133                for i, load_balancing_config in enumerate(load_balancing_configs[:]):134                    if load_balancing_config.name == "__inherit__":135                        inherit_config = load_balancing_configs.pop(i)136                        load_balancing_configs.insert(0, inherit_config)137 138        # Get credential form schemas from model credential schema or provider credential schema139        credential_schemas = self._get_credential_schema(provider_configuration)140 141        # Get decoding rsa key and cipher for decrypting credentials142        decoding_rsa_key, decoding_cipher_rsa = encrypter.get_decrypt_decoding(tenant_id)143 144        # fetch status and ttl for each config145        datas = []146        for load_balancing_config in load_balancing_configs:147            in_cooldown, ttl = LBModelManager.get_config_in_cooldown_and_ttl(148                tenant_id=tenant_id,149                provider=provider,150                model=model,151                model_type=model_type,152                config_id=load_balancing_config.id,153            )154 155            try:156                if load_balancing_config.encrypted_config:157                    credentials = json.loads(load_balancing_config.encrypted_config)158                else:159                    credentials = {}160            except JSONDecodeError:161                credentials = {}162 163            # Get provider credential secret variables164            credential_secret_variables = provider_configuration.extract_secret_variables(165                credential_schemas.credential_form_schemas166            )167 168            # decrypt credentials169            for variable in credential_secret_variables:170                if variable in credentials:171                    try:172                        credentials[variable] = encrypter.decrypt_token_with_decoding(173                            credentials.get(variable), decoding_rsa_key, decoding_cipher_rsa174                        )175                    except ValueError:176                        pass177 178            # Obfuscate credentials179            credentials = provider_configuration.obfuscated_credentials(180                credentials=credentials, credential_form_schemas=credential_schemas.credential_form_schemas181            )182 183            datas.append(184                {185                    "id": load_balancing_config.id,186                    "name": load_balancing_config.name,187                    "credentials": credentials,188                    "enabled": load_balancing_config.enabled,189                    "in_cooldown": in_cooldown,190                    "ttl": ttl,191                }192            )193 194        return is_load_balancing_enabled, datas195 196    def get_load_balancing_config(197        self, tenant_id: str, provider: str, model: str, model_type: str, config_id: str198    ) -> Optional[dict]:199        """200        Get load balancing configuration.201        :param tenant_id: workspace id202        :param provider: provider name203        :param model: model name204        :param model_type: model type205        :param config_id: load balancing config id206        :return:207        """208        # Get all provider configurations of the current workspace209        provider_configurations = self.provider_manager.get_configurations(tenant_id)210 211        # Get provider configuration212        provider_configuration = provider_configurations.get(provider)213        if not provider_configuration:214            raise ValueError(f"Provider {provider} does not exist.")215 216        # Convert model type to ModelType217        model_type = ModelType.value_of(model_type)218 219        # Get load balancing configurations220        load_balancing_model_config = (221            db.session.query(LoadBalancingModelConfig)222            .filter(223                LoadBalancingModelConfig.tenant_id == tenant_id,224                LoadBalancingModelConfig.provider_name == provider_configuration.provider.provider,225                LoadBalancingModelConfig.model_type == model_type.to_origin_model_type(),226                LoadBalancingModelConfig.model_name == model,227                LoadBalancingModelConfig.id == config_id,228            )229            .first()230        )231 232        if not load_balancing_model_config:233            return None234 235        try:236            if load_balancing_model_config.encrypted_config:237                credentials = json.loads(load_balancing_model_config.encrypted_config)238            else:239                credentials = {}240        except JSONDecodeError:241            credentials = {}242 243        # Get credential form schemas from model credential schema or provider credential schema244        credential_schemas = self._get_credential_schema(provider_configuration)245 246        # Obfuscate credentials247        credentials = provider_configuration.obfuscated_credentials(248            credentials=credentials, credential_form_schemas=credential_schemas.credential_form_schemas249        )250 251        return {252            "id": load_balancing_model_config.id,253            "name": load_balancing_model_config.name,254            "credentials": credentials,255            "enabled": load_balancing_model_config.enabled,256        }257 258    def _init_inherit_config(259        self, tenant_id: str, provider: str, model: str, model_type: ModelType260    ) -> LoadBalancingModelConfig:261        """262        Initialize the inherit configuration.263        :param tenant_id: workspace id264        :param provider: provider name265        :param model: model name266        :param model_type: model type267        :return:268        """269        # Initialize the inherit configuration270        inherit_config = LoadBalancingModelConfig(271            tenant_id=tenant_id,272            provider_name=provider,273            model_type=model_type.to_origin_model_type(),274            model_name=model,275            name="__inherit__",276        )277        db.session.add(inherit_config)278        db.session.commit()279 280        return inherit_config281 282    def update_load_balancing_configs(283        self, tenant_id: str, provider: str, model: str, model_type: str, configs: list[dict]284    ) -> None:285        """286        Update load balancing configurations.287        :param tenant_id: workspace id288        :param provider: provider name289        :param model: model name290        :param model_type: model type291        :param configs: load balancing configs292        :return:293        """294        # Get all provider configurations of the current workspace295        provider_configurations = self.provider_manager.get_configurations(tenant_id)296 297        # Get provider configuration298        provider_configuration = provider_configurations.get(provider)299        if not provider_configuration:300            raise ValueError(f"Provider {provider} does not exist.")301 302        # Convert model type to ModelType303        model_type = ModelType.value_of(model_type)304 305        if not isinstance(configs, list):306            raise ValueError("Invalid load balancing configs")307 308        current_load_balancing_configs = (309            db.session.query(LoadBalancingModelConfig)310            .filter(311                LoadBalancingModelConfig.tenant_id == tenant_id,312                LoadBalancingModelConfig.provider_name == provider_configuration.provider.provider,313                LoadBalancingModelConfig.model_type == model_type.to_origin_model_type(),314                LoadBalancingModelConfig.model_name == model,315            )316            .all()317        )318 319        # id as key, config as value320        current_load_balancing_configs_dict = {config.id: config for config in current_load_balancing_configs}321        updated_config_ids = set()322 323        for config in configs:324            if not isinstance(config, dict):325                raise ValueError("Invalid load balancing config")326 327            config_id = config.get("id")328            name = config.get("name")329            credentials = config.get("credentials")330            enabled = config.get("enabled")331 332            if not name:333                raise ValueError("Invalid load balancing config name")334 335            if enabled is None:336                raise ValueError("Invalid load balancing config enabled")337 338            # is config exists339            if config_id:340                config_id = str(config_id)341 342                if config_id not in current_load_balancing_configs_dict:343                    raise ValueError("Invalid load balancing config id: {}".format(config_id))344 345                updated_config_ids.add(config_id)346 347                load_balancing_config = current_load_balancing_configs_dict[config_id]348 349                # check duplicate name350                for current_load_balancing_config in current_load_balancing_configs:351                    if current_load_balancing_config.id != config_id and current_load_balancing_config.name == name:352                        raise ValueError("Load balancing config name {} already exists".format(name))353 354                if credentials:355                    if not isinstance(credentials, dict):356                        raise ValueError("Invalid load balancing config credentials")357 358                    # validate custom provider config359                    credentials = self._custom_credentials_validate(360                        tenant_id=tenant_id,361                        provider_configuration=provider_configuration,362                        model_type=model_type,363                        model=model,364                        credentials=credentials,365                        load_balancing_model_config=load_balancing_config,366                        validate=False,367                    )368 369                    # update load balancing config370                    load_balancing_config.encrypted_config = json.dumps(credentials)371 372                load_balancing_config.name = name373                load_balancing_config.enabled = enabled374                load_balancing_config.updated_at = datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None)375                db.session.commit()376 377                self._clear_credentials_cache(tenant_id, config_id)378            else:379                # create load balancing config380                if name == "__inherit__":381                    raise ValueError("Invalid load balancing config name")382 383                # check duplicate name384                for current_load_balancing_config in current_load_balancing_configs:385                    if current_load_balancing_config.name == name:386                        raise ValueError("Load balancing config name {} already exists".format(name))387 388                if not credentials:389                    raise ValueError("Invalid load balancing config credentials")390 391                if not isinstance(credentials, dict):392                    raise ValueError("Invalid load balancing config credentials")393 394                # validate custom provider config395                credentials = self._custom_credentials_validate(396                    tenant_id=tenant_id,397                    provider_configuration=provider_configuration,398                    model_type=model_type,399                    model=model,400                    credentials=credentials,401                    validate=False,402                )403 404                # create load balancing config405                load_balancing_model_config = LoadBalancingModelConfig(406                    tenant_id=tenant_id,407                    provider_name=provider_configuration.provider.provider,408                    model_type=model_type.to_origin_model_type(),409                    model_name=model,410                    name=name,411                    encrypted_config=json.dumps(credentials),412                )413 414                db.session.add(load_balancing_model_config)415                db.session.commit()416 417        # get deleted config ids418        deleted_config_ids = set(current_load_balancing_configs_dict.keys()) - updated_config_ids419        for config_id in deleted_config_ids:420            db.session.delete(current_load_balancing_configs_dict[config_id])421            db.session.commit()422 423            self._clear_credentials_cache(tenant_id, config_id)424 425    def validate_load_balancing_credentials(426        self,427        tenant_id: str,428        provider: str,429        model: str,430        model_type: str,431        credentials: dict,432        config_id: Optional[str] = None,433    ) -> None:434        """435        Validate load balancing credentials.436        :param tenant_id: workspace id437        :param provider: provider name438        :param model_type: model type439        :param model: model name440        :param credentials: credentials441        :param config_id: load balancing config id442        :return:443        """444        # Get all provider configurations of the current workspace445        provider_configurations = self.provider_manager.get_configurations(tenant_id)446 447        # Get provider configuration448        provider_configuration = provider_configurations.get(provider)449        if not provider_configuration:450            raise ValueError(f"Provider {provider} does not exist.")451 452        # Convert model type to ModelType453        model_type = ModelType.value_of(model_type)454 455        load_balancing_model_config = None456        if config_id:457            # Get load balancing config458            load_balancing_model_config = (459                db.session.query(LoadBalancingModelConfig)460                .filter(461                    LoadBalancingModelConfig.tenant_id == tenant_id,462                    LoadBalancingModelConfig.provider_name == provider,463                    LoadBalancingModelConfig.model_type == model_type.to_origin_model_type(),464                    LoadBalancingModelConfig.model_name == model,465                    LoadBalancingModelConfig.id == config_id,466                )467                .first()468            )469 470            if not load_balancing_model_config:471                raise ValueError(f"Load balancing config {config_id} does not exist.")472 473        # Validate custom provider config474        self._custom_credentials_validate(475            tenant_id=tenant_id,476            provider_configuration=provider_configuration,477            model_type=model_type,478            model=model,479            credentials=credentials,480            load_balancing_model_config=load_balancing_model_config,481        )482 483    def _custom_credentials_validate(484        self,485        tenant_id: str,486        provider_configuration: ProviderConfiguration,487        model_type: ModelType,488        model: str,489        credentials: dict,490        load_balancing_model_config: Optional[LoadBalancingModelConfig] = None,491        validate: bool = True,492    ) -> dict:493        """494        Validate custom credentials.495        :param tenant_id: workspace id496        :param provider_configuration: provider configuration497        :param model_type: model type498        :param model: model name499        :param credentials: credentials500        :param load_balancing_model_config: load balancing model config501        :param validate: validate credentials502        :return:503        """504        # Get credential form schemas from model credential schema or provider credential schema505        credential_schemas = self._get_credential_schema(provider_configuration)506 507        # Get provider credential secret variables508        provider_credential_secret_variables = provider_configuration.extract_secret_variables(509            credential_schemas.credential_form_schemas510        )511 512        if load_balancing_model_config:513            try:514                # fix origin data515                if load_balancing_model_config.encrypted_config:516                    original_credentials = json.loads(load_balancing_model_config.encrypted_config)517                else:518                    original_credentials = {}519            except JSONDecodeError:520                original_credentials = {}521 522            # encrypt credentials523            for key, value in credentials.items():524                if key in provider_credential_secret_variables:525                    # if send [__HIDDEN__] in secret input, it will be same as original value526                    if value == HIDDEN_VALUE and key in original_credentials:527                        credentials[key] = encrypter.decrypt_token(tenant_id, original_credentials[key])528 529        if validate:530            if isinstance(credential_schemas, ModelCredentialSchema):531                credentials = model_provider_factory.model_credentials_validate(532                    provider=provider_configuration.provider.provider,533                    model_type=model_type,534                    model=model,535                    credentials=credentials,536                )537            else:538                credentials = model_provider_factory.provider_credentials_validate(539                    provider=provider_configuration.provider.provider, credentials=credentials540                )541 542        for key, value in credentials.items():543            if key in provider_credential_secret_variables:544                credentials[key] = encrypter.encrypt_token(tenant_id, value)545 546        return credentials547 548    def _get_credential_schema(549        self, provider_configuration: ProviderConfiguration550    ) -> ModelCredentialSchema | ProviderCredentialSchema:551        """552        Get form schemas.553        :param provider_configuration: provider configuration554        :return:555        """556        # Get credential form schemas from model credential schema or provider credential schema557        if provider_configuration.provider.model_credential_schema:558            credential_schema = provider_configuration.provider.model_credential_schema559        else:560            credential_schema = provider_configuration.provider.provider_credential_schema561 562        return credential_schema563 564    def _clear_credentials_cache(self, tenant_id: str, config_id: str) -> None:565        """566        Clear credentials cache.567        :param tenant_id: workspace id568        :param config_id: load balancing config id569        :return:570        """571        provider_model_credentials_cache = ProviderCredentialsCache(572            tenant_id=tenant_id, identity_id=config_id, cache_type=ProviderCredentialsCacheType.LOAD_BALANCING_MODEL573        )574 575        provider_model_credentials_cache.delete()576