improve typing of lib/galaxy/tools/wrappers.py

This commit is contained in:
Michael R. Crusoe
2021-11-11 12:09:00 +01:00
parent 1c171e4463
commit bfa60a039d
11 changed files with 362 additions and 154 deletions
+11 -3
View File
@@ -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):
+8 -3
View File
@@ -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.
+7 -2
View File
@@ -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
+18 -19
View File
@@ -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')
+8 -3
View File
@@ -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.
+3 -1
View File
@@ -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)
+10 -3
View File
@@ -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,
+275 -111
View File
@@ -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:
@@ -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))
+11 -1
View File
@@ -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]
+10 -7
View File
@@ -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<path1>,Rewrite<path2>"
@@ -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