improve typing of lib/galaxy/util/compression_utils.py

This commit is contained in:
Michael R. Crusoe
2021-11-11 11:55:06 +01:00
parent 0f185683b7
commit 1c171e4463
2 changed files with 88 additions and 35 deletions
+76 -34
View File
@@ -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
+12 -1
View File
@@ -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]