diff --git a/multiurl/base.py b/multiurl/base.py index f73c4c5..6e56e0d 100644 --- a/multiurl/base.py +++ b/multiurl/base.py @@ -34,7 +34,7 @@ def close(self, *args, **kwargs): pass -def progress_bar(total, initial=0, desc=None): +def progress_bar(total, initial=0, desc=None, disable=False): try: # There is a bug in tqdm that expects ipywidgets # to be installed if running in a notebook @@ -45,7 +45,8 @@ def progress_bar(total, initial=0, desc=None): from tqdm import tqdm # noqa F401 except ImportError: tqdm = NoBar - + if disable: + tqdm = NoBar return tqdm( total=total, initial=initial, @@ -107,7 +108,7 @@ def extension(self, url=None): extensions.append(".unknown") return "".join(reversed(extensions)) - def download(self, target): + def download(self, target, disable_progress_bar=False): if os.path.exists(target) and not self.override_target_file: return @@ -124,6 +125,7 @@ def download(self, target): total=size, initial=skip, desc=self.title(), + disable=disable_progress_bar, ) as pbar: with open(download, mode) as f: total = self.transfer(f, pbar) diff --git a/multiurl/downloader.py b/multiurl/downloader.py index 67b4f56..e56345c 100644 --- a/multiurl/downloader.py +++ b/multiurl/downloader.py @@ -107,5 +107,5 @@ def Downloader(url, **kwargs): return MultiDownloader(downloaders) -def download(url, target, **kwargs): - return Downloader(url, **kwargs).download(target) +def download(url, target, disable_progress_bar=False, **kwargs): + return Downloader(url, **kwargs).download(target, disable_progress_bar) diff --git a/tests/test_downloader.py b/tests/test_downloader.py index 1d12f70..331c3f2 100644 --- a/tests/test_downloader.py +++ b/tests/test_downloader.py @@ -11,9 +11,11 @@ import os import pytest +from tqdm.std import tqdm from multiurl import Downloader, download from multiurl.http import FullHTTPDownloader +from multiurl.base import progress_bar, NoBar def test_http(): @@ -109,6 +111,13 @@ def test_ftp_download(tmp_path, ftpserver): assert original.read() == downloaded.read() +def test_progress_bar(): + bar = progress_bar(10, False, 5) + assert isinstance(bar, tqdm) + + nobar = progress_bar(10, True, 5) + assert isinstance(nobar, NoBar) + if __name__ == "__main__": logging.basicConfig(level=logging.DEBUG) # test_order()