Merge pull request #14971 from mr-c/mypy_0.990

bonus Python typing
This commit is contained in:
John Chilton
2022-11-28 20:50:49 +01:00
committed by GitHub
48 changed files with 969 additions and 537 deletions
+12 -12
View File
@@ -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
+3 -3
View File
@@ -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)
+3 -2
View File
@@ -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):
+1 -1
View File
@@ -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)
-1
View File
@@ -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
+3 -3
View File
@@ -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):
+3 -3
View File
@@ -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
+1 -1
View File
@@ -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,
+6 -6
View File
@@ -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
+1 -1
View File
@@ -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.
"""
+5 -17
View File
@@ -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}"
+1
View File
@@ -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"
+1 -1
View File
@@ -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
File diff suppressed because it is too large Load Diff
+10 -4
View File
@@ -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
+3 -2
View File
@@ -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'
@@ -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,
@@ -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.
@@ -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)
+6 -5
View File
@@ -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,
+180 -109
View File
@@ -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 <tables>, 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 <tables>, 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():
+12 -3
View File
@@ -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)
+1 -1
View File
@@ -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):
+1 -1
View File
@@ -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:
+1 -1
View File
@@ -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://") :]
+1 -1
View File
@@ -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
+22 -12
View File
@@ -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)
@@ -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]
+57 -41
View File
@@ -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"],
+86 -85
View File
@@ -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")
+4 -4
View File
@@ -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):
+5 -1
View File
@@ -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
+53 -21
View File
@@ -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:
+9 -2
View File
@@ -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:
+22 -8
View File
@@ -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.
@@ -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'
+5 -1
View File
@@ -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
+5 -1
View File
@@ -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
+5 -1
View File
@@ -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,
@@ -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,
-4
View File
@@ -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]
+3 -3
View File
@@ -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):
@@ -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):
+9 -4
View File
@@ -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
+27 -22
View File
@@ -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)
+17
View File
@@ -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]
+4 -1
View File
@@ -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)
+17 -10
View File
@@ -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")