diff --git a/server/app/app.py b/server/app/app.py index 61b9b7ff..e290ca35 100644 --- a/server/app/app.py +++ b/server/app/app.py @@ -33,7 +33,7 @@ def _cache_control(always, **cache_kwargs): @wraps(f) def wrapper(*args, **kwargs): response = make_response(f(*args, **kwargs)) - if not always and not current_app.app_config.server__generate_cache_control_headers: + if not always and not current_app.app_config.server_config.app__generate_cache_control_headers: return response if response.status_code >= 400: return response @@ -61,25 +61,31 @@ def cache_control_always(**cache_kwargs): @webbp.route("/", methods=["GET"]) @cache_control_always(public=True, max_age=0, no_store=True, no_cache=True, must_revalidate=True) def dataset_index(url_dataroot=None, dataset=None): - config = current_app.app_config + app_config = current_app.app_config + server_config = app_config.server_config if dataset is None: - if config.single_dataset__datapath: - location = config.single_dataset__datapath - else: + if app_config.is_multi_dataset(): return dataroot_index() + else: + location = server_config.single_dataset__datapath else: - dataroot = config.multi_dataset__dataroot.get(url_dataroot) + dataroot = None + for key, dataroot_dict in server_config.multi_dataset__dataroot.items(): + if dataroot_dict["base_url"] == url_dataroot: + dataroot = dataroot_dict["dataroot"] + break if dataroot is None: abort(HTTPStatus.NOT_FOUND) location = path_join(dataroot, dataset) - scripts = config.server__scripts - inline_scripts = config.server__inline_scripts + dataset_config = app_config.get_dataset_config(url_dataroot) + scripts = dataset_config.app__scripts + inline_scripts = dataset_config.app__inline_scripts try: cache_manager = current_app.matrix_data_cache_manager - with cache_manager.data_adaptor(location, config) as data_adaptor: - dataset_title = config.get_title(data_adaptor) + with cache_manager.data_adaptor(url_dataroot, location, app_config) as data_adaptor: + dataset_title = app_config.get_title(data_adaptor) return render_template( "index.html", datasetTitle=dataset_title, SCRIPTS=scripts, INLINE_SCRIPTS=inline_scripts ) @@ -103,11 +109,19 @@ def handle_request_exception(error): def get_data_adaptor(url_dataroot=None, dataset=None): config = current_app.app_config + server_config = config.server_config + dataset_key = None if dataset is None: - datapath = config.single_dataset__datapath + datapath = server_config.single_dataset__datapath else: - dataroot = config.multi_dataset__dataroot.get(url_dataroot) + dataroot = None + for key, dataroot_dict in server_config.multi_dataset__dataroot.items(): + if dataroot_dict["base_url"] == url_dataroot: + dataroot = dataroot_dict["dataroot"] + dataset_key = key + break + if dataroot is None: raise DatasetAccessError(f"Invalid dataset {url_dataroot}/{dataset}") datapath = path_join(dataroot, dataset) @@ -121,7 +135,7 @@ def get_data_adaptor(url_dataroot=None, dataset=None): return common_rest.abort_and_log(HTTPStatus.BAD_REQUEST, "Invalid dataset NONE", loglevel=logging.INFO) cache_manager = current_app.matrix_data_cache_manager - return cache_manager.data_adaptor(datapath, config) + return cache_manager.data_adaptor(dataset_key, datapath, config) def rest_get_data_adaptor(func): @@ -145,9 +159,12 @@ def dataroot_test_index(): data += "

Welcome to cellxgene

" config = current_app.app_config + server_config = config.server_config datasets = [] - for url_dataroot, dataroot in config.multi_dataset__dataroot.items(): - locator = DataLocator(dataroot, region_name=config.data_locator__s3__region_name) + for dataroot_dict in server_config.multi_dataset__dataroot.values(): + dataroot = dataroot_dict["dataroot"] + url_dataroot = dataroot_dict["base_url"] + locator = DataLocator(dataroot, region_name=server_config.data_locator__s3__region_name) for fname in locator.ls(): location = path_join(dataroot, fname) try: @@ -171,12 +188,12 @@ def dataroot_test_index(): def dataroot_index(): # Handle the base url for the cellxgene server when running in multi dataset mode config = current_app.app_config - if not config.multi_dataset__index: + if not config.server_config.multi_dataset__index: abort(HTTPStatus.NOT_FOUND) - elif config.multi_dataset__index is True: + elif config.server_config.multi_dataset__index is True: return dataroot_test_index() else: - return redirect(config.multi_dataset__index) + return redirect(config.server_config.multi_dataset__index) class DatasetResource(Resource): @@ -191,33 +208,33 @@ class SchemaAPI(DatasetResource): @cache_control(public=True, max_age=ONE_WEEK) @rest_get_data_adaptor def get(self, data_adaptor): - return common_rest.schema_get(data_adaptor, current_app.annotations) + return common_rest.schema_get(data_adaptor) class ConfigAPI(DatasetResource): @cache_control(public=True, max_age=ONE_WEEK) @rest_get_data_adaptor def get(self, data_adaptor): - return common_rest.config_get(current_app.app_config, data_adaptor, current_app.annotations) + return common_rest.config_get(current_app.app_config, data_adaptor) class AnnotationsObsAPI(DatasetResource): @cache_control(public=True, max_age=ONE_WEEK) @rest_get_data_adaptor def get(self, data_adaptor): - return common_rest.annotations_obs_get(request, data_adaptor, current_app.annotations) + return common_rest.annotations_obs_get(request, data_adaptor) @cache_control(no_store=True) @rest_get_data_adaptor def put(self, data_adaptor): - return common_rest.annotations_obs_put(request, data_adaptor, current_app.annotations) + return common_rest.annotations_obs_put(request, data_adaptor) class AnnotationsVarAPI(DatasetResource): @cache_control(public=True, max_age=ONE_WEEK) @rest_get_data_adaptor def get(self, data_adaptor): - return common_rest.annotations_var_get(request, data_adaptor, current_app.annotations) + return common_rest.annotations_var_get(request, data_adaptor) class DataVarAPI(DatasetResource): @@ -290,29 +307,26 @@ class Server: self.app = Flask(__name__, static_folder="../common/web/static") self._before_adding_routes(self.app, app_config) self.app.json_encoder = Float32JSONEncoder - if app_config.server__server_timing_headers: + server_config = app_config.server_config + if server_config.app__server_timing_headers: ServerTiming(self.app, force_debug=True) # enable session data self.app.permanent_session_lifetime = datetime.timedelta(days=50 * 365) # Config - secret_key = app_config.server__flask_secret_key + secret_key = server_config.app__flask_secret_key self.app.config.update(SECRET_KEY=secret_key) self.app.register_blueprint(webbp) api_version = "/api/v0.2" - if app_config.single_dataset__datapath: - bp_api = Blueprint("api", __name__, url_prefix=api_version) - resources = get_api_resources(bp_api) - self.app.register_blueprint(resources.blueprint) - - else: + if app_config.is_multi_dataset(): # NOTE: These routes only allow the dataset to be in the directory # of the dataroot, and not a subdirectory. We may want to change # the route format at some point - for url_dataroot in app_config.multi_dataset__dataroot.keys(): + for dataroot_dict in server_config.multi_dataset__dataroot.values(): + url_dataroot = dataroot_dict["base_url"] bp_api = Blueprint( f"api_dataset_{url_dataroot}", __name__, url_prefix=f"/{url_dataroot}/" + api_version ) @@ -321,9 +335,13 @@ class Server: self.app.add_url_rule( f"/{url_dataroot}//", f"dataset_index_{url_dataroot}", - lambda dataset: dataset_index(url_dataroot, dataset), + lambda dataset, url_dataroot=url_dataroot: dataset_index(url_dataroot, dataset), methods=["GET"], ) - self.app.matrix_data_cache_manager = app_config.matrix_data_cache_manager - self.app.annotations = app_config.user_annotations + else: + bp_api = Blueprint("api", __name__, url_prefix=api_version) + resources = get_api_resources(bp_api) + self.app.register_blueprint(resources.blueprint) + + self.app.matrix_data_cache_manager = server_config.matrix_data_cache_manager self.app.app_config = app_config diff --git a/server/cli/launch.py b/server/cli/launch.py index a0cf1643..00900092 100644 --- a/server/cli/launch.py +++ b/server/cli/launch.py @@ -22,13 +22,13 @@ def annotation_args(func): @click.option( "--disable-annotations", is_flag=True, - default=not DEFAULT_CONFIG.user_annotations__enable, + default=not DEFAULT_CONFIG.default_dataset_config.user_annotations__enable, show_default=True, help="Disable user annotation of data.", ) @click.option( "--annotations-file", - default=DEFAULT_CONFIG.user_annotations__local_file_csv__file, + default=DEFAULT_CONFIG.default_dataset_config.user_annotations__local_file_csv__file, show_default=True, multiple=False, metavar="", @@ -37,7 +37,7 @@ def annotation_args(func): ) @click.option( "--annotations-dir", - default=DEFAULT_CONFIG.user_annotations__local_file_csv__directory, + default=DEFAULT_CONFIG.default_dataset_config.user_annotations__local_file_csv__directory, show_default=False, multiple=False, metavar="", @@ -47,13 +47,13 @@ def annotation_args(func): @click.option( "--experimental-annotations-ontology", is_flag=True, - default=DEFAULT_CONFIG.user_annotations__ontology__enable, + default=DEFAULT_CONFIG.default_dataset_config.user_annotations__ontology__enable, show_default=True, help="When creating annotations, optionally autocomplete names from ontology terms.", ) @click.option( "--experimental-annotations-ontology-obo", - default=DEFAULT_CONFIG.user_annotations__ontology__obo_location, + default=DEFAULT_CONFIG.default_dataset_config.user_annotations__ontology__obo_location, show_default=True, metavar="", help="Location of OBO file defining cell annotation autosuggest terms.", @@ -68,7 +68,7 @@ def annotation_args(func): def config_args(func): @click.option( "--max-category-items", - default=DEFAULT_CONFIG.presentation__max_categories, + default=DEFAULT_CONFIG.default_dataset_config.presentation__max_categories, metavar="", show_default=True, help="Will not display categories with more distinct values than specified.", @@ -83,7 +83,7 @@ def config_args(func): @click.option( "--diffexp-lfc-cutoff", "-de", - default=DEFAULT_CONFIG.diffexp__lfc_cutoff, + default=DEFAULT_CONFIG.default_dataset_config.diffexp__lfc_cutoff, show_default=True, metavar="", help="Minimum log fold change threshold for differential expression.", @@ -91,14 +91,14 @@ def config_args(func): @click.option( "--disable-diffexp", is_flag=True, - default=not DEFAULT_CONFIG.diffexp__enable, + default=not DEFAULT_CONFIG.default_dataset_config.diffexp__enable, show_default=False, help="Disable on-demand differential expression.", ) @click.option( "--embedding", "-e", - default=DEFAULT_CONFIG.embeddings__names, + default=DEFAULT_CONFIG.default_dataset_config.embeddings__names, multiple=True, show_default=False, metavar="", @@ -107,7 +107,7 @@ def config_args(func): @click.option( "--experimental-enable-reembedding", is_flag=True, - default=DEFAULT_CONFIG.embeddings__enable_reembedding, + default=DEFAULT_CONFIG.default_dataset_config.embeddings__enable_reembedding, show_default=False, hidden=True, help="Enable experimental on-demand re-embedding using UMAP. WARNING: may be very slow.", @@ -123,14 +123,14 @@ def dataset_args(func): @click.option( "--obs-names", "-obs", - default=DEFAULT_CONFIG.single_dataset__obs_names, + default=DEFAULT_CONFIG.server_config.single_dataset__obs_names, metavar="", help="Name of annotation field to use for observations. If not specified cellxgene will use the the obs index.", ) @click.option( "--var-names", "-var", - default=DEFAULT_CONFIG.single_dataset__var_names, + default=DEFAULT_CONFIG.server_config.single_dataset__var_names, metavar="", help="Name of annotation to use for variables. If not specified cellxgene will use the the var index.", ) @@ -138,20 +138,20 @@ def dataset_args(func): "--backed", "-b", is_flag=True, - default=DEFAULT_CONFIG.adaptor__anndata_adaptor__backed, + default=DEFAULT_CONFIG.server_config.adaptor__anndata_adaptor__backed, show_default=False, help="Load anndata in file-backed mode. " "This may save memory, but may result in slower overall performance.", ) @click.option( "--title", "-t", - default=DEFAULT_CONFIG.single_dataset__title, + default=DEFAULT_CONFIG.server_config.single_dataset__title, metavar="", help="Title to display. If omitted will use file name.", ) @click.option( "--about", - default=DEFAULT_CONFIG.single_dataset__about, + default=DEFAULT_CONFIG.server_config.single_dataset__about, metavar="", help="URL providing more information about the dataset (hint: must be a fully specified absolute URL).", ) @@ -167,7 +167,7 @@ def server_args(func): "--debug", "-d", is_flag=True, - default=DEFAULT_CONFIG.server__debug, + default=DEFAULT_CONFIG.server_config.app__debug, show_default=True, help="Run in debug mode. This is helpful for cellxgene developers, " "or when you want more information about an error condition.", @@ -176,7 +176,7 @@ def server_args(func): "--verbose", "-v", is_flag=True, - default=DEFAULT_CONFIG.server__verbose, + default=DEFAULT_CONFIG.server_config.app__verbose, show_default=True, help="Provide verbose output, including warnings and all server requests.", ) @@ -184,7 +184,7 @@ def server_args(func): "--port", "-p", metavar="", - default=DEFAULT_CONFIG.server__port, + default=DEFAULT_CONFIG.server_config.app__port, type=int, show_default=True, help="Port to run server on. If not specified cellxgene will find an available port.", @@ -192,14 +192,14 @@ def server_args(func): @click.option( "--host", metavar="", - default=DEFAULT_CONFIG.server__host, + default=DEFAULT_CONFIG.server_config.app__host, show_default=False, help="Host IP address. By default cellxgene will use localhost (e.g. 127.0.0.1).", ) @click.option( "--scripts", "-s", - default=DEFAULT_CONFIG.server__scripts, + default=DEFAULT_CONFIG.default_dataset_config.app__scripts, multiple=True, metavar="", help="Additional script files to include in HTML page. If not specified, " @@ -220,7 +220,7 @@ def launch_args(func): @server_args @click.option( "--dataroot", - default=DEFAULT_CONFIG.multi_dataset__dataroot, + default=DEFAULT_CONFIG.server_config.multi_dataset__dataroot, metavar="", help="Enable cellxgene to serve multiple files. Supply path (local directory or URL)" " to folder containing H5AD and/or CXG datasets.", @@ -232,7 +232,7 @@ def launch_args(func): "-o", "open_browser", is_flag=True, - default=DEFAULT_CONFIG.server__open_browser, + default=DEFAULT_CONFIG.server_config.app__open_browser, show_default=True, help="Open web browser after launch.", ) @@ -296,7 +296,7 @@ class CliLaunchServer(Server): "application/octet-stream", ] Compress(app) - if app_config.server__debug: + if app_config.server_config.app__debug: CORS(app, supports_credentials=True) @@ -362,6 +362,7 @@ def launch( # app config app_config = AppConfig() + server_config = app_config.server_config try: if config_file: @@ -370,19 +371,22 @@ def launch( # Determine which config options were give on the command line. # Those will override the ones provided in the config file (if provided). cli_config = AppConfig() - cli_config.update( - server__verbose=verbose, - server__debug=debug, - server__host=host, - server__port=port, - server__scripts=scripts, - server__open_browser=open_browser, + cli_config.update_server_config( + app__verbose=verbose, + app__debug=debug, + app__host=host, + app__port=port, + app__open_browser=open_browser, single_dataset__datapath=datapath, single_dataset__title=title, single_dataset__about=about, single_dataset__obs_names=obs_names, single_dataset__var_names=var_names, multi_dataset__dataroot=dataroot, + adaptor__anndata_adaptor__backed=backed, + ) + cli_config.update_default_dataset_config( + app__scripts=scripts, user_annotations__enable=not disable_annotations, user_annotations__local_file_csv__file=annotations_file, user_annotations__local_file_csv__directory=annotations_dir, @@ -394,13 +398,15 @@ def launch( embeddings__enable_reembedding=experimental_enable_reembedding, diffexp__enable=not disable_diffexp, diffexp__lfc_cutoff=diffexp_lfc_cutoff, - adaptor__anndata_adaptor__backed=backed, ) - diff = cli_config.changes_from_default() - changes = {} - for key, val, defval in diff: - changes[key] = val - app_config.update(**changes) + + diff = cli_config.server_config.changes_from_default() + changes = {key: val for key, val, _ in diff} + app_config.update_server_config(**changes) + + diff = cli_config.default_dataset_config.changes_from_default() + changes = {key: val for key, val, _ in diff} + app_config.update_default_dataset_config(**changes) # process the configuration # any errors will be thrown as an exception. @@ -410,8 +416,8 @@ def launch( click.echo("[cellxgene] " + message) # Use a default secret if one is not provided - if not app_config.server__flask_secret_key: - app_config.update(server__flask_secret_key="SparkleAndShine") + if not server_config.app__flask_secret_key: + app_config.update_server_config(app__flask_secret_key="SparkleAndShine") app_config.complete_config(messagefn) @@ -423,12 +429,12 @@ def launch( # create the server server = CliLaunchServer(app_config) - if not app_config.server__verbose: + if not server_config.app__verbose: log = logging.getLogger("werkzeug") log.setLevel(logging.ERROR) - cellxgene_url = f"http://{app_config.server__host}:{app_config.server__port}" - if app_config.server__open_browser: + cellxgene_url = f"http://{app_config.server_config.app__host}:{app_config.server_config.app__port}" + if server_config.app__open_browser: click.echo(f"[cellxgene] Launching! Opening your browser to {cellxgene_url} now.") webbrowser.open(cellxgene_url) else: @@ -436,16 +442,16 @@ def launch( click.echo("[cellxgene] Type CTRL-C at any time to exit.") - if not app_config.server__verbose: + if not server_config.app__verbose: f = open(devnull, "w") sys.stdout = f try: server.app.run( - host=app_config.server__host, - debug=app_config.server__debug, - port=app_config.server__port, - threaded=not app_config.server__debug, + host=server_config.app__host, + debug=server_config.app__debug, + port=server_config.app__port, + threaded=not server_config.app__debug, use_debugger=False, use_reloader=False, ) diff --git a/server/common/app_config.py b/server/common/app_config.py index d0510b6a..38b47fc2 100644 --- a/server/common/app_config.py +++ b/server/common/app_config.py @@ -38,170 +38,110 @@ class AppFeature(object): class AppConfig(object): + """AppConfig stores all the configuration for cellxgene. The configuration is divided into two main parts: + server attributes, and dataset attributes. The server_config contains attributes that refer to the server process + as a whole. The default_dataset_config referes to attributes that are associated with the features and + presentations of a dataset. The dataset config attributes can be overridden depending on the url by which the + dataset was accessed. These are stored in dataroot_config. + AppConfig has methods to initialize, modify, and access the configuration. + """ + def __init__(self): + # the default configuration (see default_config.py) self.default_config = get_default_config() - self.attr_checked = {k: False for k in self.__mapping(self.default_config).keys()} - - dc = self.default_config - try: - self.server__verbose = dc["server"]["verbose"] - self.server__debug = dc["server"]["debug"] - self.server__host = dc["server"]["host"] - self.server__port = dc["server"]["port"] - self.server__scripts = dc["server"]["scripts"] - self.server__inline_scripts = dc["server"]["inline_scripts"] - self.server__open_browser = dc["server"]["open_browser"] - self.server__about_legal_tos = dc["server"]["about_legal_tos"] - self.server__about_legal_privacy = dc["server"]["about_legal_privacy"] - self.server__force_https = dc["server"]["force_https"] - self.server__flask_secret_key = dc["server"]["flask_secret_key"] - self.server__generate_cache_control_headers = dc["server"]["generate_cache_control_headers"] - self.server__server_timing_headers = dc["server"]["server_timing_headers"] - self.server__csp_directives = dc["server"]["csp_directives"] - - self.multi_dataset__dataroot = dc["multi_dataset"]["dataroot"] - self.multi_dataset__index = dc["multi_dataset"]["index"] - self.multi_dataset__allowed_matrix_types = dc["multi_dataset"]["allowed_matrix_types"] - self.multi_dataset__matrix_cache__max_datasets = dc["multi_dataset"]["matrix_cache"]["max_datasets"] - self.multi_dataset__matrix_cache__timelimit_s = dc["multi_dataset"]["matrix_cache"]["timelimit_s"] - - self.single_dataset__datapath = dc["single_dataset"]["datapath"] - self.single_dataset__obs_names = dc["single_dataset"]["obs_names"] - self.single_dataset__var_names = dc["single_dataset"]["var_names"] - self.single_dataset__about = dc["single_dataset"]["about"] - self.single_dataset__title = dc["single_dataset"]["title"] - - self.user_annotations__enable = dc["user_annotations"]["enable"] - self.user_annotations__type = dc["user_annotations"]["type"] - self.user_annotations__local_file_csv__directory = dc["user_annotations"]["local_file_csv"]["directory"] - self.user_annotations__local_file_csv__file = dc["user_annotations"]["local_file_csv"]["file"] - self.user_annotations__ontology__enable = dc["user_annotations"]["ontology"]["enable"] - self.user_annotations__ontology__obo_location = dc["user_annotations"]["ontology"]["obo_location"] - - self.presentation__max_categories = dc["presentation"]["max_categories"] - self.presentation__custom_colors = dc["presentation"]["custom_colors"] - - self.embeddings__names = dc["embeddings"]["names"] - self.embeddings__enable_reembedding = dc["embeddings"]["enable_reembedding"] - - self.diffexp__enable = dc["diffexp"]["enable"] - self.diffexp__lfc_cutoff = dc["diffexp"]["lfc_cutoff"] - self.diffexp__top_n = dc["diffexp"]["top_n"] - self.diffexp__alg_cxg__max_workers = dc["diffexp"]["alg_cxg"]["max_workers"] - self.diffexp__alg_cxg__cpu_multiplier = dc["diffexp"]["alg_cxg"]["cpu_multiplier"] - self.diffexp__alg_cxg__target_workunit = dc["diffexp"]["alg_cxg"]["target_workunit"] - - self.data_locator__s3__region_name = dc["data_locator"]["s3"]["region_name"] - - self.adaptor__cxg_adaptor__tiledb_ctx = dc["adaptor"]["cxg_adaptor"]["tiledb_ctx"] - self.adaptor__anndata_adaptor__backed = dc["adaptor"]["anndata_adaptor"]["backed"] - - self.limits__diffexp_cellcount_max = dc["limits"]["diffexp_cellcount_max"] - self.limits__column_request_max = dc["limits"]["column_request_max"] - - except KeyError as e: - raise ConfigurationError(f"Unexpected config: {str(e)}") - - # The annotation object is created during complete_config and stored here. - self.user_annotations = None - - # The matrix data cache manager is created during the complete_config and stored here. - self.matrix_data_cache_manager = None + # the server configuration + self.server_config = ServerConfig(self, self.default_config["server"]) + # the dataset config, unless overridden by an entry in dataroot_config + self.default_dataset_config = DatasetConfig(None, self, self.default_config["dataset"]) + # a dictionary of keys to DatasetConfig objects. Each key must exist in the multi_dataset__dataroot + # attribute of the server_config. + self.dataroot_config = {} # Set to true when config_completed is called self.is_completed = False + def get_dataset_config(self, dataroot_key): + if self.server_config.single_dataset__datapath: + return self.default_dataset_config + else: + return self.dataroot_config.get(dataroot_key, self.default_dataset_config) + def check_config(self): + """Verify all the attributes have been checked""" if not self.is_completed: raise ConfigurationError("The configuration has not been completed") - mapping = self.__mapping(self.default_config) - for key in mapping.keys(): - if not self.attr_checked[key]: - raise ConfigurationError(f"The attr '{key}' has not been checked") + self.server_config.check_config() + self.default_dataset_config.check_config() + for dataset_config in self.dataroot_config.values(): + dataset_config.check_config() - def __mapping(self, config): - """Create a mapping from attribute names to (location in the config tree, value)""" - dc = copy.deepcopy(config) - mapping = {} + def update_server_config(self, **kw): + self.server_config.update(**kw) + self.is_complete = False - # special cases where the value could be a dict. - # If its value is not None, the entry is added to the mapping, and not included - # in the flattening below. - dictval_cases = [ - ("adaptor", "cxg_adaptor", "tiledb_ctx"), - ("server", "csp_directives"), - ("multi_dataset", "dataroot"), - ] - for dictval_case in dictval_cases: - cur = dc - for part in dictval_case[:-1]: - cur = cur.get(part, {}) - val = cur.get(dictval_case[-1]) - if val is not None: - key = "__".join(dictval_case) - mapping[key] = (dictval_case, val) - del cur[dictval_case[-1]] - - flat_config = flatten(dc) - for key, value in flat_config.items(): - # name of the attribute - attr = "__".join(key) - mapping[attr] = (key, value) - - return mapping + def update_default_dataset_config(self, **kw): + self.default_dataset_config.update(**kw) + # update all the other dataset configs, if any + for value in self.dataroot_config.values(): + value.update(**kw) + self.is_complete = False def update_from_config_file(self, config_file): with open(config_file) as fyaml: config = yaml.load(fyaml, Loader=yaml.FullLoader) - mapping = self.__mapping(config) - for attr, (key, value) in mapping.items(): - if not hasattr(self, attr): - raise ConfigurationError(f"Unknown key from config file: {key}") - try: - setattr(self, attr, value) - except KeyError: - raise ConfigurationError(f"Unable to set config attribute: {key}") + self.server_config.update_from_config(config["server"], "server") + self.default_dataset_config.update_from_config(config["dataset"], "dataset") - self.attr_checked[attr] = False + per_dataset_config = config.get("per_dataset_config", {}) + for key, dataroot_config in per_dataset_config.items(): + self.add_dataroot_config(key, **dataroot_config) - self.is_completed = False + self.is_complete = False def write_config(self, config_file): """output the config to a yaml file""" - mapping = self.__mapping(self.default_config) - for attrname in mapping.keys(): - mapping[attrname] = getattr(self, attrname) - config = unflatten(mapping, splitter=lambda key: key.split("__")) + server = self.server_config.create_mapping(self.server_config.default_config) + dataset = self.default_dataset_config.create_mapping(self.default_dataset_config.default_config) + config = dict(server={}, dataset={}) + for attrname in server.keys(): + config["server__" + attrname] = getattr(self.server_config, attrname) + for attrname in dataset.keys(): + config["dataset__" + attrname] = getattr(self.default_dataset_config, attrname) + if self.dataroot_config: + config["per_dataset_config"] = {} + for dataroot_tag, dataroot_config in self.dataroot_config.items(): + dataset = dataroot_config.create_mapping(dataroot_config.default_config) + for attrname in dataset.keys(): + config[f"per_dataset_config__{dataroot_tag}__" + attrname] = getattr(dataroot_config, attrname) + + config = unflatten(config, splitter=lambda key: key.split("__")) yaml.dump(config, open(config_file, "w")) - def update(self, **kw): - for key, value in kw.items(): - if not hasattr(self, key): - raise ConfigurationError(f"unknown config parameter {key}.") - try: - if type(value) == tuple: - # convert tuple values to list values - value = list(value) - setattr(self, key, value) - except KeyError: - raise ConfigurationError(f"Unable to set config parameter {key}.") - - self.attr_checked[key] = False - - self.is_completed = False - def changes_from_default(self): """Return all the attribute that are different from the default""" - mapping = self.__mapping(self.default_config) - diff = [] - for attrname, (key, defval) in mapping.items(): - curval = getattr(self, attrname) - if curval != defval: - diff.append((attrname, curval, defval)) + diff_server = self.server_config.changes_from_default() + diff_dataset = self.default_dataset_config.changes_from_default() + diff = dict(server=diff_server, dataset=diff_dataset) return diff + def add_dataroot_config(self, dataroot_tag, **kw): + """Create a new dataset config object based on the default dataset config, and kw parameters""" + if dataroot_tag in self.dataroot_config: + raise ConfigurationError(f"dataroot config already exists: {dataroot_tag}") + if type(self.server_config.multi_dataset__dataroot) != dict: + raise ConfigurationError("The server__multi_dataset__dataroot must be a dictionary") + if dataroot_tag not in self.server_config.multi_dataset__dataroot: + raise ConfigurationError(f"The dataroot_tag ({dataroot_tag}) not found in server__multi_dataset__dataroot") + + self.is_completed = False + self.dataroot_config[dataroot_tag] = DatasetConfig(dataroot_tag, self, self.default_config["dataset"]) + flat_config = self.default_dataset_config.create_mapping(self.default_dataset_config.default_config) + config = {key: value[1] for key, value in flat_config.items()} + self.dataroot_config[dataroot_tag].update(**config) + self.dataroot_config[dataroot_tag].update_from_config(kw, dataroot_tag) + def complete_config(self, messagefn=None): """The configure options are checked, and any additional setup based on the config parameters is done""" @@ -218,22 +158,149 @@ class AppConfig(object): # messages we can give correct context for attributes with bad value. context = dict(messagefn=messagefn) - self.handle_server(context) - self.handle_adaptor(context) - self.handle_data_locator(context) - self.handle_adaptor(context) # may depend on data_locator - self.handle_presentation(context) - self.handle_single_dataset(context) # may depend on adaptor - self.handle_multi_dataset(context) # may depend on adaptor - self.handle_user_annotations(context) - self.handle_embeddings(context) - self.handle_diffexp(context) - self.handle_limits(context) + self.server_config.complete_config(context) + self.default_dataset_config.complete_config(context) + for dataroot_config in self.dataroot_config.values(): + dataroot_config.complete_config(context) self.is_completed = True self.check_config() - def __check_attr(self, attrname, vtype): + def get_matrix_data_cache_manager(self): + return self.server_config.matrix_data_cache_manager + + def is_multi_dataset(self): + return self.server_config.multi_dataset__dataroot is not None + + def get_title(self, data_adaptor): + return ( + self.server_config.single_dataset__title + if self.server_config.single_dataset__title + else data_adaptor.get_title() + ) + + def get_about(self, data_adaptor): + return ( + self.server_config.single_dataset__about + if self.server_config.single_dataset__about + else data_adaptor.get_about() + ) + + def get_client_config(self, data_adaptor): + """ + Return the configuration as required by the /config REST route + """ + + server_config = self.server_config + dataset_config = data_adaptor.dataset_config + annotation = dataset_config.user_annotations + + # FIXME The current set of config is not consistently presented: + # we have camalCase, hyphen-text, and underscore_text + + # make sure the configuration has been checked. + self.check_config() + + # features + features = [f.todict() for f in data_adaptor.get_features(annotation)] + + # display_names + title = self.get_title(data_adaptor) + about = self.get_about(data_adaptor) + + display_names = dict(engine=data_adaptor.get_name(), dataset=title) + + # library_versions + library_versions = {} + library_versions.update(data_adaptor.get_library_versions()) + library_versions["cellxgene"] = cellxgene_version + + # links + links = {"about-dataset": about} + + # parameters + parameters = { + "layout": dataset_config.embeddings__names, + "max-category-items": dataset_config.presentation__max_categories, + "obs_names": server_config.single_dataset__obs_names, + "var_names": server_config.single_dataset__var_names, + "diffexp_lfc_cutoff": dataset_config.diffexp__lfc_cutoff, + "backed": server_config.adaptor__anndata_adaptor__backed, + "disable-diffexp": not dataset_config.diffexp__enable, + "enable-reembedding": dataset_config.embeddings__enable_reembedding, + "annotations": False, + "annotations_file": None, + "annotations_dir": None, + "annotations_cell_ontology_enabled": False, + "annotations_cell_ontology_obopath": None, + "annotations_cell_ontology_terms": None, + "custom_colors": dataset_config.presentation__custom_colors, + "diffexp-may-be-slow": False, + "about_legal_tos": dataset_config.app__about_legal_tos, + "about_legal_privacy": dataset_config.app__about_legal_privacy, + } + + data_adaptor.update_parameters(parameters) + if annotation: + annotation.update_parameters(parameters, data_adaptor) + + # gather it all together + c = {} + config = c["config"] = {} + config["features"] = features + config["displayNames"] = display_names + config["library_versions"] = library_versions + config["links"] = links + config["parameters"] = parameters + config["limits"] = { + "column_request_max": server_config.limits__column_request_max, + "diffexp_cellcount_max": server_config.limits__diffexp_cellcount_max, + } + + return c + + +class BaseConfig(object): + """This class handles the mechanics of updating and checking attributes. + Derived classes are expected to store the actual attributes""" + + def __init__(self, app_config, default_config, dictval_cases={}): + # reference back to the app_config + self.app_config = app_config + # the complete set of attribute and their default values (unflattened) + self.default_config = default_config + # attributes where the value may be a dict (and therefore are not flattened) + self.dictval_cases = dictval_cases + # used to make sure every attribute value is checked + self.attr_checked = {k: False for k in self.create_mapping(default_config).keys()} + + def create_mapping(self, config): + """Create a mapping from attribute names to (location in the config tree, value)""" + dc = copy.deepcopy(config) + mapping = {} + + # special cases where the value could be a dict. + # If its value is not None, the entry is added to the mapping, and not included + # in the flattening below. + for dictval_case in self.dictval_cases: + cur = dc + for part in dictval_case[:-1]: + cur = cur.get(part, {}) + val = cur.get(dictval_case[-1]) + if val is not None: + key = "__".join(dictval_case) + mapping[key] = (dictval_case, val) + del cur[dictval_case[-1]] + + flat_config = flatten(dc) + for key, value in flat_config.items(): + # name of the attribute + attr = "__".join(key) + mapping[attr] = (key, value) + + return mapping + + def check_attr(self, attrname, vtype): val = getattr(self, attrname) if type(vtype) in (list, tuple): if type(val) not in vtype: @@ -250,48 +317,152 @@ class AppConfig(object): self.attr_checked[attrname] = True - def handle_server(self, context): - self.__check_attr("server__verbose", bool) - self.__check_attr("server__debug", bool) - self.__check_attr("server__host", str) - self.__check_attr("server__port", (type(None), int)) - self.__check_attr("server__scripts", list) - self.__check_attr("server__inline_scripts", list) - self.__check_attr("server__open_browser", bool) - self.__check_attr("server__force_https", bool) - self.__check_attr("server__flask_secret_key", (type(None), str)) - self.__check_attr("server__generate_cache_control_headers", bool) - self.__check_attr("server__about_legal_tos", (type(None), str)) - self.__check_attr("server__about_legal_privacy", (type(None), str)) - self.__check_attr("server__server_timing_headers", bool) - self.__check_attr("server__csp_directives", (type(None), dict)) + def check_config(self): + mapping = self.create_mapping(self.default_config) + for key in mapping.keys(): + if not self.attr_checked[key]: + raise ConfigurationError(f"The attr '{key}' has not been checked") - if self.server__port: - if not is_port_available(self.server__host, self.server__port): + def update(self, **kw): + for key, value in kw.items(): + if not hasattr(self, key): + raise ConfigurationError(f"unknown config parameter {key}.") + try: + if type(value) == tuple: + # convert tuple values to list values + value = list(value) + setattr(self, key, value) + except KeyError: + raise ConfigurationError(f"Unable to set config parameter {key}.") + + self.attr_checked[key] = False + + def update_from_config(self, config, prefix): + mapping = self.create_mapping(config) + for attr, (key, value) in mapping.items(): + if not hasattr(self, attr): + raise ConfigurationError(f"Unknown key from config file: {prefix}__{attr}") + try: + setattr(self, attr, value) + except KeyError: + raise ConfigurationError(f"Unable to set config attribute: {prefix}__{attr}") + + self.attr_checked[attr] = False + + def changes_from_default(self): + """Return all the attribute that are different from the default""" + mapping = self.create_mapping(self.default_config) + diff = [] + for attrname, (key, defval) in mapping.items(): + curval = getattr(self, attrname) + if curval != defval: + diff.append((attrname, curval, defval)) + return diff + + +class ServerConfig(BaseConfig): + """Manages the config attribute associated with the server.""" + + def __init__(self, app_config, default_config): + dictval_cases = [ + ("app", "csp_directives"), + ("adaptor", "cxg_adaptor", "tiledb_ctx"), + ("multi_dataset", "dataroot"), + ] + super().__init__(app_config, default_config, dictval_cases) + + dc = default_config + try: + self.app__verbose = dc["app"]["verbose"] + self.app__debug = dc["app"]["debug"] + self.app__host = dc["app"]["host"] + self.app__port = dc["app"]["port"] + self.app__open_browser = dc["app"]["open_browser"] + self.app__force_https = dc["app"]["force_https"] + self.app__flask_secret_key = dc["app"]["flask_secret_key"] + self.app__generate_cache_control_headers = dc["app"]["generate_cache_control_headers"] + self.app__server_timing_headers = dc["app"]["server_timing_headers"] + self.app__csp_directives = dc["app"]["csp_directives"] + + self.multi_dataset__dataroot = dc["multi_dataset"]["dataroot"] + self.multi_dataset__index = dc["multi_dataset"]["index"] + self.multi_dataset__allowed_matrix_types = dc["multi_dataset"]["allowed_matrix_types"] + self.multi_dataset__matrix_cache__max_datasets = dc["multi_dataset"]["matrix_cache"]["max_datasets"] + self.multi_dataset__matrix_cache__timelimit_s = dc["multi_dataset"]["matrix_cache"]["timelimit_s"] + + self.single_dataset__datapath = dc["single_dataset"]["datapath"] + self.single_dataset__obs_names = dc["single_dataset"]["obs_names"] + self.single_dataset__var_names = dc["single_dataset"]["var_names"] + self.single_dataset__about = dc["single_dataset"]["about"] + self.single_dataset__title = dc["single_dataset"]["title"] + + self.diffexp__alg_cxg__max_workers = dc["diffexp"]["alg_cxg"]["max_workers"] + self.diffexp__alg_cxg__cpu_multiplier = dc["diffexp"]["alg_cxg"]["cpu_multiplier"] + self.diffexp__alg_cxg__target_workunit = dc["diffexp"]["alg_cxg"]["target_workunit"] + + self.data_locator__s3__region_name = dc["data_locator"]["s3"]["region_name"] + + self.adaptor__cxg_adaptor__tiledb_ctx = dc["adaptor"]["cxg_adaptor"]["tiledb_ctx"] + self.adaptor__anndata_adaptor__backed = dc["adaptor"]["anndata_adaptor"]["backed"] + + self.limits__diffexp_cellcount_max = dc["limits"]["diffexp_cellcount_max"] + self.limits__column_request_max = dc["limits"]["column_request_max"] + + except KeyError as e: + raise ConfigurationError(f"Unexpected config: {str(e)}") + + # The matrix data cache manager is created during the complete_config and stored here. + self.matrix_data_cache_manager = None + + def complete_config(self, context): + self.handle_app(context) + self.handle_data_locator(context) + self.handle_adaptor(context) # may depend on data_locator + self.handle_single_dataset(context) # may depend on adaptor + self.handle_multi_dataset(context) # may depend on adaptor + self.handle_diffexp(context) + self.handle_limits(context) + + self.check_config() + + def handle_app(self, context): + self.check_attr("app__verbose", bool) + self.check_attr("app__debug", bool) + self.check_attr("app__host", str) + self.check_attr("app__port", (type(None), int)) + self.check_attr("app__open_browser", bool) + self.check_attr("app__force_https", bool) + self.check_attr("app__flask_secret_key", (type(None), str)) + self.check_attr("app__generate_cache_control_headers", bool) + self.check_attr("app__server_timing_headers", bool) + self.check_attr("app__csp_directives", (type(None), dict)) + + if self.app__port: + if not is_port_available(self.app__host, self.app__port): raise ConfigurationError( - f"The port selected {self.server__port} is in use, please configure an open port." + f"The port selected {self.app__port} is in use, please configure an open port." ) else: - self.server__port = find_available_port(self.server__host, DEFAULT_SERVER_PORT) + self.app__port = find_available_port(self.app__host, DEFAULT_SERVER_PORT) - if self.server__debug: + if self.app__debug: context["messagefn"]("in debug mode, setting verbose=True and open_browser=False") - self.server__verbose = True - self.server__open_browser = False + self.app__verbose = True + self.app__open_browser = False else: warnings.formatwarning = custom_format_warning - if not self.server__verbose: + if not self.app__verbose: sys.tracebacklimit = 0 # secret key: # first, from CXG_SECRET_KEY environment variable # second, from config file - self.server__flask_secret_key = os.environ.get("CXG_SECRET_KEY", self.server__flask_secret_key) + self.app__flask_secret_key = os.environ.get("CXG_SECRET_KEY", self.app__flask_secret_key) # CSP Directives are a dict of string: list(string) or string: string - if self.server__csp_directives is not None: - for k, v in self.server__csp_directives.items(): + if self.app__csp_directives is not None: + for k, v in self.app__csp_directives.items(): if not isinstance(k, str): raise ConfigurationError("CSP directive names must be a string.") if isinstance(v, list): @@ -301,26 +472,15 @@ class AppConfig(object): elif not isinstance(v, str): raise ConfigurationError("CSP directive value must be a string or list of strings.") - # scripts can be string (filename) or dict (attributes). Convert string to dict. - scripts = [] - for s in self.server__scripts: - if isinstance(s, str): - scripts.append({"src": s}) - elif isinstance(s, dict) and isinstance(s["src"], str): - scripts.append(s) - else: - raise ConfigurationError("Scripts must be string or dict") - self.server__scripts = scripts - def handle_data_locator(self, context): - self.__check_attr("data_locator__s3__region_name", (type(None), bool, str)) + self.check_attr("data_locator__s3__region_name", (type(None), bool, str)) if self.data_locator__s3__region_name is True: path = self.single_dataset__datapath or self.multi_dataset__dataroot if type(path) == dict: # if multi_dataset__dataroot is a dict, then use the first key # that is in s3. NOTE: it is not supported to have dataroots # in different regions. - paths = path.values() + paths = [val.get("dataroot") for val in path.values()] for path in paths: if path.startswith("s3://"): break @@ -332,16 +492,12 @@ class AppConfig(object): region_name = None self.data_locator__s3__region_name = region_name - def handle_presentation(self, context): - self.__check_attr("presentation__max_categories", int) - self.__check_attr("presentation__custom_colors", bool) - def handle_single_dataset(self, context): - self.__check_attr("single_dataset__datapath", (str, type(None))) - self.__check_attr("single_dataset__title", (str, type(None))) - self.__check_attr("single_dataset__about", (str, type(None))) - self.__check_attr("single_dataset__obs_names", (str, type(None))) - self.__check_attr("single_dataset__var_names", (str, type(None))) + self.check_attr("single_dataset__datapath", (str, type(None))) + self.check_attr("single_dataset__title", (str, type(None))) + self.check_attr("single_dataset__about", (str, type(None))) + self.check_attr("single_dataset__obs_names", (str, type(None))) + self.check_attr("single_dataset__var_names", (str, type(None))) if self.single_dataset__datapath is None: if self.multi_dataset__dataroot is None: @@ -357,7 +513,7 @@ class AppConfig(object): self.matrix_data_cache_manager = MatrixDataCacheManager(max_cached=1, timelimit_s=None) # preload this data set - matrix_data_loader = MatrixDataLoader(self.single_dataset__datapath, app_config=self) + matrix_data_loader = MatrixDataLoader(self.single_dataset__datapath, app_config=self.app_config) try: matrix_data_loader.pre_load_validation() except DatasetAccessError as e: @@ -388,26 +544,46 @@ class AppConfig(object): ) def handle_multi_dataset(self, context): - self.__check_attr("multi_dataset__dataroot", (type(None), dict, str)) - self.__check_attr("multi_dataset__index", (type(None), bool, str)) - self.__check_attr("multi_dataset__allowed_matrix_types", list) - self.__check_attr("multi_dataset__matrix_cache__max_datasets", int) - self.__check_attr("multi_dataset__matrix_cache__timelimit_s", (type(None), int, float)) + self.check_attr("multi_dataset__dataroot", (type(None), dict, str)) + self.check_attr("multi_dataset__index", (type(None), bool, str)) + self.check_attr("multi_dataset__allowed_matrix_types", list) + self.check_attr("multi_dataset__matrix_cache__max_datasets", int) + self.check_attr("multi_dataset__matrix_cache__timelimit_s", (type(None), int, float)) if self.multi_dataset__dataroot is None: return if type(self.multi_dataset__dataroot) == str: - self.multi_dataset__dataroot = dict(d=self.multi_dataset__dataroot) + default_dict = dict(base_url="d", dataroot=self.multi_dataset__dataroot) + self.multi_dataset__dataroot = dict(d=default_dict) - for key in self.multi_dataset__dataroot.keys(): - # sanity check for well formed keys - if type(key) != str: - raise ConfigurationError(f"error in multi_dataset__dataroot {key}") - if quote_plus(key) != key: - raise ConfigurationError(f"error in multi_dataset__dataroot {key}") - if os.path.split(os.path.normpath(key))[-1] != key: - raise ConfigurationError(f"error in multi_dataset__dataroot {key}") + for tag, dataroot_dict in self.multi_dataset__dataroot.items(): + if "base_url" not in dataroot_dict: + raise ConfigurationError(f"error in multi_dataset__dataroot: missing base_url for tag {tag}") + if "dataroot" not in dataroot_dict: + raise ConfigurationError(f"error in multi_dataset__dataroot: missing dataroot, for tag {tag}") + + base_url = dataroot_dict["base_url"] + + # sanity check for well formed base urls + bad = False + if type(base_url) != str: + bad = True + elif os.path.normpath(base_url) != base_url: + bad = True + else: + base_url_parts = base_url.split("/") + if [quote_plus(part) for part in base_url_parts] != base_url_parts: + bad = True + if ".." in base_url_parts: + bad = True + if bad: + raise ConfigurationError(f"error in multi_dataset__dataroot base_url {base_url} for tag {tag}") + + # verify all the base_urls are unique + base_urls = [d["base_url"] for d in self.multi_dataset__dataroot.values()] + if len(base_urls) > len(set(base_urls)): + raise ConfigurationError("error in multi_dataset__dataroot: base_urls must be unique") # error checking for mtype in self.multi_dataset__allowed_matrix_types: @@ -423,13 +599,114 @@ class AppConfig(object): timelimit_s=self.multi_dataset__matrix_cache__timelimit_s, ) + def handle_diffexp(self, context): + self.check_attr("diffexp__alg_cxg__max_workers", (str, int)) + self.check_attr("diffexp__alg_cxg__cpu_multiplier", int) + self.check_attr("diffexp__alg_cxg__target_workunit", int) + + max_workers = self.diffexp__alg_cxg__max_workers + cpu_multiplier = self.diffexp__alg_cxg__cpu_multiplier + cpu_count = os.cpu_count() + max_workers = min(max_workers, cpu_multiplier * cpu_count) + diffexp_tiledb.set_config(max_workers, self.diffexp__alg_cxg__target_workunit) + + def handle_adaptor(self, context): + # cxg + self.check_attr("adaptor__cxg_adaptor__tiledb_ctx", dict) + regionkey = "vfs.s3.region" + if regionkey not in self.adaptor__cxg_adaptor__tiledb_ctx: + if type(self.data_locator__s3__region_name) == str: + self.adaptor__cxg_adaptor__tiledb_ctx[regionkey] = self.data_locator__s3__region_name + + from server.data_cxg.cxg_adaptor import CxgAdaptor + + CxgAdaptor.set_tiledb_context(self.adaptor__cxg_adaptor__tiledb_ctx) + + # anndata + self.check_attr("adaptor__anndata_adaptor__backed", bool) + + def handle_limits(self, context): + self.check_attr("limits__diffexp_cellcount_max", (type(None), int)) + self.check_attr("limits__column_request_max", (type(None), int)) + + def exceeds_limit(self, limit_name, value): + limit_value = getattr(self, "limits__" + limit_name, None) + if limit_value is None: # disabled + return False + return value > limit_value + + +class DatasetConfig(BaseConfig): + """Manages the config attribute associated with a dataset.""" + + def __init__(self, tag, app_config, default_config): + super().__init__(app_config, default_config) + self.tag = tag + dc = default_config + try: + self.app__scripts = dc["app"]["scripts"] + self.app__inline_scripts = dc["app"]["inline_scripts"] + self.app__about_legal_tos = dc["app"]["about_legal_tos"] + self.app__about_legal_privacy = dc["app"]["about_legal_privacy"] + + self.presentation__max_categories = dc["presentation"]["max_categories"] + self.presentation__custom_colors = dc["presentation"]["custom_colors"] + + self.user_annotations__enable = dc["user_annotations"]["enable"] + self.user_annotations__type = dc["user_annotations"]["type"] + self.user_annotations__local_file_csv__directory = dc["user_annotations"]["local_file_csv"]["directory"] + self.user_annotations__local_file_csv__file = dc["user_annotations"]["local_file_csv"]["file"] + self.user_annotations__ontology__enable = dc["user_annotations"]["ontology"]["enable"] + self.user_annotations__ontology__obo_location = dc["user_annotations"]["ontology"]["obo_location"] + + self.embeddings__names = dc["embeddings"]["names"] + self.embeddings__enable_reembedding = dc["embeddings"]["enable_reembedding"] + + self.diffexp__enable = dc["diffexp"]["enable"] + self.diffexp__lfc_cutoff = dc["diffexp"]["lfc_cutoff"] + self.diffexp__top_n = dc["diffexp"]["top_n"] + + except KeyError as e: + raise ConfigurationError(f"Unexpected config: {str(e)}") + + # The annotation object is created during complete_config and stored here. + self.user_annotations = None + + def complete_config(self, context): + self.handle_app(context) + self.handle_presentation(context) + self.handle_user_annotations(context) + self.handle_embeddings(context) + self.handle_diffexp(context) + + def handle_app(self, context): + self.check_attr("app__scripts", list) + self.check_attr("app__inline_scripts", list) + self.check_attr("app__about_legal_tos", (type(None), str)) + self.check_attr("app__about_legal_privacy", (type(None), str)) + + # scripts can be string (filename) or dict (attributes). Convert string to dict. + scripts = [] + for s in self.app__scripts: + if isinstance(s, str): + scripts.append({"src": s}) + elif isinstance(s, dict) and isinstance(s["src"], str): + scripts.append(s) + else: + raise ConfigurationError("Scripts must be string or dict") + self.app__scripts = scripts + + def handle_presentation(self, context): + self.check_attr("presentation__max_categories", int) + self.check_attr("presentation__custom_colors", bool) + def handle_user_annotations(self, context): - self.__check_attr("user_annotations__enable", bool) - self.__check_attr("user_annotations__type", str) - self.__check_attr("user_annotations__local_file_csv__directory", (type(None), str)) - self.__check_attr("user_annotations__local_file_csv__file", (type(None), str)) - self.__check_attr("user_annotations__ontology__enable", bool) - self.__check_attr("user_annotations__ontology__obo_location", (type(None), str)) + self.check_attr("user_annotations__enable", bool) + self.check_attr("user_annotations__type", str) + self.check_attr("user_annotations__local_file_csv__directory", (type(None), str)) + self.check_attr("user_annotations__local_file_csv__file", (type(None), str)) + self.check_attr("user_annotations__ontology__enable", bool) + self.check_attr("user_annotations__ontology__obo_location", (type(None), str)) if self.user_annotations__enable: # TODO, replace this with a factory pattern once we have more than one way @@ -458,8 +735,11 @@ class AppConfig(object): # if the user has specified a fixed label file, go ahead and validate it # so that we can remove errors early in the process. - if self.single_dataset__datapath and self.user_annotations__local_file_csv__file: - with self.matrix_data_cache_manager.data_adaptor(self.single_dataset__datapath, self) as data_adaptor: + server_config = self.app_config.server_config + if server_config.single_dataset__datapath and self.user_annotations__local_file_csv__file: + with server_config.matrix_data_cache_manager.data_adaptor( + self.tag, server_config.single_dataset__datapath, self.app_config + ) as data_adaptor: data_adaptor.check_new_labels(self.user_annotations.read_labels(data_adaptor)) if self.user_annotations__ontology__enable or self.user_annotations__ontology__obo_location: @@ -487,135 +767,29 @@ class AppConfig(object): ) def handle_embeddings(self, context): - self.__check_attr("embeddings__names", list) - self.__check_attr("embeddings__enable_reembedding", bool) + self.check_attr("embeddings__names", list) + self.check_attr("embeddings__enable_reembedding", bool) - if self.single_dataset__datapath: + if self.app_config.server_config.single_dataset__datapath: if self.embeddings__enable_reembedding: - matrix_data_loader = MatrixDataLoader(self.single_dataset__datapath, app_config=self) + matrix_data_loader = MatrixDataLoader(self.single_dataset__datapath, app_config=self.app_config) if matrix_data_loader.matrix_data_type() != MatrixDataType.H5AD: raise ConfigurationError("'enable-reembedding is only supported with H5AD files.") if self.adaptor__anndata_adaptor__backed: raise ConfigurationError("enable-reembedding is not supported when run in --backed mode.") def handle_diffexp(self, context): - self.__check_attr("diffexp__enable", bool) - self.__check_attr("diffexp__lfc_cutoff", float) - self.__check_attr("diffexp__top_n", int) - self.__check_attr("diffexp__alg_cxg__max_workers", (str, int)) - self.__check_attr("diffexp__alg_cxg__cpu_multiplier", int) - self.__check_attr("diffexp__alg_cxg__target_workunit", int) + self.check_attr("diffexp__enable", bool) + self.check_attr("diffexp__lfc_cutoff", float) + self.check_attr("diffexp__top_n", int) - if self.single_dataset__datapath: - with self.matrix_data_cache_manager.data_adaptor(self.single_dataset__datapath, self) as data_adaptor: + server_config = self.app_config.server_config + if server_config.single_dataset__datapath: + with server_config.matrix_data_cache_manager.data_adaptor( + self.tag, server_config.single_dataset__datapath, self.app_config + ) as data_adaptor: if self.diffexp__enable and data_adaptor.parameters.get("diffexp_may_be_slow", False): context["messagefn"]( "CAUTION: due to the size of your dataset, " "running differential expression may take longer or fail." ) - - max_workers = self.diffexp__alg_cxg__max_workers - cpu_multiplier = self.diffexp__alg_cxg__cpu_multiplier - cpu_count = os.cpu_count() - max_workers = min(max_workers, cpu_multiplier * cpu_count) - diffexp_tiledb.set_config(max_workers, self.diffexp__alg_cxg__target_workunit) - - def handle_adaptor(self, context): - # cxg - self.__check_attr("adaptor__cxg_adaptor__tiledb_ctx", dict) - regionkey = "vfs.s3.region" - if regionkey not in self.adaptor__cxg_adaptor__tiledb_ctx: - if type(self.data_locator__s3__region_name) == str: - self.adaptor__cxg_adaptor__tiledb_ctx[regionkey] = self.data_locator__s3__region_name - - from server.data_cxg.cxg_adaptor import CxgAdaptor - - CxgAdaptor.set_tiledb_context(self.adaptor__cxg_adaptor__tiledb_ctx) - - # anndata - self.__check_attr("adaptor__anndata_adaptor__backed", bool) - - def handle_limits(self, context): - self.__check_attr("limits__diffexp_cellcount_max", (type(None), int)) - self.__check_attr("limits__column_request_max", (type(None), int)) - - def get_title(self, data_adaptor): - return self.single_dataset__title if self.single_dataset__title else data_adaptor.get_title() - - def get_about(self, data_adaptor): - return self.single_dataset__about if self.single_dataset__about else data_adaptor.get_about() - - def get_client_config(self, data_adaptor, annotation=None): - """ - Return the configuration as required by the /config REST route - """ - - # FIXME The current set of config is not consistently presented: - # we have camalCase, hyphen-text, and underscore_text - - # make sure the configuration has been checked. - self.check_config() - - # features - features = [f.todict() for f in data_adaptor.get_features(annotation)] - - # display_names - title = self.get_title(data_adaptor) - about = self.get_about(data_adaptor) - - display_names = dict(engine=data_adaptor.get_name(), dataset=title) - - # library_versions - library_versions = {} - library_versions.update(data_adaptor.get_library_versions()) - library_versions["cellxgene"] = cellxgene_version - - # links - links = {"about-dataset": about} - - # parameters - parameters = { - "layout": self.embeddings__names, - "max-category-items": self.presentation__max_categories, - "obs_names": self.single_dataset__obs_names, - "var_names": self.single_dataset__var_names, - "diffexp_lfc_cutoff": self.diffexp__lfc_cutoff, - "backed": self.adaptor__anndata_adaptor__backed, - "disable-diffexp": not self.diffexp__enable, - "enable-reembedding": self.embeddings__enable_reembedding, - "annotations": False, - "annotations_file": None, - "annotations_dir": None, - "annotations_cell_ontology_enabled": False, - "annotations_cell_ontology_obopath": None, - "annotations_cell_ontology_terms": None, - "custom_colors": self.presentation__custom_colors, - "diffexp-may-be-slow": False, - "about_legal_tos": self.server__about_legal_tos, - "about_legal_privacy": self.server__about_legal_privacy, - } - - data_adaptor.update_parameters(parameters) - if annotation: - annotation.update_parameters(parameters, data_adaptor) - - # gather it all together - c = {} - config = c["config"] = {} - config["features"] = features - config["displayNames"] = display_names - config["library_versions"] = library_versions - config["links"] = links - config["parameters"] = parameters - config["limits"] = { - "column_request_max": self.limits__column_request_max, - "diffexp_cellcount_max": self.limits__diffexp_cellcount_max, - } - - return c - - def exceeds_limit(self, limit_name, value): - limit_value = getattr(self, "limits__" + limit_name, None) - if limit_value is None: # disabled - return False - return value > limit_value diff --git a/server/common/default_config.py b/server/common/default_config.py index 5e93149c..0b0c5118 100644 --- a/server/common/default_config.py +++ b/server/common/default_config.py @@ -1,129 +1,159 @@ import yaml default_config = """ -# cellxgene configuration - server: - verbose: false - debug: false - host: "127.0.0.1" - port : null + app: + verbose: false + debug: false + host: "127.0.0.1" + port : null + open_browser: false + force_https: false + flask_secret_key: null + generate_cache_control_headers: false + server_timing_headers: false + csp_directives: null - # Scripts can be a list of either file names (string) or dicts containing keys src, integrity and crossorigin. - # these will be injected into the index template as script tags with these attributes set. - scripts: [] - # Inline scripts are a list of file names, where the contents of the file will be injected into the index. - inline_scripts: [] + multi_dataset: + # If dataroot is set, then cellxgene may serve multiple datasets. This parameter is not + # compatible with single_dataset/datapath. + # dataroot may be a string, representing the path to a directory or S3 prefix. In this + # case the datasets in that location are accessed from /d/. + # example: + # dataroot: /path/to/datasets/ + # or + # dataroot: s3://bucket/prefix/ + # + # As an alternative, dataroot can be a dictionary, where a dataset key is associated with a base_url + # and a dataroot. + # example: + # dataroot: + # d1: + # base_url: set1 + # dataroot: /path/to/set1_datasets/ + # d2: + # base_url: set2/subdir + # dataroot: /path/to/set2_datasets/ + # + # In this case, datasets can be accessed from /set1/ or + # /set2/subdir/. It is possible to have different dataset configurations + # for datasets accessed through different dataroots. For example, in one dataroot, the + # user annotations could be enabled, and in another dataroot they could be disabled. + # To specify dataroot configurations, add a new top level dictionary to the config named + # per_dataset_config. Within per_dataset_config create a dictionary for each dataroot to specialize + # ("d1" or "d2" from the example). Each of these dictionaries has the exact same form as the "dataset" + # dictionary (see below). + # When this approach is used, the values for each configuration option are checked in + # this order: per_dataset_config/, dataset, then the default values. + # + # example: + # + # per_dataset_config: + # d1: + # user_annotations: + # enable: false + # d2: + # user_annotations: + # enable: true - open_browser: false - about_legal_tos: null - about_legal_privacy: null - force_https: false - flask_secret_key: null - generate_cache_control_headers: false - server_timing_headers: false - csp_directives: null + dataroot: null -presentation: - max_categories: 1000 - custom_colors: true + # The index page when in multi-dataset mode: + # false or null: this returns a 404 code + # true: loads a test index page, which links to the datasets that are available in the dataroot + # string/URL: redirect to this URL: flask.redirect(config.multi_dataset__index) + index: false -multi_dataset: - # If dataroot is set, then cellxgene may serve multiple datasets. This parameter is not - # compatable with single_dataset/datapath. - # dataroot may be a string, representing the path to a directory or S3 prefix. In this - # case the datasets in that location are accessed from /d/. - # example: - # dataroot: /path/to/datasets/ - # or - # dataroot: s3://bucket/prefix/ - # - # As an alternative, dataroot can be a dictionary, mapping url prefixes to dataroot paths. - # example: - # dataroot: - # set1 : /path/to/set1_datasets/ - # set2 : /path/to/set2_datasets/ - # In this case, datasets can be accessed from /set1/ or - # /set2/. + # A list of allowed matrix types. If an empty list, then all matrix types are allowed + allowed_matrix_types: [] - dataroot: null + matrix_cache: + # The maximum number of datasets that may be opened at one time. The least recently used dataset + # is evicted from the cache first. + max_datasets: 5 - # The index page when in multi-dataset mode: - # false or null: this returns a 404 code - # true: loads a test index page, which links to the datasets that are available in the dataroot - # string/URL: redirect to this URL: flask.redirect(config.multi_dataset__index) - index: false + # A matrix is automatically removed from the cache after timelimit_s number of seconds. + # If timelimit_s is set to None, then there is no time limit. + timelimit_s: 30 - # A list of allowed matrix types. If an empty list, then all matrix types are allowed - allowed_matrix_types: [] + single_dataset: + # If datapath is set, then cellxgene with serve a single dataset located at datapath. This parameter is not + # compatible with multi_dataset/dataroot. + datapath: null + obs_names: null + var_names: null + about: null + title: null - matrix_cache: - # The maximum number of datasets that may be opened at one time. The least recently used dataset - # is evicted from the cache first. - max_datasets: 5 + diffexp: + alg_cxg: + # The number of threads to use is computed from: min(max_workers, cpu_multipler * cpu_count). + # Where cpu_count is determined at runtime. + max_workers: 64 + cpu_multiplier: 4 - # A matrix is automatically removed from the cache after timelimit_s number of seconds. - # If timelimit_s is set to None, then there is no time limit. - timelimit_s: 30 + # The target number of matrix elements that are evaluated + # together in one thread. + target_workunit: 16_000_000 -single_dataset: - datapath: null - obs_names: null - var_names: null - about: null - title: null + data_locator: + s3: + # s3 region name. + # if true, then the s3 location is automatically determined from the datapath or dataroot. + # if false/null, then do not set. + # if a string, then use that value (e.g. us-east-1). + region_name: true -user_annotations: - enable: true - type: local_file_csv - local_file_csv: - directory: null - file: null - ontology: - enable: false - obo_location: null + adaptor: + cxg_adaptor: + # The key/values under tiledb_ctx will be used to initialize the tiledb Context. + # If 'vfs.s3.region' is not set, then it will automatically use the setting from + # data_locator / s3 / region_name. + tiledb_ctx: + sm.tile_cache_size: 8589934592 + sm.num_reader_threads: 32 -embeddings: - names : [] - enable_reembedding: false - -diffexp: - enable: true - lfc_cutoff: 0.01 - top_n: 10 - alg_cxg: - # The number of threads to use is computed from: min(max_workers, cpu_multipler * cpu_count). - # Where cpu_count is determined at runtime. - max_workers: 64 - cpu_multiplier: 4 - - # The target number of matrix elements that are evaluated - # together in one thread. - target_workunit: 16_000_000 - -data_locator: - s3: - # s3 region name. - # if true, then the s3 location is automatically determined from the datapath or dataroot. - # if false/null, then do not set. - # if a string, then use that value (e.g. us-east-1). - region_name: true - -adaptor: - cxg_adaptor: - # The key/values under tiledb_ctx will be used to initialize the tiledb Context. - # If 'vfs.s3.region' is not set, then it will automatically use the setting from - # data_locator / s3 / region_name. - tiledb_ctx: - sm.tile_cache_size: 8589934592 - sm.num_reader_threads: 32 - - anndata_adaptor: + anndata_adaptor: backed: false -limits: - column_request_max: 32 - diffexp_cellcount_max: null + limits: + column_request_max: 32 + diffexp_cellcount_max: null + + +dataset: + app: + # Scripts can be a list of either file names (string) or dicts containing keys src, integrity and crossorigin. + # these will be injected into the index template as script tags with these attributes set. + scripts: [] + # Inline scripts are a list of file names, where the contents of the file will be injected into the index. + inline_scripts: [] + + about_legal_tos: null + about_legal_privacy: null + + presentation: + max_categories: 1000 + custom_colors: true + + user_annotations: + enable: true + type: local_file_csv + local_file_csv: + directory: null + file: null + ontology: + enable: false + obo_location: null + + embeddings: + names : [] + enable_reembedding: false + + diffexp: + enable: true + lfc_cutoff: 0.01 + top_n: 10 """ diff --git a/server/common/health.py b/server/common/health.py index 939a670f..5edaec59 100644 --- a/server/common/health.py +++ b/server/common/health.py @@ -24,10 +24,12 @@ def health_check(config): health = {"status": None, "version": "1", "releaseID": cellxgene_version} checks = False - if config.single_dataset__datapath is not None: - checks = _is_accessible(config.single_dataset__datapath, config) - elif config.multi_dataset__dataroot is not None: - checks = all([_is_accessible(datapath, config) for datapath in config.multi_dataset__dataroot.values()]) + server_config = config.server_config + if config.is_multi_dataset(): + dataroots = [datapath_dict["dataroot"] for datapath_dict in server_config.multi_dataset__dataroot.values()] + checks = all([_is_accessible(dataroot, server_config) for dataroot in dataroots]) + else: + checks = _is_accessible(server_config.single_dataset__datapath, server_config) health["status"] = "pass" if checks else "fail" code = HTTPStatus.OK if health["status"] == "pass" else HTTPStatus.BAD_REQUEST diff --git a/server/common/rest.py b/server/common/rest.py index 60a2bedf..a1e1a9ae 100644 --- a/server/common/rest.py +++ b/server/common/rest.py @@ -97,12 +97,13 @@ def _query_parameter_to_filter(args): return result -def schema_get_helper(data_adaptor, annotations): +def schema_get_helper(data_adaptor): """helper function to gather the schema from the data source and annotations""" schema = data_adaptor.get_schema() schema = copy.deepcopy(schema) # add label obs annotations as needed + annotations = data_adaptor.dataset_config.user_annotations if annotations is not None: label_schema = annotations.get_schema(data_adaptor) schema["annotations"]["obs"]["columns"].extend(label_schema) @@ -110,20 +111,20 @@ def schema_get_helper(data_adaptor, annotations): return schema -def schema_get(data_adaptor, annotations): - schema = schema_get_helper(data_adaptor, annotations) +def schema_get(data_adaptor): + schema = schema_get_helper(data_adaptor) return make_response(jsonify({"schema": schema}), HTTPStatus.OK) -def config_get(app_config, data_adaptor, annotations): - config = app_config.get_client_config(data_adaptor, annotations) +def config_get(app_config, data_adaptor): + config = app_config.get_client_config(data_adaptor) return make_response(jsonify(config), HTTPStatus.OK) -def annotations_obs_get(request, data_adaptor, annotations): +def annotations_obs_get(request, data_adaptor): fields = request.args.getlist("annotation-name", None) num_columns_requested = len(data_adaptor.get_obs_keys()) if len(fields) == 0 else len(fields) - if data_adaptor.config.exceeds_limit("column_request_max", num_columns_requested): + if data_adaptor.server_config.exceeds_limit("column_request_max", num_columns_requested): return abort(HTTPStatus.BAD_REQUEST) preferred_mimetype = request.accept_mimetypes.best_match(["application/octet-stream"]) if preferred_mimetype != "application/octet-stream": @@ -131,6 +132,7 @@ def annotations_obs_get(request, data_adaptor, annotations): try: labels = None + annotations = data_adaptor.dataset_config.user_annotations if annotations: labels = annotations.read_labels(data_adaptor) fbs = data_adaptor.annotation_to_fbs_matrix(Axis.OBS, fields, labels) @@ -139,8 +141,9 @@ def annotations_obs_get(request, data_adaptor, annotations): return abort_and_log(HTTPStatus.BAD_REQUEST, str(e), include_exc_info=True) -def annotations_put_fbs_helper(data_adaptor, annotations, fbs): +def annotations_put_fbs_helper(data_adaptor, fbs): """helper function to write annotations from fbs""" + annotations = data_adaptor.dataset_config.user_annotations if annotations is None: raise DisabledFeatureError("Writable annotations are not enabled") @@ -150,7 +153,8 @@ def annotations_put_fbs_helper(data_adaptor, annotations, fbs): annotations.write_labels(new_label_df, data_adaptor) -def annotations_obs_put(request, data_adaptor, annotations): +def annotations_obs_put(request, data_adaptor): + annotations = data_adaptor.dataset_config.user_annotations if annotations is None: return abort(HTTPStatus.NOT_IMPLEMENTED) @@ -163,17 +167,17 @@ def annotations_obs_put(request, data_adaptor, annotations): annotations.set_collection(anno_collection) try: - annotations_put_fbs_helper(data_adaptor, annotations, fbs) + annotations_put_fbs_helper(data_adaptor, fbs) res = json.dumps({"status": "OK"}) return make_response(res, HTTPStatus.OK, {"Content-Type": "application/json"}) except (ValueError, DisabledFeatureError, KeyError) as e: return abort_and_log(HTTPStatus.BAD_REQUEST, str(e), include_exc_info=True) -def annotations_var_get(request, data_adaptor, annotations): +def annotations_var_get(request, data_adaptor): fields = request.args.getlist("annotation-name", None) num_columns_requested = len(data_adaptor.get_var_keys()) if len(fields) == 0 else len(fields) - if data_adaptor.config.exceeds_limit("column_request_max", num_columns_requested): + if data_adaptor.server_config.exceeds_limit("column_request_max", num_columns_requested): return abort(HTTPStatus.BAD_REQUEST) preferred_mimetype = request.accept_mimetypes.best_match(["application/octet-stream"]) if preferred_mimetype != "application/octet-stream": @@ -181,6 +185,7 @@ def annotations_var_get(request, data_adaptor, annotations): try: labels = None + annotations = data_adaptor.dataset_config.user_annotations if annotations is not None: labels = annotations.read_labels(data_adaptor) return make_response( @@ -226,7 +231,7 @@ def data_var_get(request, data_adaptor): def colors_get(data_adaptor): - if not data_adaptor.config.presentation__custom_colors: + if not data_adaptor.dataset_config.presentation__custom_colors: return make_response(jsonify({}), HTTPStatus.OK) try: return make_response(jsonify(data_adaptor.get_colors()), HTTPStatus.OK) @@ -235,7 +240,7 @@ def colors_get(data_adaptor): def diffexp_obs_post(request, data_adaptor): - if not data_adaptor.config.diffexp__enable: + if not data_adaptor.dataset_config.diffexp__enable: return abort(HTTPStatus.NOT_IMPLEMENTED) args = request.get_json() @@ -273,7 +278,7 @@ def diffexp_obs_post(request, data_adaptor): def layout_obs_get(request, data_adaptor): fields = request.args.getlist("layout-name", None) num_columns_requested = len(data_adaptor.get_embedding_names()) if len(fields) == 0 else len(fields) - if data_adaptor.config.exceeds_limit("column_request_max", num_columns_requested): + if data_adaptor.server_config.exceeds_limit("column_request_max", num_columns_requested): return abort(HTTPStatus.BAD_REQUEST) preferred_mimetype = request.accept_mimetypes.best_match(["application/octet-stream"]) @@ -296,7 +301,7 @@ def layout_obs_get(request, data_adaptor): def layout_obs_put(request, data_adaptor): - if not data_adaptor.config.embedding__enable_reembedding: + if not data_adaptor.dataset_config.embedding__enable_reembedding: return abort(HTTPStatus.NOT_IMPLEMENTED) preferred_mimetype = request.accept_mimetypes.best_match(["application/octet-stream"]) diff --git a/server/data_anndata/anndata_adaptor.py b/server/data_anndata/anndata_adaptor.py index 60113b6d..1a9b0d17 100644 --- a/server/data_anndata/anndata_adaptor.py +++ b/server/data_anndata/anndata_adaptor.py @@ -28,10 +28,9 @@ def anndata_version_is_pre_070(): class AnndataAdaptor(DataAdaptor): - def __init__(self, data_locator, config=None): - super().__init__(config) + def __init__(self, data_locator, app_config=None, dataset_config=None): + super().__init__(data_locator, app_config, dataset_config) self.data = None - self.data_locator = data_locator self._load_data(data_locator) self._validate_and_initialize() @@ -55,14 +54,8 @@ class AnndataAdaptor(DataAdaptor): return data_locator.size() if data_locator.islocal() else 0 @staticmethod - def open(data_locator, config): - return AnndataAdaptor(data_locator, config) - - def get_location(self): - return self.data_locator.uri_or_path - - def get_data_locator(self): - return self.data_locator + def open(data_locator, app_config, dataset_config=None): + return AnndataAdaptor(data_locator, app_config, dataset_config) def get_name(self): return "cellxgene anndata adaptor version" @@ -100,7 +93,7 @@ class AnndataAdaptor(DataAdaptor): for (ax_name, var_name) in ((Axis.OBS, "obs"), (Axis.VAR, "var")): config_name = f"single_dataset__{var_name}_names" parameter_name = f"{var_name}_names" - name = getattr(self.config, config_name) + name = getattr(self.server_config, config_name) df_axis = getattr(self.data, str(ax_name)) if name is None: # Default: create unique names from index @@ -161,7 +154,7 @@ class AnndataAdaptor(DataAdaptor): with data_locator.local_handle() as lh: # as of AnnData 0.6.19, backed mode performs initial load fast, but at the # cost of significantly slower access to X data. - backed = "r" if self.config.adaptor__anndata_adaptor__backed else None + backed = "r" if self.server_config.adaptor__anndata_adaptor__backed else None self.data = anndata.read_h5ad(lh, backed=backed) except ValueError: @@ -181,7 +174,7 @@ class AnndataAdaptor(DataAdaptor): ) def _validate_and_initialize(self): - if anndata_version_is_pre_070() and self.config.adaptor__anndata_adaptor__backed: + if anndata_version_is_pre_070() and self.server_config.adaptor__anndata_adaptor__backed: warnings.warn( "Use of --backed mode with anndata versions older than 0.7 will have serious " "performance issues. Please update to at least anndata 0.7 or later." @@ -199,7 +192,7 @@ class AnndataAdaptor(DataAdaptor): # heuristic n_values = self.data.shape[0] * self.data.shape[1] - if (n_values > 1e8 and self.config.adaptor__anndata_adaptor__backed is True) or (n_values > 5e8): + if (n_values > 1e8 and self.server_config.adaptor__anndata_adaptor__backed is True) or (n_values > 5e8): self.parameters.update({"diffexp_may_be_slow": True}) def _is_valid_layout(self, arr): @@ -246,7 +239,7 @@ class AnndataAdaptor(DataAdaptor): ) if isinstance(datatype, CategoricalDtype): category_num = len(curr_axis[ann].dtype.categories) - if category_num > 500 and category_num > self.config.presentation__max_categories: + if category_num > 500 and category_num > self.dataset_config.presentation__max_categories: warnings.warn( f"{str(ax).title()} annotation '{ann}' has {category_num} categories, this may be " f"cumbersome or slow to display. We recommend setting the " @@ -277,7 +270,7 @@ class AnndataAdaptor(DataAdaptor): c) cap total list of layouts at global const MAX_LAYOUTS """ # load default layouts from the data. - layouts = self.config.embeddings__names + layouts = self.dataset_config.embeddings__names if layouts is None or len(layouts) == 0: layouts = [key[2:] for key in self.data.obsm_keys() if type(key) == str and key.startswith("X_")] @@ -329,9 +322,9 @@ class AnndataAdaptor(DataAdaptor): def compute_diffexp_ttest(self, maskA, maskB, top_n=None, lfc_cutoff=None): if top_n is None: - top_n = self.config.diffexp__top_n + top_n = self.dataset_config.diffexp__top_n if lfc_cutoff is None: - lfc_cutoff = self.config.diffexp__lfc_cutoff + lfc_cutoff = self.dataset_config.diffexp__lfc_cutoff return diffexp_generic.diffexp_ttest(self, maskA, maskB, top_n, lfc_cutoff) def get_colors(self): @@ -355,7 +348,7 @@ class AnndataAdaptor(DataAdaptor): return getattr(self.data.obs, term_name) def get_obs_index(self): - name = self.config.single_dataset__obs_names + name = self.server_config.single_dataset__obs_names if name is None: return self.original_obs_index else: diff --git a/server/data_common/data_adaptor.py b/server/data_common/data_adaptor.py index 617addad..d0eddf7e 100644 --- a/server/data_common/data_adaptor.py +++ b/server/data_common/data_adaptor.py @@ -14,12 +14,17 @@ from server.common.app_config import AppFeature, AppConfig class DataAdaptor(metaclass=ABCMeta): """Base class for loading and accessing matrix data""" - def __init__(self, config): - if type(config) != AppConfig: + def __init__(self, data_locator, app_config, dataset_config=None): + if type(app_config) != AppConfig: raise TypeError("config expected to be of type AppConfig") + # location to the dataset + self.data_locator = data_locator + # config is the application configuration - self.config = config + self.app_config = app_config + self.server_config = self.app_config.server_config + self.dataset_config = dataset_config or app_config.default_dataset_config # parameters set by this data adaptor based on the data. self.parameters = {} @@ -31,7 +36,7 @@ class DataAdaptor(metaclass=ABCMeta): @staticmethod @abstractmethod - def open(data_locator, config): + def open(data_locator, app_config, dataset_config): pass @staticmethod @@ -109,13 +114,11 @@ class DataAdaptor(metaclass=ABCMeta): def cleanup(self): pass - @abstractmethod - def get_location(self): - pass - - @abstractmethod def get_data_locator(self): - pass + return self.data_locator + + def get_location(self): + return self.data_locator.uri_or_path def get_about(self): return None @@ -149,8 +152,8 @@ class DataAdaptor(metaclass=ABCMeta): features = [ AppFeature("/cluster/", method="POST", available=False), AppFeature("/layout/obs", method="GET", available=self.get_embedding_names() is not None), - AppFeature("/layout/obs", method="PUT", available=self.config.embeddings__enable_reembedding), - AppFeature("/diffexp/", method="POST", available=self.config.diffexp__enable), + AppFeature("/layout/obs", method="PUT", available=self.dataset_config.embeddings__enable_reembedding), + AppFeature("/diffexp/", method="POST", available=self.dataset_config.diffexp__enable), AppFeature("/annotations/obs", method="PUT", available=annotations is not None), ] return features @@ -260,7 +263,6 @@ class DataAdaptor(metaclass=ABCMeta): * currently only supports access on VAR axis * currently only supports filtering on VAR axis """ - if axis != Axis.VAR: raise ValueError("Only VAR dimension access is supported") @@ -273,7 +275,7 @@ class DataAdaptor(metaclass=ABCMeta): raise FilterError("filtering on obs unsupported") num_columns = self.get_shape()[1] if var_selector is None else np.count_nonzero(var_selector) - if self.config.exceeds_limit("column_request_max", num_columns): + if self.server_config.exceeds_limit("column_request_max", num_columns): raise ExceedsLimitError("Requested dataframe columns exceed column request limit") X = self.get_X_array(obs_selector, var_selector) @@ -301,14 +303,14 @@ class DataAdaptor(metaclass=ABCMeta): except (KeyError, IndexError): raise FilterError("Error parsing filter") if top_n is None: - top_n = self.config.diffexp__top_n + top_n = self.dataset_config.diffexp__top_n - if self.config.exceeds_limit( + if self.server_config.exceeds_limit( "diffexp_cellcount_max", np.count_nonzero(obs_mask_A) + np.count_nonzero(obs_mask_B) ): raise ExceedsLimitError("Diffexp request exceeds max cell count limit") - result = self.compute_diffexp_ttest(obs_mask_A, obs_mask_B, top_n, self.config.diffexp__lfc_cutoff) + result = self.compute_diffexp_ttest(obs_mask_A, obs_mask_B, top_n, self.dataset_config.diffexp__lfc_cutoff) try: return jsonify_numpy(result) diff --git a/server/data_common/matrix_loader.py b/server/data_common/matrix_loader.py index 4f7bfaec..233e7b86 100644 --- a/server/data_common/matrix_loader.py +++ b/server/data_common/matrix_loader.py @@ -29,7 +29,7 @@ class MatrixDataCacheItem(object): self.data_lock.r_release() return None - def acquire_and_open(self, app_config): + def acquire_and_open(self, app_config, dataset_config=None): """returns the data_adaptor if cached. opens the data_adaptor if not. In either case, the a reader lock is taken. Must call release when the data_adaptor is no longer needed""" @@ -43,7 +43,7 @@ class MatrixDataCacheItem(object): if not self.data_adaptor: try: self.loader.pre_load_validation() - self.data_adaptor = self.loader.open(app_config) + self.data_adaptor = self.loader.open(app_config, dataset_config) except Exception as e: # necessary to hold the reader lock after an exception, since # the release will occur when the context exits. @@ -115,7 +115,7 @@ class MatrixDataCacheManager(object): # will automatically be refreshed. def __init__(self, max_cached, timelimit_s=None): - # key is location, value is a MatrixDataCacheInfo + # key is tuple(url_dataroot, location), value is a MatrixDataCacheInfo self.datasets = {} # lock to protect the datasets @@ -131,20 +131,21 @@ class MatrixDataCacheManager(object): self.timelimit_s = timelimit_s @contextmanager - def data_adaptor(self, location, app_config): + def data_adaptor(self, url_dataroot, location, app_config): # create a loader for to this location if it does not already exist delete_adaptor = None data_adaptor = None cache_item = None + key = (url_dataroot, location) with self.lock: self.evict_old_datasets() - info = self.datasets.get(location) + info = self.datasets.get(key) if info is not None: info.last_access = time.time() info.num_access += 1 - self.datasets[location] = info + self.datasets[key] = info data_adaptor = info.cache_item.acquire_existing() cache_item = info.cache_item @@ -165,19 +166,20 @@ class MatrixDataCacheManager(object): loader = MatrixDataLoader(location, app_config=app_config) cache_item = MatrixDataCacheItem(loader) item = MatrixDataCacheInfo(cache_item, time.time()) - self.datasets[location] = item + self.datasets[key] = item try: assert cache_item if delete_adaptor: delete_adaptor.delete() if data_adaptor is None: - data_adaptor = cache_item.acquire_and_open(app_config) + dataset_config = app_config.get_dataset_config(url_dataroot) + data_adaptor = cache_item.acquire_and_open(app_config, dataset_config) yield data_adaptor except DatasetAccessError: cache_item.release() with self.lock: - del self.datasets[location] + del self.datasets[key] cache_item.delete() cache_item = None raise @@ -215,7 +217,7 @@ class MatrixDataType(Enum): class MatrixDataLoader(object): def __init__(self, location, matrix_data_type=None, app_config=None): """ location can be a string or DataLocator """ - region_name = None if app_config is None else app_config.data_locator__s3__region_name + region_name = None if app_config is None else app_config.server_config.data_locator__s3__region_name self.location = DataLocator(location, region_name=region_name) if not self.location.exists(): raise DatasetAccessError("Dataset does not exist.", HTTPStatus.NOT_FOUND) @@ -254,12 +256,12 @@ class MatrixDataLoader(object): if not app_config: return True - if not app_config.multi_dataset__dataroot: + if not app_config.is_multi_dataset(): return True - if len(app_config.multi_dataset__allowed_matrix_types) == 0: + if len(app_config.server_config.multi_dataset__allowed_matrix_types) == 0: return True - for val in app_config.multi_dataset__allowed_matrix_types: + for val in app_config.server_config.multi_dataset__allowed_matrix_types: try: if self.matrix_data_type == MatrixDataType(val): return True @@ -279,6 +281,6 @@ class MatrixDataLoader(object): def file_size(self): return self.matrix_type.file_size(self.location) - def open(self, app_config): + def open(self, app_config, dataset_config=None): # create and return a DataAdaptor object - return self.matrix_type.open(self.location, app_config) + return self.matrix_type.open(self.location, app_config, dataset_config) diff --git a/server/data_cxg/cxg_adaptor.py b/server/data_cxg/cxg_adaptor.py index 14607967..2705f43b 100644 --- a/server/data_cxg/cxg_adaptor.py +++ b/server/data_cxg/cxg_adaptor.py @@ -24,11 +24,10 @@ class CxgAdaptor(DataAdaptor): {"sm.tile_cache_size": 8 * 1024 * 1024 * 1024, "sm.num_reader_threads": 32, "vfs.s3.region": "us-east-1"} ) - def __init__(self, data_locator, config=None): - super().__init__(config) + def __init__(self, data_locator, app_config=None, dataset_config=None): + super().__init__(data_locator, app_config, dataset_config) self.lock = threading.Lock() - self.data_locator = data_locator self.url = data_locator.uri_or_path if self.url[-1] != "/": self.url += "/" @@ -66,8 +65,8 @@ class CxgAdaptor(DataAdaptor): return 0 @staticmethod - def open(data_locator, args): - return CxgAdaptor(data_locator, args) + def open(data_locator, app_config, dataset_config=None): + return CxgAdaptor(data_locator, app_config, dataset_config) def get_about(self): return self.about if self.about else super().get_about() @@ -75,12 +74,6 @@ class CxgAdaptor(DataAdaptor): def get_title(self): return self.title if self.title else super().get_title() - def get_location(self): - return self.url - - def get_data_locator(self): - return self.data_locator - def get_name(self): return "cellxgene cxg adaptor version" @@ -196,9 +189,9 @@ class CxgAdaptor(DataAdaptor): def compute_diffexp_ttest(self, maskA, maskB, top_n=None, lfc_cutoff=None): if top_n is None: - top_n = self.config.diffexp__top_n + top_n = self.dataset_config.diffexp__top_n if lfc_cutoff is None: - lfc_cutoff = self.config.diffexp__lfc_cutoff + lfc_cutoff = self.dataset_config.diffexp__lfc_cutoff return diffexp_cxg.diffexp_ttest(self, maskA, maskB, top_n, lfc_cutoff) def get_colors(self): diff --git a/server/eb/Makefile b/server/eb/Makefile index 5e5a8a66..f6cc0d0f 100644 --- a/server/eb/Makefile +++ b/server/eb/Makefile @@ -11,11 +11,13 @@ clean: # Presumes that a top-level `make build-client` has been done to # create the client static assets. +cwd := $(shell pwd) + .PHONY: build build: clean mkdir artifact.dir; \ (cd ../.. ; \ - git ls-files server/ | cpio -pdm server/eb/artifact.dir ; ); \ + git ls-files server/ | cpio -pdm $(cwd)/artifact.dir ; ); \ $(call copy_client_assets,../../client/build,artifact.dir/server) ; \ set -e ; \ cp app.py artifact.dir/application.py; \ diff --git a/server/eb/app.py b/server/eb/app.py index dcc91641..438cd9bb 100644 --- a/server/eb/app.py +++ b/server/eb/app.py @@ -58,6 +58,7 @@ class WSGIServer(Server): @staticmethod def _before_adding_routes(app, app_config): script_hashes, style_hashes = WSGIServer.get_csp_hashes(app, app_config) + server_config = app_config.server_config csp = { "default-src": ["'self'"], "connect-src": ["'self'"], @@ -72,14 +73,14 @@ class WSGIServer(Server): if not app.debug: csp["upgrade-insecure-requests"] = "" - if app_config.server__csp_directives: - for k, v in app_config.server__csp_directives.items(): + if server_config.app__csp_directives: + for k, v in server_config.app__csp_directives.items(): if not isinstance(v, list): v = [v] csp[k] = csp.get(k, []) + v Talisman( - app, force_https=app_config.server__force_https, frame_options="DENY", content_security_policy=csp, + app, force_https=server_config.app__force_https, frame_options="DENY", content_security_policy=csp, ) @staticmethod @@ -102,16 +103,18 @@ class WSGIServer(Server): @staticmethod def compute_inline_scp_hashes(app, app_config): - inline_scripts = app_config.server__inline_scripts + dataset_configs = [app_config.default_dataset_config] + list(app_config.dataroot_config.values()) hashes = [] - for script in inline_scripts: - with app.open_resource(f"../common/web/templates/{script}") as f: - content = f.read() - # we use jinja2 template include, which trims final newline if present. - if content[-1] == 0x0A: - content = content[0:-1] - hash = base64.b64encode(hashlib.sha256(content).digest()) - hashes.append(f"'sha256-{hash.decode('utf-8')}'") + for dataset_config in dataset_configs: + inline_scripts = dataset_config.app__inline_scripts + for script in inline_scripts: + with app.open_resource(f"../common/web/templates/{script}") as f: + content = f.read() + # we use jinja2 template include, which trims final newline if present. + if content[-1] == 0x0A: + content = content[0:-1] + hash = base64.b64encode(hashlib.sha256(content).digest()) + hashes.append(f"'sha256-{hash.decode('utf-8')}'") return hashes @staticmethod @@ -156,7 +159,7 @@ try: dataroot = os.getenv("CXG_DATAROOT") if dataroot: logging.info("Configuration from CXG_DATAROOT") - app_config.update(multi_dataset__dataroot=dataroot) + app_config.update_server_config(multi_dataset__dataroot=dataroot) secret_name = os.getenv("CXG_AWS_SECRET_NAME") if secret_name: @@ -174,26 +177,23 @@ try: sys.exit(1) flask_secret_key = get_flask_secret_key(secret_region_name, secret_name) - app_config.update(server__flask_secret_key=flask_secret_key) + app_config.update_server_config(app__flask_secret_key=flask_secret_key) # features are unsupported in the current hosted server - app_config.update( - user_annotations__enable=False, - embeddings__enable_reembedding=False, - multi_dataset__allowed_matrix_types=["cxg"], + app_config.update_default_dataset_config( + user_annotations__enable=False, embeddings__enable_reembedding=False, ) + app_config.update_server_config(multi_dataset__allowed_matrix_types=["cxg"],) app_config.complete_config(logging.info) - if not app_config.server__flask_secret_key: + if not app_config.server_config.app__flask_secret_key: logging.critical( "flask_secret_key is not provided. Either set in config file, CXG_SECRET_KEY environment variable, " "or in AWS Secret Manager" ) sys.exit(1) - user_annotations = app_config.user_annotations - server = WSGIServer(app_config) debug = False @@ -203,10 +203,10 @@ except Exception: logging.critical("Caught exception during initialization", exc_info=True) sys.exit(1) -if app_config.multi_dataset__dataroot: - logging.info(f"starting server with multi_dataset__dataroot={app_config.multi_dataset__dataroot}") -elif app_config.single_dataset__datapath: - logging.info(f"starting server with single_dataset__datapath={app_config.single_dataset__datapath}") +if app_config.is_multi_dataset(): + logging.info(f"starting server with multi_dataset__dataroot={app_config.server_config.multi_dataset__dataroot}") +else: + logging.info(f"starting server with single_dataset__datapath={app_config.server_config.single_dataset__datapath}") if __name__ == "__main__": try: diff --git a/server/test/__init__.py b/server/test/__init__.py index c84e381b..48d999ba 100644 --- a/server/test/__init__.py +++ b/server/test/__init__.py @@ -27,21 +27,18 @@ def data_with_tmp_annotations(ext: MatrixDataType, annotations_fixture=False): annotations_file = path.join(tmp_dir, "test_annotations.csv") if annotations_fixture: shutil.copyfile(f"{PROJECT_ROOT}/server/test/test_datasets/pbmc3k-annotations.csv", annotations_file) - args = { - "embeddings__names": ["umap"], - "presentation__max_categories": 100, - "single_dataset__obs_names": None, - "single_dataset__var_names": None, - "diffexp__lfc_cutoff": 0.01, - } fname = { MatrixDataType.H5AD: f"{PROJECT_ROOT}/example-dataset/pbmc3k.h5ad", MatrixDataType.CXG: "test/test_datasets/pbmc3k.cxg", }[ext] data_locator = DataLocator(fname) config = AppConfig() - config.update(**args) - config.update(single_dataset__datapath=data_locator.path) + config.update_server_config( + single_dataset__obs_names=None, single_dataset__var_names=None, single_dataset__datapath=data_locator.path + ) + config.update_default_dataset_config( + embeddings__names=["umap"], presentation__max_categories=100, diffexp__lfc_cutoff=0.01, + ) config.complete_config() data = MatrixDataLoader(data_locator.abspath()).open(config) annotations = AnnotationsLocalFile(None, annotations_file) @@ -66,21 +63,21 @@ def skip_if(condition, reason: str): return decorator -def app_config(data_locator, backed=False, extra={}): - args = { - "embeddings__names": ["umap", "tsne", "pca"], - "presentation__max_categories": 100, - "single_dataset__obs_names": None, - "single_dataset__var_names": None, - "diffexp__lfc_cutoff": 0.01, - "adaptor__anndata_adaptor__backed": backed, - "single_dataset__datapath": data_locator, - "limits__diffexp_cellcount_max": None, - "limits__column_request_max": None, - } +def app_config(data_locator, backed=False, extra_server_config={}, extra_dataset_config={}): config = AppConfig() - config.update(**args) - config.update(**extra) + config.update_server_config( + single_dataset__obs_names=None, + single_dataset__var_names=None, + adaptor__anndata_adaptor__backed=backed, + single_dataset__datapath=data_locator, + limits__diffexp_cellcount_max=None, + limits__column_request_max=None, + ) + config.update_default_dataset_config( + embeddings__names=["umap", "tsne", "pca"], presentation__max_categories=100, diffexp__lfc_cutoff=0.01 + ) + config.update_server_config(**extra_server_config) + config.update_default_dataset_config(**extra_dataset_config) config.complete_config() return config diff --git a/server/test/run_diffexp.py b/server/test/run_diffexp.py index d3fad73a..bbbb3094 100644 --- a/server/test/run_diffexp.py +++ b/server/test/run_diffexp.py @@ -32,8 +32,8 @@ def main(): args = parser.parse_args() app_config = AppConfig() - app_config.single_dataset__datapath = args.dataset - app_config.server__verbose = True + app_config.update_server_config(single_dataset__datapath=args.dataset) + app_config.update_server_config(app__verbose=True) app_config.complete_config() loader = MatrixDataLoader(args.dataset) diff --git a/server/test/test_anndata_adaptor.py b/server/test/test_anndata_adaptor.py index 1cec773c..3afe2ea3 100644 --- a/server/test/test_anndata_adaptor.py +++ b/server/test/test_anndata_adaptor.py @@ -107,9 +107,9 @@ class AdaptorTest(unittest.TestCase): self.assertEqual(len(feature), 1) check_feature("POST", "/cluster/", False) - check_feature("POST", "/diffexp/", self.data.config.diffexp__enable) + check_feature("POST", "/diffexp/", self.data.dataset_config.diffexp__enable) check_feature("GET", "/layout/obs", True) - check_feature("PUT", "/layout/obs", self.data.config.embeddings__enable_reembedding) + check_feature("PUT", "/layout/obs", self.data.dataset_config.embeddings__enable_reembedding) check_feature("PUT", "/annotations/obs", False) def test_layout(self): diff --git a/server/test/test_anndata_adaptor_data_load.py b/server/test/test_anndata_adaptor_data_load.py index c80935f3..e5c01a03 100644 --- a/server/test/test_anndata_adaptor_data_load.py +++ b/server/test/test_anndata_adaptor_data_load.py @@ -15,7 +15,7 @@ class DataLoadAdaptorTest(unittest.TestCase): def setUp(self): self.data_file = DataLocator(f"{PROJECT_ROOT}/example-dataset/pbmc3k.h5ad") config = AppConfig() - config.update(single_dataset__datapath=self.data_file.path) + config.update_server_config(single_dataset__datapath=self.data_file.path) config.complete_config() self.data = AnndataAdaptor(self.data_file, config) @@ -40,14 +40,18 @@ class DataLocatorAdaptorTest(unittest.TestCase): Test various types of data locators we expect to consume """ - def setUp(self): - self.args = { - "embeddings__names": ["umap"], - "presentation__max_categories": 100, - "single_dataset__obs_names": None, - "single_dataset__var_names": None, - "diffexp__lfc_cutoff": 0.01, - } + def get_basic_config(self): + config = AppConfig() + config.update_server_config( + single_dataset__obs_names=None, + single_dataset__var_names=None, + ) + config.update_default_dataset_config( + embeddings__names=["umap"], + presentation__max_categories=100, + diffexp__lfc_cutoff=0.01, + ) + return config def stdAsserts(self, data): """ run these each time we load the data """ @@ -57,9 +61,8 @@ class DataLocatorAdaptorTest(unittest.TestCase): def test_posix_file(self): locator = DataLocator("../example-dataset/pbmc3k.h5ad") - config = AppConfig() - config.update(**self.args) - config.update(single_dataset__datapath=locator.path) + config = self.get_basic_config() + config.update_server_config(single_dataset__datapath=locator.path) config.complete_config() data = AnndataAdaptor(locator, config) self.stdAsserts(data) @@ -67,15 +70,13 @@ class DataLocatorAdaptorTest(unittest.TestCase): def test_url_https(self): url = "https://raw.githubusercontent.com/chanzuckerberg/cellxgene/main/example-dataset/pbmc3k.h5ad" locator = DataLocator(url) - config = AppConfig() - config.update(**self.args) + config = self.get_basic_config() data = AnndataAdaptor(locator, config) self.stdAsserts(data) def test_url_http(self): url = "http://raw.githubusercontent.com/chanzuckerberg/cellxgene/main/example-dataset/pbmc3k.h5ad" locator = DataLocator(url) - config = AppConfig() - config.update(**self.args) + config = self.get_basic_config() data = AnndataAdaptor(locator, config) self.stdAsserts(data) diff --git a/server/test/test_app_config.py b/server/test/test_app_config.py index 7c436a44..c55f927a 100644 --- a/server/test/test_app_config.py +++ b/server/test/test_app_config.py @@ -11,60 +11,90 @@ import requests class AppConfigTest(unittest.TestCase): def test_update(self): c = AppConfig() - c.update(server__verbose=True, multi_dataset__dataroot="datadir") - v = c.changes_from_default() - self.assertCountEqual(v, [("server__verbose", True, False), ("multi_dataset__dataroot", "datadir", None)]) + c.update_server_config(app__verbose=True, multi_dataset__dataroot="datadir") + v = c.server_config.changes_from_default() + self.assertCountEqual(v, [("app__verbose", True, False), ("multi_dataset__dataroot", "datadir", None)]) c = AppConfig() - c.update(server__scripts=(), server__inline_scripts=()) - v = c.changes_from_default() + c.update_default_dataset_config(app__scripts=(), app__inline_scripts=()) + v = c.server_config.changes_from_default() self.assertCountEqual(v, []) c = AppConfig() - c.update(server__scripts=[], server__inline_scripts=[]) - v = c.changes_from_default() + c.update_default_dataset_config(app__scripts=[], app__inline_scripts=[]) + v = c.default_dataset_config.changes_from_default() self.assertCountEqual(v, []) c = AppConfig() - c.update(server__scripts=("a", "b"), server__inline_scripts=["c", "d"]) - v = c.changes_from_default() - self.assertCountEqual(v, [("server__scripts", ["a", "b"], []), ("server__inline_scripts", ["c", "d"], [])]) + c.update_default_dataset_config(app__scripts=("a", "b"), app__inline_scripts=["c", "d"]) + v = c.default_dataset_config.changes_from_default() + self.assertCountEqual(v, [("app__scripts", ["a", "b"], []), ("app__inline_scripts", ["c", "d"], [])]) def test_multi_dataset(self): c = AppConfig() # test for illegal url_dataroots - for illegal in ("a/b", "../b", "!$*", "\\n", "", "(bad)"): - c.update(multi_dataset__dataroot={illegal: f"{PROJECT_ROOT}/example-dataset"}) + for illegal in ("../b", "!$*", "\\n", "", "(bad)"): + c.update_server_config( + multi_dataset__dataroot={"tag": {"base_url": illegal, "dataroot": "{PROJECT_ROOT}/example-dataset"}} + ) with self.assertRaises(ConfigurationError): c.complete_config() # test for legal url_dataroots - for legal in ( - "d", - "this.is-okay_", - ): - c.update(multi_dataset__dataroot={legal: f"{PROJECT_ROOT}/example-dataset"}) + for legal in ("d", "this.is-okay_", "a/b"): + c.update_server_config( + multi_dataset__dataroot={"tag": {"base_url": legal, "dataroot": "{PROJECT_ROOT}/example-dataset"}} + ) c.complete_config() # test that multi dataroots work end to end - c.update( + c.update_server_config( multi_dataset__dataroot=dict( - set1=f"{PROJECT_ROOT}/example-dataset", set2=f"{PROJECT_ROOT}/server/test/test_datasets" + s1=dict(dataroot=f"{PROJECT_ROOT}/example-dataset", base_url="set1/1/2"), + s2=dict(dataroot=f"{PROJECT_ROOT}/server/test/test_datasets", base_url="set2"), + s3=dict(dataroot=f"{PROJECT_ROOT}/server/test/test_datasets", base_url="set3"), ) ) + + # Change this default to test if the dataroot overrides below work. + c.update_default_dataset_config(app__about_legal_tos="tos_default.html") + + # specialize the configs for set1 + c.add_dataroot_config( + "s1", user_annotations__enable=False, diffexp__enable=True, app__about_legal_tos="tos_set1.html" + ) + + # specialize the configs for set2 + c.add_dataroot_config( + "s2", user_annotations__enable=True, diffexp__enable=False, app__about_legal_tos="tos_set2.html" + ) + + # no specializations for set3 (they get the default dataset config) c.complete_config() with test_server(app_config=c) as server: session = requests.Session() - r = session.get(f"{server}/set1/pbmc3k.h5ad/api/v0.2/config") + r = session.get(f"{server}/set1/1/2/pbmc3k.h5ad/api/v0.2/config") data_config = r.json() assert data_config["config"]["displayNames"]["dataset"] == "pbmc3k" + assert data_config["config"]["parameters"]["annotations"] is False + assert data_config["config"]["parameters"]["disable-diffexp"] is False + assert data_config["config"]["parameters"]["about_legal_tos"] == "tos_set1.html" r = session.get(f"{server}/set2/pbmc3k.cxg/api/v0.2/config") data_config = r.json() assert data_config["config"]["displayNames"]["dataset"] == "pbmc3k" + assert data_config["config"]["parameters"]["annotations"] is True + assert data_config["config"]["parameters"]["about_legal_tos"] == "tos_set2.html" + + r = session.get(f"{server}/set3/pbmc3k.cxg/api/v0.2/config") + data_config = r.json() + assert data_config["config"]["displayNames"]["dataset"] == "pbmc3k" + assert data_config["config"]["parameters"]["annotations"] is True + assert data_config["config"]["parameters"]["disable-diffexp"] is False + assert data_config["config"]["parameters"]["about_legal_tos"] == "tos_default.html" r = session.get(f"{server}/health") assert r.json()["status"] == "pass" diff --git a/server/test/test_diffexp.py b/server/test/test_diffexp.py index c69365ac..abe1fc57 100644 --- a/server/test/test_diffexp.py +++ b/server/test/test_diffexp.py @@ -15,8 +15,9 @@ class DiffExpTest(unittest.TestCase): """Tests the diffexp returns the expected results for one test case, using different adaptor types and different algorithms.""" - def load_dataset(self, path, extra={}): - config = app_config(path, extra=extra) + def load_dataset(self, path, extra_server_config={}, extra_dataset_config={}): + config = app_config(path, extra_server_config=extra_server_config, + extra_dataset_config=extra_dataset_config) loader = MatrixDataLoader(path) adaptor = loader.open(config) return adaptor @@ -100,7 +101,7 @@ class DiffExpTest(unittest.TestCase): # create a sparse matrix h5adfile = os.path.join(dirname, "sparse.h5ad") create_test_h5ad(h5adfile, 2000, 2000, 10, apply_col_shift) - adaptor_anndata = self.load_dataset(h5adfile, extra=dict(embeddings__names=[])) + adaptor_anndata = self.load_dataset(h5adfile, extra_dataset_config=dict(embeddings__names=[])) adata = adaptor_anndata.data sparsename = os.path.join(dirname, "sparse.cxg") diff --git a/server/test/test_eb.py b/server/test/test_eb.py new file mode 100644 index 00000000..f1e1ef9d --- /dev/null +++ b/server/test/test_eb.py @@ -0,0 +1,54 @@ +import unittest +import tempfile +import requests +import subprocess +from server.test import PROJECT_ROOT +from server.common.app_config import AppConfig +from contextlib import contextmanager +import time + + +@contextmanager +def run_eb_app(tempdirname): + ps = subprocess.Popen(["python", "artifact.dir/application.py"], cwd=tempdirname) + server = "http://localhost:5000" + for _ in range(10): + try: + requests.get(f"{server}/health") + break + except requests.exceptions.ConnectionError: + time.sleep(1) + + try: + yield server + finally: + try: + ps.terminate() + except ProcessLookupError: + pass + + +class Elastic_Beanstalk_Test(unittest.TestCase): + def test_run(self): + + tempdir = tempfile.TemporaryDirectory(dir=f"{PROJECT_ROOT}/server") + tempdirname = tempdir.name + + c = AppConfig() + # test that eb works + c.update_server_config( + multi_dataset__dataroot=f"{PROJECT_ROOT}/server/test/test_datasets", app__flask_secret_key="open sesame" + ) + + c.complete_config() + c.write_config(f"{tempdirname}/config.yaml") + + subprocess.check_call(f"git ls-files . | cpio -pdm {tempdirname}", cwd=f"{PROJECT_ROOT}/server/eb", shell=True) + subprocess.check_call(["make", "build"], cwd=tempdirname) + + with run_eb_app(tempdirname) as server: + session = requests.Session() + + r = session.get(f"{server}/d/pbmc3k.cxg/api/v0.2/config") + data_config = r.json() + assert data_config["config"]["displayNames"]["dataset"] == "pbmc3k" diff --git a/server/test/test_matrixcache.py b/server/test/test_matrixcache.py index 7a1362ae..6c68cfc3 100644 --- a/server/test/test_matrixcache.py +++ b/server/test/test_matrixcache.py @@ -21,13 +21,13 @@ class MatrixCacheTest(unittest.TestCase): shutil.copytree(source, target) def use_dataset(self, matrix_cache, dirname, app_config, dataset_index): - with matrix_cache.data_adaptor(os.path.join(dirname, str(dataset_index) + ".cxg"), app_config) as adaptor: + with matrix_cache.data_adaptor(None, os.path.join(dirname, str(dataset_index) + ".cxg"), app_config) as adaptor: pass return adaptor def use_dataset_with_error(self, matrix_cache, dirname, app_config, dataset_index): try: - with matrix_cache.data_adaptor(os.path.join(dirname, str(dataset_index) + ".cxg"), app_config): + with matrix_cache.data_adaptor(None, os.path.join(dirname, str(dataset_index) + ".cxg"), app_config): raise DatasetAccessError("something bad happened") except DatasetAccessError: # the MatrixDataCacheManager rethrows the exception, so catch and ignore @@ -38,7 +38,7 @@ class MatrixCacheTest(unittest.TestCase): result = {} for k, v in datasets.items(): # filter out the dirname and the .cxg from the name - newk = int(k[len(dirname) + 1 : -4]) + newk = int(k[1][len(dirname) + 1 : -4]) result[newk] = v return result diff --git a/server/test/test_writable_annotation.py b/server/test/test_writable_annotation.py index d4cf15b4..6ca0419d 100644 --- a/server/test/test_writable_annotation.py +++ b/server/test/test_writable_annotation.py @@ -15,12 +15,13 @@ from server.data_common.matrix_loader import MatrixDataType class WritableAnnotationTest(unittest.TestCase): def setUp(self): self.data, self.tmp_dir, self.annotations = data_with_tmp_annotations(MatrixDataType.H5AD) + self.data.dataset_config.user_annotations = self.annotations def tearDown(self): shutil.rmtree(self.tmp_dir) def annotation_put_fbs(self, fbs): - annotations_put_fbs_helper(self.data, self.annotations, fbs) + annotations_put_fbs_helper(self.data, fbs) res = json.dumps({"status": "OK"}) return res @@ -112,7 +113,7 @@ class WritableAnnotationTest(unittest.TestCase): # get labels = self.annotations.read_labels(None) fbsAll = self.data.annotation_to_fbs_matrix("obs", None, labels) - schema = schema_get_helper(self.data, self.annotations) + schema = schema_get_helper(self.data) annotations = decode_fbs.decode_matrix_FBS(fbsAll) obs_index_col_name = schema["annotations"]["obs"]["index"] self.assertEqual(annotations["n_rows"], n_rows) @@ -149,7 +150,7 @@ class WritableAnnotationTest(unittest.TestCase): self.assertEqual(len(feature), 1) check_feature("POST", "/cluster/", False) - check_feature("POST", "/diffexp/", self.data.config.diffexp__enable) + check_feature("POST", "/diffexp/", self.data.dataset_config.diffexp__enable) check_feature("GET", "/layout/obs", True) - check_feature("PUT", "/layout/obs", self.data.config.embeddings__enable_reembedding) + check_feature("PUT", "/layout/obs", self.data.dataset_config.embeddings__enable_reembedding) check_feature("PUT", "/annotations/obs", True)