diff --git a/sqlite_utils/cli.py b/sqlite_utils/cli.py index 0e56c9f..fe42ec2 100644 --- a/sqlite_utils/cli.py +++ b/sqlite_utils/cli.py @@ -13,7 +13,7 @@ import os import sys import csv as csv_std import tabulate -from .utils import find_spatialite, sqlite3, decode_base64_values +from .utils import file_progress, find_spatialite, sqlite3, decode_base64_values VALID_COLUMN_TYPES = ("INTEGER", "TEXT", "FLOAT", "BLOB") @@ -584,6 +584,7 @@ def insert_upsert_options(fn): help="Character encoding for input, defaults to utf-8", ), load_extension_option, + click.option("--silent", is_flag=True, help="Do not show progress bar"), ) ): fn = decorator(fn) @@ -608,6 +609,7 @@ def insert_upsert_implementation( default=None, encoding=None, load_extension=None, + silent=False, ): db = sqlite_utils.Database(path) _load_extensions(db, load_extension) @@ -621,9 +623,10 @@ def insert_upsert_implementation( pk = pk[0] if csv or tsv: dialect = "excel-tab" if tsv else "excel" - reader = csv_std.reader(json_file, dialect=dialect) - headers = next(reader) - docs = (dict(zip(headers, row)) for row in reader) + with file_progress(json_file, silent=silent) as json_file: + reader = csv_std.reader(json_file, dialect=dialect) + headers = next(reader) + docs = (dict(zip(headers, row)) for row in reader) elif nl: docs = (json.loads(line) for line in json_file) else: @@ -673,6 +676,7 @@ def insert( alter, encoding, load_extension, + silent, ignore, replace, truncate, @@ -702,6 +706,7 @@ def insert( truncate=truncate, encoding=encoding, load_extension=load_extension, + silent=silent, not_null=not_null, default=default, ) @@ -725,6 +730,7 @@ def upsert( default, encoding, load_extension, + silent, ): """ Upsert records based on their primary key. Works like 'insert' but if @@ -747,6 +753,7 @@ def upsert( default=default, encoding=encoding, load_extension=load_extension, + silent=silent, ) except UnicodeDecodeError as ex: raise click.ClickException(UNICODE_ERROR.format(ex)) diff --git a/sqlite_utils/utils.py b/sqlite_utils/utils.py index 89b9241..aab72ca 100644 --- a/sqlite_utils/utils.py +++ b/sqlite_utils/utils.py @@ -1,4 +1,7 @@ import base64 +import click +import contextlib +import io import os try: @@ -87,3 +90,29 @@ def find_spatialite(): if os.path.exists(path): return path return None + + +class UpdateReader(io.TextIOWrapper): + def __init__(self, raw, update): + super().__init__(raw) + self._update = update + + def read(self, size=-1): + bytes = super().read(size) + self._update(len(bytes)) + return bytes + + def readline(self, size=-1): + bytes = super().readline(size) + self._update(len(bytes)) + return bytes + + +@contextlib.contextmanager +def file_progress(file, silent=False, **kwargs): + if silent or file.raw.fileno() == 0: # 0 = stdin + yield file + else: + file_length = os.path.getsize(file.raw.name) + with click.progressbar(length=file_length, **kwargs) as bar: + yield UpdateReader(file, update=bar.update) diff --git a/tests/test_cli.py b/tests/test_cli.py index 062f319..5c3056e 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -1640,7 +1640,7 @@ def test_insert_encoding(tmpdir): open(csv_path, "wb").write(latin1_csv) # First attempt should error: bad_result = CliRunner().invoke( - cli.cli, ["insert", db_path, "places", csv_path, "--csv"] + cli.cli, ["insert", db_path, "places", csv_path, "--csv"], catch_exceptions=False ) assert bad_result.exit_code == 1 assert ( @@ -1651,6 +1651,7 @@ def test_insert_encoding(tmpdir): good_result = CliRunner().invoke( cli.cli, ["insert", db_path, "places", csv_path, "--encoding", "latin-1", "--csv"], + catch_exceptions=False ) assert good_result.exit_code == 0 db = Database(db_path)