mirror of
https://github.com/galaxyproject/galaxy.git
synced 2026-09-24 16:30:27 +08:00
improve typing of lib/galaxy/util/compression_utils.py
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user