diff --git a/lib/galaxy/app.py b/lib/galaxy/app.py index b33f8ef335d..7043d7a3fea 100644 --- a/lib/galaxy/app.py +++ b/lib/galaxy/app.py @@ -232,7 +232,7 @@ class MinimalGalaxyApplication(BasicSharedApp, HaltableContainer, SentryClientMi self.name = "galaxy" self.is_webapp = False # Read config file and check for errors - self.config: Any = self._register_singleton(config.Configuration, config.Configuration(**kwargs)) + self.config = self._register_singleton(config.GalaxyAppConfiguration, config.GalaxyAppConfiguration(**kwargs)) self.config.check() self._configure_object_store(fsmon=True) self._register_singleton(BaseObjectStore, self.object_store) @@ -393,13 +393,6 @@ class MinimalGalaxyApplication(BasicSharedApp, HaltableContainer, SentryClientMi self.security = IdEncodingHelper(id_secret=self.config.id_secret) BaseDatabaseIdField.security = self.security - def _configure_tool_shed_registry(self): - # Set up the tool sheds registry - if os.path.isfile(self.config.tool_sheds_config_file): - self.tool_shed_registry = tool_shed_registry.Registry(self.config.tool_sheds_config_file) - else: - self.tool_shed_registry = tool_shed_registry.Registry() - def _configure_engines(self, db_url, install_db_url, combined_install_database): trace_logger = getattr(self, "trace_logger", None) engine = build_engine( @@ -574,6 +567,16 @@ class GalaxyManagerApplication(MinimalManagerApp, MinimalGalaxyApplication): galaxy.model.set_datatypes_registry(self.datatypes_registry) self.configure_sentry_client() + self._configure_tool_shed_registry() + self._register_singleton(tool_shed_registry.Registry, self.tool_shed_registry) + + def _configure_tool_shed_registry(self) -> None: + # Set up the tool sheds registry + if os.path.isfile(self.config.tool_sheds_config_file): + self.tool_shed_registry = tool_shed_registry.Registry(self.config.tool_sheds_config_file) + else: + self.tool_shed_registry = tool_shed_registry.Registry() + @property def is_job_handler(self) -> bool: return ( @@ -610,9 +613,6 @@ class UniverseApplication(StructuredApp, GalaxyManagerApplication): # want to and we'll allow postfork to bind and start it. self.queue_worker = self._register_singleton(GalaxyQueueWorker, GalaxyQueueWorker(self)) - self._configure_tool_shed_registry() - self._register_singleton(tool_shed_registry.Registry, self.tool_shed_registry) - self.dependency_resolvers_view = self._register_singleton( DependencyResolversView, DependencyResolversView(self) ) @@ -672,7 +672,7 @@ class UniverseApplication(StructuredApp, GalaxyManagerApplication): # Tours registry tour_registry = build_tours_registry(self.config.tour_config_dir) self.tour_registry = tour_registry - self[ToursRegistry] = tour_registry # type: ignore[misc] + self[ToursRegistry] = tour_registry # type: ignore[type-abstract] # Webhooks registry self.webhooks_registry = self._register_singleton(WebhooksRegistry, WebhooksRegistry(self.config.webhooks_dir)) # Heartbeat for thread profiling diff --git a/lib/galaxy/app_unittest_utils/galaxy_mock.py b/lib/galaxy/app_unittest_utils/galaxy_mock.py index 8030fb8d4dc..6351b73fb9f 100644 --- a/lib/galaxy/app_unittest_utils/galaxy_mock.py +++ b/lib/galaxy/app_unittest_utils/galaxy_mock.py @@ -104,7 +104,7 @@ class MockApp(di.Container, GalaxyDataTestApp): job_metrics: JobMetrics stop: bool - def __init__(self, config=None, **kwargs): + def __init__(self, config=None, **kwargs) -> None: super().__init__() config = config or MockAppConfig(**kwargs) GalaxyDataTestApp.__init__(self, config=config, **kwargs) @@ -117,8 +117,8 @@ class MockApp(di.Container, GalaxyDataTestApp): self[GalaxyModelMapping] = self.model sts_config = ShortTermStorageConfiguration(short_term_storage_directory=os.path.join(config.data_dir, "sts")) sts_manager = ShortTermStorageManager(sts_config) - self[ShortTermStorageAllocator] = sts_manager # type: ignore[misc] - self[ShortTermStorageMonitor] = sts_manager # type: ignore[misc] + self[ShortTermStorageAllocator] = sts_manager # type: ignore[type-abstract] + self[ShortTermStorageMonitor] = sts_manager # type: ignore[type-abstract] self[galaxy_scoped_session] = self.model.context self.visualizations_registry = MockVisualizationsRegistry() self.tag_handler = tags.GalaxyTagHandler(self.model.context) diff --git a/lib/galaxy/managers/annotatable.py b/lib/galaxy/managers/annotatable.py index d6ec6aaa6af..65b76ea6f0a 100644 --- a/lib/galaxy/managers/annotatable.py +++ b/lib/galaxy/managers/annotatable.py @@ -6,7 +6,8 @@ import abc import logging from typing import Dict -from galaxy.model.scoped_session import galaxy_scoped_session +from sqlalchemy.orm import scoped_session + from .base import ( Deserializer, FunctionFilterParsersType, @@ -34,7 +35,7 @@ class AnnotatableManagerMixin: annotation_assoc: type @abc.abstractmethod - def session(self) -> galaxy_scoped_session: + def session(self) -> scoped_session: ... def annotation(self, item): diff --git a/lib/galaxy/managers/base.py b/lib/galaxy/managers/base.py index 3a92353b995..6dcb861cbbb 100644 --- a/lib/galaxy/managers/base.py +++ b/lib/galaxy/managers/base.py @@ -227,7 +227,7 @@ class ModelManager(Generic[U]): def session(self) -> scoped_session: return self.app.model.context - def _session_setattr(self, item: model._HasTable, attr: str, val: Any, flush: bool = True): + def _session_setattr(self, item: model.Base, attr: str, val: Any, flush: bool = True): setattr(item, attr, val) self.session().add(item) diff --git a/lib/galaxy/managers/datasets.py b/lib/galaxy/managers/datasets.py index 7f4036a76e0..21086287909 100644 --- a/lib/galaxy/managers/datasets.py +++ b/lib/galaxy/managers/datasets.py @@ -295,7 +295,6 @@ class DatasetAssociationManager( # DA's were meant to be proxies - but were never fully implemented as them # Instead, a dataset association HAS a dataset but contains metadata specific to a library (lda) or user (hda) - model_class: Type[model.DatasetInstance] app: MinimalManagerApp # NOTE: model_manager_class should be set in HDA/LDA subclasses diff --git a/lib/galaxy/managers/deletable.py b/lib/galaxy/managers/deletable.py index b3aaf3802f9..cbfa49eabec 100644 --- a/lib/galaxy/managers/deletable.py +++ b/lib/galaxy/managers/deletable.py @@ -14,7 +14,7 @@ from typing import ( Set, ) -from galaxy.model import _HasTable +from galaxy.model import Base from .base import ( Deserializer, ModelValidator, @@ -32,7 +32,7 @@ class DeletableManagerMixin: removed by an admin/script. """ - def _session_setattr(self, item: _HasTable, attr: str, val: Any, flush: bool = True): + def _session_setattr(self, item: Base, attr: str, val: Any, flush: bool = True): ... def delete(self, item, flush=True, **kwargs): @@ -91,7 +91,7 @@ class PurgableManagerMixin(DeletableManagerMixin): file). """ - def _session_setattr(self, item: _HasTable, attr: str, val: Any, flush: bool = True): + def _session_setattr(self, item: Base, attr: str, val: Any, flush: bool = True): ... def purge(self, item, flush=True, **kwargs): diff --git a/lib/galaxy/managers/hdas.py b/lib/galaxy/managers/hdas.py index e37f04c4ce1..ab4fc15df73 100644 --- a/lib/galaxy/managers/hdas.py +++ b/lib/galaxy/managers/hdas.py @@ -101,14 +101,14 @@ class HDAManager( # return True return super().is_accessible(item, user, **kwargs) - def is_owner(self, item: model._HasTable, user: Optional[model.User], current_history=None, **kwargs: Any) -> bool: + def is_owner(self, item: model.Base, user: Optional[model.User], current_history=None, **kwargs: Any) -> bool: """ Use history to see if current user owns HDA. """ - if self.user_manager.is_admin(user, trans=kwargs.get("trans", None)): - return True if not isinstance(item, model.HistoryDatasetAssociation): raise TypeError('"item" must be of type HistoryDatasetAssociation.') + if self.user_manager.is_admin(user, trans=kwargs.get("trans", None)): + return True history = item.history if history is None: raise HistoryDatasetAssociationNoHistoryException diff --git a/lib/galaxy/managers/histories.py b/lib/galaxy/managers/histories.py index 007f2bfe2b6..2e66215e47c 100644 --- a/lib/galaxy/managers/histories.py +++ b/lib/galaxy/managers/histories.py @@ -92,7 +92,7 @@ class HistoryManager(sharable.SharableModelManager, deletable.PurgableManagerMix def is_owner( self, - item: model._HasTable, + item: model.Base, user: Optional[model.User], current_history: Optional[model.History] = None, **kwargs: Any, diff --git a/lib/galaxy/managers/model_stores.py b/lib/galaxy/managers/model_stores.py index 60cfe228709..63a1eb6684f 100644 --- a/lib/galaxy/managers/model_stores.py +++ b/lib/galaxy/managers/model_stores.py @@ -187,8 +187,8 @@ class ModelStoreManager: import_options, model_store_format=request.model_store_format, ) - new_history = history is None and not request.for_library - if new_history: + create_new_history = history is None and not request.for_library + if create_new_history: if not model_import_store.defines_new_history(): raise RequestParameterInvalidException("Supplied model store doesn't define new history to import.") with model_import_store.target_history(legacy_history_naming=False) as new_history: @@ -197,7 +197,7 @@ class ModelStoreManager: else: object_tracker = model_import_store.perform_import( history=history, - new_history=new_history, + new_history=create_new_history, ) return object_tracker @@ -220,8 +220,8 @@ def create_objects_from_store( import_options=import_options, model_store_format=payload.model_store_format, ) - new_history = history is None and not for_library - if new_history: + create_new_history = history is None and not for_library + if create_new_history: if not model_import_store.defines_new_history(): raise RequestParameterInvalidException("Supplied model store doesn't define new history to import.") with model_import_store.target_history(legacy_history_naming=False) as new_history: @@ -230,6 +230,6 @@ def create_objects_from_store( else: object_tracker = model_import_store.perform_import( history=history, - new_history=new_history, + new_history=create_new_history, ) return object_tracker diff --git a/lib/galaxy/managers/secured.py b/lib/galaxy/managers/secured.py index 8b41092b5c9..f3e739c7b94 100644 --- a/lib/galaxy/managers/secured.py +++ b/lib/galaxy/managers/secured.py @@ -99,7 +99,7 @@ class OwnableManagerMixin: def by_id(self, id: int): ... - def is_owner(self, item: model._HasTable, user: Optional[model.User], **kwargs: Any) -> bool: + def is_owner(self, item: model.Base, user: Optional[model.User], **kwargs: Any) -> bool: """ Return True if user owns the item. """ diff --git a/lib/galaxy/metadata/set_metadata.py b/lib/galaxy/metadata/set_metadata.py index e6495ce67a6..24812559909 100644 --- a/lib/galaxy/metadata/set_metadata.py +++ b/lib/galaxy/metadata/set_metadata.py @@ -275,17 +275,13 @@ def set_metadata_portable( strip_metadata_files=False, serialize_jobs=True, ) - try: - import_model_store = store.imported_store_for_metadata( - tool_job_working_directory / "metadata/outputs_new", object_store=object_store - ) - except AssertionError: - # Remove in 21.09, this should only happen for jobs that started on <= 20.09 and finish now - import_model_store = None + import_model_store = store.imported_store_for_metadata( + tool_job_working_directory / "metadata/outputs_new", object_store=object_store + ) tool_script_file = tool_job_working_directory / "tool_script.sh" job = None - if import_model_store and export_store: + if export_store: job = next(iter(import_model_store.sa_session.objects[Job].values())) job_context = SessionlessJobContext( @@ -361,15 +357,7 @@ def set_metadata_portable( for output_name, output_dict in outputs.items(): dataset_instance_id = output_dict["id"] klass = getattr(galaxy.model, output_dict.get("model_class", "HistoryDatasetAssociation")) - dataset = None - if import_model_store: - dataset = import_model_store.sa_session.query(klass).find(dataset_instance_id) - if dataset is None: - # legacy check for jobs that started before 21.01, remove on 21.05 - filename_in = os.path.join(f"metadata/metadata_in_{output_name}") - import pickle - - dataset = pickle.load(open(filename_in, "rb")) # load DatasetInstance + dataset = import_model_store.sa_session.query(klass).find(dataset_instance_id) assert dataset is not None filename_kwds = tool_job_working_directory / f"metadata/metadata_kwds_{output_name}" diff --git a/lib/galaxy/model/__init__.py b/lib/galaxy/model/__init__.py index 8ab322eb12d..98dbb0ffe4c 100644 --- a/lib/galaxy/model/__init__.py +++ b/lib/galaxy/model/__init__.py @@ -3830,6 +3830,7 @@ class DatasetInstance(UsesCreateAndUpdateTime, _HasTable): conversion_messages = Dataset.conversion_messages permitted_actions = Dataset.permitted_actions purged: bool + creating_job_associations: List[Union[JobToOutputDatasetCollectionAssociation, JobToOutputDatasetAssociation]] class validated_states(str, Enum): UNKNOWN = "unknown" diff --git a/lib/galaxy/model/base.py b/lib/galaxy/model/base.py index 8c7a2efa06c..5d9b4a3bfd5 100644 --- a/lib/galaxy/model/base.py +++ b/lib/galaxy/model/base.py @@ -80,7 +80,7 @@ class ModelMapping(Bunch): del self.scoped_registry.registry[request_id] @property - def context(self): + def context(self) -> scoped_session: return self.session @property diff --git a/lib/galaxy/model/store/__init__.py b/lib/galaxy/model/store/__init__.py index d235c898f7c..3e2b6e9f370 100644 --- a/lib/galaxy/model/store/__init__.py +++ b/lib/galaxy/model/store/__init__.py @@ -14,12 +14,15 @@ from json import ( dumps, load, ) +from pathlib import PosixPath from tempfile import mkdtemp from typing import ( Any, Callable, cast, Dict, + Iterable, + Iterator, List, Optional, Set, @@ -37,6 +40,7 @@ from rocrate.model.computationalworkflow import ( ) from rocrate.rocrate import ROCrate from sqlalchemy.orm import joinedload +from sqlalchemy.orm.scoping import scoped_session from sqlalchemy.sql import expression from typing_extensions import Protocol @@ -56,7 +60,10 @@ from galaxy.model.orm.util import ( get_object_session, ) from galaxy.model.tags import GalaxyTagHandler -from galaxy.objectstore import ObjectStore +from galaxy.objectstore import ( + BaseObjectStore, + ObjectStore, +) from galaxy.schema.bco import ( BioComputeObjectCore, DescriptionDomain, @@ -102,7 +109,10 @@ from ..item_attrs import ( from ... import model if TYPE_CHECKING: + from tags import GalaxyTagHandlerSession + from galaxy.managers.workflows import WorkflowContentsManager + from galaxy.model import ImplicitCollectionJobs log = logging.getLogger(__name__) @@ -136,7 +146,7 @@ class StoreAppProtocol(Protocol): """Define the parts of a Galaxy-like app consumed by model store.""" datatypes_registry: Registry - object_store: ObjectStore + object_store: BaseObjectStore security: IdEncodingHelper tag_handler: GalaxyTagHandler model: GalaxyModelMapping @@ -164,11 +174,11 @@ class ImportOptions: def __init__( self, - allow_edit=False, - allow_library_creation=False, - allow_dataset_object_edit=None, - discarded_data=DEFAULT_DISCARDED_DATA_TYPE, - ): + allow_edit: bool = False, + allow_library_creation: bool = False, + allow_dataset_object_edit: Optional[bool] = None, + discarded_data: ImportDiscardedDataType = DEFAULT_DISCARDED_DATA_TYPE, + ) -> None: self.allow_edit = allow_edit self.allow_library_creation = allow_library_creation if allow_dataset_object_edit is None: @@ -179,16 +189,16 @@ class ImportOptions: class SessionlessContext: - def __init__(self): + def __init__(self) -> None: self.objects: Dict[Any, Any] = defaultdict(dict) - def flush(self): + def flush(self) -> None: pass - def add(self, obj): + def add(self, obj: Union[model.HistoryDatasetAssociation, model.Job]) -> None: self.objects[obj.__class__][obj.id] = obj - def query(self, model_class): + def query(self, model_class: Type[model.HistoryDatasetAssociation]) -> Bunch: def find(obj_id): return self.objects.get(model_class, {}).get(obj_id) or None @@ -199,7 +209,11 @@ class SessionlessContext: return Bunch(find=find, get=find, filter_by=filter_by) -def replace_metadata_file(metadata: Dict[str, Any], dataset_instance: model.DatasetInstance, sa_session): +def replace_metadata_file( + metadata: Dict[str, Any], + dataset_instance: model.DatasetInstance, + sa_session: Union[SessionlessContext, scoped_session], +) -> Dict[str, Optional[Union[model.MetadataFile, List[str], int, str]]]: def remap_objects(p, k, obj): if isinstance(obj, dict) and "model_class" in obj and obj["model_class"] == "MetadataFile": metadata_file = model.MetadataFile(dataset=dataset_instance, uuid=obj["uuid"]) @@ -212,14 +226,15 @@ def replace_metadata_file(metadata: Dict[str, Any], dataset_instance: model.Data class ModelImportStore(metaclass=abc.ABCMeta): app: Optional[StoreAppProtocol] + archive_dir: str def __init__( self, - import_options=None, + import_options: Optional[ImportOptions] = None, app: Optional[StoreAppProtocol] = None, - user=None, - object_store=None, - tag_handler=None, + user: Optional[model.User] = None, + object_store: Optional[ObjectStore] = None, + tag_handler: Optional["GalaxyTagHandlerSession"] = None, ) -> None: if object_store is None: if app is not None: @@ -241,6 +256,10 @@ class ModelImportStore(metaclass=abc.ABCMeta): else: self.import_history_encoded_id = None + @abc.abstractmethod + def workflow_paths(self) -> Iterator[Tuple[str, str]]: + pass + @abc.abstractmethod def defines_new_history(self) -> bool: """Does this store define a new history to create.""" @@ -257,6 +276,10 @@ class ModelImportStore(metaclass=abc.ABCMeta): """Return a list of library properties.""" return [] + @abc.abstractmethod + def invocations_properties(self) -> List[Any]: + pass + @abc.abstractmethod def collections_properties(self) -> List[Dict[str, Any]]: """Return a list of HDCA properties.""" @@ -265,6 +288,10 @@ class ModelImportStore(metaclass=abc.ABCMeta): def jobs_properties(self) -> List[Dict[str, Any]]: """Return a list of jobs properties.""" + @abc.abstractmethod + def implicit_collection_jobs_properties(self) -> List[Any]: + pass + @abc.abstractproperty def object_key(self) -> str: """Key used to connect objects in metadata. @@ -286,7 +313,9 @@ class ModelImportStore(metaclass=abc.ABCMeta): ) @contextlib.contextmanager - def target_history(self, default_history=None, legacy_history_naming=True): + def target_history( + self, default_history: Optional[model.History] = None, legacy_history_naming: bool = True + ) -> Iterator[Optional[model.History]]: new_history = None if self.defines_new_history(): @@ -318,7 +347,7 @@ class ModelImportStore(metaclass=abc.ABCMeta): if self.user: add_item_annotation(self.sa_session, self.user, new_history, history_properties.get("annotation")) - history = new_history + history: Optional[model.History] = new_history else: history = default_history @@ -329,7 +358,9 @@ class ModelImportStore(metaclass=abc.ABCMeta): new_history.importing = False self._flush() - def perform_import(self, history=None, new_history=False, job=None): + def perform_import( + self, history: Optional[model.History] = None, new_history: bool = False, job: Optional[model.Job] = None + ) -> "ObjectImportTracker": object_import_tracker = ObjectImportTracker() datasets_attrs = self.datasets_properties() @@ -348,7 +379,13 @@ class ModelImportStore(metaclass=abc.ABCMeta): self._flush() return object_import_tracker - def _attach_dataset_hashes(self, dataset_or_file_attrs, dataset_instance): + def _attach_dataset_hashes( + self, + dataset_or_file_attrs: Dict[str, Any], + dataset_instance: Union[ + model.LibraryDatasetDatasetAssociation, model.HistoryDatasetAssociation, model.DatasetInstance + ], + ) -> None: if "hashes" in dataset_or_file_attrs: for hash_attrs in dataset_or_file_attrs["hashes"]: hash_obj = model.DatasetHash() @@ -357,7 +394,13 @@ class ModelImportStore(metaclass=abc.ABCMeta): hash_obj.extra_files_path = hash_attrs["extra_files_path"] dataset_instance.dataset.hashes.append(hash_obj) - def _attach_dataset_sources(self, dataset_or_file_attrs, dataset_instance): + def _attach_dataset_sources( + self, + dataset_or_file_attrs: Dict[str, Any], + dataset_instance: Union[ + model.LibraryDatasetDatasetAssociation, model.HistoryDatasetAssociation, model.DatasetInstance + ], + ) -> None: if "sources" in dataset_or_file_attrs: for source_attrs in dataset_or_file_attrs["sources"]: source_obj = model.DatasetSource() @@ -372,7 +415,14 @@ class ModelImportStore(metaclass=abc.ABCMeta): dataset_instance.dataset.sources.append(source_obj) - def _import_datasets(self, object_import_tracker, datasets_attrs, history, new_history, job): + def _import_datasets( + self, + object_import_tracker: "ObjectImportTracker", + datasets_attrs: List[Any], + history: Optional[model.History], + new_history: bool, + job: Optional[model.Job], + ) -> None: object_key = self.object_key def handle_dataset_object_edit(dataset_instance, dataset_attrs): @@ -569,6 +619,8 @@ class ModelImportStore(metaclass=abc.ABCMeta): dataset_instance.dataset.purged = deleted else: dataset_instance.state = dataset_state + if not self.object_store: + raise Exception(f"self.object_store is missing from {self}.") self.object_store.update_from_file( dataset_instance.dataset, file_name=temp_dataset_file_name, create=True ) @@ -608,6 +660,8 @@ class ModelImportStore(metaclass=abc.ABCMeta): add_item_annotation(self.sa_session, self.user, dataset_instance, dataset_attrs["annotation"]) tag_list = dataset_attrs.get("tags") if tag_list: + if not self.tag_handler: + raise Exception(f"Missing self.tag_handler on {self}.") self.tag_handler.set_tags_from_list( user=self.user, item=dataset_instance, new_tags_list=tag_list, flush=False ) @@ -646,19 +700,29 @@ class ModelImportStore(metaclass=abc.ABCMeta): dataset_instance.dataset.state = dataset_instance.dataset.states.FAILED_METADATA if model_class == "HistoryDatasetAssociation": + if not isinstance(dataset_instance, model.HistoryDatasetAssociation): + raise Exception( + "Mismatch between model class and Python class, " + f"expected HistoryDatasetAssociation, got a {type(dataset_instance)}: {dataset_instance}" + ) if object_key in dataset_attrs: object_import_tracker.hdas_by_key[dataset_attrs[object_key]] = dataset_instance else: assert "id" in dataset_attrs object_import_tracker.hdas_by_id[dataset_attrs["id"]] = dataset_instance else: + if not isinstance(dataset_instance, model.LibraryDatasetDatasetAssociation): + raise Exception( + "Mismatch between model class and Python class, " + f"expected LibraryDatasetDatasetAssociation, got a {type(dataset_instance)}: {dataset_instance}" + ) if object_key in dataset_attrs: object_import_tracker.lddas_by_key[dataset_attrs[object_key]] = dataset_instance else: assert "id" in dataset_attrs object_import_tracker.lddas_by_key[dataset_attrs["id"]] = dataset_instance - def _import_libraries(self, object_import_tracker): + def _import_libraries(self, object_import_tracker: "ObjectImportTracker") -> None: object_key = self.object_key def import_folder(folder_attrs, root_folder=None): @@ -712,7 +776,13 @@ class ModelImportStore(metaclass=abc.ABCMeta): if "root_folder" in library_attrs: library.root_folder = import_folder(library_attrs["root_folder"]) - def _import_collection_instances(self, object_import_tracker, collections_attrs, history, new_history): + def _import_collection_instances( + self, + object_import_tracker: "ObjectImportTracker", + collections_attrs: List[Any], + history: Optional[model.History], + new_history: bool, + ) -> None: object_key = self.object_key def import_collection(collection_attrs): @@ -809,11 +879,17 @@ class ModelImportStore(metaclass=abc.ABCMeta): else: import_collection(collection_attrs) - def _attach_raw_id_if_editing(self, obj, attrs): + def _attach_raw_id_if_editing( + self, + obj: model.DatasetInstance, + attrs: Dict[str, Any], + ) -> None: if self.sessionless and "id" in attrs and self.import_options.allow_edit: obj.id = attrs["id"] - def _import_collection_implicit_input_associations(self, object_import_tracker, collections_attrs): + def _import_collection_implicit_input_associations( + self, object_import_tracker: "ObjectImportTracker", collections_attrs: List[Any] + ) -> None: object_key = self.object_key for collection_attrs in collections_attrs: @@ -831,7 +907,9 @@ class ModelImportStore(metaclass=abc.ABCMeta): input_dataset_collection = object_import_tracker.hdcas_by_key[input_collection_identifier] hdca.add_implicit_input_collection(name, input_dataset_collection) - def _import_dataset_copied_associations(self, object_import_tracker, datasets_attrs): + def _import_dataset_copied_associations( + self, object_import_tracker: "ObjectImportTracker", datasets_attrs: List[Any] + ) -> None: object_key = self.object_key # Re-establish copied_from_history_dataset_association relationships so history extraction @@ -869,7 +947,9 @@ class ModelImportStore(metaclass=abc.ABCMeta): else: hda_copied_from_sinks[copied_from_object_key] = dataset_key - def _import_collection_copied_associations(self, object_import_tracker, collections_attrs): + def _import_collection_copied_associations( + self, object_import_tracker: "ObjectImportTracker", collections_attrs: List[Any] + ) -> None: object_key = self.object_key # Re-establish copied_from_history_dataset_collection_association relationships so history extraction @@ -899,22 +979,26 @@ class ModelImportStore(metaclass=abc.ABCMeta): ] else: if copied_from_object_key in hdca_copied_from_sinks: - hdca.copied_from_history_dataset_association = object_import_tracker.hdcas_by_key[ + hdca.copied_from_history_dataset_collection_association = object_import_tracker.hdcas_by_key[ hdca_copied_from_sinks[copied_from_object_key] ] else: hdca_copied_from_sinks[copied_from_object_key] = dataset_collection_key - def _reassign_hids(self, object_import_tracker, history): + def _reassign_hids(self, object_import_tracker: "ObjectImportTracker", history: Optional[model.History]) -> None: # assign HIDs for newly created objects that didn't match original history requires_hid = object_import_tracker.requires_hid requires_hid_len = len(requires_hid) if requires_hid_len > 0 and not self.sessionless: + if not history: + raise Exception("Optional history is required here.") for obj in requires_hid: history.stage_addition(obj) history.add_pending_items() - def _import_workflow_invocations(self, object_import_tracker, history): + def _import_workflow_invocations( + self, object_import_tracker: "ObjectImportTracker", history: Optional[model.History] + ) -> None: # # Create jobs. # @@ -922,6 +1006,8 @@ class ModelImportStore(metaclass=abc.ABCMeta): for workflow_key, workflow_path in self.workflow_paths(): workflows_directory = os.path.join(self.archive_dir, "workflows") + if not self.app: + raise Exception(f"Missing require self.app in {self}.") workflow = self.app.workflow_contents_manager.read_workflow_from_path( self.app, self.user, workflow_path, allow_in_directory=workflows_directory ) @@ -1100,7 +1186,7 @@ class ModelImportStore(metaclass=abc.ABCMeta): if object_key in invocation_attrs: object_import_tracker.invocations_by_key[invocation_attrs[object_key]] = imported_invocation - def _import_jobs(self, object_import_tracker, history): + def _import_jobs(self, object_import_tracker: "ObjectImportTracker", history: Optional[model.History]) -> None: self._flush() object_key = self.object_key @@ -1119,8 +1205,8 @@ class ModelImportStore(metaclass=abc.ABCMeta): # only thing we allow editing currently is associations for incoming jobs. assert self.import_options.allow_edit job = self.sa_session.query(model.Job).get(job_attrs["id"]) - self._connect_job_io(job, job_attrs, _find_hda, _find_hdca, _find_dce) - self._set_job_attributes(job, job_attrs, force_terminal=False) + self._connect_job_io(job, job_attrs, _find_hda, _find_hdca, _find_dce) # type: ignore[attr-defined] + self._set_job_attributes(job, job_attrs, force_terminal=False) # type: ignore[attr-defined] # Don't edit job continue @@ -1132,23 +1218,23 @@ class ModelImportStore(metaclass=abc.ABCMeta): imported_job.imported = True imported_job.tool_id = job_attrs["tool_id"] imported_job.tool_version = job_attrs["tool_version"] - self._set_job_attributes(imported_job, job_attrs, force_terminal=True) + self._set_job_attributes(imported_job, job_attrs, force_terminal=True) # type: ignore[attr-defined] restore_times(imported_job, job_attrs) self._session_add(imported_job) # Connect jobs to input and output datasets. - params = self._normalize_job_parameters(imported_job, job_attrs, _find_hda, _find_hdca, _find_dce) + params = self._normalize_job_parameters(imported_job, job_attrs, _find_hda, _find_hdca, _find_dce) # type: ignore[attr-defined] for name, value in params.items(): # Transform parameter values when necessary. imported_job.add_parameter(name, dumps(value)) - self._connect_job_io(imported_job, job_attrs, _find_hda, _find_hdca, _find_dce) + self._connect_job_io(imported_job, job_attrs, _find_hda, _find_hdca, _find_dce) # type: ignore[attr-defined] if object_key in job_attrs: object_import_tracker.jobs_by_key[job_attrs[object_key]] = imported_job - def _import_implicit_collection_jobs(self, object_import_tracker): + def _import_implicit_collection_jobs(self, object_import_tracker: "ObjectImportTracker") -> None: object_key = self.object_key implicit_collection_jobs_attrs = self.implicit_collection_jobs_properties() @@ -1173,14 +1259,20 @@ class ModelImportStore(metaclass=abc.ABCMeta): self._session_add(icj) - def _session_add(self, obj): + def _session_add(self, obj: Union[model.DatasetInstance, model.RepresentById]) -> None: self.sa_session.add(obj) - def _flush(self): + def _flush(self) -> None: self.sa_session.flush() -def _copied_from_object_key(copied_from_chain, objects_by_key): +def _copied_from_object_key( + copied_from_chain: List[Union[Any, str]], + objects_by_key: Union[ + Dict[Union[int, str], model.HistoryDatasetAssociation], + Dict[Union[int, str], model.HistoryDatasetCollectionAssociation], + ], +) -> Optional[str]: if len(copied_from_chain) == 0: return None @@ -1221,7 +1313,7 @@ class ObjectImportTracker: jobs_by_key: Dict[ObjectKeyType, model.Job] requires_hid: List[Union[model.HistoryDatasetAssociation, model.HistoryDatasetCollectionAssociation]] - def __init__(self): + def __init__(self) -> None: self.libraries_by_key = {} self.hdas_by_key = {} self.hdas_by_id = {} @@ -1233,12 +1325,12 @@ class ObjectImportTracker: self.hda_copied_from_sinks = {} self.hdca_copied_from_sinks = {} self.jobs_by_key = {} - self.invocations_by_key = {} - self.implicit_collection_jobs_by_key = {} - self.workflows_by_key = {} + self.invocations_by_key: Dict[str, str] = {} + self.implicit_collection_jobs_by_key: Dict[str, "ImplicitCollectionJobs"] = {} + self.workflows_by_key: Dict[str, str] = {} self.requires_hid = [] - self.new_history = None + self.new_history: Optional[model.History] = None def find_hda( self, input_key: ObjectKeyType, hda_id: Optional[int] = None @@ -1273,11 +1365,13 @@ class ObjectImportTracker: class FileTracebackException(Exception): - def __init__(self, traceback, *args, **kwargs): + def __init__(self, traceback: str, *args, **kwargs) -> None: self.traceback = traceback -def get_import_model_store_for_directory(archive_dir, **kwd): +def get_import_model_store_for_directory( + archive_dir: str, **kwd +) -> Union["DirectoryImportModelStore1901", "DirectoryImportModelStoreLatest"]: traceback_file = os.path.join(archive_dir, TRACEBACK) if not os.path.isdir(archive_dir): raise Exception( @@ -1295,9 +1389,14 @@ def get_import_model_store_for_directory(archive_dir, **kwd): class DictImportModelStore(ModelImportStore): object_key = "encoded_id" - def __init__(self, store_as_dict, **kwd): + def __init__( + self, + store_as_dict: Dict[str, Any], + **kwd, + ) -> None: self._store_as_dict = store_as_dict super().__init__(**kwd) + self.archive_dir = "" def defines_new_history(self) -> bool: return DICT_STORE_ATTRS_KEY_HISTORY in self._store_as_dict @@ -1305,49 +1404,77 @@ class DictImportModelStore(ModelImportStore): def new_history_properties(self): return self._store_as_dict.get(DICT_STORE_ATTRS_KEY_HISTORY) or {} - def datasets_properties(self): + def datasets_properties( + self, + ) -> List[Dict[str, Any]]: return self._store_as_dict.get(DICT_STORE_ATTRS_KEY_DATASETS) or [] - def collections_properties(self): + def collections_properties(self) -> List[Any]: return self._store_as_dict.get(DICT_STORE_ATTRS_KEY_COLLECTIONS) or [] - def library_properties(self): + def library_properties( + self, + ) -> List[Dict[str, Any]]: return self._store_as_dict.get(DICT_STORE_ATTRS_KEY_LIBRARIES) or [] - def jobs_properties(self): + def jobs_properties(self) -> List[Any]: return self._store_as_dict.get(DICT_STORE_ATTRS_KEY_JOBS) or [] - def implicit_collection_jobs_properties(self): + def implicit_collection_jobs_properties(self) -> List[Any]: return self._store_as_dict.get(DICT_STORE_ATTRS_KEY_IMPLICIT_COLLECTION_JOBS) or [] - def invocations_properties(self): + def invocations_properties(self) -> List[Any]: return self._store_as_dict.get(DICT_STORE_ATTRS_KEY_INVOCATIONS) or [] - def workflow_paths(self): - return [] + def workflow_paths(self) -> Iterator[Tuple[str, str]]: + return + yield -def get_import_model_store_for_dict(as_dict, **kwd): +def get_import_model_store_for_dict( + as_dict: Dict[str, Any], + **kwd, +) -> DictImportModelStore: return DictImportModelStore(as_dict, **kwd) class BaseDirectoryImportModelStore(ModelImportStore): - archive_dir: str + @abc.abstractmethod + def _normalize_job_parameters( + self, + imported_job: model.Job, + job_attrs: Dict[str, Any], + _find_hda: Callable, + _find_hdca: Callable, + _find_dce: Callable, + ) -> Dict[str, Any]: + pass + + @abc.abstractmethod + def _connect_job_io( + self, + imported_job: model.Job, + job_attrs: Dict[str, Any], + _find_hda: Callable, + _find_hdca: Callable, + _find_dce: Callable, + ) -> None: + pass @property - def file_source_root(self): + def file_source_root(self) -> str: return self.archive_dir def defines_new_history(self) -> bool: new_history_attributes = os.path.join(self.archive_dir, ATTRS_FILENAME_HISTORY) return os.path.exists(new_history_attributes) - def new_history_properties(self): + def new_history_properties(self) -> Dict[str, Optional[Union[str, int, bool]]]: new_history_attributes = os.path.join(self.archive_dir, ATTRS_FILENAME_HISTORY) history_properties = load(open(new_history_attributes)) return history_properties - def datasets_properties(self): + def datasets_properties(self) -> List[Any]: datasets_attrs_file_name = os.path.join(self.archive_dir, ATTRS_FILENAME_DATASETS) datasets_attrs = load(open(datasets_attrs_file_name)) provenance_file_name = f"{datasets_attrs_file_name}.provenance" @@ -1358,18 +1485,22 @@ class BaseDirectoryImportModelStore(ModelImportStore): return datasets_attrs - def collections_properties(self): + def collections_properties(self) -> List[Any]: return self._read_list_if_exists(ATTRS_FILENAME_COLLECTIONS) - def library_properties(self): + def library_properties( + self, + ) -> List[Dict[str, Any]]: libraries_attrs = self._read_list_if_exists(ATTRS_FILENAME_LIBRARIES) libraries_attrs.extend(self._read_list_if_exists(ATTRS_FILENAME_LIBRARY_FOLDERS)) return libraries_attrs - def jobs_properties(self): + def jobs_properties( + self, + ) -> List[Dict[str, Any]]: return self._read_list_if_exists(ATTRS_FILENAME_JOBS) - def implicit_collection_jobs_properties(self): + def implicit_collection_jobs_properties(self) -> List[Union[Any, Dict[str, Union[str, List[str]]]]]: implicit_collection_jobs_attrs_file_name = os.path.join( self.archive_dir, ATTRS_FILENAME_IMPLICIT_COLLECTION_JOBS ) @@ -1378,10 +1509,12 @@ class BaseDirectoryImportModelStore(ModelImportStore): except FileNotFoundError: return [] - def invocations_properties(self): + def invocations_properties( + self, + ) -> List[Dict[str, Any]]: return self._read_list_if_exists(ATTRS_FILENAME_INVOCATIONS) - def workflow_paths(self): + def workflow_paths(self) -> Iterator[Tuple[str, str]]: workflows_directory = os.path.join(self.archive_dir, "workflows") if not os.path.exists(workflows_directory): return [] @@ -1393,7 +1526,9 @@ class BaseDirectoryImportModelStore(ModelImportStore): workflow_key = name[0 : -len(".gxwf.yml")] yield workflow_key, os.path.join(workflows_directory, name) - def _set_job_attributes(self, imported_job, job_attrs, force_terminal=False): + def _set_job_attributes( + self, imported_job: model.Job, job_attrs: Dict[str, Any], force_terminal: bool = False + ) -> None: ATTRIBUTES = ( "info", "exit_code", @@ -1417,7 +1552,7 @@ class BaseDirectoryImportModelStore(ModelImportStore): if raw_state: imported_job.set_state(raw_state) - def _read_list_if_exists(self, file_name, required=False): + def _read_list_if_exists(self, file_name: str, required: bool = False) -> List[Any]: file_name = os.path.join(self.archive_dir, file_name) if os.path.exists(file_name): attrs = load(open(file_name)) @@ -1428,7 +1563,9 @@ class BaseDirectoryImportModelStore(ModelImportStore): return attrs -def restore_times(model_object, attrs): +def restore_times( + model_object: Union[model.Job, model.WorkflowInvocation, model.WorkflowInvocationStep], attrs: Dict[str, Any] +) -> None: try: model_object.create_time = datetime.datetime.strptime(attrs["create_time"], "%Y-%m-%dT%H:%M:%S.%f") except Exception: @@ -1442,7 +1579,7 @@ def restore_times(model_object, attrs): class DirectoryImportModelStore1901(BaseDirectoryImportModelStore): object_key = "hid" - def __init__(self, archive_dir, **kwd): + def __init__(self, archive_dir: str, **kwd) -> None: archive_dir = os.path.realpath(archive_dir) # BioBlend previous to 17.01 exported histories with an extra subdir. if not os.path.exists(os.path.join(archive_dir, ATTRS_FILENAME_HISTORY)): @@ -1453,7 +1590,14 @@ class DirectoryImportModelStore1901(BaseDirectoryImportModelStore): self.archive_dir = archive_dir super().__init__(**kwd) - def _connect_job_io(self, imported_job, job_attrs, _find_hda, _find_hdca, _find_dce): + def _connect_job_io( + self, + imported_job: model.Job, + job_attrs: Dict[str, Any], + _find_hda: Callable, + _find_hdca: Callable, + _find_dce: Callable, + ) -> None: for output_key in job_attrs["output_datasets"]: output_hda = _find_hda(output_key) if output_hda: @@ -1468,7 +1612,14 @@ class DirectoryImportModelStore1901(BaseDirectoryImportModelStore): if input_hda: imported_job.add_input_dataset(input_name, input_hda) - def _normalize_job_parameters(self, imported_job, job_attrs, _find_hda, _find_hdca, _find_dce): + def _normalize_job_parameters( + self, + imported_job: model.Job, + job_attrs: Dict[str, Any], + _find_hda: Callable, + _find_hdca: Callable, + _find_dce: Callable, + ) -> Dict[str, Any]: def remap_objects(p, k, obj): if isinstance(obj, dict) and obj.get("__HistoryDatasetAssociation__", False): imported_hda = _find_hda(obj[self.object_key]) @@ -1480,7 +1631,7 @@ class DirectoryImportModelStore1901(BaseDirectoryImportModelStore): params = remap(params, remap_objects) return params - def trust_hid(self, obj_attrs): + def trust_hid(self, obj_attrs: Dict[str, Optional[Union[str, bool, int, Dict[str, Union[str, int]]]]]) -> bool: # We didn't do object tracking so we pretty much have to trust the HID and accept # that it will be wrong a lot. return True @@ -1489,12 +1640,19 @@ class DirectoryImportModelStore1901(BaseDirectoryImportModelStore): class DirectoryImportModelStoreLatest(BaseDirectoryImportModelStore): object_key = "encoded_id" - def __init__(self, archive_dir, **kwd): + def __init__(self, archive_dir: str, **kwd) -> None: archive_dir = os.path.realpath(archive_dir) self.archive_dir = archive_dir super().__init__(**kwd) - def _connect_job_io(self, imported_job, job_attrs, _find_hda, _find_hdca, _find_dce): + def _connect_job_io( + self, + imported_job: model.Job, + job_attrs: Dict[str, Any], + _find_hda: Callable, + _find_hdca: Callable, + _find_dce: Callable, + ) -> None: if imported_job.command_line is None: imported_job.command_line = job_attrs.get("command_line") @@ -1541,7 +1699,14 @@ class DirectoryImportModelStoreLatest(BaseDirectoryImportModelStore): if output_hdca: imported_job.add_output_dataset_collection(output_name, output_hdca) - def _normalize_job_parameters(self, imported_job, job_attrs, _find_hda, _find_hdca, _find_dce): + def _normalize_job_parameters( + self, + imported_job: model.Job, + job_attrs: Dict[str, Optional[Union[str, Dict[str, List[str]], int, Dict[str, List[int]]]]], + _find_hda: Callable, + _find_hdca: Callable, + _find_dce: Callable, + ) -> Dict[str, Any]: def remap_objects(p, k, obj): if isinstance(obj, dict) and "src" in obj and obj["src"] in ["hda", "hdca", "dce"]: if obj["src"] == "hda": @@ -1576,11 +1741,11 @@ class DirectoryImportModelStoreLatest(BaseDirectoryImportModelStore): params = job_attrs["params"] params = remap(params, remap_objects) - return params + return cast(Dict[str, Any], params) class BagArchiveImportModelStore(DirectoryImportModelStoreLatest): - def __init__(self, bag_archive, **kwd): + def __init__(self, bag_archive: str, **kwd) -> None: archive_dir = tempfile.mkdtemp() bdb.extract_bag(bag_archive, output_path=archive_dir) # Why this line though...? @@ -1641,15 +1806,15 @@ class DirectoryModelExportStore(ModelExportStore): def __init__( self, - export_directory: str, + export_directory: Union[PosixPath, str], app: Optional[StoreAppProtocol] = None, file_sources: Optional[ConfiguredFileSources] = None, for_edit: bool = False, - serialize_dataset_objects=None, + serialize_dataset_objects: Optional[bool] = None, export_files: Optional[str] = None, strip_metadata_files: bool = True, serialize_jobs: bool = True, - ): + ) -> None: """ :param export_directory: path to export directory. Will be created if it does not exist. :param app: Galaxy App or app-like object. Must be provided if `for_edit` and/or `serialize_dataset_objects` are True @@ -1698,7 +1863,7 @@ class DirectoryModelExportStore(ModelExportStore): self.job_output_dataset_associations: Dict[int, Dict[str, model.DatasetInstance]] = {} @property - def workflows_directory(self): + def workflows_directory(self) -> str: return os.path.join(self.export_directory, "workflows") def serialize_files(self, dataset: model.DatasetInstance, as_dict: JsonDictT) -> None: @@ -1778,10 +1943,21 @@ class DirectoryModelExportStore(ModelExportStore): self.dataset_id_to_path[dataset.dataset.id] = (as_dict.get("file_name"), as_dict.get("extra_files_path")) - def exported_key(self, obj): + def exported_key( + self, + obj: Union[model.HistoryDatasetCollectionAssociation, model.HistoryDatasetAssociation, model.DatasetInstance], + ) -> Union[str, int]: return self.serialization_options.get_identifier(self.security, obj) - def __enter__(self): + def __enter__( + self, + ) -> Union[ + "BagArchiveModelExportStore", + "DirectoryModelExportStore", + "TarModelExportStore", + "ROCrateArchiveModelExportStore", + "ROCrateModelExportStore", + ]: return self def push_metadata_files(self): @@ -1797,7 +1973,12 @@ class DirectoryModelExportStore(ModelExportStore): with open(os.path.join(self.export_directory, "tool.xml"), "w") as out: out.write(tool_source.to_string()) - def export_jobs(self, jobs: List[model.Job], jobs_attrs=None, include_job_data=True): + def export_jobs( + self, + jobs: Iterable[model.Job], + jobs_attrs: Optional[List[Dict[str, Any]]] = None, + include_job_data: bool = True, + ) -> List[Dict[str, Any]]: """ Export jobs. @@ -1812,12 +1993,12 @@ class DirectoryModelExportStore(ModelExportStore): if include_job_data: # -- Get input, output datasets. -- - input_dataset_mapping: Dict[str, List[model.DatasetInstance]] = {} - output_dataset_mapping: Dict[str, List[model.DatasetInstance]] = {} - input_dataset_collection_mapping: Dict[str, List[model.DatasetCollectionInstance]] = {} - input_dataset_collection_element_mapping: Dict[str, List[model.DatasetCollectionElement]] = {} - output_dataset_collection_mapping: Dict[str, List[model.DatasetCollectionInstance]] = {} - implicit_output_dataset_collection_mapping: Dict[str, List[model.DatasetCollection]] = {} + input_dataset_mapping: Dict[str, List[Union[str, int]]] = {} + output_dataset_mapping: Dict[str, List[Union[str, int]]] = {} + input_dataset_collection_mapping: Dict[str, List[Union[str, int]]] = {} + input_dataset_collection_element_mapping: Dict[str, List[Union[str, int]]] = {} + output_dataset_collection_mapping: Dict[str, List[Union[str, int]]] = {} + implicit_output_dataset_collection_mapping: Dict[str, List[Union[str, int]]] = {} for assoc in job.input_datasets: # Optional data inputs will not have a dataset. @@ -1957,7 +2138,9 @@ class DirectoryModelExportStore(ModelExportStore): if dataset not in self.included_datasets: self.add_dataset(dataset, include_files=add_dataset) - def export_library(self, library: model.Library, include_hidden=False, include_deleted=False): + def export_library( + self, library: model.Library, include_hidden: bool = False, include_deleted: bool = False + ) -> None: self.included_libraries.append(library) root_folder = library.root_folder self.export_library_folder_contents(root_folder, include_hidden=include_hidden, include_deleted=include_deleted) @@ -1969,8 +2152,8 @@ class DirectoryModelExportStore(ModelExportStore): ) def export_library_folder_contents( - self, library_folder: model.LibraryFolder, include_hidden=False, include_deleted=False - ): + self, library_folder: model.LibraryFolder, include_hidden: bool = False, include_deleted: bool = False + ) -> None: for library_dataset in library_folder.datasets: ldda = library_dataset.library_dataset_dataset_association add_dataset = (not ldda.visible or not include_hidden) and (not ldda.deleted or include_deleted) @@ -1979,8 +2162,8 @@ class DirectoryModelExportStore(ModelExportStore): self.export_library_folder_contents(folder, include_hidden=include_hidden, include_deleted=include_deleted) def export_workflow_invocation( - self, workflow_invocation: model.WorkflowInvocation, include_hidden=False, include_deleted=False - ): + self, workflow_invocation: model.WorkflowInvocation, include_hidden: bool = False, include_deleted: bool = False + ) -> None: self.included_invocations.append(workflow_invocation) for input_dataset in workflow_invocation.input_datasets: self.add_dataset(input_dataset.dataset) @@ -2035,7 +2218,7 @@ class DirectoryModelExportStore(ModelExportStore): def add_dataset(self, dataset: model.DatasetInstance, include_files: bool = True) -> None: self.included_datasets[dataset] = (dataset, include_files) - def _finalize(self): + def _finalize(self) -> None: export_directory = self.export_directory datasets_attrs = [] @@ -2070,7 +2253,7 @@ class DirectoryModelExportStore(ModelExportStore): jobs_attrs = [] for job_id, job_output_dataset_associations in self.job_output_dataset_associations.items(): - output_dataset_mapping = {} + output_dataset_mapping: Dict[str, List[Union[str, int]]] = {} for name, dataset in job_output_dataset_associations.items(): if name not in output_dataset_mapping: output_dataset_mapping[name] = [] @@ -2084,7 +2267,7 @@ class DirectoryModelExportStore(ModelExportStore): # # Get all jobs associated with included HDAs. - jobs_dict = {} + jobs_dict: Dict[str, model.Job] = {} implicit_collection_jobs_dict = {} def record_job(job): @@ -2110,9 +2293,15 @@ class DirectoryModelExportStore(ModelExportStore): for hda, _include_files in self.included_datasets.values(): # Get the associated job, if any. If this hda was copied from another, # we need to find the job that created the original hda + if not isinstance(hda, (model.HistoryDatasetAssociation, model.LibraryDatasetDatasetAssociation)): + raise Exception( + f"Expected a HistoryDatasetAssociation or LibraryDatasetDatasetAssociation, but got a {type(hda)}: {hda}" + ) job_hda = hda - while job_hda.copied_from_history_dataset_association: # should this check library datasets as well? - job_hda = job_hda.copied_from_history_dataset_association + while getattr( + job_hda, "copied_from_history_dataset_association", None + ): # should this check library datasets as well? + job_hda = job_hda.copied_from_history_dataset_association # type: ignore[union-attr] if not job_hda.creating_job_associations: # No viable HDA found. continue @@ -2157,6 +2346,8 @@ class DirectoryModelExportStore(ModelExportStore): assert invocation_attrs invocation_attrs["workflow"] = workflow_key + if not self.app: + raise Exception(f"Missing self.app in {self}.") self.app.workflow_contents_manager.store_workflow_artifacts( workflows_directory, workflow_key, workflow, user=history.user, history=history ) @@ -2170,7 +2361,7 @@ class DirectoryModelExportStore(ModelExportStore): with open(export_attrs_filename, "w") as export_attrs_out: dump({"galaxy_export_version": GALAXY_EXPORT_VERSION}, export_attrs_out) - def __exit__(self, exc_type, exc_val, exc_tb): + def __exit__(self, exc_type: None, exc_val: None, exc_tb: None) -> bool: if exc_type is None: self._finalize() # http://effbot.org/zone/python-with-statement.htm @@ -2179,7 +2370,17 @@ class DirectoryModelExportStore(ModelExportStore): class WriteCrates: - def _generate_markdown_readme(self): + included_invocations: List[model.WorkflowInvocation] + export_directory: Union[PosixPath, str] + included_datasets: Dict[model.DatasetInstance, Tuple[model.DatasetInstance, bool]] + dataset_id_to_path: Dict[int, Tuple[Optional[str], Optional[str]]] + + @property + @abc.abstractmethod + def workflows_directory(self) -> str: + pass + + def _generate_markdown_readme(self) -> str: markdown_parts: List[str] = [] if self._is_single_invocation_export(): invocation = self.included_invocations[0] @@ -2193,10 +2394,10 @@ class WriteCrates: return "\n".join(markdown_parts) - def _is_single_invocation_export(self): + def _is_single_invocation_export(self) -> bool: return len(self.included_invocations) == 1 - def _init_crate(self): + def _init_crate(self) -> ROCrate: ro_crate = ROCrate() markdown_path = os.path.join(self.export_directory, "README.md") @@ -2487,18 +2688,21 @@ class BcoModelExportStore(WorkflowInvocationOnlyExportStore): class ROCrateModelExportStore(DirectoryModelExportStore, WriteCrates): - def __init__(self, crate_directory, **kwds): + def __init__(self, crate_directory: Union[PosixPath, str], **kwds) -> None: self.crate_directory = crate_directory super().__init__(crate_directory, export_files="symlink", **kwds) - def _finalize(self): + def _finalize(self) -> None: super()._finalize() ro_crate = self._init_crate() ro_crate.write(self.crate_directory) class ROCrateArchiveModelExportStore(DirectoryModelExportStore, WriteCrates): - def __init__(self, uri, **kwds): + out_file: Union[PosixPath, str] + file_source_uri: Optional[Union[PosixPath, str]] + + def __init__(self, uri: Union[PosixPath, str], **kwds) -> None: temp_output_dir = tempfile.mkdtemp() self.temp_output_dir = temp_output_dir if "://" in str(uri): @@ -2511,7 +2715,7 @@ class ROCrateArchiveModelExportStore(DirectoryModelExportStore, WriteCrates): export_directory = temp_output_dir super().__init__(export_directory, **kwds) - def _finalize(self): + def _finalize(self) -> None: super()._finalize() ro_crate = self._init_crate() ro_crate.write(self.export_directory) @@ -2524,6 +2728,8 @@ class ROCrateArchiveModelExportStore(DirectoryModelExportStore, WriteCrates): if not self.file_source_uri: shutil.move(rval, self.out_file) else: + if not self.file_sources: + raise Exception(f"Need self.file_sources but {type(self)} is missing it: {self.file_sources}.") file_source_path = self.file_sources.get_file_source_path(self.file_source_uri) file_source = file_source_path.file_source assert os.path.exists(rval), rval @@ -2532,7 +2738,10 @@ class ROCrateArchiveModelExportStore(DirectoryModelExportStore, WriteCrates): class TarModelExportStore(DirectoryModelExportStore): - def __init__(self, uri, gzip=True, **kwds): + file_source_uri: Optional[Union[PosixPath, str]] + out_file: Union[PosixPath, str] + + def __init__(self, uri: Union[PosixPath, str], gzip: bool = True, **kwds) -> None: self.gzip = gzip temp_output_dir = tempfile.mkdtemp() self.temp_output_dir = temp_output_dir @@ -2546,10 +2755,12 @@ class TarModelExportStore(DirectoryModelExportStore): export_directory = temp_output_dir super().__init__(export_directory, **kwds) - def _finalize(self): + def _finalize(self) -> None: super()._finalize() tar_export_directory(self.export_directory, self.out_file, self.gzip) if self.file_source_uri: + if not self.file_sources: + raise Exception(f"Need self.file_sources but {type(self)} is missing it: {self.file_sources}.") file_source_path = self.file_sources.get_file_source_path(self.file_source_uri) file_source = file_source_path.file_source assert os.path.exists(self.out_file) @@ -2558,17 +2769,19 @@ class TarModelExportStore(DirectoryModelExportStore): class BagDirectoryModelExportStore(DirectoryModelExportStore): - def __init__(self, out_directory, **kwds): + def __init__(self, out_directory: str, **kwds) -> None: self.out_directory = out_directory super().__init__(out_directory, **kwds) - def _finalize(self): + def _finalize(self) -> None: super()._finalize() bdb.make_bag(self.out_directory) class BagArchiveModelExportStore(BagDirectoryModelExportStore): - def __init__(self, uri, bag_archiver="tgz", **kwds): + file_source_uri: Optional[Union[PosixPath, str]] + + def __init__(self, uri: Union[PosixPath, str], bag_archiver: str = "tgz", **kwds) -> None: # bag_archiver in tgz, zip, tar self.bag_archiver = bag_archiver temp_output_dir = tempfile.mkdtemp() @@ -2583,12 +2796,14 @@ class BagArchiveModelExportStore(BagDirectoryModelExportStore): export_directory = temp_output_dir super().__init__(export_directory, **kwds) - def _finalize(self): + def _finalize(self) -> None: super()._finalize() rval = bdb.archive_bag(self.export_directory, self.bag_archiver) if not self.file_source_uri: shutil.move(rval, self.out_file) else: + if not self.file_sources: + raise Exception(f"Need self.file_sources but {type(self)} is missing it: {self.file_sources}.") file_source_path = self.file_sources.get_file_source_path(self.file_source_uri) file_source = file_source_path.file_source assert os.path.exists(rval) @@ -2598,7 +2813,7 @@ class BagArchiveModelExportStore(BagDirectoryModelExportStore): def get_export_store_factory( app, download_format: str, export_files=None, bco_export_options: Optional[BcoExportOptions] = None -) -> Callable[[str], ModelExportStore]: +) -> Callable[[Union[PosixPath, str]], ModelExportStore]: export_store_class: Union[ Type[TarModelExportStore], Type[BagArchiveModelExportStore], @@ -2632,7 +2847,7 @@ def get_export_store_factory( return lambda path: export_store_class(path, **export_store_class_kwds) -def tar_export_directory(export_directory: str, out_file: str, gzip: bool) -> None: +def tar_export_directory(export_directory: Union[PosixPath, str], out_file: Union[PosixPath, str], gzip: bool) -> None: tarfile_mode = "w" if gzip: tarfile_mode += ":gz" @@ -2642,7 +2857,7 @@ def tar_export_directory(export_directory: str, out_file: str, gzip: bool) -> No store_archive.add(os.path.join(export_directory, export_path), arcname=export_path) -def get_export_dataset_filename(name, ext, hid): +def get_export_dataset_filename(name: str, ext: str, hid: int) -> str: """ Builds a filename for a dataset using its name an extension. """ @@ -2650,7 +2865,9 @@ def get_export_dataset_filename(name, ext, hid): return f"{base}_{hid}.{ext}" -def imported_store_for_metadata(directory, object_store=None): +def imported_store_for_metadata( + directory: str, object_store: Optional[ObjectStore] = None +) -> BaseDirectoryImportModelStore: import_options = ImportOptions(allow_dataset_object_edit=True, allow_edit=True) import_model_store = get_import_model_store_for_directory( directory, import_options=import_options, object_store=object_store @@ -2671,7 +2888,7 @@ def source_to_import_store( raise Exception( "Can only specify a model_store_format as an argument to source_to_import_store in conjuction with URIs" ) - model_import_store = get_import_model_store_for_dict( + model_import_store: ModelImportStore = get_import_model_store_for_dict( source, import_options=import_options, app=app, diff --git a/lib/galaxy/model/tool_shed_install/__init__.py b/lib/galaxy/model/tool_shed_install/__init__.py index 33081223b47..8bd5f916906 100644 --- a/lib/galaxy/model/tool_shed_install/__init__.py +++ b/lib/galaxy/model/tool_shed_install/__init__.py @@ -1,7 +1,13 @@ import logging import os from enum import Enum -from typing import TYPE_CHECKING +from typing import ( + Any, + Callable, + Dict, + Optional, + TYPE_CHECKING, +) from sqlalchemy import ( Boolean, @@ -170,7 +176,7 @@ class ToolShedRepository(Base): self.status = status self.error_message = error_message - def as_dict(self, value_mapper=None): + def as_dict(self, value_mapper: Optional[Dict[str, Callable]] = None) -> Dict[str, Any]: return self.to_dict(view="element", value_mapper=value_mapper) @property @@ -515,7 +521,7 @@ class ToolShedRepository(Base): return asbool(self.tool_shed_status.get("revision_update", False)) return False - def to_dict(self, view="collection", value_mapper=None): + def to_dict(self, view="collection", value_mapper: Optional[Dict[str, Callable]] = None) -> Dict[str, Any]: if value_mapper is None: value_mapper = {} rval = {} @@ -527,7 +533,7 @@ class ToolShedRepository(Base): try: rval[key] = self.__getattribute__(key) if key in value_mapper: - rval[key] = value_mapper.get(key, rval[key]) + rval[key] = value_mapper[key](rval[key]) except AttributeError: rval[key] = None return rval diff --git a/lib/galaxy/structured_app.py b/lib/galaxy/structured_app.py index 78a0bbe4a63..fc55a642de0 100644 --- a/lib/galaxy/structured_app.py +++ b/lib/galaxy/structured_app.py @@ -47,6 +47,7 @@ if TYPE_CHECKING: from galaxy.managers.hdas import HDAManager from galaxy.managers.histories import HistoryManager from galaxy.managers.workflows import WorkflowsManager + from galaxy.tool_shed.galaxy_install.installed_repository_manager import InstalledRepositoryManager from galaxy.tools import ToolBox from galaxy.tools.cache import ToolCache from galaxy.tools.data import ToolDataTableManager @@ -104,7 +105,7 @@ class MinimalManagerApp(MinimalApp): library_folder_manager: Any # 'galaxy.managers.folders.FolderManager' library_manager: Any # 'galaxy.managers.libraries.LibraryManager' role_manager: Any # 'galaxy.managers.roles.RoleManager' - installed_repository_manager: Any # 'galaxy.tool_shed.galaxy_install.installed_repository_manager.InstalledRepositoryManager' + installed_repository_manager: "InstalledRepositoryManager" user_manager: Any job_config: "JobConfiguration" job_manager: Any # galaxy.jobs.manager.JobManager @@ -113,6 +114,7 @@ class MinimalManagerApp(MinimalApp): genomes: "Genomes" error_reports: "ErrorReports" object_store: BaseObjectStore + tool_shed_registry: ToolShedRegistry @property @abc.abstractmethod @@ -146,7 +148,6 @@ class StructuredApp(MinimalManagerApp): data_provider_registry: Any # 'galaxy.visualization.data_providers.registry.DataProviderRegistry' tool_data_tables: "ToolDataTableManager" tool_cache: "ToolCache" - tool_shed_registry: ToolShedRegistry tool_shed_repository_cache: Optional[ToolShedRepositoryCache] watchers: "ConfigWatchers" workflow_scheduling_manager: Any # 'galaxy.workflow.scheduling_manager.WorkflowSchedulingManager' diff --git a/lib/galaxy/tool_shed/galaxy_install/install_manager.py b/lib/galaxy/tool_shed/galaxy_install/install_manager.py index 95cd8be7eff..d34e8c4ba5c 100644 --- a/lib/galaxy/tool_shed/galaxy_install/install_manager.py +++ b/lib/galaxy/tool_shed/galaxy_install/install_manager.py @@ -126,7 +126,7 @@ class InstallRepositoryManager: shed_tool_conf=None, reinstalling=False, tool_panel_section_mapping=None, - ): + ) -> None: """ Generate the metadata for the installed tool shed repository, among other things. This method is called when an administrator is installing a new repository or @@ -206,6 +206,7 @@ class InstallRepositoryManager: ) if "data_manager" in irmm_metadata_dict: dmh = data_manager.DataManagerHandler(self.app) + assert shed_config_dict dmh.install_data_managers( self.app.config.shed_data_manager_config_file, irmm_metadata_dict, diff --git a/lib/galaxy/tool_shed/galaxy_install/installed_repository_manager.py b/lib/galaxy/tool_shed/galaxy_install/installed_repository_manager.py index 0ee698e62d6..11a08b6d65f 100644 --- a/lib/galaxy/tool_shed/galaxy_install/installed_repository_manager.py +++ b/lib/galaxy/tool_shed/galaxy_install/installed_repository_manager.py @@ -82,7 +82,7 @@ class InstalledRepositoryManager: self.installed_dependent_repositories_of_installed_repositories = {} @property - def tool_paths(self): + def tool_paths(self) -> List[str]: """Return all possible tool_path attributes of all tool config files.""" if len(self._tool_paths) != len(self.tool_configs): # This could be happen at startup or after the creation of a new shed_tool_conf.xml file @@ -93,6 +93,7 @@ class InstalledRepositoryManager: if error_message: log.error(error_message) else: + assert tree tool_path = tree.getroot().get("tool_path") if tool_path: tool_paths.append(tool_path) @@ -402,7 +403,9 @@ class InstalledRepositoryManager: missing_repository_dependencies["description"] = description return installed_repository_dependencies, missing_repository_dependencies - def get_installed_and_missing_repository_dependencies_for_new_or_updated_install(self, repo_info_tuple): + def get_installed_and_missing_repository_dependencies_for_new_or_updated_install( + self, repo_info_tuple + ) -> Tuple[Dict[str, Any], Dict[str, Any]]: """ Parse the received repository_dependencies dictionary that is associated with a repository being installed into Galaxy for the first time and attempt to determine repository dependencies that are @@ -572,7 +575,9 @@ class InstalledRepositoryManager: missing_tool_dependencies[td_key] = val return installed_tool_dependencies, missing_tool_dependencies - def get_repository_dependency_tups_for_installed_repository(self, repository, dependency_tups=None, status=None): + def get_repository_dependency_tups_for_installed_repository( + self, repository, dependency_tups=None, status=None + ) -> List[RepositoryTupleT]: """ Return a list of of tuples defining tool_shed_repository objects (whose status can be anything) required by the received repository. The returned list defines the entire repository dependency tree. This method is called @@ -646,7 +651,7 @@ class InstalledRepositoryManager: deleted_tool_dependency_names.append(original_dependency_val_dict["name"]) return updated_tool_dependency_names, deleted_tool_dependency_names - def uninstall_repository(self, repository: ToolShedRepository, remove_from_disk=True): + def uninstall_repository(self, repository: ToolShedRepository, remove_from_disk=True) -> str: errors = "" shed_tool_conf, tool_path, relative_install_dir = suc.get_tool_panel_config_tool_path_install_dir( app=self.app, repository=repository @@ -699,7 +704,7 @@ class InstalledRepositoryManager: def remove_entry_from_installed_repository_dependencies_of_installed_repositories( self, repository: ToolShedRepository - ): + ) -> None: """ Remove an entry from self.installed_repository_dependencies_of_installed_repositories. A side-effect of this method is removal of appropriate value items from self.installed_dependent_repositories_of_installed_repositories. @@ -795,7 +800,7 @@ class InstalledRepositoryManager: return "True" return "False" - def set_prior_installation_required(self, repository, required_repository): + def set_prior_installation_required(self, repository, required_repository) -> str: """ Return True if the received required_repository must be installed before the received repository. diff --git a/lib/galaxy/tool_shed/galaxy_install/tools/data_manager.py b/lib/galaxy/tool_shed/galaxy_install/tools/data_manager.py index 53f1c3b78f0..991fdd70a2b 100644 --- a/lib/galaxy/tool_shed/galaxy_install/tools/data_manager.py +++ b/lib/galaxy/tool_shed/galaxy_install/tools/data_manager.py @@ -73,7 +73,7 @@ class DataManagerHandler: relative_install_dir: StrPath, repository, repository_tools_tups, - ): + ) -> List["DataManager"]: rval: List["DataManager"] = [] if "data_manager" in metadata_dict: tpm = tool_panel_manager.ToolPanelManager(self.app) diff --git a/lib/galaxy/tool_shed/util/container_util.py b/lib/galaxy/tool_shed/util/container_util.py index a30f9fb504b..efb5988ced7 100644 --- a/lib/galaxy/tool_shed/util/container_util.py +++ b/lib/galaxy/tool_shed/util/container_util.py @@ -1,4 +1,5 @@ import logging +from typing import Union from galaxy.util.tool_shed.common_util import remove_protocol_from_tool_shed_url @@ -13,8 +14,8 @@ def generate_repository_dependencies_key_for_repository( repository_name: str, repository_owner: str, changeset_revision: str, - prior_installation_required: bool, - only_if_compiling_contained_td: bool, + prior_installation_required: Union[bool, str], + only_if_compiling_contained_td: Union[bool, str], ) -> str: """ Assumes tool shed is current tool shed since repository dependencies across tool sheds @@ -27,11 +28,11 @@ def generate_repository_dependencies_key_for_repository( return "{}{}{}{}{}{}{}{}{}{}{}".format( tool_shed, STRSEP, - str(repository_name), + repository_name, STRSEP, - str(repository_owner), + repository_owner, STRSEP, - str(changeset_revision), + changeset_revision, STRSEP, str(prior_installation_required), STRSEP, diff --git a/lib/galaxy/tool_util/data/__init__.py b/lib/galaxy/tool_util/data/__init__.py index 983ebeb0c1e..425f6e46253 100644 --- a/lib/galaxy/tool_util/data/__init__.py +++ b/lib/galaxy/tool_util/data/__init__.py @@ -18,12 +18,16 @@ import time from glob import glob from tempfile import NamedTemporaryFile from typing import ( + Any, BinaryIO, + Callable, Dict, List, Optional, Set, + Tuple, Type, + TYPE_CHECKING, Union, ) @@ -31,15 +35,23 @@ import requests from galaxy import util from galaxy.exceptions import MessageException -from galaxy.util import RW_R__R__ +from galaxy.util import ( + Element, + RW_R__R__, +) from galaxy.util.dictifiable import Dictifiable from galaxy.util.filelock import FileLock +from galaxy.util.path import StrPath from galaxy.util.renamed_temporary_file import RenamedTemporaryFile from ._schema import ( ToolDataEntry, ToolDataEntryList, ) +if TYPE_CHECKING: + from galaxy.config import GalaxyAppConfiguration + from galaxy.tools.data_manager.manager import DataManager + log = logging.getLogger(__name__) DEFAULT_TABLE_TYPE = "tabular" @@ -90,12 +102,9 @@ class ToolDataPathFiles: return os.path.exists(path) -ConfigFilesT = Union[str, os.PathLike, List[Union[str, os.PathLike]]] - - class ToolDataTable(Dictifiable): type_key: str - data: List + data: List[List[str]] @classmethod def from_dict(cls, d): @@ -109,20 +118,20 @@ class ToolDataTable(Dictifiable): def __init__( self, - config_element, - tool_data_path, - from_shed_config=False, - filename=None, - tool_data_path_files=None, - other_config_dict=None, - ): + config_element: Element, + tool_data_path: Optional[StrPath], + tool_data_path_files: ToolDataPathFiles, + from_shed_config: bool = False, + filename: Optional[StrPath] = None, + other_config_dict: Optional[Union["GalaxyAppConfiguration", Dict[str, Any]]] = None, + ) -> None: self.name = config_element.get("name") self.comment_char = config_element.get("comment_char") self.empty_field_value = config_element.get("empty_field_value", "") - self.empty_field_values = {} + self.empty_field_values: Dict[str, str] = {} self.allow_duplicate_entries = util.asbool(config_element.get("allow_duplicate_entries", True)) - self.here = filename and os.path.dirname(filename) - self.filenames = {} + self.here = os.path.dirname(filename) if filename else None + self.filenames: Dict[str, Dict[str, Any]] = {} self.tool_data_path = tool_data_path self.tool_data_path_files = tool_data_path_files self.other_config_dict = other_config_dict or {} @@ -131,7 +140,7 @@ class ToolDataTable(Dictifiable): # This value has no external meaning, and does not represent an abstract version of the underlying data self._loaded_content_version = 1 self._load_info = ( - [config_element, tool_data_path], + (config_element, tool_data_path), { "from_shed_config": from_shed_config, "tool_data_path_files": self.tool_data_path_files, @@ -139,9 +148,9 @@ class ToolDataTable(Dictifiable): "filename": filename, }, ) - self._merged_load_info = [] + self._merged_load_info: List[Tuple[Type[ToolDataTable], Tuple[Tuple[Element, StrPath], Dict[str, Any]]]] = [] - def _update_version(self, version=None): + def _update_version(self, version: Optional[int] = None) -> int: if version is not None: self._loaded_content_version = version else: @@ -151,14 +160,30 @@ class ToolDataTable(Dictifiable): def get_empty_field_by_name(self, name): return self.empty_field_values.get(name, self.empty_field_value) - def _add_entry(self, entry, allow_duplicates=True, persist=False, entry_source=None, **kwd): + def _add_entry( + self, + entry: Union[List[str], Dict[str, str]], + allow_duplicates: bool = True, + persist: bool = False, + entry_source=None, + **kwd, + ) -> None: raise NotImplementedError("Abstract method") - def add_entry(self, entry, allow_duplicates=True, persist=False, entry_source=None, **kwd): + def add_entry( + self, + entry: Union[List[str], Dict[str, str]], + allow_duplicates: bool = True, + persist: bool = False, + entry_source=None, + **kwd, + ) -> int: self._add_entry(entry, allow_duplicates=allow_duplicates, persist=persist, entry_source=entry_source, **kwd) return self._update_version() - def add_entries(self, entries, allow_duplicates=True, persist=False, entry_source=None, **kwd): + def add_entries( + self, entries: List[List[str]], allow_duplicates: bool = True, persist: bool = False, entry_source=None, **kwd + ) -> int: for entry in entries: try: self.add_entry( @@ -178,13 +203,20 @@ class ToolDataTable(Dictifiable): def is_current_version(self, other_version): return self._loaded_content_version == other_version - def merge_tool_data_table(self, other_table, allow_duplicates=True, persist=False, entry_source=None, **kwd): + def merge_tool_data_table( + self, + other_table: "ToolDataTable", + allow_duplicates: bool = True, + persist: bool = False, + entry_source=None, + **kwd, + ) -> int: raise NotImplementedError("Abstract method") - def reload_from_files(self): + def reload_from_files(self) -> int: new_version = self._update_version() merged_info = self._merged_load_info - self.__init__(*self._load_info[0], **self._load_info[1]) + self.__init__(*self._load_info[0], **self._load_info[1]) # type: ignore[misc] self._update_version(version=new_version) for (tool_data_table_class, load_info) in merged_info: self.merge_tool_data_table(tool_data_table_class(*load_info[0], **load_info[1]), allow_duplicates=False) @@ -214,26 +246,32 @@ class TabularToolDataTable(ToolDataTable): def __init__( self, - config_element, - tool_data_path, - from_shed_config=False, - filename=None, - tool_data_path_files=None, - other_config_dict=None, - ): + config_element: Element, + tool_data_path: Optional[StrPath], + tool_data_path_files: ToolDataPathFiles, + from_shed_config: bool = False, + filename: Optional[StrPath] = None, + other_config_dict: Optional[Union["GalaxyAppConfiguration", Dict[str, Any]]] = None, + ) -> None: super().__init__( config_element, tool_data_path, + tool_data_path_files, from_shed_config, filename, - tool_data_path_files, other_config_dict=other_config_dict, ) self.config_element = config_element self.data = [] self.configure_and_load(config_element, tool_data_path, from_shed_config) - def configure_and_load(self, config_element, tool_data_path, from_shed_config=False, url_timeout=10): + def configure_and_load( + self, + config_element: Element, + tool_data_path: Optional[StrPath], + from_shed_config: bool = False, + url_timeout: float = 10, + ) -> None: """ Configure and load table from an XML element. """ @@ -307,6 +345,7 @@ class TabularToolDataTable(ToolDataTable): # in self.tool_data_path. file_path, file_name = os.path.split(filename) if file_path != self.tool_data_path: + assert self.tool_data_path corrected_filename = os.path.join(self.tool_data_path, file_name) if self.tool_data_path_files.exists(corrected_filename): filename = corrected_filename @@ -375,7 +414,8 @@ class TabularToolDataTable(ToolDataTable): self.missing_index_file = None self.extend_data_with(filename) - def get_fields(self): + # This method is used in tools, so need to keep its API stable + def get_fields(self) -> List[List[str]]: return self.data def get_field(self, value): @@ -385,7 +425,8 @@ class TabularToolDataTable(ToolDataTable): rval = TabularToolDataField(i) return rval - def get_named_fields_list(self): + # This method is used in tools, so need to keep its API stable + def get_named_fields_list(self) -> List[Dict[Union[str, int], str]]: rval = [] named_columns = self.get_column_name_list() for fields in self.get_fields(): @@ -393,7 +434,7 @@ class TabularToolDataTable(ToolDataTable): for i, field in enumerate(fields): if i == len(named_columns): break - field_name = named_columns[i] + field_name: Optional[Union[str, int]] = named_columns[i] if field_name is None: field_name = i # check that this is supposed to be 0 based. field_dict[field_name] = field @@ -403,7 +444,7 @@ class TabularToolDataTable(ToolDataTable): def get_version_fields(self): return (self._loaded_content_version, self.get_fields()) - def parse_column_spec(self, config_element): + def parse_column_spec(self, config_element: Element) -> None: """ Parse column definitions, which can either be a set of 'column' elements with a name and index (as in dynamic options config), or a shorthand @@ -412,7 +453,7 @@ class TabularToolDataTable(ToolDataTable): A column named 'value' is required. """ - self.columns = {} + self.columns: Dict[str, int] = {} if config_element.find("columns") is not None: column_names = util.xml_text(config_element.find("columns")) column_names = [n.strip() for n in column_names.split(",")] @@ -437,13 +478,15 @@ class TabularToolDataTable(ToolDataTable): if "name" not in self.columns: self.columns["name"] = self.columns["value"] - def extend_data_with(self, filename, errors=None): + def extend_data_with(self, filename: str, errors: Optional[List[str]] = None) -> None: here = os.path.dirname(os.path.abspath(filename)) self.data.extend(self.parse_file_fields(filename, errors=errors, here=here)) if not self.allow_duplicate_entries: self._deduplicate_data() - def parse_file_fields(self, filename, errors: Optional[List[str]] = None, here="__HERE__"): + def parse_file_fields( + self, filename: str, errors: Optional[List[str]] = None, here: str = "__HERE__" + ) -> List[List[str]]: """ Parse separated lines from file and return a list of tuples. @@ -472,8 +515,9 @@ class TabularToolDataTable(ToolDataTable): log.debug("Loaded %i lines from '%s' for '%s'", len(rval), filename, self.name) return rval - def get_column_name_list(self): - rval = [] + # This method is used in tools, so need to keep its API stable + def get_column_name_list(self) -> List[Union[str, None]]: + rval: List[Union[str, None]] = [] for i in range(self.largest_index + 1): found_column = False for name, index in self.columns.items(): @@ -488,26 +532,27 @@ class TabularToolDataTable(ToolDataTable): rval.append(None) return rval - def get_entry(self, query_attr, query_val, return_attr, default=None): + # This method is used in tools, so need to keep its API stable + def get_entry(self, query_attr: str, query_val: str, return_attr: str, default: None = None): """ Returns table entry associated with a col/val pair. """ - rval = self.get_entries(query_attr, query_val, return_attr, default=default, limit=1) + rval = self.get_entries(query_attr, query_val, return_attr, limit=1) if rval: return rval[0] return default - def get_entries(self, query_attr, query_val, return_attr, default=None, limit=None): + def get_entries(self, query_attr: str, query_val: str, return_attr: str, limit=None) -> List: """ - Returns table entry associated with a col/val pair. + Returns table entries associated with a col/val pair. """ query_col = self.columns.get(query_attr, None) if query_col is None: - return default + return [] if return_attr is not None: return_col = self.columns.get(return_attr, None) if return_col is None: - return default + return [] rval = [] # Look for table entry. for fields in self.get_fields(): @@ -521,9 +566,12 @@ class TabularToolDataTable(ToolDataTable): rval.append(fields[return_col]) if limit is not None and len(rval) == limit: break - return rval or default + return rval - def get_filename_for_source(self, source, default=None): + # This method is used in tools, so need to keep its API stable + def get_filename_for_source( + self, source: Optional[Union[Dict, "DataManager"]], default: Optional[str] = None + ) -> Optional[str]: if source: # if dict, assume is compatible info dict, otherwise call method if isinstance(source, dict): @@ -534,7 +582,7 @@ class TabularToolDataTable(ToolDataTable): source_repo_info = None filename = default for name, value in self.filenames.items(): - repo_info = value.get("tool_shed_repository", None) + repo_info = value.get("tool_shed_repository") if (not source_repo_info and not repo_info) or ( source_repo_info and repo_info and source_repo_info == repo_info ): @@ -542,7 +590,14 @@ class TabularToolDataTable(ToolDataTable): break return filename - def _add_entry(self, entry, allow_duplicates=True, persist=False, entry_source=None, **kwd): + def _add_entry( + self, + entry: Union[List[str], Dict[str, str]], + allow_duplicates: bool = True, + persist: bool = False, + entry_source=None, + **kwd, + ) -> None: # accepts dict or list of columns if isinstance(entry, dict): fields = [] @@ -677,8 +732,8 @@ class TabularToolDataTable(ToolDataTable): def xml_string(self): return util.xml_to_string(self.config_element) - def to_dict(self, view="collection"): - rval = super().to_dict(view=view) + def to_dict(self, view: str = "collection", value_mapper: Optional[Dict[str, Callable]] = None) -> Dict[str, Any]: + rval = super().to_dict(view, value_mapper) if view == "element": rval["columns"] = sorted(self.columns.keys(), key=lambda x: self.columns[x]) rval["fields"] = self.get_fields() @@ -689,7 +744,7 @@ class TabularToolDataField(Dictifiable): dict_collection_visible_keys: List[str] = [] - def __init__(self, data): + def __init__(self, data: Dict): self.data = data def __getitem__(self, key): @@ -727,8 +782,8 @@ class TabularToolDataField(Dictifiable): sha1.update(util.smart_str(fmap[k])) return sha1.hexdigest() - def to_dict(self): - rval = super().to_dict() + def to_dict(self, view: str = "collection", value_mapper: Optional[Dict[str, Callable]] = None) -> Dict[str, Any]: + rval = super().to_dict(view, value_mapper) rval["name"] = self.data["value"] rval["fields"] = self.data rval["base_dir"] = (self.get_base_dir(),) @@ -737,7 +792,7 @@ class TabularToolDataField(Dictifiable): return rval -def _expand_here_template(content, here=None): +def _expand_here_template(content: str, here: Optional[str]) -> str: if here and content: content = string.Template(content).safe_substitute({"__HERE__": here}) return content @@ -750,16 +805,16 @@ tool_data_table_types_list: List[Type[ToolDataTable]] = [TabularToolDataTable] class ToolDataTableManager(Dictifiable): """Manages a collection of tool data tables""" - data_tables: Dict[str, "ToolDataTable"] + data_tables: Dict[str, ToolDataTable] tool_data_table_types = {cls.type_key: cls for cls in tool_data_table_types_list} def __init__( self, tool_data_path: str, - config_filename: Optional[ConfigFilesT] = None, + config_filename: Optional[Union[StrPath, List[StrPath]]] = None, tool_data_table_config_path_set=None, - other_config_dict=None, - ): + other_config_dict: Optional[Union["GalaxyAppConfiguration", Dict[str, Any]]] = None, + ) -> None: self.tool_data_path = tool_data_path # This stores all defined data table entries from both the tool_data_table_conf.xml file and the shed_tool_data_table_conf.xml file # at server startup. If tool shed repositories are installed that contain a valid file named tool_data_table_conf.xml.sample, entries @@ -776,13 +831,13 @@ class ToolDataTableManager(Dictifiable): data_tables = [ToolDataEntry(**table.to_dict()) for table in self.data_tables.values()] return ToolDataEntryList.construct(__root__=data_tables) - def __getitem__(self, key: str): + def __getitem__(self, key: str) -> ToolDataTable: return self.data_tables.__getitem__(key) - def __setitem__(self, key: str, value): + def __setitem__(self, key: str, value) -> None: return self.data_tables.__setitem__(key, value) - def __contains__(self, key: str): + def __contains__(self, key: str) -> bool: return self.data_tables.__contains__(key) def get(self, name: str, default=None): @@ -791,22 +846,27 @@ class ToolDataTableManager(Dictifiable): except KeyError: return default - def set(self, name: str, value): + def set(self, name: str, value: ToolDataTable) -> None: self[name] = value def get_tables(self) -> Dict[str, "ToolDataTable"]: return self.data_tables - def to_dict(self, view: str = "collection", value_mapper=None): - return {name: data_table.to_dict(view="export") for name, data_table in self.data_tables.items()} + def to_dict( + self, view: str = "collection", value_mapper: Optional[Dict[str, Callable]] = None + ) -> Dict[str, Dict[str, Any]]: + return { + name: data_table.to_dict(view="export", value_mapper=value_mapper) + for name, data_table in self.data_tables.items() + } - def to_json(self, path: Union[str, os.PathLike]) -> None: + def to_json(self, path: StrPath) -> None: with open(path, "w") as out: out.write(json.dumps(self.to_dict())) def load_from_config_file( - self, config_filename: ConfigFilesT, tool_data_path: Union[str, os.PathLike], from_shed_config: bool = False - ): + self, config_filename: StrPath, tool_data_path: Optional[StrPath], from_shed_config: bool = False + ) -> List[Element]: """ This method is called under 3 conditions: @@ -817,56 +877,60 @@ class ToolDataTableManager(Dictifiable): Galaxy instance. In this case, we have 2 entry types to handle, files whose root tag is , for example: """ table_elems = [] - config_filenames: List[Union[str, os.PathLike]] - if not isinstance(config_filename, list): - config_filenames = [config_filename] - else: - config_filenames = config_filename - for filename in config_filenames: - tree = util.parse_xml(filename) - root = tree.getroot() - for table_elem in root.findall("table"): - table = self.from_elem( - table_elem, - tool_data_path, - from_shed_config, - filename=filename, - tool_data_path_files=self.tool_data_path_files, - other_config_dict=self.other_config_dict, + tree = util.parse_xml(config_filename) + root = tree.getroot() + for table_elem in root.findall("table"): + table = self.from_elem( + table_elem, + tool_data_path, + from_shed_config, + filename=config_filename, + tool_data_path_files=self.tool_data_path_files, + other_config_dict=self.other_config_dict, + ) + table_elems.append(table_elem) + if table.name not in self.data_tables: + self.data_tables[table.name] = table + log.debug("Loaded tool data table '%s' from file '%s'", table.name, config_filename) + else: + log.debug( + "Loading another instance of data table '%s' from file '%s', attempting to merge content.", + table.name, + config_filename, ) - table_elems.append(table_elem) - if table.name not in self.data_tables: - self.data_tables[table.name] = table - log.debug("Loaded tool data table '%s' from file '%s'", table.name, filename) - else: - log.debug( - "Loading another instance of data table '%s' from file '%s', attempting to merge content.", - table.name, - filename, - ) - self.data_tables[table.name].merge_tool_data_table( - table, allow_duplicates=False - ) # only merge content, do not persist to disk, do not allow duplicate rows when merging - # FIXME: This does not account for an entry with the same unique build ID, but a different path. + self.data_tables[table.name].merge_tool_data_table( + table, allow_duplicates=False + ) # only merge content, do not persist to disk, do not allow duplicate rows when merging + # FIXME: This does not account for an entry with the same unique build ID, but a different path. return table_elems def from_elem( - self, table_elem, tool_data_path, from_shed_config, filename, tool_data_path_files, other_config_dict=None - ): + self, + table_elem: Element, + tool_data_path: Optional[StrPath], + from_shed_config: bool, + filename: StrPath, + tool_data_path_files: ToolDataPathFiles, + other_config_dict: Optional[Union["GalaxyAppConfiguration", Dict[str, Any]]] = None, + ) -> ToolDataTable: table_type = table_elem.get("type", "tabular") assert table_type in self.tool_data_table_types, f"Unknown data table type '{table_type}'" return self.tool_data_table_types[table_type]( table_elem, tool_data_path, + tool_data_path_files=tool_data_path_files, from_shed_config=from_shed_config, filename=filename, - tool_data_path_files=tool_data_path_files, other_config_dict=other_config_dict, ) def add_new_entries_from_config_file( - self, config_filename, tool_data_path, shed_tool_data_table_config, persist=False - ): + self, + config_filename: StrPath, + tool_data_path: Optional[StrPath], + shed_tool_data_table_config: StrPath, + persist: bool = False, + ) -> Tuple[List[Element], str]: """ This method is called when a tool shed repository that includes a tool_data_table_conf.xml.sample file is being installed into a local galaxy instance. We have 2 cases to handle, files whose root tag is , for example:: @@ -904,7 +968,12 @@ class ToolDataTableManager(Dictifiable): self.to_xml_file(shed_tool_data_table_config, table_elems) return table_elems, error_message - def to_xml_file(self, shed_tool_data_table_config, new_elems=None, remove_elems=None): + def to_xml_file( + self, + shed_tool_data_table_config: StrPath, + new_elems: Optional[List[Element]] = None, + remove_elems: Optional[List[Element]] = None, + ) -> None: """ Write the current in-memory version of the shed_tool_data_table_conf.xml file to disk. remove_elems are removed before new_elems are added. @@ -950,7 +1019,9 @@ class ToolDataTableManager(Dictifiable): if out_path_is_new: self.tool_data_path_files.update_files() - def reload_tables(self, table_names=None, path=None): + def reload_tables( + self, table_names: Optional[Union[List[str], str]] = None, path: Optional[str] = None + ) -> List[str]: """ Reload tool data tables. If neither table_names nor path is given, reloads all tool data tables. """ @@ -967,7 +1038,7 @@ class ToolDataTableManager(Dictifiable): log.debug("Reloaded tool data table '%s' from files.", table_name) return table_names - def get_table_names_by_path(self, path): + def get_table_names_by_path(self, path: str) -> List[str]: """Returns a list of table names given a path""" table_names = set() for name, data_table in self.data_tables.items(): diff --git a/lib/galaxy/tool_util/fetcher.py b/lib/galaxy/tool_util/fetcher.py index c9874e017f2..3c9bedaaaf9 100644 --- a/lib/galaxy/tool_util/fetcher.py +++ b/lib/galaxy/tool_util/fetcher.py @@ -1,23 +1,32 @@ +from typing import ( + Dict, + Type, + TYPE_CHECKING, +) + from galaxy.util import plugin_config +if TYPE_CHECKING: + from galaxy.tool_util.locations import ToolLocationResolver + class ToolLocationFetcher: def __init__(self): self.resolver_classes = self.__resolvers_dict() - def __resolvers_dict(self): + def __resolvers_dict(self) -> Dict[str, Type["ToolLocationResolver"]]: import galaxy.tool_util.locations return plugin_config.plugins_dict(galaxy.tool_util.locations, "scheme") - def to_tool_path(self, path_or_uri_like, **kwds): + def to_tool_path(self, path_or_uri_like: str, **kwds) -> str: if "://" not in path_or_uri_like: path = path_or_uri_like else: uri_like = path_or_uri_like if ":" not in path_or_uri_like: raise Exception("Invalid URI passed to get_tool_source") - scheme, rest = uri_like.split(":", 2) + scheme = uri_like.split(":", 2)[0] if scheme not in self.resolver_classes: raise Exception(f"Unknown tool scheme [{scheme}] for URI [{uri_like}]") path = self.resolver_classes[scheme]().get_tool_source_path(uri_like) diff --git a/lib/galaxy/tool_util/locations/__init__.py b/lib/galaxy/tool_util/locations/__init__.py index b3e0d9534e5..b9fe699de50 100644 --- a/lib/galaxy/tool_util/locations/__init__.py +++ b/lib/galaxy/tool_util/locations/__init__.py @@ -14,7 +14,7 @@ class ToolLocationResolver(metaclass=ABCMeta): """Short label for the type of location resolver and URI scheme.""" @abstractmethod - def get_tool_source_path(self, uri_like): + def get_tool_source_path(self, uri_like: str) -> str: """Return a local path for the uri_like string.""" def _temp_path(self, uri_like): diff --git a/lib/galaxy/tool_util/locations/dockstore.py b/lib/galaxy/tool_util/locations/dockstore.py index 991fbdc098b..3ef9c929889 100644 --- a/lib/galaxy/tool_util/locations/dockstore.py +++ b/lib/galaxy/tool_util/locations/dockstore.py @@ -11,7 +11,7 @@ class DockStoreResolver(ToolLocationResolver): scheme = "dockstore" - def get_tool_source_path(self, uri_like): + def get_tool_source_path(self, uri_like: str) -> str: assert uri_like.startswith("dockstore://") tool_id = uri_like[len("dockstore://") :] if ":" in tool_id: diff --git a/lib/galaxy/tool_util/locations/file.py b/lib/galaxy/tool_util/locations/file.py index 2d13ef2a7d4..cdc09981f17 100644 --- a/lib/galaxy/tool_util/locations/file.py +++ b/lib/galaxy/tool_util/locations/file.py @@ -5,6 +5,6 @@ class HttpToolResolver(ToolLocationResolver): scheme = "file" - def get_tool_source_path(self, uri_like): + def get_tool_source_path(self, uri_like: str) -> str: assert uri_like.startswith("file://") return uri_like[len("file://") :] diff --git a/lib/galaxy/tool_util/locations/http.py b/lib/galaxy/tool_util/locations/http.py index dea0c868e42..e8fb7372cb0 100644 --- a/lib/galaxy/tool_util/locations/http.py +++ b/lib/galaxy/tool_util/locations/http.py @@ -9,7 +9,7 @@ class HttpToolResolver(ToolLocationResolver): def __init__(self, **kwds): pass - def get_tool_source_path(self, uri_like): + def get_tool_source_path(self, uri_like: str) -> str: tmp_path = self._temp_path(uri_like) download_to_file(uri_like, tmp_path) return tmp_path diff --git a/lib/galaxy/tool_util/parser/factory.py b/lib/galaxy/tool_util/parser/factory.py index 5093fc6e7d4..6a88f7b9d96 100644 --- a/lib/galaxy/tool_util/parser/factory.py +++ b/lib/galaxy/tool_util/parser/factory.py @@ -1,11 +1,20 @@ """Constructors for concrete tool and input source objects.""" import logging +from typing import ( + Callable, + Dict, + List, + Optional, +) from yaml import safe_load from galaxy.tool_util.loader import load_tool_with_refereces -from galaxy.util import parse_xml_string_to_etree +from galaxy.util import ( + ElementTree, + parse_xml_string_to_etree, +) from galaxy.util.yaml_util import ordered_load from .cwl import ( CwlToolSource, @@ -28,21 +37,21 @@ from ..fetcher import ToolLocationFetcher log = logging.getLogger(__name__) -def build_xml_tool_source(xml_string): +def build_xml_tool_source(xml_string: str) -> XmlToolSource: return XmlToolSource(parse_xml_string_to_etree(xml_string)) -def build_cwl_tool_source(yaml_string): +def build_cwl_tool_source(yaml_string: str) -> CwlToolSource: proxy = tool_proxy(tool_object=safe_load(yaml_string)) # regular CwlToolSource sets basename as tool id, but that's not going to cut it in production return CwlToolSource(tool_proxy=proxy) -def build_yaml_tool_source(yaml_string): +def build_yaml_tool_source(yaml_string: str) -> YamlToolSource: return YamlToolSource(safe_load(yaml_string)) -TOOL_SOURCE_FACTORIES = { +TOOL_SOURCE_FACTORIES: Dict[str, Callable[[str], ToolSource]] = { "XmlToolSource": build_xml_tool_source, "YamlToolSource": build_yaml_tool_source, "CwlToolSource": build_cwl_tool_source, @@ -50,13 +59,13 @@ TOOL_SOURCE_FACTORIES = { def get_tool_source( - config_file=None, - xml_tree=None, - enable_beta_formats=True, - tool_location_fetcher=None, - macro_paths=None, - tool_source_class=None, - raw_tool_source=None, + config_file: Optional[str] = None, + xml_tree: Optional[ElementTree] = None, + enable_beta_formats: bool = True, + tool_location_fetcher: Optional[ToolLocationFetcher] = None, + macro_paths: Optional[List[str]] = None, + tool_source_class: Optional[str] = None, + raw_tool_source: Optional[str] = None, ) -> ToolSource: """Return a ToolSource object corresponding to supplied source. @@ -75,6 +84,7 @@ def get_tool_source( if tool_location_fetcher is None: tool_location_fetcher = ToolLocationFetcher() + assert config_file config_file = tool_location_fetcher.to_tool_path(config_file) if not enable_beta_formats: tree, macro_paths = load_tool_with_refereces(config_file) diff --git a/lib/galaxy/tool_util/parser/output_collection_def.py b/lib/galaxy/tool_util/parser/output_collection_def.py index 8e5481feef2..b2c76f51fcd 100644 --- a/lib/galaxy/tool_util/parser/output_collection_def.py +++ b/lib/galaxy/tool_util/parser/output_collection_def.py @@ -118,7 +118,7 @@ class FilePatternDatasetCollectionDescription(DatasetCollectionDescription): self.recurse = asbool(kwargs.get("recurse", False)) self.match_relative_path = asbool(kwargs.get("match_relative_path", False)) if pattern in NAMED_PATTERNS: - pattern = NAMED_PATTERNS.get(pattern) + pattern = NAMED_PATTERNS[pattern] self.pattern = pattern self.sort_by = sort_by = kwargs.get("sort_by", DEFAULT_SORT_BY) if sort_by.startswith("reverse_"): @@ -149,7 +149,7 @@ class FilePatternDatasetCollectionDescription(DatasetCollectionDescription): return as_dict @property - def discover_patterns(self): + def discover_patterns(self) -> List[str]: return [self.pattern] diff --git a/lib/galaxy/tool_util/parser/output_objects.py b/lib/galaxy/tool_util/parser/output_objects.py index 48314fb5090..9cf29a42077 100644 --- a/lib/galaxy/tool_util/parser/output_objects.py +++ b/lib/galaxy/tool_util/parser/output_objects.py @@ -1,5 +1,11 @@ -from typing import List +from typing import ( + Any, + Dict, + List, + Optional, +) +from galaxy.util import Element from galaxy.util.dictifiable import Dictifiable from .output_actions import ToolOutputActionGroup from .output_collection_def import ( @@ -9,7 +15,14 @@ from .output_collection_def import ( class ToolOutputBase(Dictifiable): - def __init__(self, name, label=None, filters=None, hidden=False, from_expression=None): + def __init__( + self, + name: str, + label: Optional[str] = None, + filters: Optional[List[Element]] = None, + hidden: bool = False, + from_expression: Optional[str] = None, + ) -> None: super().__init__() self.name = name self.label = label @@ -50,18 +63,18 @@ class ToolOutput(ToolOutputBase): def __init__( self, - name, - format=None, - format_source=None, - metadata_source=None, - parent=None, - label=None, - filters=None, - actions=None, - hidden=False, - implicit=False, - from_expression=None, - ): + name: str, + format: Optional[str] = None, + format_source: Optional[str] = None, + metadata_source: Optional[str] = None, + parent: Optional[str] = None, + label: Optional[str] = None, + filters: Optional[List[Element]] = None, + actions: Optional[ToolOutputActionGroup] = None, + hidden: bool = False, + implicit: bool = False, + from_expression: Optional[str] = None, + ) -> None: super().__init__(name, label=label, filters=filters, hidden=hidden, from_expression=from_expression) self.output_type = "data" self.format = format @@ -71,14 +84,17 @@ class ToolOutput(ToolOutputBase): self.actions = actions # Initialize default values - self.change_format = [] + self.change_format: List[Element] = [] self.implicit = implicit - self.from_work_dir = None + self.from_work_dir: Optional[str] = None self.dataset_collector_descriptions: List[DatasetCollectionDescription] = [] + self.default_identifier_source: Optional[str] = None + self.count: Optional[int] = None + self.tool: Optional[Any] # Tuple emulation - def __len__(self): + def __len__(self) -> int: return 3 def __getitem__(self, index): @@ -106,7 +122,7 @@ class ToolOutput(ToolOutputBase): return as_dict @staticmethod - def from_dict(name, output_dict, tool=None): + def from_dict(name: str, output_dict: Dict[str, Any], tool: Optional[object] = None) -> "ToolOutput": output = ToolOutput(name) output.format = output_dict.get("format", "data") output.change_format = [] @@ -119,7 +135,7 @@ class ToolOutput(ToolOutputBase): output.filters = [] output.tool = tool output.from_work_dir = output_dict.get("from_work_dir", None) - output.hidden = output_dict.get("hidden", "") + output.hidden = output_dict.get("hidden", False) # TODO: implement tool output action group fixes output.actions = ToolOutputActionGroup(output, None) output.dataset_collector_descriptions = dataset_collector_descriptions_from_output_dict(output_dict) @@ -183,30 +199,30 @@ class ToolOutputCollection(ToolOutputBase): def __init__( self, - name, - structure, - label=None, - filters=None, - hidden=False, - default_format="data", - default_format_source=None, - default_metadata_source=None, - inherit_format=False, - inherit_metadata=False, - ): + name: str, + structure: "ToolOutputCollectionStructure", + label: Optional[str] = None, + filters: Optional[List[Element]] = None, + hidden: bool = False, + default_format: str = "data", + default_format_source: Optional[str] = None, + default_metadata_source: Optional[str] = None, + inherit_format: bool = False, + inherit_metadata: bool = False, + ) -> None: super().__init__(name, label=label, filters=filters, hidden=hidden) self.output_type = "collection" self.collection = True self.default_format = default_format self.structure = structure - self.outputs = {} + self.outputs: Dict[str, str] = {} self.inherit_format = inherit_format self.inherit_metadata = inherit_metadata self.metadata_source = default_metadata_source self.format_source = default_format_source - self.change_format = [] # TODO + self.change_format: List = [] # TODO: not implemented def known_outputs(self, inputs, type_registry): if self.dynamic_structure: @@ -275,13 +291,13 @@ class ToolOutputCollection(ToolOutputBase): return as_dict @staticmethod - def from_dict(name, output_dict, tool=None): + def from_dict(name, output_dict, tool=None) -> "ToolOutputCollection": structure = ToolOutputCollectionStructure.from_dict(output_dict["structure"]) rval = ToolOutputCollection( name, structure=structure, label=output_dict.get("label", None), - filters=None, + filters=[], hidden=output_dict.get("hidden", False), default_format=output_dict.get("default_format", "data"), default_format_source=output_dict.get("default_format_source", None), @@ -299,12 +315,12 @@ class ToolOutputCollection(ToolOutputBase): class ToolOutputCollectionStructure: def __init__( self, - collection_type, - collection_type_source=None, - collection_type_from_rules=None, - structured_like=None, - dataset_collector_descriptions=None, - ): + collection_type: Optional[str], + collection_type_source: Optional[str] = None, + collection_type_from_rules: Optional[str] = None, + structured_like: Optional[str] = None, + dataset_collector_descriptions: Optional[List[DatasetCollectionDescription]] = None, + ) -> None: self.collection_type = collection_type self.collection_type_source = collection_type_source self.collection_type_from_rules = collection_type_from_rules @@ -349,7 +365,7 @@ class ToolOutputCollectionStructure: } @staticmethod - def from_dict(as_dict): + def from_dict(as_dict) -> "ToolOutputCollectionStructure": structure = ToolOutputCollectionStructure( collection_type=as_dict["collection_type"], collection_type_source=as_dict["collection_type_source"], diff --git a/lib/galaxy/tool_util/verify/script.py b/lib/galaxy/tool_util/verify/script.py index e8852fa3832..619ed8f87fe 100644 --- a/lib/galaxy/tool_util/verify/script.py +++ b/lib/galaxy/tool_util/verify/script.py @@ -1,18 +1,25 @@ #!/usr/bin/env python import argparse +import concurrent.futures.thread import datetime as dt import json import logging import os import sys import tempfile -from collections import namedtuple from concurrent.futures import ( thread, ThreadPoolExecutor, ) -from typing import List +from typing import ( + Any, + Callable, + Dict, + List, + NamedTuple, + Optional, +) import yaml @@ -30,18 +37,27 @@ ALL_VERSION = "*" LATEST_VERSION = None -TestReference = namedtuple("TestReference", ["tool_id", "tool_version", "test_index"]) -TestException = namedtuple("TestException", ["tool_id", "exception", "was_recorded"]) +class TestReference(NamedTuple): + tool_id: str + tool_version: Optional[str] + test_index: int + + +class TestException(NamedTuple): + tool_id: str + exception: Exception + was_recorded: bool class Results: - test_exceptions: List[Exception] + test_exceptions: List[TestException] - def __init__(self, default_suitename, test_json, append=False, galaxy_url=None): + def __init__( + self, default_suitename: str, test_json: str, append: bool = False, galaxy_url: Optional[str] = None + ) -> None: self.test_json = test_json or "-" self.galaxy_url = galaxy_url test_results = [] - test_exceptions: List[Exception] = [] suitename = default_suitename if append: assert test_json != "-" @@ -51,16 +67,16 @@ class Results: if "suitename" in previous_results: suitename = previous_results["suitename"] self.test_results = test_results - self.test_exceptions = test_exceptions + self.test_exceptions = [] self.suitename = suitename - def register_result(self, result): + def register_result(self, result: Dict[str, Any]) -> None: self.test_results.append(result) - def register_exception(self, test_exception): + def register_exception(self, test_exception: TestException) -> None: self.test_exceptions.append(test_exception) - def already_successful(self, test_reference): + def already_successful(self, test_reference: TestReference) -> bool: test_data = self._previous_test_data(test_reference) if test_data: if "status" in test_data and test_data["status"] == "success": @@ -68,7 +84,7 @@ class Results: return False - def already_executed(self, test_reference): + def already_executed(self, test_reference: TestReference) -> bool: test_data = self._previous_test_data(test_reference) if test_data: if "status" in test_data and test_data["status"] != "skipped": @@ -76,7 +92,7 @@ class Results: return False - def _previous_test_data(self, test_reference): + def _previous_test_data(self, test_reference: TestReference) -> Optional[Dict[str, Any]]: test_id = _test_id_for_reference(test_reference) for test_result in self.test_results: if test_result.get("id") != test_id: @@ -89,7 +105,7 @@ class Results: return None - def write(self): + def write(self) -> None: tests = sorted(self.test_results, key=lambda el: el["id"]) n_passed, n_failures, n_skips = 0, 0, 0 n_errors = len([e for e in self.test_exceptions if not e.was_recorded]) @@ -127,56 +143,37 @@ class Results: with open(self.test_json, "w") as f: json.dump(report_obj, f) - def info_message(self): + def info_message(self) -> str: messages = [] passed_tests = self._tests_with_status("success") messages.append("Passed tool tests ({}): {}".format(len(passed_tests), [t["id"] for t in passed_tests])) failed_tests = self._tests_with_status("failure") messages.append("Failed tool tests ({}): {}".format(len(failed_tests), [t["id"] for t in failed_tests])) - skiped_tests = self._tests_with_status("skip") - messages.append("Skipped tool tests ({}): {}".format(len(skiped_tests), [t["id"] for t in skiped_tests])) + skipped_tests = self._tests_with_status("skip") + messages.append("Skipped tool tests ({}): {}".format(len(skipped_tests), [t["id"] for t in skipped_tests])) errored_tests = self._tests_with_status("error") messages.append("Errored tool tests ({}): {}".format(len(errored_tests), [t["id"] for t in errored_tests])) return "\n".join(messages) - @property - def success_count(self): - self._tests_with_status("success") - - @property - def skip_count(self): - self._tests_with_status("skip") - - @property - def error_count(self): - return self._tests_with_status("error") + len(self.test_exceptions) - - @property - def failure_count(self): - return self._tests_with_status("failure") - - def _tests_with_status(self, status): + def _tests_with_status(self, status: str) -> List[Dict[str, Any]]: return [t for t in self.test_results if t.get("data", {}).get("status") == status] def test_tools( - galaxy_interactor, - test_references, - results, - log=None, - parallel_tests=1, - history_per_test_case=False, - history_name=None, - no_history_reuse=False, - no_history_cleanup=False, - publish_history=False, - retries=0, - verify_kwds=None, -): - """Run through tool tests and write report. - - Refactor this into Galaxy in 21.01. - """ + galaxy_interactor: GalaxyInteractorApi, + test_references: List[TestReference], + results: Results, + log: Optional[logging.Logger] = None, + parallel_tests: int = 1, + history_per_test_case: bool = False, + history_name: Optional[str] = None, + no_history_reuse: bool = False, + no_history_cleanup: bool = False, + publish_history: bool = False, + retries: int = 0, + verify_kwds: Optional[Dict[str, Any]] = None, +) -> None: + """Run through tool tests and write report.""" verify_kwds = (verify_kwds or {}).copy() tool_test_start = dt.datetime.now() history_created = False @@ -221,8 +218,8 @@ def test_tools( try: executor.shutdown(wait=True) except KeyboardInterrupt: - executor._threads.clear() - thread._threads_queues.clear() + executor._threads.clear() # type: ignore[attr-defined] + thread._threads_queues.clear() # type: ignore[attr-defined] results.write() if log: if results.test_json == "-": @@ -236,7 +233,7 @@ def test_tools( galaxy_interactor.delete_history(test_history) -def _test_id_for_reference(test_reference): +def _test_id_for_reference(test_reference: "TestReference") -> str: tool_id = test_reference.tool_id tool_version = test_reference.tool_version test_index = test_reference.test_index @@ -253,15 +250,15 @@ def _test_id_for_reference(test_reference): def _test_tool( - executor, - test_reference, - results, - galaxy_interactor, - log, - retries, - publish_history, - verify_kwds, -): + executor: concurrent.futures.thread.ThreadPoolExecutor, + test_reference: "TestReference", + results: Results, + galaxy_interactor: GalaxyInteractorApi, + log: Optional[logging.Logger], + retries: int, + publish_history: bool, + verify_kwds: Dict[str, Any], +) -> None: tool_id = test_reference.tool_id tool_version = test_reference.tool_version test_index = test_reference.test_index @@ -272,7 +269,7 @@ def _test_tool( test_id = _test_id_for_reference(test_reference) - def run_test(): + def run_test() -> None: run_retries = retries job_data = None job_exception = None @@ -323,16 +320,16 @@ def _test_tool( def build_case_references( - galaxy_interactor, - tool_id=ALL_TOOLS, - tool_version=LATEST_VERSION, - test_index=ALL_TESTS, - page_size=0, - page_number=0, - test_filters=None, - log=None, -): - test_references = [] + galaxy_interactor: GalaxyInteractorApi, + tool_id: str = ALL_TOOLS, + tool_version: Optional[str] = LATEST_VERSION, + test_index: int = ALL_TESTS, + page_size: int = 0, + page_number: int = 0, + test_filters: Optional[List[Callable[[TestReference], bool]]] = None, + log: Optional[logging.Logger] = None, +) -> List[TestReference]: + test_references: List[TestReference] = [] if tool_id == ALL_TOOLS: tests_summary = galaxy_interactor.get_tests_summary() for tool_id, tool_versions_dict in tests_summary.items(): @@ -341,8 +338,7 @@ def build_case_references( test_reference = TestReference(tool_id, tool_version, test_index) test_references.append(test_reference) else: - assert tool_id - tool_test_dicts = galaxy_interactor.get_tool_tests(tool_id, tool_version=tool_version) or {} + tool_test_dicts = galaxy_interactor.get_tool_tests(tool_id, tool_version=tool_version) for i, tool_test_dict in enumerate(tool_test_dicts): this_tool_version = tool_test_dict.get("tool_version", tool_version) this_test_index = i @@ -351,7 +347,7 @@ def build_case_references( test_references.append(test_reference) if test_filters is not None and len(test_filters) > 0: - filtered_test_references = [] + filtered_test_references: List[TestReference] = [] for test_reference in test_references: skip_test = False for test_filter in test_filters: @@ -361,7 +357,10 @@ def build_case_references( skip_test = True if not skip_test: filtered_test_references.append(test_reference) - log.info(f"Skipping {len(test_references)-len(filtered_test_references)} out of {len(test_references)} tests.") + if log is not None: + log.info( + f"Skipping {len(test_references)-len(filtered_test_references)} out of {len(test_references)} tests." + ) test_references = filtered_test_references if page_size > 0: @@ -372,7 +371,7 @@ def build_case_references( return test_references -def main(argv=None): +def main(argv=None) -> None: if argv is None: argv = sys.argv[1:] @@ -384,7 +383,11 @@ def main(argv=None): sys.exit(1) -def run_tests(args, test_filters=None, log=None): +def run_tests( + args: argparse.Namespace, + test_filters: Optional[List[Callable[[TestReference], bool]]] = None, + log: Optional[logging.Logger] = None, +) -> None: # Split out argument parsing so we can quickly build other scripts - such as a script # to run all tool tests for a workflow by just passing in a custom test_filters. test_filters = test_filters or [] @@ -464,12 +467,10 @@ def run_tests(args, test_filters=None, log=None): exceptions = results.test_exceptions if exceptions: exception = exceptions[0] - if hasattr(exception, "exception"): - exception = exception.exception - raise exception + raise exception.exception -def setup_global_logger(name, log_file=None, verbose=False): +def setup_global_logger(name: str, log_file: Optional[str] = None, verbose: bool = False) -> logging.Logger: formatter = logging.Formatter("%(asctime)s %(levelname)-5s - %(message)s") console = logging.StreamHandler() console.setFormatter(formatter) @@ -490,7 +491,7 @@ def setup_global_logger(name, log_file=None, verbose=False): return logger -def arg_parser(): +def arg_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(description=DESCRIPTION) parser.add_argument("-u", "--galaxy-url", default="http://localhost:8080", help="Galaxy URL") parser.add_argument("-k", "--key", default=None, help="Galaxy User API Key") diff --git a/lib/galaxy/tools/data/__init__.py b/lib/galaxy/tools/data/__init__.py index 0bdbbabd96e..97454ca1413 100644 --- a/lib/galaxy/tools/data/__init__.py +++ b/lib/galaxy/tools/data/__init__.py @@ -68,21 +68,21 @@ class RefgenieToolDataTable(TabularToolDataTable): self, config_element, tool_data_path, + tool_data_path_files, from_shed_config=False, filename=None, - tool_data_path_files=None, other_config_dict=None, - ): + ) -> None: super().__init__( config_element, tool_data_path, + tool_data_path_files, from_shed_config, filename, - tool_data_path_files, other_config_dict=other_config_dict, ) self.config_element = config_element - self.data = [] + self.data: List[List[str]] = [] self.configure_and_load(config_element, tool_data_path, from_shed_config) def configure_and_load(self, config_element, tool_data_path, from_shed_config=False, url_timeout=10): diff --git a/lib/galaxy/tools/evaluation.py b/lib/galaxy/tools/evaluation.py index ae211ecf0e2..2ae9cd70ca2 100644 --- a/lib/galaxy/tools/evaluation.py +++ b/lib/galaxy/tools/evaluation.py @@ -27,6 +27,7 @@ from galaxy.structured_app import ( BasicSharedApp, MinimalToolApp, ) +from galaxy.tool_util.data import TabularToolDataTable from galaxy.tools.parameters import ( visit_input_values, wrapped_json, @@ -463,7 +464,10 @@ class ToolEvaluator: Queries and returns an entry in a data table. """ if table_name in self.app.tool_data_tables: - return self.app.tool_data_tables[table_name].get_entry(query_attr, query_val, return_attr) + table = self.app.tool_data_tables[table_name] + if not isinstance(table, TabularToolDataTable): + raise Exception(f"Expected a TabularToolDataTable but got a {type(table)}: {table}.") + return table.get_entry(query_attr, query_val, return_attr) param_dict["__tool_directory__"] = self.compute_environment.tool_directory() param_dict["__get_data_table_entry__"] = get_data_table_entry diff --git a/lib/galaxy/util/__init__.py b/lib/galaxy/util/__init__.py index 2df7c31c6a1..d72e90d924f 100644 --- a/lib/galaxy/util/__init__.py +++ b/lib/galaxy/util/__init__.py @@ -37,6 +37,9 @@ from typing import ( List, Optional, overload, + Tuple, + TypeVar, + Union, ) from urllib.parse import ( urlencode, @@ -62,14 +65,20 @@ except ImportError: LXML_AVAILABLE = True try: from lxml import etree - from lxml.etree import ( - _Element as Element, - ElementTree, - ) + from lxml.etree import _Element as Element + + # lxml.etree.ElementTree is a function that returns a new instance of the + # lxml.etree._ElementTree class. This class doesn't have a proper + # __init__() method, so we can add a __new__() constructor that mimicks + # xml.etree.ElementTree.ElementTree initialization. + class ElementTree(etree._ElementTree): + def __new__(cls, element=None, file=None) -> etree.ElementTree: + return etree.ElementTree(element, file=file) + except ImportError: LXML_AVAILABLE = False import xml.etree.ElementTree as etree # type: ignore[assignment,no-redef] - from xml.etree.ElementTree import ( # noqa: F401 + from xml.etree.ElementTree import ( # type: ignore[assignment] Element, ElementTree, ) @@ -286,7 +295,7 @@ def unique_id(KEY_SIZE=128): return md5(random_bits).hexdigest() -def parse_xml(fname: StrPath, strip_whitespace=True, remove_comments=True) -> etree.ElementTree: +def parse_xml(fname: StrPath, strip_whitespace=True, remove_comments=True) -> ElementTree: """Returns a parsed xml tree""" parser = None if remove_comments and LXML_AVAILABLE: @@ -313,28 +322,28 @@ def parse_xml(fname: StrPath, strip_whitespace=True, remove_comments=True) -> et return tree -def parse_xml_string(xml_string, strip_whitespace=True) -> etree.Element: +def parse_xml_string(xml_string: str, strip_whitespace: bool = True) -> Element: try: - tree = etree.fromstring(xml_string) + elem = etree.fromstring(xml_string) except ValueError as e: if "strings with encoding declaration are not supported" in unicodify(e): - tree = etree.fromstring(xml_string.encode("utf-8")) + elem = etree.fromstring(xml_string.encode("utf-8")) else: raise e if strip_whitespace: - for elem in tree.iter("*"): - if elem.text is not None: - elem.text = elem.text.strip() - if elem.tail is not None: - elem.tail = elem.tail.strip() - return tree + for sub_elem in elem.iter("*"): + if sub_elem.text is not None: + sub_elem.text = sub_elem.text.strip() + if sub_elem.tail is not None: + sub_elem.tail = sub_elem.tail.strip() + return elem -def parse_xml_string_to_etree(xml_string, strip_whitespace=True): +def parse_xml_string_to_etree(xml_string: str, strip_whitespace: bool = True) -> ElementTree: return ElementTree(parse_xml_string(xml_string=xml_string, strip_whitespace=strip_whitespace)) -def xml_to_string(elem, pretty=False) -> str: +def xml_to_string(elem: Element, pretty: bool = False) -> str: """ Returns a string from an xml tree. """ @@ -1046,7 +1055,32 @@ def string_as_bool_or_none(string): return False -def listify(item, do_strip: bool = False) -> List[Any]: +ItemType = TypeVar("ItemType") + + +@overload +def listify(item: Union[None, Literal[False]], do_strip: bool = False) -> List: + ... + + +@overload +def listify(item: str, do_strip: bool = False) -> List[str]: + ... + + +@overload +def listify(item: Union[List[ItemType], Tuple[ItemType, ...]], do_strip: bool = False) -> List[ItemType]: + ... + + +# Unfortunately we cannot use ItemType .. -> List[ItemType] in the next overload +# because then that would also match Union types. +@overload +def listify(item: Any, do_strip: bool = False) -> List: + ... + + +def listify(item: Any, do_strip: bool = False) -> List: """ Make a single item a single item list. @@ -1065,9 +1099,7 @@ def listify(item, do_strip: bool = False) -> List[Any]: """ if not item: return [] - elif isinstance(item, list): - return item - elif isinstance(item, tuple): + elif isinstance(item, (list, tuple)): return list(item) elif isinstance(item, str) and item.count(","): if do_strip: diff --git a/lib/galaxy/util/dictifiable.py b/lib/galaxy/util/dictifiable.py index 551638dc3a1..1d1ec228262 100644 --- a/lib/galaxy/util/dictifiable.py +++ b/lib/galaxy/util/dictifiable.py @@ -1,5 +1,11 @@ import datetime import uuid +from typing import ( + Any, + Callable, + Dict, + Optional, +) def dict_for(obj, **kwds): @@ -12,7 +18,7 @@ class Dictifiable: when for sharing objects across boundaries, such as the API, tool scripts, and JavaScript code.""" - def to_dict(self, view="collection", value_mapper=None): + def to_dict(self, view: str = "collection", value_mapper: Optional[Dict[str, Callable]] = None) -> Dict[str, Any]: """ Return item dictionary. """ @@ -29,8 +35,9 @@ class Dictifiable: try: return item.to_dict(view=view, value_mapper=value_mapper) except Exception: + assert value_mapper is not None if key in value_mapper: - return value_mapper.get(key)(item) + return value_mapper[key](item) if type(item) == datetime.datetime: return item.isoformat() elif type(item) == uuid.UUID: diff --git a/lib/galaxy/webapps/base/webapp.py b/lib/galaxy/webapps/base/webapp.py index f68aab1d5a3..45df16c83af 100644 --- a/lib/galaxy/webapps/base/webapp.py +++ b/lib/galaxy/webapps/base/webapp.py @@ -13,6 +13,7 @@ from http.cookies import CookieError from typing import ( Any, Dict, + Optional, ) from urllib.parse import urlparse @@ -36,6 +37,10 @@ from galaxy.exceptions import ( from galaxy.managers import context from galaxy.managers.session import GalaxySessionManager from galaxy.managers.users import UserManager +from galaxy.structured_app import ( + BasicSharedApp, + MinimalApp, +) from galaxy.util import ( asbool, safe_makedirs, @@ -96,7 +101,9 @@ class WebApplication(base.WebApplication): injection_aware: bool = False - def __init__(self, galaxy_app, session_cookie="galaxysession", name=None): + def __init__( + self, galaxy_app: MinimalApp, session_cookie: str = "galaxysession", name: Optional[str] = None + ) -> None: super().__init__() self.name = name galaxy_app.is_webapp = True @@ -188,7 +195,7 @@ class WebApplication(base.WebApplication): def make_body_iterable(self, trans, body): return base.WebApplication.make_body_iterable(self, trans, body) - def transaction_chooser(self, environ, galaxy_app, session_cookie): + def transaction_chooser(self, environ, galaxy_app: BasicSharedApp, session_cookie: str): return GalaxyWebTransaction(environ, galaxy_app, self, session_cookie) def add_ui_controllers(self, package_name, app): @@ -275,12 +282,14 @@ class GalaxyWebTransaction(base.DefaultWebTransaction, context.ProvidesHistoryCo (specifically the user's "cookie" session and history) """ - def __init__(self, environ: Dict[str, Any], app, webapp, session_cookie=None) -> None: + def __init__( + self, environ: Dict[str, Any], app: BasicSharedApp, webapp: WebApplication, session_cookie: Optional[str] = None + ) -> None: self._app = app self.webapp = webapp self.user_manager = app[UserManager] self.session_manager = app[GalaxySessionManager] - base.DefaultWebTransaction.__init__(self, environ) + super().__init__(environ) self.expunge_all() config = self.app.config self.debug = asbool(config.get("debug", False)) @@ -305,11 +314,13 @@ class GalaxyWebTransaction(base.DefaultWebTransaction, context.ProvidesHistoryCo # If not, check for an active session but do not create one. # If an error message is set here, it's sent back using # trans.show_error in the response -- in expose_api. + assert session_cookie self.error_message = self._authenticate_api(session_cookie) elif self.app.name == "reports": self.galaxy_session = None else: # This is a web request, get or create session. + assert session_cookie self._ensure_valid_session(session_cookie) if self.galaxy_session: # When we've authenticated by session, we have to check the @@ -318,7 +329,7 @@ class GalaxyWebTransaction(base.DefaultWebTransaction, context.ProvidesHistoryCo if config.use_remote_user and self.galaxy_session.user.deleted: self.response.send_redirect(url_for("/static/user_disabled.html")) if config.require_login: - self._ensure_logged_in_user(environ, session_cookie) + self._ensure_logged_in_user(session_cookie) if config.session_duration: # TODO DBTODO All ajax calls from the client need to go through # a single point of control where we can do things like @@ -494,7 +505,7 @@ class GalaxyWebTransaction(base.DefaultWebTransaction, context.ProvidesHistoryCo if self.app.config.cookie_domain is not None: self.response.cookies[name]["domain"] = self.app.config.cookie_domain - def _authenticate_api(self, session_cookie): + def _authenticate_api(self, session_cookie: str) -> Optional[str]: """ Authenticate for the API via key or session (if available). """ @@ -524,8 +535,9 @@ class GalaxyWebTransaction(base.DefaultWebTransaction, context.ProvidesHistoryCo # Anonymous API interaction -- anything but @expose_api_anonymous will fail past here. self.user = None self.galaxy_session = None + return None - def _ensure_valid_session(self, session_cookie, create=True): + def _ensure_valid_session(self, session_cookie: str, create: bool = True) -> None: """ Ensure that a valid Galaxy session exists and is available as trans.session (part of initialization) @@ -606,6 +618,7 @@ class GalaxyWebTransaction(base.DefaultWebTransaction, context.ProvidesHistoryCo log.warning(f"User '{galaxy_session.user.email}' is marked deleted, invalidating session") # Do we need to invalidate the session for some reason? if invalidate_existing_session: + assert galaxy_session prev_galaxy_session = galaxy_session prev_galaxy_session.is_valid = False galaxy_session = None @@ -630,10 +643,11 @@ class GalaxyWebTransaction(base.DefaultWebTransaction, context.ProvidesHistoryCo if invalidate_existing_session: self.get_or_create_default_history() - def _ensure_logged_in_user(self, environ, session_cookie): + def _ensure_logged_in_user(self, session_cookie: str) -> None: # The value of session_cookie can be one of # 'galaxysession' or 'galaxycommunitysession' # Currently this method does nothing unless session_cookie is 'galaxysession' + assert self.galaxy_session if session_cookie == "galaxysession" and self.galaxy_session.user is None: # TODO: re-engineer to eliminate the use of allowed_paths # as maintenance overhead is far too high. diff --git a/lib/galaxy/webapps/galaxy/api/history_contents.py b/lib/galaxy/webapps/galaxy/api/history_contents.py index a47bdad8e45..3cec18b0fe1 100644 --- a/lib/galaxy/webapps/galaxy/api/history_contents.py +++ b/lib/galaxy/webapps/galaxy/api/history_contents.py @@ -73,7 +73,6 @@ from galaxy.webapps.galaxy.api.common import ( from galaxy.webapps.galaxy.services.history_contents import ( CreateHistoryContentFromStore, CreateHistoryContentPayload, - DatasetDetailsType, DirectionOptions, HistoriesContentsService, HistoryContentsFilterList, @@ -234,7 +233,7 @@ def parse_legacy_index_query_params( else: content_types = [e.value for e in HistoryContentType] - id_list: Optional[List[DecodedDatabaseIdField]] = None + id_list = None if ids: id_list = util.listify(ids) # If explicit ids given, always used detailed result. @@ -257,7 +256,7 @@ def parse_legacy_index_query_params( def parse_dataset_details(details: Optional[str]): """Parses the different values that the `dataset_details` parameter can have from a string.""" - dataset_details: Optional[DatasetDetailsType] = None + dataset_details = None if details is not None and details != "all": dataset_details = set(util.listify(details)) else: # either None or 'all' diff --git a/lib/tool_shed/webapp/api/categories.py b/lib/tool_shed/webapp/api/categories.py index 859b856edc9..63448b7e69c 100644 --- a/lib/tool_shed/webapp/api/categories.py +++ b/lib/tool_shed/webapp/api/categories.py @@ -1,4 +1,8 @@ import logging +from typing import ( + Callable, + Dict, +) import tool_shed.util.shed_util_common as suc from galaxy import ( @@ -20,7 +24,7 @@ log = logging.getLogger(__name__) class CategoriesController(BaseAPIController): """RESTful controller for interactions with categories in the Tool Shed.""" - def __get_value_mapper(self, trans): + def __get_value_mapper(self, trans) -> Dict[str, Callable]: value_mapper = {"id": trans.security.encode_id} return value_mapper diff --git a/lib/tool_shed/webapp/api/groups.py b/lib/tool_shed/webapp/api/groups.py index 512b542979b..1f1cbff1114 100644 --- a/lib/tool_shed/webapp/api/groups.py +++ b/lib/tool_shed/webapp/api/groups.py @@ -1,4 +1,8 @@ import logging +from typing import ( + Callable, + Dict, +) from galaxy import ( util, @@ -28,7 +32,7 @@ class GroupsController(BaseAPIController): super().__init__(app) self.group_manager = groups.GroupManager() - def __get_value_mapper(self, trans): + def __get_value_mapper(self, trans) -> Dict[str, Callable]: value_mapper = {"id": trans.security.encode_id} return value_mapper diff --git a/lib/tool_shed/webapp/api/repositories.py b/lib/tool_shed/webapp/api/repositories.py index 2102b167619..5dfc86d2ecf 100644 --- a/lib/tool_shed/webapp/api/repositories.py +++ b/lib/tool_shed/webapp/api/repositories.py @@ -5,6 +5,10 @@ import tarfile from collections import namedtuple from io import StringIO from time import strftime +from typing import ( + Callable, + Dict, +) from sqlalchemy import ( and_, @@ -289,7 +293,7 @@ class RepositoriesController(BaseAPIController): return [] return repository.installable_revisions(self.app) - def __get_value_mapper(self, trans): + def __get_value_mapper(self, trans) -> Dict[str, Callable]: value_mapper = { "id": trans.security.encode_id, "repository_id": trans.security.encode_id, diff --git a/lib/tool_shed/webapp/api/repository_revisions.py b/lib/tool_shed/webapp/api/repository_revisions.py index da37cfd71e8..791d94189e3 100644 --- a/lib/tool_shed/webapp/api/repository_revisions.py +++ b/lib/tool_shed/webapp/api/repository_revisions.py @@ -1,4 +1,8 @@ import logging +from typing import ( + Callable, + Dict, +) from sqlalchemy import and_ @@ -21,7 +25,7 @@ log = logging.getLogger(__name__) class RepositoryRevisionsController(BaseAPIController): """RESTful controller for interactions with tool shed repository revisions.""" - def __get_value_mapper(self, trans): + def __get_value_mapper(self, trans) -> Dict[str, Callable]: value_mapper = { "id": trans.security.encode_id, "repository_id": trans.security.encode_id, diff --git a/mypy.ini b/mypy.ini index 6340a19a970..bbefdbab609 100644 --- a/mypy.ini +++ b/mypy.ini @@ -210,8 +210,6 @@ check_untyped_defs = False check_untyped_defs = False [mypy-galaxy.tools.bundled.extract.extract_genomic_dna] check_untyped_defs = False -[mypy-galaxy.tool_util.parser.output_objects] -check_untyped_defs = False [mypy-galaxy.tool_util.deps.resolvers] check_untyped_defs = False [mypy-galaxy.tool_util.deps.mulled.mulled_update_singularity_containers] @@ -248,8 +246,6 @@ check_untyped_defs = False check_untyped_defs = False [mypy-galaxy.tools.expressions.evaluation] check_untyped_defs = False -[mypy-galaxy.tool_util.data] -check_untyped_defs = False [mypy-galaxy.tool_util.verify] check_untyped_defs = False [mypy-galaxy.tool_util.toolbox.watcher] diff --git a/test/integration/test_data_manager.py b/test/integration/test_data_manager.py index f093101e122..d6a42b00bae 100644 --- a/test/integration/test_data_manager.py +++ b/test/integration/test_data_manager.py @@ -146,7 +146,7 @@ class TestDataManagerIntegration(integration_util.IntegrationTestCase, UsesShed) self._app.tool_data_tables.get("all_fasta").remove_entry(table_content["dm6"]) entries = self._app.tool_data_tables.get("all_fasta").get_entries("dbkey", "dm6", "dbkey") - assert entries is None + assert not entries def test_data_manager_manual_multiple(self): """ @@ -190,14 +190,14 @@ class TestDataManagerIntegration(integration_util.IntegrationTestCase, UsesShed) self._app.tool_data_tables.get("all_fasta").remove_entry(table_content["dm6"]) entries = self._app.tool_data_tables.get("all_fasta").get_entries("dbkey", "dm6", "dbkey") - assert entries is None + assert not entries self._app.tool_data_tables.get("all_fasta").remove_entry(table_content["NC_001617.1"]) entries = self._app.tool_data_tables.get("all_fasta").get_entries( "dbkey", "another_unique_dbkey_value", "dbkey" ) - assert entries is None + assert not entries @classmethod def get_secure_ascii_digits(cls, n=12): diff --git a/test/integration/test_data_manager_refgenie.py b/test/integration/test_data_manager_refgenie.py index 73b828dd8f8..3af08162ab4 100644 --- a/test/integration/test_data_manager_refgenie.py +++ b/test/integration/test_data_manager_refgenie.py @@ -96,7 +96,7 @@ class TestDataManagerIntegration(integration_util.IntegrationTestCase, UsesShed) self._app.tool_data_tables.get("all_fasta").to_dict(view="element")["fields"][0] ) entries = self._app.tool_data_tables.get("all_fasta").get_entries("dbkey", "dm6", "dbkey") - assert entries is None + assert not entries def test_data_manager_manual_refgenie_dbkeys(self): """ @@ -121,7 +121,7 @@ class TestDataManagerIntegration(integration_util.IntegrationTestCase, UsesShed) self._app.tool_data_tables.get("__dbkeys__").to_dict(view="element")["fields"][0] ) entries = self._app.tool_data_tables.get("all_fasta").get_entries("name", "dm7", "name") - assert entries is None + assert not entries @classmethod def get_secure_ascii_digits(cls, n=12): diff --git a/test/unit/app/tools/test_actions.py b/test/unit/app/tools/test_actions.py index 8b4b8898eeb..a111018a845 100644 --- a/test/unit/app/tools/test_actions.py +++ b/test/unit/app/tools/test_actions.py @@ -1,5 +1,8 @@ import string -from typing import cast +from typing import ( + cast, + Optional, +) from galaxy import model from galaxy.app_unittest_utils import tools_support @@ -228,14 +231,16 @@ def __assert_output_format_is(expected, output, input_extensions=None, param_con assert actual_format == expected, f"Actual format {actual_format}, does not match expected {expected}" -def quick_output(format, format_source=None, change_format_xml=None): +def quick_output( + format: str, format_source: Optional[str] = None, change_format_xml: Optional[str] = None +) -> ToolOutput: test_output = ToolOutput("test_output") test_output.format = format test_output.format_source = format_source if change_format_xml: - test_output.change_format = XML(change_format_xml) + test_output.change_format = XML(change_format_xml).findall("change_format") else: - test_output.change_format = None + test_output.change_format = [] return test_output diff --git a/test/unit/tool_util/test_verify_script.py b/test/unit/tool_util/test_verify_script.py index 6d6a9484bd0..4c8f42c26a7 100644 --- a/test/unit/tool_util/test_verify_script.py +++ b/test/unit/tool_util/test_verify_script.py @@ -1,6 +1,10 @@ import json import os from tempfile import NamedTemporaryFile +from typing import ( + cast, + NoReturn, +) from unittest import mock from galaxy.tool_util.unittest_utils.interactor import ( @@ -10,6 +14,7 @@ from galaxy.tool_util.unittest_utils.interactor import ( MockGalaxyInteractor, NEW_HISTORY_ID, ) +from galaxy.tool_util.verify.interactor import GalaxyInteractorApi from galaxy.tool_util.verify.script import ( arg_parser, build_case_references, @@ -21,7 +26,7 @@ from galaxy.tool_util.verify.script import ( VT_PATH = "galaxy.tool_util.verify.script.verify_tool" -def test_arg_parse(): +def test_arg_parse() -> None: parser = arg_parser() # defaults @@ -61,7 +66,7 @@ def test_arg_parse(): assert args.skip == "executed" -def test_test_tools(): +def test_test_tools() -> None: interactor = MockGalaxyInteractor() f = NamedTemporaryFile() results = Results("my suite", f.name) @@ -73,7 +78,7 @@ def test_test_tools(): with mock.patch(VT_PATH) as mock_verify: assert_results_not_written(results) run( - interactor, + cast(GalaxyInteractorApi, interactor), test_references, results, ) @@ -85,7 +90,7 @@ def test_test_tools(): assert interactor.history_deleted -def test_test_tools_no_history_cleanup(): +def test_test_tools_no_history_cleanup() -> None: interactor = MockGalaxyInteractor() f = NamedTemporaryFile() results = Results("my suite", f.name) @@ -95,7 +100,7 @@ def test_test_tools_no_history_cleanup(): with mock.patch(VT_PATH) as mock_verify: assert_results_not_written(results) run( - interactor, + cast(GalaxyInteractorApi, interactor), test_references, results, no_history_cleanup=True, @@ -108,7 +113,7 @@ def test_test_tools_no_history_cleanup(): assert not interactor.history_deleted -def test_test_tools_history_reuse(): +def test_test_tools_history_reuse() -> None: interactor = MockGalaxyInteractor() f = NamedTemporaryFile() results = Results(EXISTING_SUITE_NAME, f.name) @@ -118,7 +123,7 @@ def test_test_tools_history_reuse(): with mock.patch(VT_PATH) as mock_verify: assert_results_not_written(results) run( - interactor, + cast(GalaxyInteractorApi, interactor), test_references, results, no_history_reuse=False, @@ -133,7 +138,7 @@ def test_test_tools_history_reuse(): assert not interactor.history_deleted -def test_test_tools_no_history_reuse(): +def test_test_tools_no_history_reuse() -> None: interactor = MockGalaxyInteractor() f = NamedTemporaryFile() results = Results("existing suite", f.name) @@ -143,7 +148,7 @@ def test_test_tools_no_history_reuse(): with mock.patch(VT_PATH) as mock_verify: assert_results_not_written(results) run( - interactor, + cast(GalaxyInteractorApi, interactor), test_references, results, no_history_reuse=True, @@ -158,7 +163,7 @@ def test_test_tools_no_history_reuse(): assert interactor.history_deleted -def test_test_tools_history_name(): +def test_test_tools_history_name() -> None: interactor = MockGalaxyInteractor() f = NamedTemporaryFile() results = Results("my suite", f.name) @@ -168,7 +173,7 @@ def test_test_tools_history_name(): with mock.patch(VT_PATH) as mock_verify: assert_results_not_written(results) run( - interactor, + cast(GalaxyInteractorApi, interactor), test_references, results, history_name="testfoo", @@ -183,7 +188,7 @@ def test_test_tools_history_name(): assert interactor.history_deleted -def test_test_tool_per_test_history(): +def test_test_tool_per_test_history() -> None: interactor = MockGalaxyInteractor() f = NamedTemporaryFile() results = Results("my suite", f.name) @@ -194,7 +199,7 @@ def test_test_tool_per_test_history(): with mock.patch(VT_PATH) as mock_verify: assert_results_not_written(results) run( - interactor, + cast(GalaxyInteractorApi, interactor), test_references, results, history_per_test_case=True, @@ -207,7 +212,7 @@ def test_test_tool_per_test_history(): assert not interactor.history_deleted -def test_test_tools_records_exception(): +def test_test_tools_records_exception() -> None: interactor = MockGalaxyInteractor() f = NamedTemporaryFile() results = Results("my suite", f.name) @@ -217,12 +222,12 @@ def test_test_tools_records_exception(): with mock.patch(VT_PATH) as mock_verify: assert_results_not_written(results) - def side_effect(*args, **kwd): + def side_effect(*args, **kwd) -> NoReturn: raise Exception("Cow") mock_verify.side_effect = side_effect run( - interactor, + cast(GalaxyInteractorApi, interactor), test_references, results, ) @@ -232,7 +237,7 @@ def test_test_tools_records_exception(): assert_results_written(results) -def test_test_tools_records_retry_exception(): +def test_test_tools_records_retry_exception() -> None: interactor = MockGalaxyInteractor() f = NamedTemporaryFile() results = Results("my suite", f.name) @@ -253,7 +258,7 @@ def test_test_tools_records_retry_exception(): mock_verify.side_effect = side_effect run( - interactor, + cast(GalaxyInteractorApi, interactor), test_references, results, retries=1, @@ -294,8 +299,8 @@ def test_results(): assert "Skipped tool tests (1)" in message -def test_build_references(): - interactor = MockGalaxyInteractor() +def test_build_references() -> None: + interactor = cast(GalaxyInteractorApi, MockGalaxyInteractor()) test_references = build_case_references(interactor) assert len(test_references) == 6 @@ -332,11 +337,11 @@ def test_build_references(): assert test_reference.test_index == 2 -def assert_results_not_written(results): +def assert_results_not_written(results: Results) -> None: assert os.stat(results.test_json).st_size == 0 -def assert_results_written(results): +def assert_results_written(results: Results) -> None: assert os.stat(results.test_json).st_size > 0 with open(results.test_json) as f: json.load(f) diff --git a/test/unit/util/test_utils.py b/test/unit/util/test_utils.py index fb3b470557d..828ae610af7 100644 --- a/test/unit/util/test_utils.py +++ b/test/unit/util/test_utils.py @@ -122,3 +122,20 @@ def test_galaxy_directory(monkeypatch): assert path1 == path2 == path3 assert os.path.isabs(path1) + + +def test_listify() -> None: + assert util.listify(None) == [] + assert util.listify(False) == [] + assert util.listify(True) == [True] + assert util.listify("foo") == ["foo"] + assert util.listify("foo, bar") == ["foo", " bar"] + assert util.listify("foo, bar", do_strip=True) == ["foo", "bar"] + assert util.listify([1, 2, 3]) == [1, 2, 3] + assert util.listify((1, 2, 3)) == [1, 2, 3] + s = {1, 2, 3} + assert util.listify(s) == [s] + d = {"a": 1, "b": 2, "c": 3} + assert util.listify(d) == [d] + o = object() + assert util.listify(o) == [o] diff --git a/test/unit/webapps/test_routes.py b/test/unit/webapps/test_routes.py index 2673adbad52..36fd7143b63 100644 --- a/test/unit/webapps/test_routes.py +++ b/test/unit/webapps/test_routes.py @@ -1,5 +1,8 @@ +from typing import cast + from routes import request_config +from galaxy.structured_app import MinimalApp from galaxy.util.bunch import Bunch from galaxy.web import url_for from galaxy.webapps.base.webapp import WebApplication @@ -22,7 +25,7 @@ class MockWebApplication(WebApplication): def test_galaxy_routes(): test_config = Bunch(template_cache_path="/tmp") - app = Bunch(config=test_config, security=object(), trace_logger=None, name="galaxy") + app = cast(MinimalApp, Bunch(config=test_config, security=object(), trace_logger=None, name="galaxy")) test_webapp = MockWebApplication(app) galaxy_buildapp.populate_api_routes(test_webapp, app) diff --git a/test/unit/webapps/test_webapp_base.py b/test/unit/webapps/test_webapp_base.py index 1a747d10001..9f329b9ea98 100644 --- a/test/unit/webapps/test_webapp_base.py +++ b/test/unit/webapps/test_webapp_base.py @@ -3,16 +3,23 @@ Unit tests for ``galaxy.web.framework.webapp`` """ import logging import re +from typing import ( + cast, + Optional, +) import galaxy.config from galaxy.app_unittest_utils import galaxy_mock -from galaxy.webapps.base import webapp as Webapp +from galaxy.webapps.base.webapp import ( + GalaxyWebTransaction, + WebApplication, +) log = logging.getLogger(__name__) -class StubGalaxyWebTransaction(Webapp.GalaxyWebTransaction): - def _ensure_valid_session(self, session_cookie, create=True): +class StubGalaxyWebTransaction(GalaxyWebTransaction): + def _ensure_valid_session(self, session_cookie: str, create: bool = True) -> None: pass @@ -28,12 +35,12 @@ class CORSParsingMockConfig(galaxy_mock.MockAppConfig): class TestGalaxyWebTransactionHeaders: - def _new_trans(self, allowed_origin_hostnames=None): + def _new_trans(self, allowed_origin_hostnames: Optional[str] = None) -> StubGalaxyWebTransaction: app = galaxy_mock.MockApp() app.config = CORSParsingMockConfig(allowed_origin_hostnames=allowed_origin_hostnames) - webapp = galaxy_mock.MockWebapp(app.security) + webapp = cast(WebApplication, galaxy_mock.MockWebapp(app.security)) environ = galaxy_mock.buildMockEnviron() - trans = StubGalaxyWebTransaction(environ, app, webapp) + trans = StubGalaxyWebTransaction(environ, app, webapp, "session_cookie") return trans def assert_cors_header_equals(self, headers, should_be): @@ -42,7 +49,7 @@ class TestGalaxyWebTransactionHeaders: def assert_cors_header_missing(self, headers): assert not ("access-control-allow-origin" in headers) - def test_parse_allowed_origin_hostnames(self): + def test_parse_allowed_origin_hostnames(self) -> None: """Should return a list of (possibly) mixed strings and regexps""" config = CORSParsingMockConfig() @@ -57,16 +64,16 @@ class TestGalaxyWebTransactionHeaders: assert isinstance(hostnames[1], str) assert isinstance(hostnames[2], str) - def test_default_set_cors_headers(self): + def test_default_set_cors_headers(self) -> None: """No CORS headers should be set (or even checked) by default""" trans = self._new_trans(allowed_origin_hostnames=None) - assert isinstance(trans, Webapp.GalaxyWebTransaction) + assert isinstance(trans, GalaxyWebTransaction) trans.request.headers["Origin"] = "http://lisaskelprecipes.pinterest.com?id=kelpcake" trans.set_cors_headers() self.assert_cors_header_missing(trans.response.headers) - def test_set_cors_headers(self): + def test_set_cors_headers(self) -> None: """Origin should be echo'd when it matches an allowed hostname""" # an asterisk is a special 'allow all' string trans = self._new_trans(allowed_origin_hostnames="*,beep.com")