From 1c171e4463fdf8638cd287099d9448066ccfe35b Mon Sep 17 00:00:00 2001 From: "Michael R. Crusoe" Date: Tue, 9 Nov 2021 12:32:50 +0100 Subject: [PATCH] improve typing of lib/galaxy/util/compression_utils.py --- lib/galaxy/util/compression_utils.py | 110 ++++++++++++++++++--------- setup.cfg | 13 +++- 2 files changed, 88 insertions(+), 35 deletions(-) diff --git a/lib/galaxy/util/compression_utils.py b/lib/galaxy/util/compression_utils.py index 4a22807dd6a..90e6813f69d 100644 --- a/lib/galaxy/util/compression_utils.py +++ b/lib/galaxy/util/compression_utils.py @@ -1,21 +1,26 @@ +import bz2 import gzip import io import logging import os import tarfile import zipfile +from typing import Any, cast, Generator, IO, Iterable, List, Optional, Tuple, Union from galaxy.util.path import safe_relpath from .checkers import ( - bz2, is_bz2, is_gzip ) log = logging.getLogger(__name__) +FileObjType = Union[gzip.GzipFile, bz2.BZ2File, IO[Any], io.TextIOWrapper] -def get_fileobj(filename, mode="r", compressed_formats=None): + +def get_fileobj( + filename: str, mode: str = "r", compressed_formats: Optional[List[str]] = None +) -> FileObjType: """ Returns a fileobj. If the file is compressed, return an appropriate file reader. In text mode, always use 'utf-8' encoding. @@ -28,7 +33,9 @@ def get_fileobj(filename, mode="r", compressed_formats=None): return get_fileobj_raw(filename, mode, compressed_formats)[1] -def get_fileobj_raw(filename, mode="r", compressed_formats=None): +def get_fileobj_raw( + filename: str, mode: str = "r", compressed_formats: Optional[List[str]] = None +) -> Tuple[Optional[str], FileObjType]: if compressed_formats is None: compressed_formats = ['bz2', 'gzip', 'zip'] # Remove 't' from mode, which may cause an error for compressed files @@ -38,7 +45,7 @@ def get_fileobj_raw(filename, mode="r", compressed_formats=None): mode = 'r' compressed_format = None if 'gzip' in compressed_formats and is_gzip(filename): - fh = gzip.GzipFile(filename, mode) + fh: Union[gzip.GzipFile, bz2.BZ2File, IO[bytes]] = gzip.GzipFile(filename, mode) compressed_format = 'gzip' elif 'bz2' in compressed_formats and is_bz2(filename): fh = bz2.BZ2File(filename, mode) @@ -56,14 +63,18 @@ def get_fileobj_raw(filename, mode="r", compressed_formats=None): elif 'b' in mode: return compressed_format, open(filename, mode) else: - return compressed_format, open(filename, mode, encoding='utf-8') - if 'b' not in mode: - return compressed_format, io.TextIOWrapper(fh, encoding='utf-8') + return compressed_format, open(filename, mode, encoding="utf-8") + if "b" not in mode: + return compressed_format, io.TextIOWrapper( + cast(IO[bytes], fh), encoding="utf-8" + ) else: return compressed_format, fh -def file_iter(fname, sep=None): +def file_iter( + fname: str, sep: Optional[Any] = None +) -> Generator[Union[List[bytes], Any, List[str]], None, None]: """ This generator iterates over a file and yields its lines splitted via the C{sep} parameter. Skips empty lines and lines starting with @@ -79,13 +90,18 @@ def file_iter(fname, sep=None): yield line.split(sep) +ArchiveMemberType = Union[tarfile.TarInfo, zipfile.ZipInfo] + + class CompressedFile: + archive: Union[tarfile.TarFile, zipfile.ZipFile] + @staticmethod - def can_decompress(file_path): + def can_decompress(file_path: str) -> bool: return tarfile.is_tarfile(file_path) or zipfile.is_zipfile(file_path) - def __init__(self, file_path, mode='r'): + def __init__(self, file_path: str, mode: str = "r") -> None: if tarfile.is_tarfile(file_path): self.file_type = 'tar' elif zipfile.is_zipfile(file_path) and not file_path.endswith('.jar'): @@ -101,7 +117,7 @@ class CompressedFile: raise NameError(f'File type {self.file_type} specified, no open method found.') @property - def common_prefix_dir(self): + def common_prefix_dir(self) -> str: """ Get the common prefix directory for all the files in the archive, if any. @@ -114,14 +130,24 @@ class CompressedFile: common_prefix = os.path.commonprefix([self.getname(item) for item in contents]) # If the common_prefix does not end with a slash, check that is a # directory and all other files are contained in it - if len(common_prefix) >= 1 and not common_prefix.endswith(os.sep) and self.isdir(self.getmember(common_prefix)) \ - and all(self.getname(item).startswith(common_prefix + os.sep) for item in contents if self.isfile(item)): + common_prefix_member = self.getmember(common_prefix) + if ( + len(common_prefix) >= 1 + and not common_prefix.endswith(os.sep) + and common_prefix_member + and self.isdir(common_prefix_member) + and all( + self.getname(item).startswith(common_prefix + os.sep) + for item in contents + if self.isfile(item) + ) + ): common_prefix += os.sep if not common_prefix.endswith(os.sep): common_prefix = '' return common_prefix - def extract(self, path): + def extract(self, path: str) -> str: '''Determine the path to which the archive should be extracted.''' contents = self.getmembers() extraction_path = path @@ -132,16 +158,27 @@ class CompressedFile: extraction_path = os.path.join(path, self.file_name) if not os.path.exists(extraction_path): os.makedirs(extraction_path) - self.archive.extractall(extraction_path, members=self.safemembers()) + if isinstance(self.archive, tarfile.TarFile): + members_t = cast(Iterable[tarfile.TarInfo], self.safemembers()) + self.archive.extractall(extraction_path, members=members_t) + else: + members_z = cast(Iterable[str], self.safemembers()) + self.archive.extractall(extraction_path, members=members_z) else: if not common_prefix_dir: extraction_path = os.path.join(path, self.file_name) if not os.path.exists(extraction_path): os.makedirs(extraction_path) - self.archive.extractall(extraction_path, members=self.safemembers()) + if isinstance(self.archive, tarfile.TarFile): + members_t = cast(Iterable[tarfile.TarInfo], self.safemembers()) + self.archive.extractall(extraction_path, members=members_t) + else: + members_z = cast(Iterable[str], self.safemembers()) + self.archive.extractall(extraction_path, members=members_z) # Since .zip files store unix permissions separately, we need to iterate through the zip file # and set permissions on extracted members. if self.file_type == 'zip': + assert isinstance(self.archive, zipfile.ZipFile) for zipped_file in contents: filename = self.getname(zipped_file) absolute_filepath = os.path.join(extraction_path, filename) @@ -155,10 +192,11 @@ class CompressedFile: log.warning(f"Unable to change permission on extracted file '{absolute_filepath}' as it does not exist") return os.path.abspath(os.path.join(extraction_path, common_prefix_dir)) - def safemembers(self): + def safemembers(self) -> Union[Iterable[tarfile.TarInfo], Iterable[str]]: members = self.archive common_prefix_dir = self.common_prefix_dir if self.file_type == "tar": + assert isinstance(members, tarfile.TarFile) for finfo in members: if not safe_relpath(finfo.name): raise Exception(f"Path '{finfo.name}' is blocked (illegal path).") @@ -168,57 +206,61 @@ class CompressedFile: raise Exception(f"Link '{finfo.name}' to '{finfo.linkname}' is blocked.") yield finfo elif self.file_type == "zip": + assert isinstance(members, zipfile.ZipFile) for name in members.namelist(): if not safe_relpath(name): raise Exception(f"{name} is blocked (illegal path).") yield name - def getmembers_tar(self): + def getmembers_tar(self) -> List[tarfile.TarInfo]: + assert isinstance(self.archive, tarfile.TarFile) return self.archive.getmembers() - def getmembers_zip(self): + def getmembers_zip(self) -> List[zipfile.ZipInfo]: + assert isinstance(self.archive, zipfile.ZipFile) return self.archive.infolist() - def getname_tar(self, item): + def getname_tar(self, item: tarfile.TarInfo) -> str: return item.name - def getname_zip(self, item): + def getname_zip(self, item: zipfile.ZipInfo) -> str: return item.filename - def getmember(self, name): + def getmember(self, name: str) -> Optional[ArchiveMemberType]: for member in self.getmembers(): if self.getname(member) == name: return member + return None - def getmembers(self): - return getattr(self, f'getmembers_{self.type}')() + def getmembers(self) -> List[ArchiveMemberType]: + return cast(List[ArchiveMemberType], getattr(self, f"getmembers_{self.type}")()) - def getname(self, member): - return getattr(self, f'getname_{self.type}')(member) + def getname(self, member: ArchiveMemberType) -> str: + return cast(str, getattr(self, f"getname_{self.type}")(member)) - def isdir(self, member): - return getattr(self, f'isdir_{self.type}')(member) + def isdir(self, member: ArchiveMemberType) -> bool: + return cast(bool, getattr(self, f"isdir_{self.type}")(member)) - def isdir_tar(self, member): + def isdir_tar(self, member: tarfile.TarInfo) -> bool: return member.isdir() - def isdir_zip(self, member): + def isdir_zip(self, member: zipfile.ZipInfo) -> bool: if member.filename.endswith(os.sep): return True return False - def isfile(self, member): + def isfile(self, member: ArchiveMemberType) -> bool: if not self.isdir(member): return True return False - def open_tar(self, filepath, mode): + def open_tar(self, filepath: str, mode: str) -> tarfile.TarFile: return tarfile.open(filepath, mode, errorlevel=0) - def open_zip(self, filepath, mode): + def open_zip(self, filepath: str, mode: str) -> zipfile.ZipFile: return zipfile.ZipFile(filepath, mode) - def zipfile_ok(self, path_to_archive): + def zipfile_ok(self, path_to_archive: str) -> bool: """ This function is a bit pedantic and not functionally necessary. It checks whether there is no file pointing outside of the extraction, because ZipFile.extractall() has some potential diff --git a/setup.cfg b/setup.cfg index 53f7bc4b492..d6001d672f8 100644 --- a/setup.cfg +++ b/setup.cfg @@ -219,7 +219,18 @@ check_untyped_defs = False [mypy-test.functional.webhooks.tour_generator] check_untyped_defs = False [mypy-galaxy.util.compression_utils] -check_untyped_defs = False +disallow_any_generics = True +disallow_subclassing_any = True +disallow_untyped_calls = False +disallow_untyped_defs = True +disallow_incomplete_defs = True +check_untyped_defs = True +disallow_untyped_decorators = True +no_implicit_optional = True +warn_unused_ignores = True +warn_return_any = True +no_implicit_reexport = True +strict_equality = True [mypy-galaxy.tools.bundled.filters.gff.gff_filter_by_feature_count] check_untyped_defs = False [mypy-galaxy.tools.bundled.extract.extract_genomic_dna]