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