Underground-Digital/Workflow-Engine
0
1import logging2 3from flask_login import current_user4from flask_restful import Resource, reqparse5from werkzeug.exceptions import Forbidden6 7from controllers.console import api8from controllers.console.wraps import account_initialization_required, setup_required9from core.model_runtime.entities.model_entities import ModelType10from core.model_runtime.errors.validate import CredentialsValidateFailedError11from core.model_runtime.utils.encoders import jsonable_encoder12from libs.login import login_required13from services.model_load_balancing_service import ModelLoadBalancingService14from services.model_provider_service import ModelProviderService15 16 17class DefaultModelApi(Resource):18 @setup_required19 @login_required20 @account_initialization_required21 def get(self):22 parser = reqparse.RequestParser()23 parser.add_argument(24 "model_type",25 type=str,26 required=True,27 nullable=False,28 choices=[mt.value for mt in ModelType],29 location="args",30 )31 args = parser.parse_args()32 33 tenant_id = current_user.current_tenant_id34 35 model_provider_service = ModelProviderService()36 default_model_entity = model_provider_service.get_default_model_of_model_type(37 tenant_id=tenant_id, model_type=args["model_type"]38 )39 40 return jsonable_encoder({"data": default_model_entity})41 42 @setup_required43 @login_required44 @account_initialization_required45 def post(self):46 if not current_user.is_admin_or_owner:47 raise Forbidden()48 49 parser = reqparse.RequestParser()50 parser.add_argument("model_settings", type=list, required=True, nullable=False, location="json")51 args = parser.parse_args()52 53 tenant_id = current_user.current_tenant_id54 55 model_provider_service = ModelProviderService()56 model_settings = args["model_settings"]57 for model_setting in model_settings:58 if "model_type" not in model_setting or model_setting["model_type"] not in [mt.value for mt in ModelType]:59 raise ValueError("invalid model type")60 61 if "provider" not in model_setting:62 continue63 64 if "model" not in model_setting:65 raise ValueError("invalid model")66 67 try:68 model_provider_service.update_default_model_of_model_type(69 tenant_id=tenant_id,70 model_type=model_setting["model_type"],71 provider=model_setting["provider"],72 model=model_setting["model"],73 )74 except Exception as ex:75 logging.exception(f"{model_setting['model_type']} save error: {ex}")76 raise ex77 78 return {"result": "success"}79 80 81class ModelProviderModelApi(Resource):82 @setup_required83 @login_required84 @account_initialization_required85 def get(self, provider):86 tenant_id = current_user.current_tenant_id87 88 model_provider_service = ModelProviderService()89 models = model_provider_service.get_models_by_provider(tenant_id=tenant_id, provider=provider)90 91 return jsonable_encoder({"data": models})92 93 @setup_required94 @login_required95 @account_initialization_required96 def post(self, provider: str):97 if not current_user.is_admin_or_owner:98 raise Forbidden()99 100 tenant_id = current_user.current_tenant_id101 102 parser = reqparse.RequestParser()103 parser.add_argument("model", type=str, required=True, nullable=False, location="json")104 parser.add_argument(105 "model_type",106 type=str,107 required=True,108 nullable=False,109 choices=[mt.value for mt in ModelType],110 location="json",111 )112 parser.add_argument("credentials", type=dict, required=False, nullable=True, location="json")113 parser.add_argument("load_balancing", type=dict, required=False, nullable=True, location="json")114 parser.add_argument("config_from", type=str, required=False, nullable=True, location="json")115 args = parser.parse_args()116 117 model_load_balancing_service = ModelLoadBalancingService()118 119 if (120 "load_balancing" in args121 and args["load_balancing"]122 and "enabled" in args["load_balancing"]123 and args["load_balancing"]["enabled"]124 ):125 if "configs" not in args["load_balancing"]:126 raise ValueError("invalid load balancing configs")127 128 # save load balancing configs129 model_load_balancing_service.update_load_balancing_configs(130 tenant_id=tenant_id,131 provider=provider,132 model=args["model"],133 model_type=args["model_type"],134 configs=args["load_balancing"]["configs"],135 )136 137 # enable load balancing138 model_load_balancing_service.enable_model_load_balancing(139 tenant_id=tenant_id, provider=provider, model=args["model"], model_type=args["model_type"]140 )141 else:142 # disable load balancing143 model_load_balancing_service.disable_model_load_balancing(144 tenant_id=tenant_id, provider=provider, model=args["model"], model_type=args["model_type"]145 )146 147 if args.get("config_from", "") != "predefined-model":148 model_provider_service = ModelProviderService()149 150 try:151 model_provider_service.save_model_credentials(152 tenant_id=tenant_id,153 provider=provider,154 model=args["model"],155 model_type=args["model_type"],156 credentials=args["credentials"],157 )158 except CredentialsValidateFailedError as ex:159 logging.exception(f"save model credentials error: {ex}")160 raise ValueError(str(ex))161 162 return {"result": "success"}, 200163 164 @setup_required165 @login_required166 @account_initialization_required167 def delete(self, provider: str):168 if not current_user.is_admin_or_owner:169 raise Forbidden()170 171 tenant_id = current_user.current_tenant_id172 173 parser = reqparse.RequestParser()174 parser.add_argument("model", type=str, required=True, nullable=False, location="json")175 parser.add_argument(176 "model_type",177 type=str,178 required=True,179 nullable=False,180 choices=[mt.value for mt in ModelType],181 location="json",182 )183 args = parser.parse_args()184 185 model_provider_service = ModelProviderService()186 model_provider_service.remove_model_credentials(187 tenant_id=tenant_id, provider=provider, model=args["model"], model_type=args["model_type"]188 )189 190 return {"result": "success"}, 204191 192 193class ModelProviderModelCredentialApi(Resource):194 @setup_required195 @login_required196 @account_initialization_required197 def get(self, provider: str):198 tenant_id = current_user.current_tenant_id199 200 parser = reqparse.RequestParser()201 parser.add_argument("model", type=str, required=True, nullable=False, location="args")202 parser.add_argument(203 "model_type",204 type=str,205 required=True,206 nullable=False,207 choices=[mt.value for mt in ModelType],208 location="args",209 )210 args = parser.parse_args()211 212 model_provider_service = ModelProviderService()213 credentials = model_provider_service.get_model_credentials(214 tenant_id=tenant_id, provider=provider, model_type=args["model_type"], model=args["model"]215 )216 217 model_load_balancing_service = ModelLoadBalancingService()218 is_load_balancing_enabled, load_balancing_configs = model_load_balancing_service.get_load_balancing_configs(219 tenant_id=tenant_id, provider=provider, model=args["model"], model_type=args["model_type"]220 )221 222 return {223 "credentials": credentials,224 "load_balancing": {"enabled": is_load_balancing_enabled, "configs": load_balancing_configs},225 }226 227 228class ModelProviderModelEnableApi(Resource):229 @setup_required230 @login_required231 @account_initialization_required232 def patch(self, provider: str):233 tenant_id = current_user.current_tenant_id234 235 parser = reqparse.RequestParser()236 parser.add_argument("model", type=str, required=True, nullable=False, location="json")237 parser.add_argument(238 "model_type",239 type=str,240 required=True,241 nullable=False,242 choices=[mt.value for mt in ModelType],243 location="json",244 )245 args = parser.parse_args()246 247 model_provider_service = ModelProviderService()248 model_provider_service.enable_model(249 tenant_id=tenant_id, provider=provider, model=args["model"], model_type=args["model_type"]250 )251 252 return {"result": "success"}253 254 255class ModelProviderModelDisableApi(Resource):256 @setup_required257 @login_required258 @account_initialization_required259 def patch(self, provider: str):260 tenant_id = current_user.current_tenant_id261 262 parser = reqparse.RequestParser()263 parser.add_argument("model", type=str, required=True, nullable=False, location="json")264 parser.add_argument(265 "model_type",266 type=str,267 required=True,268 nullable=False,269 choices=[mt.value for mt in ModelType],270 location="json",271 )272 args = parser.parse_args()273 274 model_provider_service = ModelProviderService()275 model_provider_service.disable_model(276 tenant_id=tenant_id, provider=provider, model=args["model"], model_type=args["model_type"]277 )278 279 return {"result": "success"}280 281 282class ModelProviderModelValidateApi(Resource):283 @setup_required284 @login_required285 @account_initialization_required286 def post(self, provider: str):287 tenant_id = current_user.current_tenant_id288 289 parser = reqparse.RequestParser()290 parser.add_argument("model", type=str, required=True, nullable=False, location="json")291 parser.add_argument(292 "model_type",293 type=str,294 required=True,295 nullable=False,296 choices=[mt.value for mt in ModelType],297 location="json",298 )299 parser.add_argument("credentials", type=dict, required=True, nullable=False, location="json")300 args = parser.parse_args()301 302 model_provider_service = ModelProviderService()303 304 result = True305 error = None306 307 try:308 model_provider_service.model_credentials_validate(309 tenant_id=tenant_id,310 provider=provider,311 model=args["model"],312 model_type=args["model_type"],313 credentials=args["credentials"],314 )315 except CredentialsValidateFailedError as ex:316 result = False317 error = str(ex)318 319 response = {"result": "success" if result else "error"}320 321 if not result:322 response["error"] = error323 324 return response325 326 327class ModelProviderModelParameterRuleApi(Resource):328 @setup_required329 @login_required330 @account_initialization_required331 def get(self, provider: str):332 parser = reqparse.RequestParser()333 parser.add_argument("model", type=str, required=True, nullable=False, location="args")334 args = parser.parse_args()335 336 tenant_id = current_user.current_tenant_id337 338 model_provider_service = ModelProviderService()339 parameter_rules = model_provider_service.get_model_parameter_rules(340 tenant_id=tenant_id, provider=provider, model=args["model"]341 )342 343 return jsonable_encoder({"data": parameter_rules})344 345 346class ModelProviderAvailableModelApi(Resource):347 @setup_required348 @login_required349 @account_initialization_required350 def get(self, model_type):351 tenant_id = current_user.current_tenant_id352 353 model_provider_service = ModelProviderService()354 models = model_provider_service.get_models_by_model_type(tenant_id=tenant_id, model_type=model_type)355 356 return jsonable_encoder({"data": models})357 358 359api.add_resource(ModelProviderModelApi, "/workspaces/current/model-providers/<string:provider>/models")360api.add_resource(361 ModelProviderModelEnableApi,362 "/workspaces/current/model-providers/<string:provider>/models/enable",363 endpoint="model-provider-model-enable",364)365api.add_resource(366 ModelProviderModelDisableApi,367 "/workspaces/current/model-providers/<string:provider>/models/disable",368 endpoint="model-provider-model-disable",369)370api.add_resource(371 ModelProviderModelCredentialApi, "/workspaces/current/model-providers/<string:provider>/models/credentials"372)373api.add_resource(374 ModelProviderModelValidateApi, "/workspaces/current/model-providers/<string:provider>/models/credentials/validate"375)376 377api.add_resource(378 ModelProviderModelParameterRuleApi, "/workspaces/current/model-providers/<string:provider>/models/parameter-rules"379)380api.add_resource(ModelProviderAvailableModelApi, "/workspaces/current/models/model-types/<string:model_type>")381api.add_resource(DefaultModelApi, "/workspaces/current/default-model")382 