From bfa60a039dee1d81685f685df991e6ea4c6cb812 Mon Sep 17 00:00:00 2001 From: "Michael R. Crusoe" Date: Tue, 9 Nov 2021 18:10:03 +0100 Subject: [PATCH] improve typing of lib/galaxy/tools/wrappers.py --- lib/galaxy/datatypes/binary.py | 14 +- lib/galaxy/datatypes/data.py | 11 +- lib/galaxy/datatypes/registry.py | 9 +- lib/galaxy/model/__init__.py | 37 +- lib/galaxy/model/metadata.py | 11 +- lib/galaxy/tools/parameters/basic.py | 4 +- lib/galaxy/tools/parameters/wrapped_json.py | 13 +- lib/galaxy/tools/wrappers.py | 386 +++++++++++++----- .../visualization/data_providers/genome.py | 2 +- setup.cfg | 12 +- test/unit/app/tools/test_wrappers.py | 17 +- 11 files changed, 362 insertions(+), 154 deletions(-) diff --git a/lib/galaxy/datatypes/binary.py b/lib/galaxy/datatypes/binary.py index 9de3825e558..3d171cf1fe5 100644 --- a/lib/galaxy/datatypes/binary.py +++ b/lib/galaxy/datatypes/binary.py @@ -25,6 +25,7 @@ from bx.seq.twobit import TWOBIT_MAGIC_NUMBER, TWOBIT_MAGIC_NUMBER_SWAP from galaxy import util from galaxy.datatypes import metadata from galaxy.datatypes.data import ( + Data, DatatypeValidation, get_file_peek, ) @@ -267,7 +268,10 @@ class Bref3(Binary): class DynamicCompressedArchive(CompressedArchive): - def matches_any(self, target_datatypes): + compressed_format: str + uncompressed_datatype_instance: Data + + def matches_any(self, target_datatypes) -> bool: """Treat two aspects of compressed datatypes separately. """ compressed_target_datatypes = [] @@ -281,8 +285,12 @@ class DynamicCompressedArchive(CompressedArchive): # TODO: Add gz and bz2 as proper datatypes and use those instances instead of # CompressedArchive() in the following check. - return self.uncompressed_datatype_instance.matches_any(uncompressed_target_datatypes) or \ - CompressedArchive().matches_any(compressed_target_datatypes) + if not hasattr(self, "uncompressed_datatype_instance"): + raise Exception("Missing 'uncompressed_datatype_instance' attribute.") + else: + return self.uncompressed_datatype_instance.matches_any( + uncompressed_target_datatypes + ) or CompressedArchive().matches_any(compressed_target_datatypes) class GzDynamicCompressedArchive(DynamicCompressedArchive): diff --git a/lib/galaxy/datatypes/data.py b/lib/galaxy/datatypes/data.py index ae2c9c89b7c..58926ca2e3a 100644 --- a/lib/galaxy/datatypes/data.py +++ b/lib/galaxy/datatypes/data.py @@ -6,7 +6,7 @@ import shutil import string import tempfile from inspect import isclass -from typing import Any, Dict, Optional +from typing import Any, Dict, List, Optional, Tuple, TYPE_CHECKING import webob.exc from markupsafe import escape @@ -29,6 +29,9 @@ from . import ( metadata ) +if TYPE_CHECKING: + from galaxy.model import DatasetInstance + XSS_VULNERABLE_MIME_TYPES = [ 'image/svg+xml', # Unfiltered by Galaxy and may contain JS that would be executed by some browsers. 'application/xml', # Some browsers will evalute SVG embedded JS in such XML documents. @@ -633,7 +636,9 @@ class Data(metaclass=DataMeta): """Returns available converters by type for this dataset""" return datatypes_registry.get_converters_by_datatype(original_dataset.ext) - def find_conversion_destination(self, dataset, accepted_formats, datatypes_registry, **kwd): + def find_conversion_destination( + self, dataset, accepted_formats: List[str], datatypes_registry, **kwd + ) -> Tuple[bool, Optional[str], Optional["DatasetInstance"]]: """Returns ( direct_match, converted_ext, existing converted dataset )""" return datatypes_registry.find_conversion_destination_for_dataset_by_extensions(dataset, accepted_formats, **kwd) @@ -726,7 +731,7 @@ class Data(metaclass=DataMeta): def has_resolution(self): return False - def matches_any(self, target_datatypes): + def matches_any(self, target_datatypes: List[Any]) -> bool: """ Check if this datatype is of any of the target_datatypes or is a subtype thereof. diff --git a/lib/galaxy/datatypes/registry.py b/lib/galaxy/datatypes/registry.py index a2669f872ea..cf13e051355 100644 --- a/lib/galaxy/datatypes/registry.py +++ b/lib/galaxy/datatypes/registry.py @@ -6,7 +6,7 @@ import imp import logging import os from string import Template -from typing import Dict +from typing import Dict, List, Optional, Tuple, TYPE_CHECKING import yaml @@ -28,6 +28,9 @@ from . import ( ) from .display_applications.application import DisplayApplication +if TYPE_CHECKING: + from galaxy.model import DatasetInstance + class ConfigurationError(Exception): pass @@ -885,7 +888,9 @@ class Registry: return converters[target_ext] return None - def find_conversion_destination_for_dataset_by_extensions(self, dataset_or_ext, accepted_formats, converter_safe=True): + def find_conversion_destination_for_dataset_by_extensions( + self, dataset_or_ext, accepted_formats: List[str], converter_safe: bool = True + ) -> Tuple[bool, Optional[str], Optional["DatasetInstance"]]: """ returns (direct_match, converted_ext, converted_dataset) - direct match is True iff no the data set already has an accepted format diff --git a/lib/galaxy/model/__init__.py b/lib/galaxy/model/__init__.py index 29950ebfce1..a141cb910e2 100644 --- a/lib/galaxy/model/__init__.py +++ b/lib/galaxy/model/__init__.py @@ -27,6 +27,7 @@ from typing import ( List, NamedTuple, Optional, + Tuple, Type, TYPE_CHECKING, Union, @@ -143,6 +144,8 @@ AUTO_PROPAGATED_TAGS = ["name"] if TYPE_CHECKING: + from galaxy.datatypes.data import Data + class _HasTable: table: Table __table__: Table @@ -210,8 +213,9 @@ def set_datatypes_registry(d_registry): class HasTags: - dict_collection_visible_keys = ['tags'] - dict_element_visible_keys = ['tags'] + dict_collection_visible_keys = ["tags"] + dict_element_visible_keys = ["tags"] + tags: List["ItemTagAssociation"] def to_dict(self, *args, **kwargs): rval = super().to_dict(*args, **kwargs) @@ -3449,7 +3453,7 @@ class DatasetHash(Base, Serializable): return rval -def datatype_for_extension(extension, datatypes_registry=None): +def datatype_for_extension(extension, datatypes_registry=None) -> "Data": if extension is not None: extension = extension.lower() if datatypes_registry is None: @@ -3545,12 +3549,12 @@ class DatasetInstance(_HasTable): object_session(self).flush() # flush here, because hda.flush() won't flush the Dataset object state = property(get_dataset_state, set_dataset_state) - def get_file_name(self): + def get_file_name(self) -> str: if self.dataset.purged: return "" return self.dataset.get_file_name() - def set_file_name(self, filename): + def set_file_name(self, filename: str): return self.dataset.set_file_name(filename) file_name = property(get_file_name, set_file_name) @@ -3568,7 +3572,7 @@ class DatasetInstance(_HasTable): return self.dataset.extra_files_path_exists() @property - def datatype(self): + def datatype(self) -> "Data": return datatype_for_extension(self.extension) def get_metadata(self): @@ -3595,7 +3599,7 @@ class DatasetInstance(_HasTable): meta_types.append(meta_type) return meta_types - def get_metadata_file_paths_and_extensions(self): + def get_metadata_file_paths_and_extensions(self) -> List[Tuple[str, str]]: metadata = self.metadata metadata_files = [] for metadata_name in self.metadata_file_types: @@ -3802,7 +3806,9 @@ class DatasetInstance(_HasTable): def can_convert_to(self, format): return format in self.get_converter_types() - def find_conversion_destination(self, accepted_formats, **kwd): + def find_conversion_destination( + self, accepted_formats: List[str], **kwd + ) -> Tuple[bool, Optional[str], Optional["DatasetInstance"]]: """Returns ( target_ext, existing converted dataset )""" return self.datatype.find_conversion_destination(self, accepted_formats, _get_datatypes_registry(), **kwd) @@ -5263,7 +5269,9 @@ class DatasetCollection(Base, Dictifiable, UsesAnnotations, Serializable): return [(row[:-2], row.extension, row.Dataset.file_name) for row in q] @property - def element_identifiers_extensions_paths_and_metadata_files(self): + def element_identifiers_extensions_paths_and_metadata_files( + self, + ) -> List[List[Any]]: q = self._get_nested_collection_attributes( element_attributes=('element_identifier',), hda_attributes=('extension',), @@ -8125,6 +8133,7 @@ class ItemTagAssociation(Dictifiable): dict_element_visible_keys = dict_collection_visible_keys associated_item_names: List[str] = [] user_tname: Column + user_value = Column(TrimmedString(255), index=True) def __init_subclass__(cls, **kwargs): super().__init_subclass__(**kwargs) @@ -8151,7 +8160,6 @@ class HistoryTagAssociation(Base, ItemTagAssociation, RepresentById): user_id = Column(Integer, ForeignKey('galaxy_user.id'), index=True) user_tname = Column(TrimmedString(255), index=True) value = Column(TrimmedString(255), index=True) - user_value = Column(TrimmedString(255), index=True) history = relationship('History', back_populates='tags') tag = relationship('Tag') user = relationship('User') @@ -8167,7 +8175,6 @@ class HistoryDatasetAssociationTagAssociation(Base, ItemTagAssociation, Represen user_id = Column(Integer, ForeignKey('galaxy_user.id'), index=True) user_tname = Column(TrimmedString(255), index=True) value = Column(TrimmedString(255), index=True) - user_value = Column(TrimmedString(255), index=True) history_dataset_association = relationship('HistoryDatasetAssociation', back_populates='tags') tag = relationship('Tag') user = relationship('User') @@ -8183,7 +8190,6 @@ class LibraryDatasetDatasetAssociationTagAssociation(Base, ItemTagAssociation, R user_id = Column(Integer, ForeignKey('galaxy_user.id'), index=True) user_tname = Column(TrimmedString(255), index=True) value = Column(TrimmedString(255), index=True) - user_value = Column(TrimmedString(255), index=True) library_dataset_dataset_association = relationship( 'LibraryDatasetDatasetAssociation', back_populates='tags') tag = relationship('Tag') @@ -8199,7 +8205,6 @@ class PageTagAssociation(Base, ItemTagAssociation, RepresentById): user_id = Column(Integer, ForeignKey('galaxy_user.id'), index=True) user_tname = Column(TrimmedString(255), index=True) value = Column(TrimmedString(255), index=True) - user_value = Column(TrimmedString(255), index=True) page = relationship('Page', back_populates='tags') tag = relationship('Tag') user = relationship('User') @@ -8214,7 +8219,6 @@ class WorkflowStepTagAssociation(Base, ItemTagAssociation, RepresentById): user_id = Column(Integer, ForeignKey('galaxy_user.id'), index=True) user_tname = Column(TrimmedString(255), index=True) value = Column(TrimmedString(255), index=True) - user_value = Column(TrimmedString(255), index=True) workflow_step = relationship('WorkflowStep', back_populates='tags') tag = relationship('Tag') user = relationship('User') @@ -8229,7 +8233,6 @@ class StoredWorkflowTagAssociation(Base, ItemTagAssociation, RepresentById): user_id = Column(Integer, ForeignKey('galaxy_user.id'), index=True) user_tname = Column(TrimmedString(255), index=True) value = Column(TrimmedString(255), index=True) - user_value = Column(TrimmedString(255), index=True) stored_workflow = relationship('StoredWorkflow', back_populates='tags') tag = relationship('Tag') user = relationship('User') @@ -8244,7 +8247,6 @@ class VisualizationTagAssociation(Base, ItemTagAssociation, RepresentById): user_id = Column(Integer, ForeignKey('galaxy_user.id'), index=True) user_tname = Column(TrimmedString(255), index=True) value = Column(TrimmedString(255), index=True) - user_value = Column(TrimmedString(255), index=True) visualization = relationship('Visualization', back_populates='tags') tag = relationship('Tag') user = relationship('User') @@ -8260,7 +8262,6 @@ class HistoryDatasetCollectionTagAssociation(Base, ItemTagAssociation, Represent user_id = Column(Integer, ForeignKey('galaxy_user.id'), index=True) user_tname = Column(TrimmedString(255), index=True) value = Column(TrimmedString(255), index=True) - user_value = Column(TrimmedString(255), index=True) dataset_collection = relationship('HistoryDatasetCollectionAssociation', back_populates='tags') tag = relationship('Tag') user = relationship('User') @@ -8276,7 +8277,6 @@ class LibraryDatasetCollectionTagAssociation(Base, ItemTagAssociation, Represent user_id = Column(Integer, ForeignKey('galaxy_user.id'), index=True) user_tname = Column(TrimmedString(255), index=True) value = Column(TrimmedString(255), index=True) - user_value = Column(TrimmedString(255), index=True) dataset_collection = relationship('LibraryDatasetCollectionAssociation', back_populates='tags') tag = relationship('Tag') user = relationship('User') @@ -8291,7 +8291,6 @@ class ToolTagAssociation(Base, ItemTagAssociation, RepresentById): user_id = Column(Integer, ForeignKey('galaxy_user.id'), index=True) user_tname = Column(TrimmedString(255), index=True) value = Column(TrimmedString(255), index=True) - user_value = Column(TrimmedString(255), index=True) tag = relationship('Tag') user = relationship('User') diff --git a/lib/galaxy/model/metadata.py b/lib/galaxy/model/metadata.py index 9d448d6df36..a77048b3f95 100644 --- a/lib/galaxy/model/metadata.py +++ b/lib/galaxy/model/metadata.py @@ -13,6 +13,7 @@ import weakref from collections import OrderedDict from collections.abc import Mapping from os.path import abspath +from typing import Any, Iterator, TYPE_CHECKING, Union from sqlalchemy.orm import object_session from sqlalchemy.orm.attributes import flag_modified @@ -28,6 +29,10 @@ from galaxy.util import ( ) from galaxy.util.json import safe_dumps +if TYPE_CHECKING: + from galaxy.model import DatasetInstance + from galaxy.model.none_like import NoneDataset + log = logging.getLogger(__name__) STATEMENTS = "__galaxy_statements__" # this is the name of the property in a Datatype class where new metadata spec element Statements are stored @@ -66,7 +71,7 @@ class MetadataCollection(Mapping): retrieved, returning default values in cases when metadata is not set. """ - def __init__(self, parent): + def __init__(self, parent: Union["DatasetInstance", "NoneDataset"]) -> None: self.parent = parent # initialize dict if needed if self.parent._metadata is None: @@ -88,7 +93,7 @@ class MetadataCollection(Mapping): def spec(self): return self.parent.datatype.metadata_spec - def __iter__(self): + def __iter__(self) -> Iterator[Any]: yield from self.spec.keys() def __getitem__(self, key): @@ -142,7 +147,7 @@ class MetadataCollection(Mapping): else: log.info(f"Attempted to delete invalid key '{name}' from MetadataCollection") - def element_is_set(self, name): + def element_is_set(self, name) -> bool: """ check if the meta data with the given name is set, i.e. diff --git a/lib/galaxy/tools/parameters/basic.py b/lib/galaxy/tools/parameters/basic.py index d085eeb73d4..3425f4d1828 100644 --- a/lib/galaxy/tools/parameters/basic.py +++ b/lib/galaxy/tools/parameters/basic.py @@ -256,7 +256,7 @@ class ToolParameter(Dictifiable): return str_value return "Not available." - def to_param_dict_string(self, value, other_values=None): + def to_param_dict_string(self, value, other_values=None) -> str: """Called via __str__ when used in the Cheetah template""" if value is None: value = "" @@ -845,6 +845,8 @@ class SelectToolParameter(ToolParameter): y,z """ + value_label: str + def __init__(self, tool, input_source, context=None): input_source = ensure_input_source(input_source) super().__init__(tool, input_source) diff --git a/lib/galaxy/tools/parameters/wrapped_json.py b/lib/galaxy/tools/parameters/wrapped_json.py index 1e419204437..82c37e8dccb 100644 --- a/lib/galaxy/tools/parameters/wrapped_json.py +++ b/lib/galaxy/tools/parameters/wrapped_json.py @@ -1,4 +1,5 @@ import logging +from typing import Any, Dict, List, Sequence import packaging.version @@ -33,8 +34,12 @@ def data_collection_input_to_path(v): return v.all_paths -def data_collection_input_to_staging_path_and_source_path(v, invalid_chars=('/',), include_collection_name=False): - staging_paths = v.get_all_staging_paths(invalid_chars=invalid_chars, include_collection_name=include_collection_name) +def data_collection_input_to_staging_path_and_source_path( + v, invalid_chars: Sequence[str] = ("/",), include_collection_name: bool = False +) -> List[Dict[str, Any]]: + staging_paths = v.get_all_staging_paths( + invalid_chars=invalid_chars, include_collection_name=include_collection_name + ) source_paths = v.all_paths metadata_files = v.all_metadata_files return [ @@ -44,7 +49,9 @@ def data_collection_input_to_staging_path_and_source_path(v, invalid_chars=('/', } for staging_path, source_path, metadata_files in zip(staging_paths, source_paths, metadata_files)] -def data_input_to_staging_path_and_source_path(v, invalid_chars=('/',)): +def data_input_to_staging_path_and_source_path( + v, invalid_chars: Sequence[str] = ("/",) +) -> Dict[str, Any]: staging_path = v.get_staging_path(invalid_chars=invalid_chars) return { 'staging_path': staging_path, diff --git a/lib/galaxy/tools/wrappers.py b/lib/galaxy/tools/wrappers.py index d91aeb189d8..80a575249ef 100644 --- a/lib/galaxy/tools/wrappers.py +++ b/lib/galaxy/tools/wrappers.py @@ -3,8 +3,30 @@ import os import shlex import tempfile from functools import total_ordering +from typing import ( + Any, + cast, + Dict, + Iterable, + Iterator, + KeysView, + List, + Optional, + Sequence, + Tuple, + TYPE_CHECKING, + Union, +) from galaxy import exceptions +from galaxy.model import ( + DatasetCollection, + DatasetCollectionElement, + DatasetCollectionInstance, + DatasetInstance, + HasTags, + HistoryDatasetCollectionAssociation, +) from galaxy.model.none_like import NoneDataset from galaxy.security.object_wrapper import wrap_with_safe_string from galaxy.tools.parameters.wrapped_json import ( @@ -13,6 +35,13 @@ from galaxy.tools.parameters.wrapped_json import ( ) from galaxy.util import filesystem_safe_string +if TYPE_CHECKING: + from galaxy.tools import Tool + from galaxy.tools.parameters.basic import SelectToolParameter, ToolParameter + from galaxy.datatypes.registry import Registry + from galaxy.jobs import ComputeEnvironment + from galaxy.model.metadata import MetadataCollection + log = logging.getLogger(__name__) # Fields in .log files corresponding to paths, must have one of the following @@ -27,11 +56,14 @@ class ToolParameterValueWrapper: Base class for object that Wraps a Tool Parameter and Value. """ - def __bool__(self): + value: Union[str, List[str]] + input: "ToolParameter" + + def __bool__(self) -> bool: return bool(self.value) __nonzero__ = __bool__ - def get_display_text(self, quote=True): + def get_display_text(self, quote: bool = True) -> str: """ Returns a string containing the value that would be displayed to the user in the tool interface. When quote is True (default), the string is escaped for e.g. command-line usage. @@ -47,21 +79,21 @@ class RawObjectWrapper(ToolParameterValueWrapper): Wraps an object so that __str__ returns module_name:class_name. """ - def __init__(self, obj): + def __init__(self, obj: Any): self.obj = obj - def __bool__(self): + def __bool__(self) -> bool: return bool(self.obj) # FIXME: would it be safe/backwards compatible to rename .obj to .value, so that we can just inherit this method? __nonzero__ = __bool__ - def __str__(self): + def __str__(self) -> str: try: return f"{self.obj.__module__}:{self.obj.__class__.__name__}" except Exception: # Most likely None, which lacks __module__. return str(self.obj) - def __getattr__(self, key): + def __getattr__(self, key: Any) -> Any: return getattr(self.obj, key) @@ -71,12 +103,17 @@ class InputValueWrapper(ToolParameterValueWrapper): Wraps an input so that __str__ gives the "param_dict" representation. """ - def __init__(self, input, value, other_values=None): + def __init__( + self, + input: "ToolParameter", + value: str, + other_values: Optional[Dict[str, str]] = None, + ) -> None: self.input = input self.value = value - self._other_values = other_values or {} + self._other_values: Dict[str, str] = other_values or {} - def _get_cast_value(self, other): + def _get_cast_value(self, other: Any) -> Union[str, int, float, bool, None]: if self.input.type == 'boolean' and isinstance(other, str): return str(self) # For backward compatibility, allow `$wrapper != ""` for optional non-text param @@ -85,44 +122,46 @@ class InputValueWrapper(ToolParameterValueWrapper): return str(self) else: return None - cast = { + cast_table = { 'text': str, 'integer': int, 'float': float, 'boolean': bool, } - return cast.get(self.input.type, str)(self) + return cast( + Union[str, int, float, bool], cast_table.get(self.input.type, str)(self) + ) - def __eq__(self, other): - return self._get_cast_value(other) == other + def __eq__(self, other: Any) -> bool: + return bool(self._get_cast_value(other) == other) - def __ne__(self, other): + def __ne__(self, other: Any) -> bool: return not self == other - def __str__(self): + def __str__(self) -> str: to_param_dict_string = self.input.to_param_dict_string(self.value, self._other_values) if isinstance(to_param_dict_string, list): return ','.join(to_param_dict_string) else: return to_param_dict_string - def __iter__(self): + def __iter__(self) -> Iterable[str]: to_param_dict_string = self.input.to_param_dict_string(self.value, self._other_values) if not isinstance(to_param_dict_string, list): return iter([to_param_dict_string]) else: return iter(to_param_dict_string) - def __getattr__(self, key): + def __getattr__(self, key: Any) -> Any: return getattr(self.value, key) - def __gt__(self, other): - return self._get_cast_value(other) > other + def __gt__(self, other: Any) -> bool: + return bool(self._get_cast_value(other) > other) - def __int__(self): + def __int__(self) -> int: return int(float(self)) - def __float__(self): + def __float__(self) -> float: return float(str(self)) @@ -132,20 +171,28 @@ class SelectToolParameterWrapper(ToolParameterValueWrapper): attributes are accessible. """ + input: "SelectToolParameter" + class SelectToolParameterFieldWrapper: """ Provide access to any field by name or index for this particular value. Only applicable for dynamic_options selects, which have more than simple 'options' defined (name, value, selected). """ - def __init__(self, input, value, other_values, compute_environment): + def __init__( + self, + input: "SelectToolParameter", + value: Union[str, List[str]], + other_values: Optional[Dict[str, str]], + compute_environment: Optional["ComputeEnvironment"], + ) -> None: self._input = input self._value = value self._other_values = other_values - self._fields = {} + self._fields: Dict[str, str] = {} self._compute_environment = compute_environment - def __getattr__(self, name): + def __getattr__(self, name: str) -> Any: if name not in self._fields: self._fields[name] = self._input.options.get_field_by_name_for_value(name, self._value, None, self._other_values) values = map(str, self._fields[name]) @@ -158,12 +205,17 @@ class SelectToolParameterWrapper(ToolParameterValueWrapper): new_values.append(rewrite_value) else: new_values.append(value) - - values = new_values + return self._input.separator.join(new_values) return self._input.separator.join(values) - def __init__(self, input, value, other_values=None, compute_environment=None): + def __init__( + self, + input: "SelectToolParameter", + value: Union[str, List[str]], + other_values: Optional[Dict[str, str]] = None, + compute_environment: Optional["ComputeEnvironment"] = None, + ): self.input = input self.value = value self.input.value_label = input.value_to_display_text(value) @@ -171,30 +223,32 @@ class SelectToolParameterWrapper(ToolParameterValueWrapper): self.compute_environment = compute_environment self.fields = self.SelectToolParameterFieldWrapper(input, value, other_values, self.compute_environment) - def __eq__(self, other): + def __eq__(self, other: Any) -> bool: if isinstance(other, str): if other == '' and self.value in (None, []): # Allow $wrapper == '' for select (self.value is None) and multiple select (self.value is []) params return True return str(self) == other else: - return super() == other + return super().__eq__(other) - def __ne__(self, other): + def __ne__(self, other: Any) -> bool: return not self == other - def __str__(self): + def __str__(self) -> str: # Assuming value is never a path - otherwise would need to pass # along following argument value_map=self._path_rewriter. - return self.input.to_param_dict_string(self.value, other_values=self._other_values) + return str( + self.input.to_param_dict_string(self.value, other_values=self._other_values) + ) - def __add__(self, x): + def __add__(self, x: Any) -> str: return f'{self}{x}' - def __getattr__(self, key): + def __getattr__(self, key: Any) -> Any: return getattr(self.input, key) - def __iter__(self): + def __iter__(self) -> Iterable[str]: if not self.input.multiple: raise Exception("Tried to iterate over a non-multiple parameter.") return self.value.__iter__() @@ -206,6 +260,8 @@ class DatasetFilenameWrapper(ToolParameterValueWrapper): attributes are accessible. """ + false_path: Optional[str] + class MetadataWrapper: """ Wraps a Metadata Collection to return MetadataParameters wrapped @@ -213,12 +269,16 @@ class DatasetFilenameWrapper(ToolParameterValueWrapper): of a Metadata Collection. """ - def __init__(self, dataset, compute_environment=None): + def __init__( + self, + dataset: DatasetInstance, + compute_environment: Optional["ComputeEnvironment"] = None, + ) -> None: self.dataset = dataset - self.metadata = dataset.metadata + self.metadata: "MetadataCollection" = dataset.metadata self.compute_environment = compute_environment - def __getattr__(self, name): + def __getattr__(self, name: str) -> Any: rval = self.metadata.get(name, None) if name in self.metadata.spec: if rval is None: @@ -239,33 +299,49 @@ class DatasetFilenameWrapper(ToolParameterValueWrapper): rval = wrap_with_safe_string(rval) return rval - def __bool__(self): - return self.metadata.__nonzero__() + def __bool__(self) -> bool: + return bool(self.metadata.__nonzero__()) __nonzero__ = __bool__ - def __iter__(self): + def __iter__(self) -> Iterator[Any]: return self.metadata.__iter__() - def element_is_set(self, name): + def element_is_set(self, name: str) -> bool: return self.metadata.element_is_set(name) - def get(self, key, default=None): + def get(self, key: str, default: Any = None) -> Any: try: return getattr(self, key) except Exception: return default - def items(self): + def items(self) -> Iterator[Tuple[str, Any]]: return iter((k, self.get(k)) for k, v in self.metadata.items()) - def __init__(self, dataset, datatypes_registry=None, tool=None, name=None, compute_environment=None, identifier=None, io_type="input", formats=None): + def __init__( + self, + dataset: Optional[DatasetInstance], + datatypes_registry: Optional["Registry"] = None, + tool: Optional["Tool"] = None, + name: Optional[str] = None, + compute_environment: Optional["ComputeEnvironment"] = None, + identifier: Optional[str] = None, + io_type: str = "input", + formats: Optional[List[str]] = None, + ) -> None: if not dataset: try: # TODO: allow this to work when working with grouping - ext = tool.inputs[name].extensions[0] + ext = tool.inputs[name].extensions[0] # type: ignore except Exception: - ext = 'data' - self.dataset = wrap_with_safe_string(NoneDataset(datatypes_registry=datatypes_registry, ext=ext), no_wrap_classes=ToolParameterValueWrapper) + ext = "data" + self.dataset = cast( + DatasetInstance, + wrap_with_safe_string( + NoneDataset(datatypes_registry=datatypes_registry, ext=ext), + no_wrap_classes=ToolParameterValueWrapper, + ), + ) else: # Tool wrappers should not normally be accessing .dataset directly, # so we will wrap it and keep the original around for file paths @@ -274,10 +350,11 @@ class DatasetFilenameWrapper(ToolParameterValueWrapper): direct_match, target_ext, converted_dataset = dataset.find_conversion_destination(formats) if not direct_match and target_ext and converted_dataset: dataset = converted_dataset - self.unsanitized = dataset + self.unsanitized: DatasetInstance = dataset self.dataset = wrap_with_safe_string(dataset, no_wrap_classes=ToolParameterValueWrapper) + assert dataset self.metadata = self.MetadataWrapper(dataset, compute_environment) - if hasattr(dataset, 'tags'): + if isinstance(dataset, HasTags): self.groups = {tag.user_value.lower() for tag in dataset.tags if tag.user_tname == 'group'} else: # May be a 'FakeDatasetAssociation' @@ -301,21 +378,27 @@ class DatasetFilenameWrapper(ToolParameterValueWrapper): self._element_identifier = identifier @property - def element_identifier(self): + def element_identifier(self) -> str: identifier = self._element_identifier if identifier is None: identifier = self.name return identifier @property - def file_ext(self): - return getattr(self.unsanitized.datatype, 'file_ext_export_alias', self.dataset.extension) + def file_ext(self) -> str: + return str( + getattr( + self.unsanitized.datatype, + "file_ext_export_alias", + self.dataset.extension, + ) + ) @property - def name_and_ext(self): + def name_and_ext(self) -> str: return f"{self.element_identifier}.{self.file_ext}" - def get_staging_path(self, invalid_chars=('/',)): + def get_staging_path(self, invalid_chars: Sequence[str] = ("/",)) -> str: """ Strip leading dots, unicode null chars, replace `/` with `_`, truncate at 255 characters. @@ -326,18 +409,26 @@ class DatasetFilenameWrapper(ToolParameterValueWrapper): return f"{safe_element_identifier}.{self.file_ext}" @property - def all_metadata_files(self): + def all_metadata_files(self) -> List[Tuple[str, str]]: return self.unsanitized.get_metadata_file_paths_and_extensions() if self else [] - def serialize(self, invalid_chars=('/',)): - return data_input_to_staging_path_and_source_path(self, invalid_chars=invalid_chars) if self else {} + def serialize(self, invalid_chars: Sequence[str] = ("/",)) -> Dict[str, Any]: + return ( + data_input_to_staging_path_and_source_path( + self, invalid_chars=invalid_chars + ) + if self + else {} + ) @property - def is_collection(self): + def is_collection(self) -> bool: return False - def is_of_type(self, *exts): + def is_of_type(self, *exts: str) -> bool: datatypes = [] + if not self.datatypes_registry: + raise Exception("datatypes_registry is required to use 'is_of_type'.") for e in exts: datatype = self.datatypes_registry.get_datatype_by_extension(e) if datatype is not None: @@ -346,13 +437,13 @@ class DatasetFilenameWrapper(ToolParameterValueWrapper): log.warning(f"Datatype class not found for extension '{e}', which is used as parameter of 'is_of_type()' method") return self.dataset.datatype.matches_any(datatypes) - def __str__(self): + def __str__(self) -> str: if self.false_path is not None: return self.false_path else: - return self.unsanitized.file_name + return str(self.unsanitized.file_name) - def __getattr__(self, key): + def __getattr__(self, key: Any) -> Any: if self.false_path is not None and key == 'file_name': # Path to dataset was rewritten for this job. return self.false_path @@ -385,17 +476,24 @@ class DatasetFilenameWrapper(ToolParameterValueWrapper): else: return getattr(self.dataset, key) - def __bool__(self): + def __bool__(self) -> bool: return bool(self.dataset) __nonzero__ = __bool__ class HasDatasets: - def _dataset_wrapper(self, dataset, **kwargs): + job_working_directory: Optional[str] + + def __iter__(self) -> Iterator[Any]: + pass + + def _dataset_wrapper( + self, dataset: DatasetInstance, **kwargs: Any + ) -> DatasetFilenameWrapper: return DatasetFilenameWrapper(dataset, **kwargs) - def paths_as_file(self, sep="\n"): + def paths_as_file(self, sep: str = "\n") -> str: contents = sep.join(map(str, self)) with tempfile.NamedTemporaryFile(mode='w+', prefix="gx_file_list", dir=self.job_working_directory, delete=False) as fh: fh.write(contents) @@ -403,28 +501,54 @@ class HasDatasets: return filepath -class DatasetListWrapper(list, ToolParameterValueWrapper, HasDatasets): - """ - """ +class DatasetListWrapper( + List[DatasetFilenameWrapper], ToolParameterValueWrapper, HasDatasets +): + """ """ - def __init__(self, job_working_directory, datasets, **kwargs): - self._dataset_elements_cache = {} - if not isinstance(datasets, list): + def __init__( + self, + job_working_directory: Optional[str], + datasets: Union[ + Sequence[ + Union[ + None, + DatasetInstance, + DatasetCollectionInstance, + DatasetCollectionElement, + ] + ], + DatasetInstance, + ], + **kwargs: Any, + ) -> None: + self._dataset_elements_cache: Dict[str, List[DatasetFilenameWrapper]] = {} + if not isinstance(datasets, Sequence): datasets = [datasets] - def to_wrapper(dataset): - if hasattr(dataset, "dataset_instance"): - element = dataset - dataset = element.dataset_instance - kwargs["identifier"] = element.element_identifier - return self._dataset_wrapper(dataset, **kwargs) + def to_wrapper( + dataset: Union[ + None, + DatasetInstance, + DatasetCollectionInstance, + DatasetCollectionElement, + ] + ) -> DatasetFilenameWrapper: + if isinstance(dataset, DatasetCollectionElement): + dataset2 = dataset.dataset_instance + kwargs["identifier"] = dataset.element_identifier + else: + dataset2 = dataset + return self._dataset_wrapper(dataset2, **kwargs) list.__init__(self, map(to_wrapper, datasets)) self.job_working_directory = job_working_directory @staticmethod - def to_dataset_instances(dataset_instance_sources): - dataset_instances = [] + def to_dataset_instances( + dataset_instance_sources: Any, + ) -> List[Union[None, DatasetInstance]]: + dataset_instances: List[Optional[DatasetInstance]] = [] if not isinstance(dataset_instance_sources, list): dataset_instance_sources = [dataset_instance_sources] for dataset_instance_source in dataset_instance_sources: @@ -438,7 +562,7 @@ class DatasetListWrapper(list, ToolParameterValueWrapper, HasDatasets): dataset_instances.extend(dataset_instance_source.collection.dataset_elements) return dataset_instances - def get_datasets_for_group(self, group): + def get_datasets_for_group(self, group: str) -> List[DatasetFilenameWrapper]: group = str(group).lower() if not self._dataset_elements_cache.get(group): wrappers = [] @@ -448,25 +572,37 @@ class DatasetListWrapper(list, ToolParameterValueWrapper, HasDatasets): self._dataset_elements_cache[group] = wrappers return self._dataset_elements_cache[group] - def serialize(self, invalid_chars=('/',)): + def serialize(self, invalid_chars: Sequence[str] = ("/",)) -> List[Dict[str, Any]]: return [v.serialize(invalid_chars) for v in self] - def __str__(self): + def __str__(self) -> str: return ','.join(map(str, self)) - def __bool__(self): + def __bool__(self) -> bool: # Fail `#if $param` checks in cheetah if optional input is not provided return any(self) __nonzero__ = __bool__ class DatasetCollectionWrapper(ToolParameterValueWrapper, HasDatasets): + name: Optional[str] + collection: DatasetCollection - def __init__(self, job_working_directory, has_collection, datatypes_registry, **kwargs): + def __init__( + self, + job_working_directory: Optional[str], + has_collection: Union[ + None, DatasetCollectionElement, HistoryDatasetCollectionAssociation + ], + datatypes_registry: "Registry", + **kwargs: Any, + ) -> None: super().__init__() self.job_working_directory = job_working_directory - self._dataset_elements_cache = {} - self._element_identifiers_extensions_paths_and_metadata_files = None + self._dataset_elements_cache: Dict[str, List[DatasetFilenameWrapper]] = {} + self._element_identifiers_extensions_paths_and_metadata_files: Optional[ + List[List[Any]] + ] = None self.datatypes_registry = datatypes_registry kwargs['datatypes_registry'] = datatypes_registry self.kwargs = kwargs @@ -477,12 +613,10 @@ class DatasetCollectionWrapper(ToolParameterValueWrapper, HasDatasets): else: self.__input_supplied = True - if hasattr(has_collection, "name"): - # It is a HistoryDatasetCollectionAssociation + if isinstance(has_collection, HistoryDatasetCollectionAssociation): collection = has_collection.collection self.name = has_collection.name - elif hasattr(has_collection, "child_collection"): - # It is a DatasetCollectionElement instance referencing another collection + elif isinstance(has_collection, DatasetCollectionElement): collection = has_collection.child_collection self.name = has_collection.element_identifier else: @@ -499,7 +633,11 @@ class DatasetCollectionWrapper(ToolParameterValueWrapper, HasDatasets): element_identifier = dataset_collection_element.element_identifier if dataset_collection_element.is_collection: - element_wrapper = DatasetCollectionWrapper(job_working_directory, dataset_collection_element, **kwargs) + element_wrapper: Union[ + DatasetCollectionWrapper, DatasetFilenameWrapper + ] = DatasetCollectionWrapper( + job_working_directory, dataset_collection_element, **kwargs + ) else: element_wrapper = self._dataset_wrapper(element_object, identifier=element_identifier, **kwargs) @@ -509,7 +647,7 @@ class DatasetCollectionWrapper(ToolParameterValueWrapper, HasDatasets): self.__element_instances = element_instances self.__element_instance_list = element_instance_list - def get_datasets_for_group(self, group): + def get_datasets_for_group(self, group: str) -> List[DatasetFilenameWrapper]: group = str(group).lower() if not self._dataset_elements_cache.get(group): wrappers = [] @@ -519,37 +657,47 @@ class DatasetCollectionWrapper(ToolParameterValueWrapper, HasDatasets): self._dataset_elements_cache[group] = wrappers return self._dataset_elements_cache[group] - def keys(self): + def keys(self) -> Union[List[str], KeysView[Any]]: if not self.__input_supplied: return [] return self.__element_instances.keys() @property - def is_collection(self): + def is_collection(self) -> bool: return True @property - def element_identifier(self): + def element_identifier(self) -> Optional[str]: return self.name @property - def all_paths(self): + def all_paths(self) -> List[str]: return [path for _, _, path, _ in self.element_identifiers_extensions_paths_and_metadata_files] @property - def all_metadata_files(self): + def all_metadata_files(self) -> List[List[str]]: return [metadata_files for _, _, _, metadata_files in self.element_identifiers_extensions_paths_and_metadata_files] @property - def element_identifiers_extensions_paths_and_metadata_files(self): + def element_identifiers_extensions_paths_and_metadata_files( + self, + ) -> List[List[Any]]: if self._element_identifiers_extensions_paths_and_metadata_files is None: if self.collection: - self._element_identifiers_extensions_paths_and_metadata_files = self.collection.element_identifiers_extensions_paths_and_metadata_files + result = ( + self.collection.element_identifiers_extensions_paths_and_metadata_files + ) + self._element_identifiers_extensions_paths_and_metadata_files = result + return result else: return [] return self._element_identifiers_extensions_paths_and_metadata_files - def get_all_staging_paths(self, invalid_chars=('/',), include_collection_name=False): + def get_all_staging_paths( + self, + invalid_chars: Sequence[str] = ("/",), + include_collection_name: bool = False, + ) -> List[str]: safe_element_identifiers = [] for element_identifiers, extension, *_ in self.element_identifiers_extensions_paths_and_metadata_files: datatype = self.datatypes_registry.get_datatype_by_extension(extension) @@ -559,7 +707,7 @@ class DatasetCollectionWrapper(ToolParameterValueWrapper, HasDatasets): for element_identifier in element_identifiers: max_len = 254 - len(extension) if include_collection_name: - max_len = max_len - (len(self.name) + 1) + max_len = max_len - (len(self.name or "") + 1) assert max_len >= 1, 'Could not stage element, element identifier is too long' current_element_identifier = filesystem_safe_string(element_identifier, max_len=max_len, invalid_chars=invalid_chars) if include_collection_name and self.name: @@ -569,14 +717,24 @@ class DatasetCollectionWrapper(ToolParameterValueWrapper, HasDatasets): safe_element_identifiers.append(f'{os.path.sep.join(current_element_identifiers)}.{extension}') return safe_element_identifiers - def serialize(self, invalid_chars=('/',), include_collection_name=False): - return data_collection_input_to_staging_path_and_source_path(self, invalid_chars=invalid_chars, include_collection_name=include_collection_name) + def serialize( + self, + invalid_chars: Sequence[str] = ("/",), + include_collection_name: bool = False, + ) -> List[Dict[str, Any]]: + return data_collection_input_to_staging_path_and_source_path( + self, + invalid_chars=invalid_chars, + include_collection_name=include_collection_name, + ) @property - def is_input_supplied(self): + def is_input_supplied(self) -> bool: return self.__input_supplied - def __getitem__(self, key): + def __getitem__( + self, key: Union[str, int] + ) -> Union[None, "DatasetCollectionWrapper", DatasetFilenameWrapper]: if not self.__input_supplied: return None if isinstance(key, int): @@ -584,7 +742,9 @@ class DatasetCollectionWrapper(ToolParameterValueWrapper, HasDatasets): else: return self.__element_instances[key] - def __getattr__(self, key): + def __getattr__( + self, key: str + ) -> Union[None, "DatasetCollectionWrapper", DatasetFilenameWrapper]: if not self.__input_supplied: return None try: @@ -592,12 +752,14 @@ class DatasetCollectionWrapper(ToolParameterValueWrapper, HasDatasets): except KeyError: raise AttributeError() - def __iter__(self): + def __iter__( + self, + ) -> Iterator[Union["DatasetCollectionWrapper", DatasetFilenameWrapper]]: if not self.__input_supplied: return [].__iter__() return self.__element_instance_list.__iter__() - def __bool__(self): + def __bool__(self) -> bool: # Fail `#if $param` checks in cheetah is optional input # not specified or if resulting collection is empty. return self.__input_supplied and bool(self.__element_instance_list) @@ -607,13 +769,15 @@ class DatasetCollectionWrapper(ToolParameterValueWrapper, HasDatasets): class ElementIdentifierMapper: """Track mapping of dataset collection elements datasets to element identifiers.""" - def __init__(self, input_datasets=None): + def __init__(self, input_datasets: Optional[Dict[str, Any]] = None) -> None: if input_datasets is not None: self.identifier_key_dict = {v: f"{k}|__identifier__" for k, v in input_datasets.items()} else: self.identifier_key_dict = {} - def identifier(self, dataset_value, input_values): + def identifier( + self, dataset_value: str, input_values: Dict[str, str] + ) -> Optional[str]: identifier_key = self.identifier_key_dict.get(dataset_value, None) element_identifier = None if identifier_key: diff --git a/lib/galaxy/visualization/data_providers/genome.py b/lib/galaxy/visualization/data_providers/genome.py index 0291258bc38..cd6f5d622ba 100644 --- a/lib/galaxy/visualization/data_providers/genome.py +++ b/lib/galaxy/visualization/data_providers/genome.py @@ -234,7 +234,7 @@ class GenomeDataProvider(BaseDataProvider): """ # Get column names. try: - column_names = self.original_dataset.datatype.column_names + column_names = self.original_dataset.datatype.column_names # type: ignore except AttributeError: try: column_names = list(range(self.original_dataset.metadata.columns)) diff --git a/setup.cfg b/setup.cfg index d6001d672f8..a85298c1255 100644 --- a/setup.cfg +++ b/setup.cfg @@ -414,7 +414,17 @@ check_untyped_defs = False [mypy-tool_shed.util.metadata_util] check_untyped_defs = False [mypy-galaxy.tools.wrappers] -check_untyped_defs = False +disallow_any_generics = True +disallow_subclassing_any = True +disallow_untyped_calls = False +disallow_untyped_defs = True +disallow_incomplete_defs = True +check_untyped_defs = True +disallow_untyped_decorators = True +no_implicit_optional = True +warn_unused_ignores = True +no_implicit_reexport = True +strict_equality = True [mypy-galaxy.tools.error_reports.plugins.base_git] check_untyped_defs = False [mypy-galaxy.tool_util.deps.mulled.mulled_build] diff --git a/test/unit/app/tools/test_wrappers.py b/test/unit/app/tools/test_wrappers.py index 6a025569b68..fb94a026280 100644 --- a/test/unit/app/tools/test_wrappers.py +++ b/test/unit/app/tools/test_wrappers.py @@ -1,11 +1,14 @@ import os import tempfile +from typing import cast from unittest.mock import Mock import pytest from galaxy.datatypes.metadata import MetadataSpecCollection from galaxy.job_execution.datasets import DatasetPath +from galaxy.jobs import ComputeEnvironment +from galaxy.model import DatasetInstance from galaxy.tools.parameters.basic import ( BooleanToolParameter, DrillDownSelectToolParameter, @@ -101,7 +104,7 @@ def test_select_wrapper_multiple(tool): @with_mock_tool def test_select_wrapper_with_path_rewritting(tool): parameter = _setup_blast_tool(tool, multiple=True) - compute_environment = MockComputeEnvironment(None) + compute_environment = cast(ComputeEnvironment, MockComputeEnvironment(None)) wrapper = SelectToolParameterWrapper(parameter, ["val1", "val2"], other_values={}, compute_environment=compute_environment) assert wrapper.fields.path == "Rewrite,Rewrite" @@ -190,7 +193,7 @@ def test_input_value_wrapper_input_value_wrapper_comparison(tool): def test_dataset_wrapper(): - dataset = MockDataset() + dataset = cast(DatasetInstance, MockDataset()) wrapper = DatasetFilenameWrapper(dataset) assert str(wrapper) == MOCK_DATASET_PATH assert wrapper.file_name == MOCK_DATASET_PATH @@ -199,9 +202,9 @@ def test_dataset_wrapper(): def test_dataset_wrapper_false_path(): - dataset = MockDataset() + dataset = cast(DatasetInstance, MockDataset()) new_path = "/new/path/dataset_123.dat" - wrapper = DatasetFilenameWrapper(dataset, compute_environment=MockComputeEnvironment(false_path=new_path)) + wrapper = DatasetFilenameWrapper(dataset, compute_environment=cast(ComputeEnvironment, MockComputeEnvironment(false_path=new_path))) assert str(wrapper) == new_path assert wrapper.file_name == new_path @@ -223,19 +226,19 @@ class MockComputeEnvironment: def test_dataset_false_extra_files_path(): - dataset = MockDataset() + dataset = cast(DatasetInstance, MockDataset()) wrapper = DatasetFilenameWrapper(dataset) assert wrapper.extra_files_path == MOCK_DATASET_EXTRA_FILES_PATH new_path = "/new/path/dataset_123.dat" dataset_path = DatasetPath(123, MOCK_DATASET_PATH, false_path=new_path) - wrapper = DatasetFilenameWrapper(dataset, compute_environment=MockComputeEnvironment(dataset_path)) + wrapper = DatasetFilenameWrapper(dataset, compute_environment=cast(ComputeEnvironment, MockComputeEnvironment(dataset_path))) # Setting false_path is not enough to override assert wrapper.extra_files_path == MOCK_DATASET_EXTRA_FILES_PATH new_files_path = "/new/path/dataset_123_files" - wrapper = DatasetFilenameWrapper(dataset, compute_environment=MockComputeEnvironment(false_path=new_path, false_extra_files_path=new_files_path)) + wrapper = DatasetFilenameWrapper(dataset, compute_environment=cast(ComputeEnvironment, MockComputeEnvironment(false_path=new_path, false_extra_files_path=new_files_path))) assert wrapper.extra_files_path == new_files_path