Fixed issue #433 - CLI eats cursor

The issue is that underlying iterator is not fully consumed within the body of
the `with file_progress()` block. Instead, that block creates generator
expressions like `docs = (dict(zip(headers, row)) for row in reader)`

These iterables are consumed later, outside the `with file_progress()` block,
which consumes the underlying iterator, and in turn updates the progress bar.

This means that the `ProgressBar.__exit__` method gets called before the last
time the `ProgressBar.update` method gets called. The result is that the code to
make the cursor invisible (inside the `update()` method) is called after the
cleanup code to make it visible (in the `__exit__` method).

The fix is to move consumption of the `docs` iterators within the progress bar block.

(An additional fix, to make ProgressBar more robust against this kind of misuse, would
to make it refusing to update after its `__exit__` method had been called, just
like files cannot be `read()` after they are closed. That requires a in the
click library).
This commit is contained in:
Luke Plant 2023-10-04 18:49:28 +01:00
commit 76113d1cb1

View file

@ -1024,93 +1024,97 @@ def insert_upsert_implementation(
if flatten: if flatten:
docs = (_flatten(doc) for doc in docs) docs = (_flatten(doc) for doc in docs)
if stop_after: if stop_after:
docs = itertools.islice(docs, stop_after) docs = itertools.islice(docs, stop_after)
if convert: if convert:
variable = "row" variable = "row"
if lines: if lines:
variable = "line" variable = "line"
elif text: elif text:
variable = "text" variable = "text"
fn = _compile_code(convert, imports, variable=variable) fn = _compile_code(convert, imports, variable=variable)
if lines: if lines:
docs = (fn(doc["line"]) for doc in docs) docs = (fn(doc["line"]) for doc in docs)
elif text: elif text:
# Special case: this is allowed to be an iterable # Special case: this is allowed to be an iterable
text_value = list(docs)[0]["text"] text_value = list(docs)[0]["text"]
fn_return = fn(text_value) fn_return = fn(text_value)
if isinstance(fn_return, dict): if isinstance(fn_return, dict):
docs = [fn_return] docs = [fn_return]
else:
try:
docs = iter(fn_return)
except TypeError:
raise click.ClickException(
"--convert must return dict or iterator"
)
else: else:
try: docs = (fn(doc) or doc for doc in docs)
docs = iter(fn_return)
except TypeError:
raise click.ClickException("--convert must return dict or iterator")
else:
docs = (fn(doc) or doc for doc in docs)
extra_kwargs = { extra_kwargs = {
"ignore": ignore, "ignore": ignore,
"replace": replace, "replace": replace,
"truncate": truncate, "truncate": truncate,
"analyze": analyze, "analyze": analyze,
} }
if not_null: if not_null:
extra_kwargs["not_null"] = set(not_null) extra_kwargs["not_null"] = set(not_null)
if default: if default:
extra_kwargs["defaults"] = dict(default) extra_kwargs["defaults"] = dict(default)
if upsert: if upsert:
extra_kwargs["upsert"] = upsert extra_kwargs["upsert"] = upsert
# docs should all be dictionaries # docs should all be dictionaries
docs = (verify_is_dict(doc) for doc in docs) docs = (verify_is_dict(doc) for doc in docs)
# Apply {"$base64": true, ...} decoding, if needed # Apply {"$base64": true, ...} decoding, if needed
docs = (decode_base64_values(doc) for doc in docs) docs = (decode_base64_values(doc) for doc in docs)
# For bulk_sql= we use cursor.executemany() instead # For bulk_sql= we use cursor.executemany() instead
if bulk_sql: if bulk_sql:
if batch_size: if batch_size:
doc_chunks = chunks(docs, batch_size) doc_chunks = chunks(docs, batch_size)
else: else:
doc_chunks = [docs] doc_chunks = [docs]
for doc_chunk in doc_chunks: for doc_chunk in doc_chunks:
with db.conn: with db.conn:
db.conn.cursor().executemany(bulk_sql, doc_chunk) db.conn.cursor().executemany(bulk_sql, doc_chunk)
return return
try: try:
db[table].insert_all( db[table].insert_all(
docs, pk=pk, batch_size=batch_size, alter=alter, **extra_kwargs docs, pk=pk, batch_size=batch_size, alter=alter, **extra_kwargs
)
except Exception as e:
if (
isinstance(e, OperationalError)
and e.args
and "has no column named" in e.args[0]
):
raise click.ClickException(
"{}\n\nTry using --alter to add additional columns".format(e.args[0])
) )
# If we can find sql= and parameters= arguments, show those except Exception as e:
variables = _find_variables(e.__traceback__, ["sql", "parameters"]) if (
if "sql" in variables and "parameters" in variables: isinstance(e, OperationalError)
raise click.ClickException( and e.args
"{}\n\nsql = {}\nparameters = {}".format( and "has no column named" in e.args[0]
str(e), variables["sql"], variables["parameters"] ):
raise click.ClickException(
"{}\n\nTry using --alter to add additional columns".format(
e.args[0]
)
) )
) # If we can find sql= and parameters= arguments, show those
else: variables = _find_variables(e.__traceback__, ["sql", "parameters"])
raise if "sql" in variables and "parameters" in variables:
if tracker is not None: raise click.ClickException(
db[table].transform(types=tracker.types) "{}\n\nsql = {}\nparameters = {}".format(
str(e), variables["sql"], variables["parameters"]
)
)
else:
raise
if tracker is not None:
db[table].transform(types=tracker.types)
# Clean up open file-like objects # Clean up open file-like objects
if sniff_buffer: if sniff_buffer:
sniff_buffer.close() sniff_buffer.close()
if decoded_buffer: if decoded_buffer:
decoded_buffer.close() decoded_buffer.close()
def _find_variables(tb, vars): def _find_variables(tb, vars):