Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
__init__.py941 linesDownload Raw Back to utils
1##########################################################################2#3# pgAdmin 4 - PostgreSQL Tools4#5# Copyright (C) 2013 - 2024, The pgAdmin Development Team6# This software is released under the PostgreSQL Licence7#8##########################################################################9 10import os11import sys12import json13import subprocess14from collections import defaultdict15from operator import attrgetter16 17from pathlib import Path18from flask import Blueprint, current_app, url_for19from flask_babel import gettext20from flask_security import current_user, login_required21from flask_security.utils import get_post_login_redirect, \22    get_post_logout_redirect23from threading import Lock24import config25from .paths import get_storage_directory26from .preferences import Preferences27from pgadmin.utils.constants import UTILITIES_ARRAY, USER_NOT_FOUND, \28    MY_STORAGE, ACCESS_DENIED_MESSAGE, INTERNAL29from pgadmin.utils.ajax import make_json_response30from pgadmin.model import db, User, ServerGroup, Server31from urllib.parse import unquote32 33ADD_SERVERS_MSG = "Added %d Server Group(s) and %d Server(s)."34 35 36class PgAdminModule(Blueprint):37    """38    Base class for every PgAdmin Module.39 40    This class defines a set of method and attributes that41    every module should implement.42    """43 44    def __init__(self, name, import_name, **kwargs):45        kwargs.setdefault('url_prefix', '/' + name)46        kwargs.setdefault('template_folder', 'templates')47        kwargs.setdefault('static_folder', 'static')48        self.submodules = []49        self.parentmodules = []50 51        super().__init__(name, import_name, **kwargs)52 53    def register_preferences(self):54        # To be implemented by child classes55        pass56 57    def register(self, app, options):58        """59        Override the default register function to automagically register60        sub-modules at once.61        """62 63        super().register(app, options)64 65        def create_module_preference():66            # Create preference for each module by default67            if hasattr(self, 'LABEL'):68                self.preference = Preferences(self.name, self.LABEL)69            else:70                self.preference = Preferences(self.name, None)71 72            self.register_preferences()73 74        # Create and register the module preference object and preferences for75        # it just before starting app76        app.register_before_app_start(create_module_preference)77 78        for module in self.submodules:79            module.parentmodules.append(self)80            if app.blueprints.get(module.name) is None:81                app.register_blueprint(module)82                app.register_logout_hook(module)83 84    def get_own_messages(self):85        """86        Returns:87            dict: the i18n messages used by this module, not including any88                messages needed by the submodules.89        """90        return dict()91 92    def get_own_menuitems(self):93        """94        Returns:95            dict: the menuitems for this module, not including96                any needed from the submodules.97        """98        return defaultdict(list)99 100    def get_exposed_url_endpoints(self):101        """102        Returns:103            list: a list of url endpoints exposed to the client.104        """105        return []106 107    @property108    def messages(self):109        res = self.get_own_messages()110 111        for module in self.submodules:112            res.update(module.messages)113        return res114 115    @property116    def menu_items(self):117        menu_items = self.get_own_menuitems()118        for module in self.submodules:119            for key, value in module.menu_items.items():120                menu_items[key].extend(value)121        menu_items = dict((key, sorted(value, key=attrgetter('priority')))122                          for key, value in menu_items.items())123        return menu_items124 125    @property126    def exposed_endpoints(self):127        res = self.get_exposed_url_endpoints()128 129        for module in self.submodules:130            res += module.exposed_endpoints131 132        return res133 134 135IS_WIN = (os.name == 'nt')136 137sys_encoding = sys.getdefaultencoding()138if not sys_encoding or sys_encoding == 'ascii':139    # Fall back to 'utf-8', if we couldn't determine the default encoding,140    # or 'ascii'.141    sys_encoding = 'utf-8'142 143fs_encoding = sys.getfilesystemencoding()144if not fs_encoding or fs_encoding == 'ascii':145    # Fall back to 'utf-8', if we couldn't determine the file-system encoding,146    # or 'ascii'.147    fs_encoding = 'utf-8'148 149 150def u_encode(_s, _encoding=sys_encoding):151    return _s152 153 154def file_quote(_p):155    return _p156 157 158if IS_WIN:159    import ctypes160    from ctypes import wintypes161 162    def env(name):163        if name in os.environ:164            return os.environ[name]165        return None166 167    _GetShortPathNameW = ctypes.windll.kernel32.GetShortPathNameW168    _GetShortPathNameW.argtypes = [169        wintypes.LPCWSTR, wintypes.LPWSTR, wintypes.DWORD170    ]171    _GetShortPathNameW.restype = wintypes.DWORD172 173    def fs_short_path(_path):174        """175        Gets the short path name of a given long path.176        http://stackoverflow.com/a/23598461/200291177        """178        buf_size = len(_path)179        while True:180            res = ctypes.create_unicode_buffer(buf_size)181            # Note:- _GetShortPathNameW may return empty value182            # if directory doesn't exist.183            needed = _GetShortPathNameW(_path, res, buf_size)184 185            if buf_size >= needed:186                return res.value187            else:188                buf_size += needed189 190    def document_dir():191        CSIDL_PERSONAL = 5  # My Documents192        SHGFP_TYPE_CURRENT = 0  # Get current, not default value193 194        buf = ctypes.create_unicode_buffer(wintypes.MAX_PATH)195        ctypes.windll.shell32.SHGetFolderPathW(196            None, CSIDL_PERSONAL, None, SHGFP_TYPE_CURRENT, buf197        )198 199        return buf.value200 201else:202    def env(name):203        if name in os.environ:204            return os.environ[name]205        return None206 207    def fs_short_path(_path):208        return _path209 210    def document_dir():211        return os.path.realpath(os.path.expanduser('~/'))212 213 214def get_complete_file_path(file, validate=True):215    """216    Args:217        file: File returned by file manager218 219    Returns:220         Full path for the file221    """222    if not file:223        return None224 225    # If desktop mode226    if current_app.PGADMIN_RUNTIME or not current_app.config['SERVER_MODE']:227        return file if os.path.isfile(file) else None228 229    storage_dir = get_storage_directory()230    if storage_dir:231        file = os.path.join(232            storage_dir,233            file.lstrip('/').lstrip('\\')234        )235        if IS_WIN:236            file = file.replace('\\', '/')237            file = fs_short_path(file)238 239    if validate:240        return file if os.path.isfile(file) else None241    else:242        return file243 244 245def filename_with_file_manager_path(_file, create_file=False,246                                    skip_permission_check=False):247    """248    Args:249        file: File name returned from client file manager250        create_file: Set flag to False when file creation doesn't require251        skip_permission_check:252    Returns:253        Filename to use for backup with full path taken from preference254    """255    # retrieve storage directory path256    try:257        last_storage = Preferences.module('file_manager').preference(258            'last_storage').get()259    except Exception:260        last_storage = MY_STORAGE261 262    if last_storage != MY_STORAGE:263        sel_dir_list = [sdir for sdir in current_app.config['SHARED_STORAGE']264                        if sdir['name'] == last_storage]265        selected_dir = sel_dir_list[0] if len(266            sel_dir_list) == 1 else None267 268        if selected_dir and selected_dir['restricted_access'] and \269                not current_user.has_role("Administrator"):270            return make_json_response(success=0,271                                      errormsg=ACCESS_DENIED_MESSAGE,272                                      info='ACCESS_DENIED',273                                      status=403)274        storage_dir = get_storage_directory(275            shared_storage=last_storage)276    else:277        storage_dir = get_storage_directory()278 279    from pgadmin.misc.file_manager import Filemanager280    Filemanager.check_access_permission(281        storage_dir, _file, skip_permission_check)282    if storage_dir:283        _file = os.path.join(storage_dir, _file.lstrip('/').lstrip('\\'))284    elif not os.path.isabs(_file):285        _file = os.path.join(document_dir(), _file)286 287    def short_filepath():288        short_path = fs_short_path(_file)289        # fs_short_path() function may return empty path on Windows290        # if directory doesn't exists. In that case we strip the last path291        # component and get the short path.292        if os.name == 'nt' and short_path == '':293            base_name = os.path.basename(_file)294            dir_name = os.path.dirname(_file)295            short_path = fs_short_path(dir_name) + '\\' + base_name296        return short_path297 298    if create_file:299        # Touch the file to get the short path of the file on windows.300        with open(_file, 'a'):301            return short_filepath()302 303    return short_filepath()304 305 306def does_utility_exist(file):307    """308    This function will check the utility file exists on given path.309    :return:310    """311    error_msg = None312 313    if file is None:314        error_msg = gettext("Utility file not found. Please correct the Binary"315                            " Path in the Preferences dialog")316        return error_msg317 318    if Path(config.STORAGE_DIR) == Path(file) or \319            Path(config.STORAGE_DIR) in Path(file).parents:320        error_msg = gettext("Please correct the Binary Path in the Preferences"321                            " dialog. pgAdmin storage directory can not be a"322                            " utility binary directory.")323 324    if not os.path.exists(file):325        error_msg = gettext("'%s' file not found. Please correct the Binary"326                            " Path in the Preferences dialog" % file)327    return error_msg328 329 330def get_server(sid):331    """332    # Fetch the server  etc333    :param sid:334    :return: server335    """336    server = Server.query.filter_by(id=sid).first()337    return server338 339 340def get_binary_path_versions(binary_path: str) -> dict:341    ret = {}342    binary_path = os.path.abspath(343        replace_binary_path(binary_path)344    )345 346    for utility in UTILITIES_ARRAY:347        ret[utility] = None348        full_path = os.path.join(binary_path,349                                 (utility if os.name != 'nt' else350                                  (utility + '.exe')))351 352        try:353            # if path doesn't exist raise exception354            if not os.path.isdir(binary_path):355                current_app.logger.warning('Invalid binary path.')356                raise Exception()357            # Get the output of the '--version' command358            cmd = subprocess.run(359                [full_path, '--version'],360                shell=False,361                capture_output=True,362                text=True363            )364            if cmd.returncode == 0:365                ret[utility] = cmd.stdout.split(") ", 1)[1].strip()366            else:367                raise Exception()368        except Exception as _:369            continue370 371    return ret372 373 374def set_binary_path(binary_path, bin_paths, server_type,375                    version_number=None, set_as_default=False,376                    is_fixed_path=False):377    """378    This function is used to iterate through the utilities and set the379    default binary path.380    """381    path_with_dir = binary_path if "$DIR" in binary_path else None382    binary_versions = get_binary_path_versions(binary_path)383 384    for utility, version in binary_versions.items():385        version_number = version if version_number is None else version_number386        # version will be None if binary not present387        version_number = version_number or ''388        if version_number.find('.'):389            version_number = version_number.split('.', 1)[0]390        try:391            # Get the paths array based on server type392            if 'pg_bin_paths' in bin_paths or 'as_bin_paths' in bin_paths:393                paths_array = bin_paths['pg_bin_paths']394                if server_type == 'ppas':395                    paths_array = bin_paths['as_bin_paths']396            else:397                paths_array = bin_paths398 399            for path in paths_array:400                if path['version'].find(version_number) == 0 and \401                        path['binaryPath'] is None:402                    path['binaryPath'] = path_with_dir \403                        if path_with_dir is not None else binary_path404                    if set_as_default:405                        path['isDefault'] = True406                    # Whether the fixed path in the config file exists or not407                    path['isFixed'] = is_fixed_path408                    break409            break410        except Exception:411            continue412 413 414def replace_binary_path(binary_path):415    """416    This function is used to check if $DIR is present in417    the binary path. If it is there then replace it with418    module.419    """420    if "$DIR" in binary_path:421        # When running as an WSGI application, we will not find the422        # '__file__' attribute for the '__main__' module.423        main_module_file = getattr(424            sys.modules['__main__'], '__file__', None425        )426 427        if main_module_file is not None:428            binary_path = binary_path.replace(429                "$DIR", os.path.dirname(main_module_file)430            )431 432    return binary_path433 434 435def add_value(attr_dict, key, value):436    """Add a value to the attribute dict if non-empty.437 438    Args:439        attr_dict (dict): The dictionary to add the values to440        key (str): The key for the new value441        value (str): The value to add442 443    Returns:444        The updated attribute dictionary445    """446    if value != "" and value is not None:447        attr_dict[key] = value448 449    return attr_dict450 451 452def dump_database_servers(output_file, selected_servers,453                          dump_user=current_user, from_setup=False,454                          auth_source=INTERNAL):455    """Dump the server groups and servers.456    """457    user = _does_user_exist(dump_user, from_setup, auth_source)458    if user is None:459        return False, USER_NOT_FOUND % dump_user460 461    user_id = user.id462    # Dict to collect the output463    object_dict = {}464    # Counters465    servers_dumped = 0466 467    # Dump servers468    servers = Server.query.filter_by(user_id=user_id).all()469    server_dict = {}470    for server in servers:471        if selected_servers is None or (472            isinstance(selected_servers, list) and len(selected_servers) == 0)\473                or str(server.id) in selected_servers\474                or server.id in selected_servers:475            # Get the group name476            group_name = ServerGroup.query.filter_by(477                user_id=user_id, id=server.servergroup_id).first().name478 479            attr_dict = {}480            add_value(attr_dict, "Name", server.name)481            add_value(attr_dict, "Group", group_name)482            add_value(attr_dict, "Host", server.host)483            add_value(attr_dict, "Port", server.port)484            add_value(attr_dict, "MaintenanceDB", server.maintenance_db)485            add_value(attr_dict, "Username", server.username)486            add_value(attr_dict, "Role", server.role)487            add_value(attr_dict, "Comment", server.comment)488            add_value(attr_dict, "Shared", server.shared)489            add_value(attr_dict, "SharedUsername", server.shared_username)490            add_value(attr_dict, "DBRestriction", server.db_res)491            add_value(attr_dict, "BGColor", server.bgcolor)492            add_value(attr_dict, "FGColor", server.fgcolor)493            add_value(attr_dict, "Service", server.service)494            add_value(attr_dict, "UseSSHTunnel", server.use_ssh_tunnel)495            add_value(attr_dict, "TunnelHost", server.tunnel_host)496            add_value(attr_dict, "TunnelPort", server.tunnel_port)497            add_value(attr_dict, "TunnelUsername", server.tunnel_username)498            add_value(attr_dict, "TunnelAuthentication",499                      server.tunnel_authentication)500            add_value(attr_dict, "KerberosAuthentication",501                      server.kerberos_conn),502            add_value(attr_dict, "ConnectionParameters",503                      server.connection_params)504 505            # if desktop mode or server mode with506            # ENABLE_SERVER_PASS_EXEC_CMD flag is True507            if not current_app.config['SERVER_MODE'] or \508                    current_app.config['ENABLE_SERVER_PASS_EXEC_CMD']:509                add_value(attr_dict, "PasswordExecCommand",510                          server.passexec_cmd)511                add_value(attr_dict, "PasswordExecExpiration",512                          server.passexec_expiration)513 514            servers_dumped = servers_dumped + 1515 516            server_dict[servers_dumped] = attr_dict517 518    object_dict["Servers"] = server_dict519 520    try:521        if from_setup:522            file_path = unquote(output_file)523        else:524            file_path = filename_with_file_manager_path(unquote(output_file))525    except Exception as e:526        return _handle_error(str(e), from_setup)527 528    # write to file529    file_content = json.dumps(object_dict, indent=4)530    error_str = "Error: {0}"531    try:532        with open(file_path, 'w') as output_file:533            output_file.write(file_content)534    except IOError as e:535        err_msg = error_str.format(e.strerror)536        return _handle_error(err_msg, from_setup)537    except Exception as e:538        err_msg = error_str.format(e.strerror)539        return _handle_error(err_msg, from_setup)540 541    msg = gettext("Configuration for %s servers dumped to %s" %542                  (servers_dumped, output_file.name))543    print(msg)544 545    return True, msg546 547 548def validate_json_data(data, is_admin):549    """550    Used internally by load_servers to validate servers data.551    :param data: servers data552    :param is_admin:553    :return: error message if any554    """555    skip_servers = []556    # Loop through the servers...557    if "Servers" not in data:558        return gettext("'Servers' attribute not found in the specified file.")559 560    for server in data["Servers"]:561        obj = data["Servers"][server]562 563        # Check if server is shared.Won't import if user is non-admin564        if obj.get('Shared', None) and not is_admin:565            print("Won't import the server '%s' as it is shared " %566                  obj["Name"])567            skip_servers.append(server)568            continue569 570        def check_attrib(attrib):571            if attrib not in obj:572                return gettext("'%s' attribute not found for server '%s'" %573                               (attrib, server))574            return None575 576        def check_is_integer(value):577            if not isinstance(value, int):578                return gettext("Port must be integer for server '%s'" % server)579            return None580 581        for attrib in ("Group", "Name"):582            errmsg = check_attrib(attrib)583            if errmsg:584                return errmsg585 586        is_service_attrib_available = obj.get("Service", None) is not None587 588        if not is_service_attrib_available:589            for attrib in ("Port", "Username"):590                errmsg = check_attrib(attrib)591                if errmsg:592                    return errmsg593                if attrib == 'Port':594                    errmsg = check_is_integer(obj[attrib])595                    if errmsg:596                        return errmsg597 598        errmsg = check_attrib("MaintenanceDB")599        if errmsg:600            return errmsg601 602        if "Host" not in obj and not is_service_attrib_available:603            return gettext("'Host' or 'Service' attribute not "604                           "found for server '%s'" % server)605 606    for server in skip_servers:607        del data["Servers"][server]608    return None609 610 611def load_database_servers(input_file, selected_servers,612                          load_user=current_user, from_setup=False,613                          auth_source=INTERNAL):614    """Load server groups and servers.615    """616    user = _does_user_exist(load_user, from_setup, auth_source)617    if user is None:618        return False, USER_NOT_FOUND % load_user619 620    # generate full path of file621    try:622        if from_setup:623            file_path = unquote(input_file)624        else:625            file_path = filename_with_file_manager_path(unquote(input_file))626    except Exception as e:627        return _handle_error(str(e), from_setup)628 629    try:630        with open(file_path) as f:631            data = json.load(f)632    except json.decoder.JSONDecodeError as e:633        return _handle_error(gettext("Error parsing input file %s: %s" %634                             (file_path, e)), from_setup)635    except Exception as e:636        return _handle_error(gettext("Error reading input file %s: [%d] %s" %637                             (file_path, e.errno, e.strerror)), from_setup)638 639    f.close()640 641    user_id = user.id642    # Counters643    groups_added = 0644    servers_added = 0645 646    # Get the server groups647    groups = ServerGroup.query.filter_by(user_id=user_id)648 649    # Validate server data650    error_msg = validate_json_data(data, user.has_role("Administrator"))651    if error_msg is not None and from_setup:652        print(ADD_SERVERS_MSG % (groups_added, servers_added))653        return _handle_error(error_msg, from_setup)654 655    for server in data["Servers"]:656        if selected_servers is None or str(server) in selected_servers:657            obj = data["Servers"][server]658 659            # Get the group. Create if necessary660            group_id = next(661                (g.id for g in groups if g.name == obj["Group"]), -1)662 663            if group_id == -1:664                new_group = ServerGroup()665                new_group.name = obj["Group"]666                new_group.user_id = user_id667                db.session.add(new_group)668 669                try:670                    db.session.commit()671                except Exception as e:672                    if from_setup:673                        print(ADD_SERVERS_MSG % (groups_added, servers_added))674                    return _handle_error(675                        gettext("Error creating server group '%s': %s" %676                                (new_group.name, e)), from_setup)677 678                group_id = new_group.id679                groups_added = groups_added + 1680                groups = ServerGroup.query.filter_by(user_id=user_id)681 682            # Create the server683            new_server = Server()684            new_server.name = obj["Name"]685            new_server.servergroup_id = group_id686            new_server.user_id = user_id687            new_server.maintenance_db = obj["MaintenanceDB"]688 689            new_server.host = obj.get("Host", None)690 691            new_server.port = obj.get("Port", None)692 693            new_server.username = obj.get("Username", None)694 695            new_server.role = obj.get("Role", None)696 697            new_server.comment = obj.get("Comment", None)698 699            new_server.db_res = obj.get("DBRestriction", None)700 701            if 'ConnectionParameters' in obj:702                new_server.connection_params = \703                    obj.get("ConnectionParameters", None)704            else:705                # JSON file format is old before introduction of the706                # connection parameters.707                conn_param = dict()708                for item in ['HostAddr', 'SSLMode', 'PassFile', 'SSLCert',709                             'SSLKey', 'SSLRootCert', 'SSLCrl', 'Timeout',710                             'SSLCompression']:711                    if item in obj:712                        key = item.lower()713                        if item == 'Timeout':714                            key = 'connect_timeout'715                        conn_param[key] = obj.get(item)716 717                new_server.connection_params = conn_param718 719            new_server.bgcolor = obj.get("BGColor", None)720 721            new_server.fgcolor = obj.get("FGColor", None)722 723            new_server.service = obj.get("Service", None)724 725            new_server.use_ssh_tunnel = obj.get("UseSSHTunnel", None)726 727            new_server.tunnel_host = obj.get("TunnelHost", None)728 729            new_server.tunnel_port = obj.get("TunnelPort", None)730 731            new_server.tunnel_username = obj.get("TunnelUsername", None)732 733            new_server.tunnel_authentication = \734                obj.get("TunnelAuthentication", None)735 736            new_server.shared = obj.get("Shared", None)737 738            new_server.shared_username = obj.get("SharedUsername", None)739 740            new_server.kerberos_conn = obj.get("KerberosAuthentication", None)741 742            # if desktop mode or server mode with743            # ENABLE_SERVER_PASS_EXEC_CMD flag is True744            if not current_app.config['SERVER_MODE'] or \745                    current_app.config['ENABLE_SERVER_PASS_EXEC_CMD']:746                new_server.passexec_cmd = obj.get("PasswordExecCommand", None)747                new_server.passexec_expiration = obj.get(748                    "PasswordExecExpiration", None)749 750            db.session.add(new_server)751 752            try:753                db.session.commit()754            except Exception as e:755                if from_setup:756                    print(ADD_SERVERS_MSG % (groups_added, servers_added))757                return _handle_error(gettext("Error creating server '%s': %s" %758                                             (new_server.name, e)), from_setup)759 760            servers_added = servers_added + 1761 762    msg = ADD_SERVERS_MSG % (groups_added, servers_added)763    print(msg)764 765    return True, msg766 767 768def clear_database_servers(load_user=current_user, from_setup=False,769                           auth_source=INTERNAL):770    """Clear groups and servers configurations.771    """772    user = _does_user_exist(load_user, from_setup, auth_source)773    if user is None:774        return False775 776    user_id = user.id777 778    # Remove all servers779    servers = Server.query.filter_by(user_id=user_id)780    for server in servers:781        db.session.delete(server)782 783    # Remove all servergroups except for the first784    # This matches the UI behavior in785    # web/pgadmin/browser/server_groups/__init__.py#delete786    # TODO: Investigate if we can skip the first with an `offset(1)`787    groups = ServerGroup.query.filter_by(user_id=user_id).order_by("id")788    default_sg = groups.first()789    for group in groups:790        if group.id != default_sg.id:791            db.session.delete(group)792 793    try:794        db.session.commit()795    except Exception as e:796        error_msg = \797            gettext("Error clearing server configuration with error (%s)" %798                    str(e))799        if from_setup:800            print(error_msg)801            sys.exit(1)802 803        return False, error_msg804 805 806def _does_user_exist(user, from_setup, auth_source=INTERNAL):807    """808    This function will check user is exist or not. If exist then return809    """810    if isinstance(user, User):811        auth_source = user.auth_source812        user = user.username813 814    new_user = User.query.filter_by(username=user,815                                    auth_source=auth_source).first()816 817    if new_user is None:818        print(USER_NOT_FOUND % user)819        if from_setup:820            sys.exit(1)821 822    return new_user823 824 825def _handle_error(error_msg, from_setup):826    """827    This function is used to print the error msg and exit from app if828    called from setup.py829    """830    if from_setup:831        print(error_msg)832        sys.exit(1)833 834    return False, error_msg835 836 837# Shortcut configuration for Accesskey838ACCESSKEY_FIELDS = [839    {840        'name': 'key',841        'type': 'keyCode',842        'label': gettext('Key')843    }844]845 846# Shortcut configuration847SHORTCUT_FIELDS = [848    {849        'name': 'key',850        'type': 'keyCode',851        'label': gettext('Key')852    },853    {854        'name': 'shift',855        'type': 'checkbox',856        'label': gettext('Shift')857    },858 859    {860        'name': 'control',861        'type': 'checkbox',862        'label': gettext('Ctrl')863    },864    {865        'name': 'alt',866        'type': 'checkbox',867        'label': gettext('Alt/Option')868    }869]870 871 872class KeyManager:873    def __init__(self):874        self.users = dict()875        self.lock = Lock()876 877    @login_required878    def get(self):879        user = self.users.get(current_user.id, None)880        if user is not None:881            return user.get('key', None)882 883    @login_required884    def set(self, _key, _new_login=True):885        with self.lock:886            user = self.users.get(current_user.id, None)887            if user is None:888                self.users[current_user.id] = dict(889                    session_count=1, key=_key)890            else:891                if _new_login:892                    user['session_count'] += 1893                user['key'] = _key894 895    @login_required896    def reset(self):897        with self.lock:898            user = self.users.get(current_user.id, None)899 900            if user is not None:901                # This will not decrement if session expired902                user['session_count'] -= 1903                if user['session_count'] == 0:904                    del self.users[current_user.id]905 906    @login_required907    def hard_reset(self):908        with self.lock:909            user = self.users.get(current_user.id, None)910 911            if user is not None:912                del self.users[current_user.id]913 914 915def get_safe_post_login_redirect():916    allow_list = [917        url_for('browser.index')918    ]919    if "SCRIPT_NAME" in os.environ and os.environ["SCRIPT_NAME"]:920        allow_list.append(os.environ["SCRIPT_NAME"])921 922    url = get_post_login_redirect()923    for item in allow_list:924        if url.startswith(item):925            return url926 927    return url_for('browser.index')928 929 930def get_safe_post_logout_redirect():931    allow_list = [932        url_for('security.login')933    ]934    if "SCRIPT_NAME" in os.environ and os.environ["SCRIPT_NAME"]:935        allow_list.append(os.environ["SCRIPT_NAME"])936    url = get_post_logout_redirect()937    for item in allow_list:938        if url.startswith(item):939            return url940    return url_for('security.login')941 
codekingpro/portable-devtools · Team Ai