Underground-Digital/Workflow-Engine
0
1import datetime2 3import pytz4from flask import request5from flask_login import current_user6from flask_restful import Resource, fields, marshal_with, reqparse7 8from configs import dify_config9from constants.languages import supported_language10from controllers.console import api11from controllers.console.workspace.error import (12 AccountAlreadyInitedError,13 CurrentPasswordIncorrectError,14 InvalidInvitationCodeError,15 RepeatPasswordNotMatchError,16)17from controllers.console.wraps import account_initialization_required, setup_required18from extensions.ext_database import db19from fields.member_fields import account_fields20from libs.helper import TimestampField, timezone21from libs.login import login_required22from models import AccountIntegrate, InvitationCode23from services.account_service import AccountService24from services.errors.account import CurrentPasswordIncorrectError as ServiceCurrentPasswordIncorrectError25 26 27class AccountInitApi(Resource):28 @setup_required29 @login_required30 def post(self):31 account = current_user32 33 if account.status == "active":34 raise AccountAlreadyInitedError()35 36 parser = reqparse.RequestParser()37 38 if dify_config.EDITION == "CLOUD":39 parser.add_argument("invitation_code", type=str, location="json")40 41 parser.add_argument("interface_language", type=supported_language, required=True, location="json")42 parser.add_argument("timezone", type=timezone, required=True, location="json")43 args = parser.parse_args()44 45 if dify_config.EDITION == "CLOUD":46 if not args["invitation_code"]:47 raise ValueError("invitation_code is required")48 49 # check invitation code50 invitation_code = (51 db.session.query(InvitationCode)52 .filter(53 InvitationCode.code == args["invitation_code"],54 InvitationCode.status == "unused",55 )56 .first()57 )58 59 if not invitation_code:60 raise InvalidInvitationCodeError()61 62 invitation_code.status = "used"63 invitation_code.used_at = datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None)64 invitation_code.used_by_tenant_id = account.current_tenant_id65 invitation_code.used_by_account_id = account.id66 67 account.interface_language = args["interface_language"]68 account.timezone = args["timezone"]69 account.interface_theme = "light"70 account.status = "active"71 account.initialized_at = datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None)72 db.session.commit()73 74 return {"result": "success"}75 76 77class AccountProfileApi(Resource):78 @setup_required79 @login_required80 @account_initialization_required81 @marshal_with(account_fields)82 def get(self):83 return current_user84 85 86class AccountNameApi(Resource):87 @setup_required88 @login_required89 @account_initialization_required90 @marshal_with(account_fields)91 def post(self):92 parser = reqparse.RequestParser()93 parser.add_argument("name", type=str, required=True, location="json")94 args = parser.parse_args()95 96 # Validate account name length97 if len(args["name"]) < 3 or len(args["name"]) > 30:98 raise ValueError("Account name must be between 3 and 30 characters.")99 100 updated_account = AccountService.update_account(current_user, name=args["name"])101 102 return updated_account103 104 105class AccountAvatarApi(Resource):106 @setup_required107 @login_required108 @account_initialization_required109 @marshal_with(account_fields)110 def post(self):111 parser = reqparse.RequestParser()112 parser.add_argument("avatar", type=str, required=True, location="json")113 args = parser.parse_args()114 115 updated_account = AccountService.update_account(current_user, avatar=args["avatar"])116 117 return updated_account118 119 120class AccountInterfaceLanguageApi(Resource):121 @setup_required122 @login_required123 @account_initialization_required124 @marshal_with(account_fields)125 def post(self):126 parser = reqparse.RequestParser()127 parser.add_argument("interface_language", type=supported_language, required=True, location="json")128 args = parser.parse_args()129 130 updated_account = AccountService.update_account(current_user, interface_language=args["interface_language"])131 132 return updated_account133 134 135class AccountInterfaceThemeApi(Resource):136 @setup_required137 @login_required138 @account_initialization_required139 @marshal_with(account_fields)140 def post(self):141 parser = reqparse.RequestParser()142 parser.add_argument("interface_theme", type=str, choices=["light", "dark"], required=True, location="json")143 args = parser.parse_args()144 145 updated_account = AccountService.update_account(current_user, interface_theme=args["interface_theme"])146 147 return updated_account148 149 150class AccountTimezoneApi(Resource):151 @setup_required152 @login_required153 @account_initialization_required154 @marshal_with(account_fields)155 def post(self):156 parser = reqparse.RequestParser()157 parser.add_argument("timezone", type=str, required=True, location="json")158 args = parser.parse_args()159 160 # Validate timezone string, e.g. America/New_York, Asia/Shanghai161 if args["timezone"] not in pytz.all_timezones:162 raise ValueError("Invalid timezone string.")163 164 updated_account = AccountService.update_account(current_user, timezone=args["timezone"])165 166 return updated_account167 168 169class AccountPasswordApi(Resource):170 @setup_required171 @login_required172 @account_initialization_required173 @marshal_with(account_fields)174 def post(self):175 parser = reqparse.RequestParser()176 parser.add_argument("password", type=str, required=False, location="json")177 parser.add_argument("new_password", type=str, required=True, location="json")178 parser.add_argument("repeat_new_password", type=str, required=True, location="json")179 args = parser.parse_args()180 181 if args["new_password"] != args["repeat_new_password"]:182 raise RepeatPasswordNotMatchError()183 184 try:185 AccountService.update_account_password(current_user, args["password"], args["new_password"])186 except ServiceCurrentPasswordIncorrectError:187 raise CurrentPasswordIncorrectError()188 189 return {"result": "success"}190 191 192class AccountIntegrateApi(Resource):193 integrate_fields = {194 "provider": fields.String,195 "created_at": TimestampField,196 "is_bound": fields.Boolean,197 "link": fields.String,198 }199 200 integrate_list_fields = {201 "data": fields.List(fields.Nested(integrate_fields)),202 }203 204 @setup_required205 @login_required206 @account_initialization_required207 @marshal_with(integrate_list_fields)208 def get(self):209 account = current_user210 211 account_integrates = db.session.query(AccountIntegrate).filter(AccountIntegrate.account_id == account.id).all()212 213 base_url = request.url_root.rstrip("/")214 oauth_base_path = "/console/api/oauth/login"215 providers = ["github", "google"]216 217 integrate_data = []218 for provider in providers:219 existing_integrate = next((ai for ai in account_integrates if ai.provider == provider), None)220 if existing_integrate:221 integrate_data.append(222 {223 "id": existing_integrate.id,224 "provider": provider,225 "created_at": existing_integrate.created_at,226 "is_bound": True,227 "link": None,228 }229 )230 else:231 integrate_data.append(232 {233 "id": None,234 "provider": provider,235 "created_at": None,236 "is_bound": False,237 "link": f"{base_url}{oauth_base_path}/{provider}",238 }239 )240 241 return {"data": integrate_data}242 243 244# Register API resources245api.add_resource(AccountInitApi, "/account/init")246api.add_resource(AccountProfileApi, "/account/profile")247api.add_resource(AccountNameApi, "/account/name")248api.add_resource(AccountAvatarApi, "/account/avatar")249api.add_resource(AccountInterfaceLanguageApi, "/account/interface-language")250api.add_resource(AccountInterfaceThemeApi, "/account/interface-theme")251api.add_resource(AccountTimezoneApi, "/account/timezone")252api.add_resource(AccountPasswordApi, "/account/password")253api.add_resource(AccountIntegrateApi, "/account/integrates")254# api.add_resource(AccountEmailApi, '/account/email')255# api.add_resource(AccountEmailVerifyApi, '/account/email-verify')256 