+110









Yeuoly
GitHub
takatost
kurokobo
Novice Lee
zxhlyh
AkaraChen
Yi
Joel
JzoNg
twwu
Hiroshi Fujita
AkaraChen
NFish
Wu Tianwei
非法操作
Novice
Hiroki Nagai
Gen Sato
eux
huangzhuo1949
huangzhuo
lotsik
crazywoola
nite-knite
Jyong
github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
gakkiyomi
CN-P5
CN-P5
Chuehnone
yihong
Kevin9703
-LAN-
Boris Feld
mbo
mabo
Warren Chen
JzoNgKVO
jiandanfeng
zhu-an
zhaoqingyu.1075
海狸大師
Xu Song
rayshaw001
Ding Jiatong
Bowen Liang
JasonVV
le0zh
zhuxinliang
k-zaku
luckylhb90
hobo.l
jiangbo721
刘江波
Shun Miyazawa
EricPan
crazywoola
sino
Jhvcc
lowell
Boris Polonsky
Ademílson Tonato
Ademílson Tonato
IWAI, Masaharu <iwaim.sub@gmail.com>
Yueh-Po Peng
Jason
Xin Zhang
yjc980121
heyszt
Abdullah AlOsaimi
Abdullah AlOsaimi
Yingchun Lai
Hash Brown
zuodongxu
Masashi Tomooka
aplio
Obada Khalili
Nam Vu
Kei YAMAZAKI
TechnoHouse
Riddhimaan-Senapati
MaFee921
te-chan
HQidea
Joshbly
xhe
weiwenyan-dev
ex_wenyan.wei
engchina
engchina
dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
呆萌闷油瓶
Kemal
Lazy_Frog
Yi Xiao
Steven sun
steven
Kalo Chin
Katy Tao
depy
胡春东
Junjie.M
MuYu
Naoki Takashima
Summer-Gu
Fei He
ybalbert001
Yuanbo Li
douxc
liuzhenghua
Wu Jiayang
Your Name
kimjion
AugNSo
llinvokerl
liusurong.lsr
Vasu Negi
Hundredwz
Xiyuan Chen
403e2d58b9
Signed-off-by: yihong0618 <zouzou0208@gmail.com> Signed-off-by: -LAN- <laipz8200@outlook.com> Signed-off-by: xhe <xw897002528@gmail.com> Signed-off-by: dependabot[bot] <support@github.com> Co-authored-by: takatost <takatost@gmail.com> Co-authored-by: kurokobo <kuro664@gmail.com> Co-authored-by: Novice Lee <novicelee@NoviPro.local> Co-authored-by: zxhlyh <jasonapring2015@outlook.com> Co-authored-by: AkaraChen <akarachen@outlook.com> Co-authored-by: Yi <yxiaoisme@gmail.com> Co-authored-by: Joel <iamjoel007@gmail.com> Co-authored-by: JzoNg <jzongcode@gmail.com> Co-authored-by: twwu <twwu@dify.ai> Co-authored-by: Hiroshi Fujita <fujita-h@users.noreply.github.com> Co-authored-by: AkaraChen <85140972+AkaraChen@users.noreply.github.com> Co-authored-by: NFish <douxc512@gmail.com> Co-authored-by: Wu Tianwei <30284043+WTW0313@users.noreply.github.com> Co-authored-by: 非法操作 <hjlarry@163.com> Co-authored-by: Novice <857526207@qq.com> Co-authored-by: Hiroki Nagai <82458324+nagaihiroki-git@users.noreply.github.com> Co-authored-by: Gen Sato <52241300+halogen22@users.noreply.github.com> Co-authored-by: eux <euxuuu@gmail.com> Co-authored-by: huangzhuo1949 <167434202+huangzhuo1949@users.noreply.github.com> Co-authored-by: huangzhuo <huangzhuo1@xiaomi.com> Co-authored-by: lotsik <lotsik@mail.ru> Co-authored-by: crazywoola <100913391+crazywoola@users.noreply.github.com> Co-authored-by: nite-knite <nkCoding@gmail.com> Co-authored-by: Jyong <76649700+JohnJyong@users.noreply.github.com> Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> Co-authored-by: gakkiyomi <gakkiyomi@aliyun.com> Co-authored-by: CN-P5 <heibai2006@gmail.com> Co-authored-by: CN-P5 <heibai2006@qq.com> Co-authored-by: Chuehnone <1897025+chuehnone@users.noreply.github.com> Co-authored-by: yihong <zouzou0208@gmail.com> Co-authored-by: Kevin9703 <51311316+Kevin9703@users.noreply.github.com> Co-authored-by: -LAN- <laipz8200@outlook.com> Co-authored-by: Boris Feld <lothiraldan@gmail.com> Co-authored-by: mbo <himabo@gmail.com> Co-authored-by: mabo <mabo@aeyes.ai> Co-authored-by: Warren Chen <warren.chen830@gmail.com> Co-authored-by: JzoNgKVO <27049666+JzoNgKVO@users.noreply.github.com> Co-authored-by: jiandanfeng <chenjh3@wangsu.com> Co-authored-by: zhu-an <70234959+xhdd123321@users.noreply.github.com> Co-authored-by: zhaoqingyu.1075 <zhaoqingyu.1075@bytedance.com> Co-authored-by: 海狸大師 <86974027+yenslife@users.noreply.github.com> Co-authored-by: Xu Song <xusong.vip@gmail.com> Co-authored-by: rayshaw001 <396301947@163.com> Co-authored-by: Ding Jiatong <dingjiatong@gmail.com> Co-authored-by: Bowen Liang <liangbowen@gf.com.cn> Co-authored-by: JasonVV <jasonwangiii@outlook.com> Co-authored-by: le0zh <newlight@qq.com> Co-authored-by: zhuxinliang <zhuxinliang@didiglobal.com> Co-authored-by: k-zaku <zaku99@outlook.jp> Co-authored-by: luckylhb90 <luckylhb90@gmail.com> Co-authored-by: hobo.l <hobo.l@binance.com> Co-authored-by: jiangbo721 <365065261@qq.com> Co-authored-by: 刘江波 <jiangbo721@163.com> Co-authored-by: Shun Miyazawa <34241526+miya@users.noreply.github.com> Co-authored-by: EricPan <30651140+Egfly@users.noreply.github.com> Co-authored-by: crazywoola <427733928@qq.com> Co-authored-by: sino <sino2322@gmail.com> Co-authored-by: Jhvcc <37662342+Jhvcc@users.noreply.github.com> Co-authored-by: lowell <lowell.hu@zkteco.in> Co-authored-by: Boris Polonsky <BorisPolonsky@users.noreply.github.com> Co-authored-by: Ademílson Tonato <ademilsonft@outlook.com> Co-authored-by: Ademílson Tonato <ademilson.tonato@refurbed.com> Co-authored-by: IWAI, Masaharu <iwaim.sub@gmail.com> Co-authored-by: Yueh-Po Peng (Yabi) <94939112+y10ab1@users.noreply.github.com> Co-authored-by: Jason <ggbbddjm@gmail.com> Co-authored-by: Xin Zhang <sjhpzx@gmail.com> Co-authored-by: yjc980121 <3898524+yjc980121@users.noreply.github.com> Co-authored-by: heyszt <36215648+hieheihei@users.noreply.github.com> Co-authored-by: Abdullah AlOsaimi <osaimiacc@gmail.com> Co-authored-by: Abdullah AlOsaimi <189027247+osaimi@users.noreply.github.com> Co-authored-by: Yingchun Lai <laiyingchun@apache.org> Co-authored-by: Hash Brown <hi@xzd.me> Co-authored-by: zuodongxu <192560071+zuodongxu@users.noreply.github.com> Co-authored-by: Masashi Tomooka <tmokmss@users.noreply.github.com> Co-authored-by: aplio <ryo.091219@gmail.com> Co-authored-by: Obada Khalili <54270856+obadakhalili@users.noreply.github.com> Co-authored-by: Nam Vu <zuzoovn@gmail.com> Co-authored-by: Kei YAMAZAKI <1715090+kei-yamazaki@users.noreply.github.com> Co-authored-by: TechnoHouse <13776377+deephbz@users.noreply.github.com> Co-authored-by: Riddhimaan-Senapati <114703025+Riddhimaan-Senapati@users.noreply.github.com> Co-authored-by: MaFee921 <31881301+2284730142@users.noreply.github.com> Co-authored-by: te-chan <t-nakanome@sakura-is.co.jp> Co-authored-by: HQidea <HQidea@users.noreply.github.com> Co-authored-by: Joshbly <36315710+Joshbly@users.noreply.github.com> Co-authored-by: xhe <xw897002528@gmail.com> Co-authored-by: weiwenyan-dev <154779315+weiwenyan-dev@users.noreply.github.com> Co-authored-by: ex_wenyan.wei <ex_wenyan.wei@tcl.com> Co-authored-by: engchina <12236799+engchina@users.noreply.github.com> Co-authored-by: engchina <atjapan2015@gmail.com> Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com> Co-authored-by: 呆萌闷油瓶 <253605712@qq.com> Co-authored-by: Kemal <kemalmeler@outlook.com> Co-authored-by: Lazy_Frog <4590648+lazyFrogLOL@users.noreply.github.com> Co-authored-by: Yi Xiao <54782454+YIXIAO0@users.noreply.github.com> Co-authored-by: Steven sun <98230804+Tuyohai@users.noreply.github.com> Co-authored-by: steven <sunzwj@digitalchina.com> Co-authored-by: Kalo Chin <91766386+fdb02983rhy@users.noreply.github.com> Co-authored-by: Katy Tao <34019945+KatyTao@users.noreply.github.com> Co-authored-by: depy <42985524+h4ckdepy@users.noreply.github.com> Co-authored-by: 胡春东 <gycm520@gmail.com> Co-authored-by: Junjie.M <118170653@qq.com> Co-authored-by: MuYu <mr.muzea@gmail.com> Co-authored-by: Naoki Takashima <39912547+takatea@users.noreply.github.com> Co-authored-by: Summer-Gu <37869445+gubinjie@users.noreply.github.com> Co-authored-by: Fei He <droxer.he@gmail.com> Co-authored-by: ybalbert001 <120714773+ybalbert001@users.noreply.github.com> Co-authored-by: Yuanbo Li <ybalbert@amazon.com> Co-authored-by: douxc <7553076+douxc@users.noreply.github.com> Co-authored-by: liuzhenghua <1090179900@qq.com> Co-authored-by: Wu Jiayang <62842862+Wu-Jiayang@users.noreply.github.com> Co-authored-by: Your Name <you@example.com> Co-authored-by: kimjion <45935338+kimjion@users.noreply.github.com> Co-authored-by: AugNSo <song.tiankai@icloud.com> Co-authored-by: llinvokerl <38915183+llinvokerl@users.noreply.github.com> Co-authored-by: liusurong.lsr <liusurong.lsr@alibaba-inc.com> Co-authored-by: Vasu Negi <vasu-negi@users.noreply.github.com> Co-authored-by: Hundredwz <1808096180@qq.com> Co-authored-by: Xiyuan Chen <52963600+GareArc@users.noreply.github.com>
572 lines
23 KiB
Python
572 lines
23 KiB
Python
import datetime
|
|
import json
|
|
import logging
|
|
from json import JSONDecodeError
|
|
from typing import Optional, Union
|
|
|
|
from constants import HIDDEN_VALUE
|
|
from core.entities.provider_configuration import ProviderConfiguration
|
|
from core.helper import encrypter
|
|
from core.helper.model_provider_cache import ProviderCredentialsCache, ProviderCredentialsCacheType
|
|
from core.model_manager import LBModelManager
|
|
from core.model_runtime.entities.model_entities import ModelType
|
|
from core.model_runtime.entities.provider_entities import (
|
|
ModelCredentialSchema,
|
|
ProviderCredentialSchema,
|
|
)
|
|
from core.model_runtime.model_providers.model_provider_factory import ModelProviderFactory
|
|
from core.provider_manager import ProviderManager
|
|
from extensions.ext_database import db
|
|
from models.provider import LoadBalancingModelConfig
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
class ModelLoadBalancingService:
|
|
def __init__(self) -> None:
|
|
self.provider_manager = ProviderManager()
|
|
|
|
def enable_model_load_balancing(self, tenant_id: str, provider: str, model: str, model_type: str) -> None:
|
|
"""
|
|
enable model load balancing.
|
|
|
|
:param tenant_id: workspace id
|
|
:param provider: provider name
|
|
:param model: model name
|
|
:param model_type: model type
|
|
:return:
|
|
"""
|
|
# Get all provider configurations of the current workspace
|
|
provider_configurations = self.provider_manager.get_configurations(tenant_id)
|
|
|
|
# Get provider configuration
|
|
provider_configuration = provider_configurations.get(provider)
|
|
if not provider_configuration:
|
|
raise ValueError(f"Provider {provider} does not exist.")
|
|
|
|
# Enable model load balancing
|
|
provider_configuration.enable_model_load_balancing(model=model, model_type=ModelType.value_of(model_type))
|
|
|
|
def disable_model_load_balancing(self, tenant_id: str, provider: str, model: str, model_type: str) -> None:
|
|
"""
|
|
disable model load balancing.
|
|
|
|
:param tenant_id: workspace id
|
|
:param provider: provider name
|
|
:param model: model name
|
|
:param model_type: model type
|
|
:return:
|
|
"""
|
|
# Get all provider configurations of the current workspace
|
|
provider_configurations = self.provider_manager.get_configurations(tenant_id)
|
|
|
|
# Get provider configuration
|
|
provider_configuration = provider_configurations.get(provider)
|
|
if not provider_configuration:
|
|
raise ValueError(f"Provider {provider} does not exist.")
|
|
|
|
# disable model load balancing
|
|
provider_configuration.disable_model_load_balancing(model=model, model_type=ModelType.value_of(model_type))
|
|
|
|
def get_load_balancing_configs(
|
|
self, tenant_id: str, provider: str, model: str, model_type: str
|
|
) -> tuple[bool, list[dict]]:
|
|
"""
|
|
Get load balancing configurations.
|
|
:param tenant_id: workspace id
|
|
:param provider: provider name
|
|
:param model: model name
|
|
:param model_type: model type
|
|
:return:
|
|
"""
|
|
# Get all provider configurations of the current workspace
|
|
provider_configurations = self.provider_manager.get_configurations(tenant_id)
|
|
|
|
# Get provider configuration
|
|
provider_configuration = provider_configurations.get(provider)
|
|
if not provider_configuration:
|
|
raise ValueError(f"Provider {provider} does not exist.")
|
|
|
|
# Convert model type to ModelType
|
|
model_type_enum = ModelType.value_of(model_type)
|
|
|
|
# Get provider model setting
|
|
provider_model_setting = provider_configuration.get_provider_model_setting(
|
|
model_type=model_type_enum,
|
|
model=model,
|
|
)
|
|
|
|
is_load_balancing_enabled = False
|
|
if provider_model_setting and provider_model_setting.load_balancing_enabled:
|
|
is_load_balancing_enabled = True
|
|
|
|
# Get load balancing configurations
|
|
load_balancing_configs = (
|
|
db.session.query(LoadBalancingModelConfig)
|
|
.filter(
|
|
LoadBalancingModelConfig.tenant_id == tenant_id,
|
|
LoadBalancingModelConfig.provider_name == provider_configuration.provider.provider,
|
|
LoadBalancingModelConfig.model_type == model_type_enum.to_origin_model_type(),
|
|
LoadBalancingModelConfig.model_name == model,
|
|
)
|
|
.order_by(LoadBalancingModelConfig.created_at)
|
|
.all()
|
|
)
|
|
|
|
if provider_configuration.custom_configuration.provider:
|
|
# check if the inherit configuration exists,
|
|
# inherit is represented for the provider or model custom credentials
|
|
inherit_config_exists = False
|
|
for load_balancing_config in load_balancing_configs:
|
|
if load_balancing_config.name == "__inherit__":
|
|
inherit_config_exists = True
|
|
break
|
|
|
|
if not inherit_config_exists:
|
|
# Initialize the inherit configuration
|
|
inherit_config = self._init_inherit_config(tenant_id, provider, model, model_type_enum)
|
|
|
|
# prepend the inherit configuration
|
|
load_balancing_configs.insert(0, inherit_config)
|
|
else:
|
|
# move the inherit configuration to the first
|
|
for i, load_balancing_config in enumerate(load_balancing_configs[:]):
|
|
if load_balancing_config.name == "__inherit__":
|
|
inherit_config = load_balancing_configs.pop(i)
|
|
load_balancing_configs.insert(0, inherit_config)
|
|
|
|
# Get credential form schemas from model credential schema or provider credential schema
|
|
credential_schemas = self._get_credential_schema(provider_configuration)
|
|
|
|
# Get decoding rsa key and cipher for decrypting credentials
|
|
decoding_rsa_key, decoding_cipher_rsa = encrypter.get_decrypt_decoding(tenant_id)
|
|
|
|
# fetch status and ttl for each config
|
|
datas = []
|
|
for load_balancing_config in load_balancing_configs:
|
|
in_cooldown, ttl = LBModelManager.get_config_in_cooldown_and_ttl(
|
|
tenant_id=tenant_id,
|
|
provider=provider,
|
|
model=model,
|
|
model_type=model_type_enum,
|
|
config_id=load_balancing_config.id,
|
|
)
|
|
|
|
try:
|
|
if load_balancing_config.encrypted_config:
|
|
credentials = json.loads(load_balancing_config.encrypted_config)
|
|
else:
|
|
credentials = {}
|
|
except JSONDecodeError:
|
|
credentials = {}
|
|
|
|
# Get provider credential secret variables
|
|
credential_secret_variables = provider_configuration.extract_secret_variables(
|
|
credential_schemas.credential_form_schemas
|
|
)
|
|
|
|
# decrypt credentials
|
|
for variable in credential_secret_variables:
|
|
if variable in credentials:
|
|
try:
|
|
credentials[variable] = encrypter.decrypt_token_with_decoding(
|
|
credentials.get(variable), decoding_rsa_key, decoding_cipher_rsa
|
|
)
|
|
except ValueError:
|
|
pass
|
|
|
|
# Obfuscate credentials
|
|
credentials = provider_configuration.obfuscated_credentials(
|
|
credentials=credentials, credential_form_schemas=credential_schemas.credential_form_schemas
|
|
)
|
|
|
|
datas.append(
|
|
{
|
|
"id": load_balancing_config.id,
|
|
"name": load_balancing_config.name,
|
|
"credentials": credentials,
|
|
"enabled": load_balancing_config.enabled,
|
|
"in_cooldown": in_cooldown,
|
|
"ttl": ttl,
|
|
}
|
|
)
|
|
|
|
return is_load_balancing_enabled, datas
|
|
|
|
def get_load_balancing_config(
|
|
self, tenant_id: str, provider: str, model: str, model_type: str, config_id: str
|
|
) -> Optional[dict]:
|
|
"""
|
|
Get load balancing configuration.
|
|
:param tenant_id: workspace id
|
|
:param provider: provider name
|
|
:param model: model name
|
|
:param model_type: model type
|
|
:param config_id: load balancing config id
|
|
:return:
|
|
"""
|
|
# Get all provider configurations of the current workspace
|
|
provider_configurations = self.provider_manager.get_configurations(tenant_id)
|
|
|
|
# Get provider configuration
|
|
provider_configuration = provider_configurations.get(provider)
|
|
if not provider_configuration:
|
|
raise ValueError(f"Provider {provider} does not exist.")
|
|
|
|
# Convert model type to ModelType
|
|
model_type_enum = ModelType.value_of(model_type)
|
|
|
|
# Get load balancing configurations
|
|
load_balancing_model_config = (
|
|
db.session.query(LoadBalancingModelConfig)
|
|
.filter(
|
|
LoadBalancingModelConfig.tenant_id == tenant_id,
|
|
LoadBalancingModelConfig.provider_name == provider_configuration.provider.provider,
|
|
LoadBalancingModelConfig.model_type == model_type_enum.to_origin_model_type(),
|
|
LoadBalancingModelConfig.model_name == model,
|
|
LoadBalancingModelConfig.id == config_id,
|
|
)
|
|
.first()
|
|
)
|
|
|
|
if not load_balancing_model_config:
|
|
return None
|
|
|
|
try:
|
|
if load_balancing_model_config.encrypted_config:
|
|
credentials = json.loads(load_balancing_model_config.encrypted_config)
|
|
else:
|
|
credentials = {}
|
|
except JSONDecodeError:
|
|
credentials = {}
|
|
|
|
# Get credential form schemas from model credential schema or provider credential schema
|
|
credential_schemas = self._get_credential_schema(provider_configuration)
|
|
|
|
# Obfuscate credentials
|
|
credentials = provider_configuration.obfuscated_credentials(
|
|
credentials=credentials, credential_form_schemas=credential_schemas.credential_form_schemas
|
|
)
|
|
|
|
return {
|
|
"id": load_balancing_model_config.id,
|
|
"name": load_balancing_model_config.name,
|
|
"credentials": credentials,
|
|
"enabled": load_balancing_model_config.enabled,
|
|
}
|
|
|
|
def _init_inherit_config(
|
|
self, tenant_id: str, provider: str, model: str, model_type: ModelType
|
|
) -> LoadBalancingModelConfig:
|
|
"""
|
|
Initialize the inherit configuration.
|
|
:param tenant_id: workspace id
|
|
:param provider: provider name
|
|
:param model: model name
|
|
:param model_type: model type
|
|
:return:
|
|
"""
|
|
# Initialize the inherit configuration
|
|
inherit_config = LoadBalancingModelConfig(
|
|
tenant_id=tenant_id,
|
|
provider_name=provider,
|
|
model_type=model_type.to_origin_model_type(),
|
|
model_name=model,
|
|
name="__inherit__",
|
|
)
|
|
db.session.add(inherit_config)
|
|
db.session.commit()
|
|
|
|
return inherit_config
|
|
|
|
def update_load_balancing_configs(
|
|
self, tenant_id: str, provider: str, model: str, model_type: str, configs: list[dict]
|
|
) -> None:
|
|
"""
|
|
Update load balancing configurations.
|
|
:param tenant_id: workspace id
|
|
:param provider: provider name
|
|
:param model: model name
|
|
:param model_type: model type
|
|
:param configs: load balancing configs
|
|
:return:
|
|
"""
|
|
# Get all provider configurations of the current workspace
|
|
provider_configurations = self.provider_manager.get_configurations(tenant_id)
|
|
|
|
# Get provider configuration
|
|
provider_configuration = provider_configurations.get(provider)
|
|
if not provider_configuration:
|
|
raise ValueError(f"Provider {provider} does not exist.")
|
|
|
|
# Convert model type to ModelType
|
|
model_type_enum = ModelType.value_of(model_type)
|
|
|
|
if not isinstance(configs, list):
|
|
raise ValueError("Invalid load balancing configs")
|
|
|
|
current_load_balancing_configs = (
|
|
db.session.query(LoadBalancingModelConfig)
|
|
.filter(
|
|
LoadBalancingModelConfig.tenant_id == tenant_id,
|
|
LoadBalancingModelConfig.provider_name == provider_configuration.provider.provider,
|
|
LoadBalancingModelConfig.model_type == model_type_enum.to_origin_model_type(),
|
|
LoadBalancingModelConfig.model_name == model,
|
|
)
|
|
.all()
|
|
)
|
|
|
|
# id as key, config as value
|
|
current_load_balancing_configs_dict = {config.id: config for config in current_load_balancing_configs}
|
|
updated_config_ids = set()
|
|
|
|
for config in configs:
|
|
if not isinstance(config, dict):
|
|
raise ValueError("Invalid load balancing config")
|
|
|
|
config_id = config.get("id")
|
|
name = config.get("name")
|
|
credentials = config.get("credentials")
|
|
enabled = config.get("enabled")
|
|
|
|
if not name:
|
|
raise ValueError("Invalid load balancing config name")
|
|
|
|
if enabled is None:
|
|
raise ValueError("Invalid load balancing config enabled")
|
|
|
|
# is config exists
|
|
if config_id:
|
|
config_id = str(config_id)
|
|
|
|
if config_id not in current_load_balancing_configs_dict:
|
|
raise ValueError("Invalid load balancing config id: {}".format(config_id))
|
|
|
|
updated_config_ids.add(config_id)
|
|
|
|
load_balancing_config = current_load_balancing_configs_dict[config_id]
|
|
|
|
# check duplicate name
|
|
for current_load_balancing_config in current_load_balancing_configs:
|
|
if current_load_balancing_config.id != config_id and current_load_balancing_config.name == name:
|
|
raise ValueError("Load balancing config name {} already exists".format(name))
|
|
|
|
if credentials:
|
|
if not isinstance(credentials, dict):
|
|
raise ValueError("Invalid load balancing config credentials")
|
|
|
|
# validate custom provider config
|
|
credentials = self._custom_credentials_validate(
|
|
tenant_id=tenant_id,
|
|
provider_configuration=provider_configuration,
|
|
model_type=model_type_enum,
|
|
model=model,
|
|
credentials=credentials,
|
|
load_balancing_model_config=load_balancing_config,
|
|
validate=False,
|
|
)
|
|
|
|
# update load balancing config
|
|
load_balancing_config.encrypted_config = json.dumps(credentials)
|
|
|
|
load_balancing_config.name = name
|
|
load_balancing_config.enabled = enabled
|
|
load_balancing_config.updated_at = datetime.datetime.now(datetime.UTC).replace(tzinfo=None)
|
|
db.session.commit()
|
|
|
|
self._clear_credentials_cache(tenant_id, config_id)
|
|
else:
|
|
# create load balancing config
|
|
if name == "__inherit__":
|
|
raise ValueError("Invalid load balancing config name")
|
|
|
|
# check duplicate name
|
|
for current_load_balancing_config in current_load_balancing_configs:
|
|
if current_load_balancing_config.name == name:
|
|
raise ValueError("Load balancing config name {} already exists".format(name))
|
|
|
|
if not credentials:
|
|
raise ValueError("Invalid load balancing config credentials")
|
|
|
|
if not isinstance(credentials, dict):
|
|
raise ValueError("Invalid load balancing config credentials")
|
|
|
|
# validate custom provider config
|
|
credentials = self._custom_credentials_validate(
|
|
tenant_id=tenant_id,
|
|
provider_configuration=provider_configuration,
|
|
model_type=model_type_enum,
|
|
model=model,
|
|
credentials=credentials,
|
|
validate=False,
|
|
)
|
|
|
|
# create load balancing config
|
|
load_balancing_model_config = LoadBalancingModelConfig(
|
|
tenant_id=tenant_id,
|
|
provider_name=provider_configuration.provider.provider,
|
|
model_type=model_type_enum.to_origin_model_type(),
|
|
model_name=model,
|
|
name=name,
|
|
encrypted_config=json.dumps(credentials),
|
|
)
|
|
|
|
db.session.add(load_balancing_model_config)
|
|
db.session.commit()
|
|
|
|
# get deleted config ids
|
|
deleted_config_ids = set(current_load_balancing_configs_dict.keys()) - updated_config_ids
|
|
for config_id in deleted_config_ids:
|
|
db.session.delete(current_load_balancing_configs_dict[config_id])
|
|
db.session.commit()
|
|
|
|
self._clear_credentials_cache(tenant_id, config_id)
|
|
|
|
def validate_load_balancing_credentials(
|
|
self,
|
|
tenant_id: str,
|
|
provider: str,
|
|
model: str,
|
|
model_type: str,
|
|
credentials: dict,
|
|
config_id: Optional[str] = None,
|
|
) -> None:
|
|
"""
|
|
Validate load balancing credentials.
|
|
:param tenant_id: workspace id
|
|
:param provider: provider name
|
|
:param model_type: model type
|
|
:param model: model name
|
|
:param credentials: credentials
|
|
:param config_id: load balancing config id
|
|
:return:
|
|
"""
|
|
# Get all provider configurations of the current workspace
|
|
provider_configurations = self.provider_manager.get_configurations(tenant_id)
|
|
|
|
# Get provider configuration
|
|
provider_configuration = provider_configurations.get(provider)
|
|
if not provider_configuration:
|
|
raise ValueError(f"Provider {provider} does not exist.")
|
|
|
|
# Convert model type to ModelType
|
|
model_type_enum = ModelType.value_of(model_type)
|
|
|
|
load_balancing_model_config = None
|
|
if config_id:
|
|
# Get load balancing config
|
|
load_balancing_model_config = (
|
|
db.session.query(LoadBalancingModelConfig)
|
|
.filter(
|
|
LoadBalancingModelConfig.tenant_id == tenant_id,
|
|
LoadBalancingModelConfig.provider_name == provider,
|
|
LoadBalancingModelConfig.model_type == model_type_enum.to_origin_model_type(),
|
|
LoadBalancingModelConfig.model_name == model,
|
|
LoadBalancingModelConfig.id == config_id,
|
|
)
|
|
.first()
|
|
)
|
|
|
|
if not load_balancing_model_config:
|
|
raise ValueError(f"Load balancing config {config_id} does not exist.")
|
|
|
|
# Validate custom provider config
|
|
self._custom_credentials_validate(
|
|
tenant_id=tenant_id,
|
|
provider_configuration=provider_configuration,
|
|
model_type=model_type_enum,
|
|
model=model,
|
|
credentials=credentials,
|
|
load_balancing_model_config=load_balancing_model_config,
|
|
)
|
|
|
|
def _custom_credentials_validate(
|
|
self,
|
|
tenant_id: str,
|
|
provider_configuration: ProviderConfiguration,
|
|
model_type: ModelType,
|
|
model: str,
|
|
credentials: dict,
|
|
load_balancing_model_config: Optional[LoadBalancingModelConfig] = None,
|
|
validate: bool = True,
|
|
) -> dict:
|
|
"""
|
|
Validate custom credentials.
|
|
:param tenant_id: workspace id
|
|
:param provider_configuration: provider configuration
|
|
:param model_type: model type
|
|
:param model: model name
|
|
:param credentials: credentials
|
|
:param load_balancing_model_config: load balancing model config
|
|
:param validate: validate credentials
|
|
:return:
|
|
"""
|
|
# Get credential form schemas from model credential schema or provider credential schema
|
|
credential_schemas = self._get_credential_schema(provider_configuration)
|
|
|
|
# Get provider credential secret variables
|
|
provider_credential_secret_variables = provider_configuration.extract_secret_variables(
|
|
credential_schemas.credential_form_schemas
|
|
)
|
|
|
|
if load_balancing_model_config:
|
|
try:
|
|
# fix origin data
|
|
if load_balancing_model_config.encrypted_config:
|
|
original_credentials = json.loads(load_balancing_model_config.encrypted_config)
|
|
else:
|
|
original_credentials = {}
|
|
except JSONDecodeError:
|
|
original_credentials = {}
|
|
|
|
# encrypt credentials
|
|
for key, value in credentials.items():
|
|
if key in provider_credential_secret_variables:
|
|
# if send [__HIDDEN__] in secret input, it will be same as original value
|
|
if value == HIDDEN_VALUE and key in original_credentials:
|
|
credentials[key] = encrypter.decrypt_token(tenant_id, original_credentials[key])
|
|
|
|
if validate:
|
|
model_provider_factory = ModelProviderFactory(tenant_id)
|
|
if isinstance(credential_schemas, ModelCredentialSchema):
|
|
credentials = model_provider_factory.model_credentials_validate(
|
|
provider=provider_configuration.provider.provider,
|
|
model_type=model_type,
|
|
model=model,
|
|
credentials=credentials,
|
|
)
|
|
else:
|
|
credentials = model_provider_factory.provider_credentials_validate(
|
|
provider=provider_configuration.provider.provider, credentials=credentials
|
|
)
|
|
|
|
for key, value in credentials.items():
|
|
if key in provider_credential_secret_variables:
|
|
credentials[key] = encrypter.encrypt_token(tenant_id, value)
|
|
|
|
return credentials
|
|
|
|
def _get_credential_schema(
|
|
self, provider_configuration: ProviderConfiguration
|
|
) -> Union[ModelCredentialSchema, ProviderCredentialSchema]:
|
|
"""Get form schemas."""
|
|
if provider_configuration.provider.model_credential_schema:
|
|
return provider_configuration.provider.model_credential_schema
|
|
elif provider_configuration.provider.provider_credential_schema:
|
|
return provider_configuration.provider.provider_credential_schema
|
|
else:
|
|
raise ValueError("No credential schema found")
|
|
|
|
def _clear_credentials_cache(self, tenant_id: str, config_id: str) -> None:
|
|
"""
|
|
Clear credentials cache.
|
|
:param tenant_id: workspace id
|
|
:param config_id: load balancing config id
|
|
:return:
|
|
"""
|
|
provider_model_credentials_cache = ProviderCredentialsCache(
|
|
tenant_id=tenant_id, identity_id=config_id, cache_type=ProviderCredentialsCacheType.LOAD_BALANCING_MODEL
|
|
)
|
|
|
|
provider_model_credentials_cache.delete()
|