diff --git a/lib/galaxy/model/__init__.py b/lib/galaxy/model/__init__.py index 9d79f822313..cfbf1a90717 100644 --- a/lib/galaxy/model/__init__.py +++ b/lib/galaxy/model/__init__.py @@ -18,6 +18,7 @@ from uuid import UUID, uuid4 from six import string_types from sqlalchemy import ( + alias, and_, func, inspect, @@ -1556,6 +1557,36 @@ class History(HasTags, Dictifiable, UsesAnnotations, HasName): self._active_datasets_and_roles = query.all() return self._active_datasets_and_roles + @property + def active_visible_datasets_and_roles(self): + if not hasattr(self, '_active_visible_datasets_and_roles'): + db_session = object_session(self) + query = (db_session.query(HistoryDatasetAssociation) + .filter(HistoryDatasetAssociation.table.c.history_id == self.id) + .filter(not_(HistoryDatasetAssociation.deleted)) + .filter(HistoryDatasetAssociation.visible) + .order_by(HistoryDatasetAssociation.table.c.hid.asc()) + .options(joinedload("dataset"), + joinedload("dataset.actions"), + joinedload("dataset.actions.role"), + joinedload("tags"))) + self._active_visible_datasets_and_roles = query.all() + return self._active_visible_datasets_and_roles + + @property + def active_visible_dataset_collections(self): + if not hasattr(self, '_active_visible_dataset_collections'): + db_session = object_session(self) + query = (db_session.query(HistoryDatasetCollectionAssociation) + .filter(HistoryDatasetCollectionAssociation.table.c.history_id == self.id) + .filter(not_(HistoryDatasetCollectionAssociation.deleted)) + .filter(HistoryDatasetCollectionAssociation.visible) + .order_by(HistoryDatasetCollectionAssociation.table.c.hid.asc()) + .options(joinedload("collection"), + joinedload("tags"))) + self._active_visible_dataset_collections = query.all() + return self._active_visible_dataset_collections + @property def active_contents(self): """ Return all active contents ordered by hid. @@ -1968,6 +1999,18 @@ class Dataset(StorableObject): return False +def datatype_for_extension(extension, datatypes_registry=None): + if datatypes_registry is None: + datatypes_registry = _get_datatypes_registry() + if not extension or extension == 'auto' or extension == '_sniff_': + extension = 'data' + ret = datatypes_registry.get_datatype_by_extension(extension) + if ret is None: + log.warning("Datatype class not found for extension '%s'" % extension) + return datatypes_registry.get_datatype_by_extension('data') + return ret + + class DatasetInstance(object): """A base class for all 'dataset instances', HDAs, LDAs, etc""" states = Dataset.states @@ -2045,14 +2088,7 @@ class DatasetInstance(object): @property def datatype(self): - extension = self.extension - if not extension or extension == 'auto' or extension == '_sniff_': - extension = 'data' - ret = _get_datatypes_registry().get_datatype_by_extension(extension) - if ret is None: - log.warning("Datatype class not found for extension '%s'" % extension) - return _get_datatypes_registry().get_datatype_by_extension('data') - return ret + return datatype_for_extension(self.extension) def get_metadata(self): # using weakref to store parent (to prevent circ ref), @@ -3142,6 +3178,78 @@ class DatasetCollection(object, Dictifiable, UsesAnnotations): self.populated_state = DatasetCollection.populated_states.NEW self.element_count = element_count + @property + def dataset_states_and_extensions_summary(self): + if not hasattr(self, '_dataset_states_and_extensions_summary'): + db_session = object_session(self) + + dc = alias(DatasetCollection.table) + de = alias(DatasetCollectionElement.table) + hda = alias(HistoryDatasetAssociation.table) + dataset = alias(Dataset.table) + + select_from = dc.outerjoin(de, de.c.dataset_collection_id == dc.c.id) + + depth_collection_type = self.collection_type + while ":" in depth_collection_type: + child_collection = alias(DatasetCollection.table) + child_collection_element = alias(DatasetCollectionElement.table) + select_from = select_from.outerjoin(child_collection, child_collection.c.id == de.c.child_collection_id) + select_from = select_from.outerjoin(child_collection_element, child_collection_element.c.dataset_collection_id == child_collection.c.id) + + de = child_collection_element + depth_collection_type = depth_collection_type.split(":", 1)[1] + + select_from = select_from.outerjoin(hda, hda.c.id == de.c.hda_id).outerjoin(dataset, hda.c.dataset_id == dataset.c.id) + select_stmt = select([hda.c.extension, dataset.c.state]).select_from(select_from).where(dc.c.id == self.id).distinct() + extensions = set() + states = set() + for extension, state in db_session.execute(select_stmt).fetchall(): + states.add(state) + extensions.add(extension) + + self._dataset_states_and_extensions_summary = (states, extensions) + + return self._dataset_states_and_extensions_summary + + @property + def populated_optimized(self): + if not hasattr(self, '_populated_optimized'): + _populated_optimized = True + if ":" not in self.collection_type: + _populated_optimized = self.populated_state == DatasetCollection.populated_states.OK + else: + db_session = object_session(self) + + dc = alias(DatasetCollection.table) + de = alias(DatasetCollectionElement.table) + + select_from = dc.outerjoin(de, de.c.dataset_collection_id == dc.c.id) + + collection_depth_aliases = [dc] + + depth_collection_type = self.collection_type + while ":" in depth_collection_type: + child_collection = alias(DatasetCollection.table) + child_collection_element = alias(DatasetCollectionElement.table) + select_from = select_from.outerjoin(child_collection, child_collection.c.id == de.c.child_collection_id) + select_from = select_from.outerjoin(child_collection_element, child_collection_element.c.dataset_collection_id == child_collection.c.id) + + collection_depth_aliases.append(child_collection) + + de = child_collection_element + depth_collection_type = depth_collection_type.split(":", 1)[1] + + select_stmt = select(list(map(lambda dc: dc.c.populated_state, collection_depth_aliases))).select_from(select_from).where(dc.c.id == self.id).distinct() + for populated_states in db_session.execute(select_stmt).fetchall(): + for populated_state in populated_states: + if populated_state != DatasetCollection.populated_states.OK: + _populated_optimized = False + + self._populated_optimized = _populated_optimized + + return self._populated_optimized + @property def populated(self): top_level_populated = self.populated_state == DatasetCollection.populated_states.OK diff --git a/lib/galaxy/tools/__init__.py b/lib/galaxy/tools/__init__.py index 29da3072470..fefc2be9e18 100755 --- a/lib/galaxy/tools/__init__.py +++ b/lib/galaxy/tools/__init__.py @@ -56,6 +56,10 @@ from galaxy.tools.parameters.basic import ( ToolParameter, workflow_building_modes, ) +from galaxy.tools.parameters.dataset_matcher import ( + set_dataset_matcher_factory, + unset_dataset_matcher_factory, +) from galaxy.tools.parameters.grouping import Conditional, ConditionalWhen, Repeat, Section, UploadDataset from galaxy.tools.parameters.input_translation import ToolInputTranslator from galaxy.tools.parameters.meta import expand_meta_parameters @@ -1826,7 +1830,9 @@ class Tool(object, Dictifiable): # create tool model tool_model = self.to_dict(request_context) tool_model['inputs'] = [] + set_dataset_matcher_factory(request_context, self, state_inputs) self.populate_model(request_context, self.inputs, state_inputs, tool_model['inputs']) + unset_dataset_matcher_factory(request_context) # create tool help tool_help = '' diff --git a/lib/galaxy/tools/parameters/basic.py b/lib/galaxy/tools/parameters/basic.py index e6b1b12c76f..ba905c3ec9d 100644 --- a/lib/galaxy/tools/parameters/basic.py +++ b/lib/galaxy/tools/parameters/basic.py @@ -26,8 +26,7 @@ from galaxy.util.expressions import ExpressionContext from galaxy.web import url_for from . import validation from .dataset_matcher import ( - DatasetCollectionMatcher, - DatasetMatcher + get_dataset_matcher_factory, ) from .sanitize import ToolParameterSanitizer from ..parameters import ( @@ -1484,15 +1483,16 @@ class BaseDataToolParameter(ToolParameter): return None history = trans.history if history is not None: - dataset_matcher = DatasetMatcher(trans, self, None, other_values) + dataset_matcher_factory = get_dataset_matcher_factory(trans) + dataset_matcher = dataset_matcher_factory.dataset_matcher(self, other_values) if isinstance(self, DataToolParameter): - for hda in reversed(history.active_datasets_and_roles): - match = dataset_matcher.hda_match(hda, check_security=False) + for hda in reversed(history.active_visible_datasets_and_roles): + match = dataset_matcher.hda_match(hda) if match: return match.hda else: - dataset_collection_matcher = DatasetCollectionMatcher(dataset_matcher) - for hdca in reversed(history.active_dataset_collections): + dataset_collection_matcher = dataset_matcher_factory.dataset_collection_matcher(dataset_matcher) + for hdca in reversed(history.active_visible_dataset_collections): if dataset_collection_matcher.hdca_match(hdca, reduction=self.multiple): return hdca @@ -1603,13 +1603,6 @@ class DataToolParameter(BaseDataToolParameter): raise ValueError("Datatype class not found for extension '%s', which is used as 'type' attribute in conversion of data parameter '%s'" % (conv_type, self.name)) self.conversions.append((name, conv_extension, [conv_type])) - def match_collections(self, history, dataset_matcher, reduction=True): - dataset_collection_matcher = DatasetCollectionMatcher(dataset_matcher) - - for history_dataset_collection in history.active_dataset_collections: - if dataset_collection_matcher.hdca_match(history_dataset_collection, reduction=reduction): - yield history_dataset_collection - def from_json(self, value, trans, other_values={}): if trans.workflow_building_mode is workflow_building_modes.ENABLED: return None @@ -1797,7 +1790,8 @@ class DataToolParameter(BaseDataToolParameter): return d # prepare dataset/collection matching - dataset_matcher = DatasetMatcher(trans, self, None, other_values) + dataset_matcher_factory = get_dataset_matcher_factory(trans) + dataset_matcher = dataset_matcher_factory.dataset_matcher(self, other_values) multiple = self.multiple # build and append a new select option @@ -1811,8 +1805,9 @@ class DataToolParameter(BaseDataToolParameter): # add datasets hda_list = util.listify(other_values.get(self.name)) - for hda in history.active_datasets_and_roles: - match = dataset_matcher.hda_match(hda, check_security=False) + # Prefetch all at once, big list of visible, non-deleted datasets. + for hda in history.active_visible_datasets_and_roles: + match = dataset_matcher.hda_match(hda) if match: m = match.hda hda_list = [h for h in hda_list if h != m and h != hda] @@ -1829,8 +1824,8 @@ class DataToolParameter(BaseDataToolParameter): append(d['options']['hda'], hda, '(%s) %s' % (hda_state, hda.name), 'hda', True) # add dataset collections - dataset_collection_matcher = DatasetCollectionMatcher(dataset_matcher) - for hdca in history.active_dataset_collections: + dataset_collection_matcher = dataset_matcher_factory.dataset_collection_matcher(dataset_matcher) + for hdca in history.active_visible_dataset_collections: if dataset_collection_matcher.hdca_match(hdca, reduction=multiple): append(d['options']['hdca'], hdca, hdca.name, 'hdca') @@ -1866,19 +1861,16 @@ class DataCollectionToolParameter(BaseDataToolParameter): dataset_collection_type_descriptions = trans.app.dataset_collections_service.collection_type_descriptions return history_query.HistoryQuery.from_parameter(self, dataset_collection_type_descriptions) - def match_collections(self, trans, history, dataset_matcher): + def match_collections(self, trans, history, dataset_collection_matcher): dataset_collections = trans.app.dataset_collections_service.history_dataset_collections(history, self._history_query(trans)) - dataset_collection_matcher = DatasetCollectionMatcher(dataset_matcher) for dataset_collection_instance in dataset_collections: if not dataset_collection_matcher.hdca_match(dataset_collection_instance): continue yield dataset_collection_instance - def match_multirun_collections(self, trans, history, dataset_matcher): - dataset_collection_matcher = DatasetCollectionMatcher(dataset_matcher) - - for history_dataset_collection in history.active_dataset_collections: + def match_multirun_collections(self, trans, history, dataset_collection_matcher): + for history_dataset_collection in history.active_visible_dataset_collections: if not self._history_query(trans).can_map_over(history_dataset_collection): continue @@ -1954,10 +1946,12 @@ class DataCollectionToolParameter(BaseDataToolParameter): return d # prepare dataset/collection matching - dataset_matcher = DatasetMatcher(trans, self, None, other_values) + dataset_matcher_factory = get_dataset_matcher_factory(trans) + dataset_matcher = dataset_matcher_factory.dataset_matcher(self, other_values) + dataset_collection_matcher = dataset_matcher_factory.dataset_collection_matcher(dataset_matcher) # append directly matched collections - for hdca in self.match_collections(trans, history, dataset_matcher): + for hdca in self.match_collections(trans, history, dataset_collection_matcher): d['options']['hdca'].append({ 'id' : trans.security.encode_id(hdca.id), 'hid' : hdca.hid, @@ -1967,7 +1961,7 @@ class DataCollectionToolParameter(BaseDataToolParameter): }) # append matching subcollections - for hdca in self.match_multirun_collections(trans, history, dataset_matcher): + for hdca in self.match_multirun_collections(trans, history, dataset_collection_matcher): subcollection_type = self._history_query(trans).can_map_over(hdca).collection_type d['options']['hdca'].append({ 'id' : trans.security.encode_id(hdca.id), diff --git a/lib/galaxy/tools/parameters/dataset_matcher.py b/lib/galaxy/tools/parameters/dataset_matcher.py index 2cdb6277064..7fd5b8eea81 100644 --- a/lib/galaxy/tools/parameters/dataset_matcher.py +++ b/lib/galaxy/tools/parameters/dataset_matcher.py @@ -4,8 +4,82 @@ import galaxy.model log = getLogger(__name__) -ROLES_UNSET = object() -INVALID_STATES = [galaxy.model.Dataset.states.ERROR, galaxy.model.Dataset.states.DISCARDED] + +def set_dataset_matcher_factory(trans, tool, param_values): + trans.dataset_matcher_factory = DatasetMatcherFactory(trans, tool, param_values) + + +def unset_dataset_matcher_factory(trans): + trans.dataset_matcher_factory = None + + +def get_dataset_matcher_factory(trans): + dataset_matcher_factory = getattr(trans, "dataset_matcher_factory", None) + return dataset_matcher_factory or DatasetMatcherFactory(trans) + + +class DatasetMatcherFactory(object): + """""" + + def __init__(self, trans, tool=None, param_values=None): + self._trans = trans + self._tool = tool + self._data_inputs = [] + self._matches_format_cache = {} + if tool: + valid_input_states = tool.valid_input_states + else: + valid_input_states = galaxy.model.Dataset.valid_input_states + self.valid_input_states = valid_input_states + can_process_summary = False + if tool is not None and param_values is not None: + self._collect_data_inputs(tool, param_values) + + require_public = self._tool and self._tool.tool_type == 'data_destination' + if not require_public and self._data_inputs: + can_process_summary = True + for data_input in self._data_inputs: + if data_input.options: + can_process_summary = False + break + self._can_process_summary = can_process_summary + + def matches_any_format(self, hda_extension, formats): + for format in formats: + if self.matches_format(hda_extension, format): + return True + return False + + def matches_format(self, hda_extension, format): + # cache datatype checking combinations for fast recall + if hda_extension not in self._matches_format_cache: + self._matches_format_cache[hda_extension] = {} + + formats = self._matches_format_cache[hda_extension] + if format not in formats: + datatype = galaxy.model.datatype_for_extension(hda_extension, datatypes_registry=self._trans.app.datatypes_registry) + formats[format] = datatype.matches_any([format]) + + return formats[format] + + def _collect_data_inputs(self, tool, param_values): + def visitor(input, value, prefix, parent=None, **kwargs): + type_name = type(input).__name__ + if "DataToolParameter" in type_name: + self._data_inputs.append(input) + elif "DataCollectionToolParameter" in type_name: + self._data_inputs.append(input) + + tool.visit_inputs(param_values, visitor) + + def dataset_matcher(self, param, other_values): + return DatasetMatcher(self, self._trans, param, other_values) + + def dataset_collection_matcher(self, dataset_matcher): + if self._can_process_summary: + return SummaryDatasetCollectionMatcher(self, dataset_matcher) + else: + return DatasetCollectionMatcher(dataset_matcher) class DatasetMatcher(object): @@ -17,12 +91,11 @@ class DatasetMatcher(object): and permission handling. """ - def __init__(self, trans, param, value, other_values): + def __init__(self, dataset_matcher_factory, trans, param, other_values): + self.dataset_matcher_factory = dataset_matcher_factory self.trans = trans self.param = param self.tool = param.tool - self.value = value - self.current_user_roles = ROLES_UNSET filter_value = None if param.options and other_values: try: @@ -31,20 +104,7 @@ class DatasetMatcher(object): pass # no valid options self.filter_value = filter_value - def hda_accessible(self, hda, check_security=True): - """ Does HDA correspond to dataset that is an a valid state and is - accessible to user. - """ - dataset = hda.dataset - has_tool = self.tool - if has_tool: - valid_input_states = self.tool.valid_input_states - else: - valid_input_states = galaxy.model.Dataset.valid_input_states - state_valid = dataset.state in valid_input_states - return state_valid and (not check_security or self.__can_access_dataset(dataset)) - - def valid_hda_match(self, hda, check_implicit_conversions=True, check_security=False): + def valid_hda_match(self, hda, check_implicit_conversions=True): """ Return False of this parameter can not be matched to the supplied HDA, otherwise return a description of the match (either a HdaDirectMatch describing a direct match or a HdaImplicitMatch @@ -52,7 +112,7 @@ class DatasetMatcher(object): """ rval = False formats = self.param.formats - if hda.datatype.matches_any(formats): + if self.dataset_matcher_factory.matches_any_format(hda.extension, formats): rval = HdaDirectMatch(hda) else: if not check_implicit_conversions: @@ -62,8 +122,6 @@ class DatasetMatcher(object): original_hda = hda if converted_dataset: hda = converted_dataset - if check_security and not self.__can_access_dataset(hda.dataset): - return False rval = HdaImplicitMatch(hda, target_ext, original_hda) else: return False @@ -71,31 +129,21 @@ class DatasetMatcher(object): return False return rval - def hda_match(self, hda, check_implicit_conversions=True, check_security=True, ensure_visible=True): + def hda_match(self, hda, check_implicit_conversions=True, ensure_visible=True): """ If HDA is accessible, return information about whether it could match this parameter and if so how. See valid_hda_match for more information. """ - accessible = self.hda_accessible(hda, check_security=check_security) - if accessible and (not ensure_visible or hda.visible or (self.selected(hda) and not hda.implicitly_converted_parent_datasets)): + dataset = hda.dataset + valid_state = dataset.state in self.dataset_matcher_factory.valid_input_states + if valid_state and (not ensure_visible or hda.visible): # If we are sending data to an external application, then we need to make sure there are no roles # associated with the dataset that restrict its access from "public". require_public = self.tool and self.tool.tool_type == 'data_destination' - if require_public and not self.trans.app.security_agent.dataset_is_public(hda.dataset): - return False - if self.filter(hda): + if require_public and not self.trans.app.security_agent.dataset_is_public(dataset): return False return self.valid_hda_match(hda, check_implicit_conversions=check_implicit_conversions) - def selected(self, hda): - """ Given value for DataToolParameter, is this HDA "selected". - """ - value = self.value - if value and str(value[0]).isdigit(): - return hda.id in map(int, value) - else: - return value and hda in value - def filter(self, hda): """ Filter out this value based on other values for job (if applicable). @@ -103,12 +151,6 @@ class DatasetMatcher(object): param = self.param return param.options and param.get_options_filter_attribute(hda) != self.filter_value - def __can_access_dataset(self, dataset): - # Lazily cache current_user_roles. - if self.current_user_roles is ROLES_UNSET: - self.current_user_roles = self.trans.get_current_user_roles() - return self.trans.app.security_agent.can_access_dataset(self.current_user_roles, dataset) - class HdaDirectMatch(object): """ Supplied HDA was a valid option directly (did not need to find implicit @@ -138,6 +180,33 @@ class HdaImplicitMatch(object): return True +class SummaryDatasetCollectionMatcher(object): + + def __init__(self, dataset_matcher_factory, dataset_matcher): + self.dataset_matcher_factory = dataset_matcher_factory + self.dataset_matcher = dataset_matcher + + def hdca_match(self, history_dataset_collection_association, reduction=False): + dataset_collection = history_dataset_collection_association.collection + if reduction and dataset_collection.collection_type.find(":") > 0: + return False + + if not dataset_collection.populated_optimized: + return False + + (states, extensions) = dataset_collection.dataset_states_and_extensions_summary + for state in states: + if state not in self.dataset_matcher_factory.valid_input_states: + return False + + formats = self.dataset_matcher.param.formats + for extension in extensions: + if not self.dataset_matcher_factory.matches_any_format(extension, formats): + return False + + return True + + class DatasetCollectionMatcher(object): def __init__(self, dataset_matcher): diff --git a/test/unit/tools/test_data_parameters.py b/test/unit/tools/test_data_parameters.py index c7e0aa2e4e4..abc40180a5c 100644 --- a/test/unit/tools/test_data_parameters.py +++ b/test/unit/tools/test_data_parameters.py @@ -1,5 +1,4 @@ from galaxy import model -from galaxy.util import bunch from .test_parameter_parsing import BaseParameterTestCase from ..unittest_utils import galaxy_mock @@ -39,9 +38,9 @@ class DataToolParameterTestCase(BaseParameterTestCase): assert field['options']['hda'][0]['name'] == "hda2" assert field['options']['hda'][1]['name'] == "hda1" - hda2.datatype_matches = False + hda2.extension = 'data' field = self._simple_field() - assert len(field['options']['hda']) == 1 + assert len(field['options']['hda']) == 1, field assert field['options']['hda'][0]['name'] == "hda1" def test_field_display_hidden_hdas_only_if_selected(self): @@ -66,7 +65,7 @@ class DataToolParameterTestCase(BaseParameterTestCase): def test_field_implicit_conversion_new(self): hda1 = MockHistoryDatasetAssociation(name="hda1", id=1) - hda1.datatype_matches = False + hda1.extension = 'data' hda1.conversion_destination = ("tabular", None) self.stub_active_datasets(hda1) field = self._simple_field() @@ -76,7 +75,7 @@ class DataToolParameterTestCase(BaseParameterTestCase): def test_field_implicit_conversion_existing(self): hda1 = MockHistoryDatasetAssociation(name="hda1", id=1) - hda1.datatype_matches = False + hda1.extension = 'data' hda1.conversion_destination = ("tabular", MockHistoryDatasetAssociation(name="hda1converted", id=2)) self.stub_active_datasets(hda1) field = self._simple_field() @@ -124,7 +123,7 @@ class DataToolParameterTestCase(BaseParameterTestCase): def test_get_initial_with_previously_converted_data(self): hda1 = MockHistoryDatasetAssociation(name="hda1", id=1) - hda1.datatype_matches = False + hda1.extension = 'data' converted = MockHistoryDatasetAssociation(name="hda1converted", id=2) hda1.conversion_destination = ("tabular", converted) self.stub_active_datasets(hda1) @@ -132,10 +131,10 @@ class DataToolParameterTestCase(BaseParameterTestCase): def test_get_initial_with_to_be_converted_data(self): hda1 = MockHistoryDatasetAssociation(name="hda1", id=1) - hda1.datatype_matches = False + hda1.extension = 'data' hda1.conversion_destination = ("tabular", None) self.stub_active_datasets(hda1) - assert hda1 == self.param.get_initial_value(self.trans, {}) + assert hda1 == self.param.get_initial_value(self.trans, {}), hda1 def _new_hda(self): hda = model.HistoryDatasetAssociation() @@ -157,6 +156,7 @@ class DataToolParameterTestCase(BaseParameterTestCase): def stub_active_datasets(self, *hdas): self.test_history._active_datasets_and_roles = [h for h in hdas if not h.deleted] + self.test_history._active_visible_datasets_and_roles = [h for h in hdas if not h.deleted and h.visible] def _simple_field(self, **kwds): return self.param.to_dict(trans=self.trans, **kwds) @@ -170,7 +170,7 @@ class DataToolParameterTestCase(BaseParameterTestCase): optional_text = "" if self.optional: optional_text = 'optional="True"' - template_xml = '''''' + template_xml = '''''' param_str = template_xml % (multi_text, optional_text) self._param = self._parameter_for(tool=self.mock_tool, xml=param_str) @@ -190,11 +190,8 @@ class MockHistoryDatasetAssociation(object): self.deleted = False self.dataset = test_dataset self.visible = True - self.datatype_matches = True self.conversion_destination = (None, None) - self.datatype = bunch.Bunch( - matches_any=lambda formats: self.datatype_matches, - ) + self.extension = "txt" self.dbkey = "hg19" self.implicitly_converted_parent_datasets = False self.name = name diff --git a/test/unit/tools/test_dataset_matcher.py b/test/unit/tools/test_dataset_matcher.py index 6962f1e4449..c057c717c2d 100644 --- a/test/unit/tools/test_dataset_matcher.py +++ b/test/unit/tools/test_dataset_matcher.py @@ -13,32 +13,6 @@ from ..tools_support import UsesApp class DatasetMatcherTestCase(TestCase, UsesApp): - def test_hda_accessible(self): - # Cannot access errored or discard datasets. - self.mock_hda.dataset.state = model.Dataset.states.ERROR - assert not self.test_context.hda_accessible(self.mock_hda) - - self.mock_hda.dataset.state = model.Dataset.states.DISCARDED - assert not self.test_context.hda_accessible(self.mock_hda) - - # Can access datasets in other states. - self.mock_hda.dataset.state = model.Dataset.states.OK - assert self.test_context.hda_accessible(self.mock_hda) - - self.mock_hda.dataset.state = model.Dataset.states.QUEUED - assert self.test_context.hda_accessible(self.mock_hda) - - # Cannot access dataset if security agent says no. - self.app.security_agent.can_access_dataset = lambda roles, dataset: False - assert not self.test_context.hda_accessible(self.mock_hda) - - def test_selected(self): - self.test_context.value = [] - assert not self.test_context.selected(self.mock_hda) - - self.test_context.value = [self.mock_hda] - assert self.test_context.selected(self.mock_hda) - def test_hda_mismatches(self): # Datasets not visible are not "valid" for param. self.mock_hda.visible = False @@ -46,13 +20,13 @@ class DatasetMatcherTestCase(TestCase, UsesApp): # Datasets that don't match datatype are not valid. self.mock_hda.visible = True - self.mock_hda.datatype_matches = False + self.mock_hda.extension = 'data' assert not self.test_context.hda_match(self.mock_hda) def test_valid_hda_direct_match(self): # Datasets that visible and matching are valid self.mock_hda.visible = True - self.mock_hda.datatype_matches = True + self.mock_hda.extension = 'txt' hda_match = self.test_context.hda_match(self.mock_hda, check_implicit_conversions=False) assert hda_match @@ -64,7 +38,7 @@ class DatasetMatcherTestCase(TestCase, UsesApp): def test_valid_hda_implicit_convered(self): # Find conversion returns an HDA to an already implicitly converted # dataset. - self.mock_hda.datatype_matches = False + self.mock_hda.extension = 'data' converted_hda = model.HistoryDatasetAssociation() self.mock_hda.conversion_destination = ("tabular", converted_hda) hda_match = self.test_context.hda_match(self.mock_hda) @@ -77,7 +51,7 @@ class DatasetMatcherTestCase(TestCase, UsesApp): def test_hda_match_implicit_can_convert(self): # Find conversion returns a target extension to convert to, but not # a previously implicitly converted dataset. - self.mock_hda.datatype_matches = False + self.mock_hda.extension = 'data' self.mock_hda.conversion_destination = ("tabular", None) hda_match = self.test_context.hda_match(self.mock_hda) @@ -87,7 +61,7 @@ class DatasetMatcherTestCase(TestCase, UsesApp): assert hda_match.target_ext == "tabular" def test_hda_match_properly_skips_conversion(self): - self.mock_hda.datatype_matches = False + self.mock_hda.extension = 'data' self.mock_hda.conversion_destination = ("tabular", bunch.Bunch()) hda_match = self.test_context.hda_match(self.mock_hda, check_implicit_conversions=False) assert not hda_match @@ -148,20 +122,18 @@ class DatasetMatcherTestCase(TestCase, UsesApp): option_xml = "" if self.filtered_param: option_xml = '''''' - param_xml = XML('''%s''' % option_xml) + param_xml = XML('''%s''' % option_xml) self.param = basic.DataToolParameter( self.tool, param_xml, ) - - self._test_context = dataset_matcher.DatasetMatcher( - trans=bunch.Bunch( - app=self.app, - get_current_user_roles=lambda: self.current_user_roles, - workflow_building_mode=True, - ), + trans = bunch.Bunch( + app=self.app, + get_current_user_roles=lambda: self.current_user_roles, + workflow_building_mode=True, + ) + self._test_context = dataset_matcher.get_dataset_matcher_factory(trans).dataset_matcher( param=self.param, - value=[], other_values=self.other_values )