Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
models.py382 linesDownload Raw Back to workspace
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