From b8f0cb5a6c13b158d5c02e13f9ab4059f60a0b4c Mon Sep 17 00:00:00 2001 From: ikatyal2110 <134458944+ikatyal2110@users.noreply.github.com> Date: Thu, 6 Aug 2026 21:03:31 +0530 Subject: [PATCH] Fix convert --dry-run using bracket quoting instead of quote_identifier() --- sqlite_utils/cli.py | 219 ++++++++++++++++++++++---------------------- 1 file changed, 109 insertions(+), 110 deletions(-) diff --git a/sqlite_utils/cli.py b/sqlite_utils/cli.py index 7fab72b18..0d6b035c9 100644 --- a/sqlite_utils/cli.py +++ b/sqlite_utils/cli.py @@ -1,17 +1,30 @@ import base64 +import csv as csv_std import difflib -from typing import Any -import click -from click_default_group import DefaultGroup -from datetime import datetime, timezone import hashlib +import inspect +import io +import itertools +import json +import os import pathlib +import pdb # noqa: T100 +import sys +import textwrap +from datetime import datetime, timezone from runpy import run_module +from typing import Any + +import click +import tabulate +from click_default_group import DefaultGroup + import sqlite_utils +from sqlite_utils import recipes from sqlite_utils.db import ( + DEFAULT, AlterError, BadMultiValues, - DEFAULT, DescIndex, InvalidColumns, NoTable, @@ -19,36 +32,28 @@ PrimaryKeyRequired, quote_identifier, ) -from sqlite_utils.plugins import ensure_plugins_loaded, pm, get_plugins +from sqlite_utils.plugins import ensure_plugins_loaded, get_plugins, pm from sqlite_utils.utils import maximize_csv_field_size_limit -from sqlite_utils import recipes -import textwrap -import inspect -import io -import itertools -import json -import os -import pdb -import sys -import csv as csv_std -import tabulate + from .utils import ( + Format, OperationalError, + TypeTracker, _compile_code, chunks, + decode_base64_values, dedupe_keys, file_progress, find_spatialite, - flatten as _flatten, - sqlite3, - decode_base64_values, progressbar, rows_from_file, - Format, - TypeTracker, + sqlite3, +) +from .utils import ( + flatten as _flatten, ) -CONTEXT_SETTINGS = dict(help_option_names=["-h", "--help"]) +CONTEXT_SETTINGS = {"help_option_names": ["-h", "--help"]} def _register_db_for_cleanup(db): @@ -67,7 +72,7 @@ def _close_databases(ctx): for db in ctx.meta.get("_databases_to_close", []): try: db.close() - except Exception: + except sqlite3.Error: pass @@ -174,7 +179,6 @@ def functions_option(fn): @click.version_option() def cli(): "Commands for interacting with a SQLite database" - pass @cli.command() @@ -891,7 +895,7 @@ def enable_counts(path, tables, load_extension): # Check all tables exist bad_tables = [table for table in tables if not db[table].exists()] if bad_tables: - raise click.ClickException("Invalid tables: {}".format(bad_tables)) + raise click.ClickException(f"Invalid tables: {bad_tables}") for table in tables: db.table(table).enable_counts() @@ -1140,9 +1144,7 @@ def _insert_docs(docs, tracker=None): ) ): raise click.ClickException( - "{}\n\nTry using --alter to add additional columns".format( - e.args[0] - ) + f"{e.args[0]}\n\nTry using --alter to add additional columns" ) # If we can find sql= and parameters= arguments, show those variables = _find_variables(e.__traceback__, ["sql", "parameters"]) @@ -1240,7 +1242,7 @@ def _insert_docs(docs, tracker=None): reader = csv_std.reader(decoded, **csv_reader_args) # type: ignore first_row = next(reader) if no_headers: - headers = ["untitled_{}".format(i + 1) for i in range(len(first_row))] + headers = [f"untitled_{i + 1}" for i in range(len(first_row))] reader = itertools.chain([first_row], reader) else: headers = first_row @@ -1269,9 +1271,7 @@ def _insert_docs(docs, tracker=None): docs = [docs] except json.decoder.JSONDecodeError as ex: raise click.ClickException( - "Invalid JSON - use --csv for CSV or --tsv for TSV files\n\nJSON error: {}".format( - ex - ) + f"Invalid JSON - use --csv for CSV or --tsv for TSV files\n\nJSON error: {ex}" ) if flatten: docs = (_flatten(doc) for doc in docs) @@ -1290,7 +1290,7 @@ def _insert_docs(docs, tracker=None): docs = (fn(doc["line"]) for doc in docs) elif text: # Special case: this is allowed to be an iterable - text_value = list(docs)[0]["text"] + text_value = next(iter(docs))["text"] fn_return = fn(text_value) if isinstance(fn_return, dict): docs = [fn_return] @@ -1774,17 +1774,14 @@ def create_table( ctype = columns.pop(0) if ctype.upper() not in VALID_COLUMN_TYPES: raise click.ClickException( - "column types must be one of {}".format(VALID_COLUMN_TYPES) + f"column types must be one of {VALID_COLUMN_TYPES}" ) coltypes[name] = ctype.upper() # Does table already exist? - if table in db.table_names(): - if not ignore and not replace and not transform: - raise click.ClickException( - 'Table "{}" already exists. Use --replace to delete and replace it.'.format( - table - ) - ) + if table in db.table_names() and not ignore and not replace and not transform: + raise click.ClickException( + f'Table "{table}" already exists. Use --replace to delete and replace it.' + ) db.table(table).create( coltypes, pk=pks[0] if len(pks) == 1 else pks, @@ -1819,7 +1816,7 @@ def duplicate(path, table, new_table, ignore, load_extension): db.table(table).duplicate(new_table) except NoTable: if not ignore: - raise click.ClickException('Table "{}" does not exist'.format(table)) + raise click.ClickException(f'Table "{table}" does not exist') @cli.command(name="rename-table") @@ -1843,9 +1840,7 @@ def rename_table(path, table, new_name, ignore, load_extension): db.rename_table(table, new_name) except sqlite3.OperationalError as ex: if not ignore: - raise click.ClickException( - 'Table "{}" could not be renamed. {}'.format(table, str(ex)) - ) + raise click.ClickException(f'Table "{table}" could not be renamed. {ex!s}') @cli.command(name="drop-table") @@ -1874,10 +1869,10 @@ def drop_table(path, table, ignore, load_extension): # A view exists with this name if not ignore: raise click.ClickException( - '"{}" is a view, not a table - use drop-view to drop it'.format(table) + f'"{table}" is a view, not a table - use drop-view to drop it' ) except OperationalError: - raise click.ClickException('Table "{}" does not exist'.format(table)) + raise click.ClickException(f'Table "{table}" does not exist') @cli.command(name="create-view") @@ -1919,9 +1914,7 @@ def create_view(path, view, select, ignore, replace, load_extension): db.view(view).drop() else: raise click.ClickException( - 'View "{}" already exists. Use --replace to delete and replace it.'.format( - view - ) + f'View "{view}" already exists. Use --replace to delete and replace it.' ) db.create_view(view, select) @@ -1953,9 +1946,9 @@ def drop_view(path, view, ignore, load_extension): return if view in db.table_names(): raise click.ClickException( - '"{}" is a table, not a view - use drop-table to drop it'.format(view) + f'"{view}" is a table, not a view - use drop-table to drop it' ) - raise click.ClickException('View "{}" does not exist'.format(view)) + raise click.ClickException(f'View "{view}" does not exist') @cli.command() @@ -2177,7 +2170,7 @@ def memory( file_path = pathlib.Path(path) stem = file_path.stem if stem_counts.get(stem): - file_table = "{}_{}".format(stem, stem_counts[stem]) + file_table = f"{stem}_{stem_counts[stem]}" else: file_table = stem stem_counts[stem] = stem_counts.get(stem, 1) + 1 @@ -2196,14 +2189,14 @@ def memory( if tracker is not None and db.table(file_table).exists(): db.table(file_table).transform(types=tracker.types) # Add convenient t / t1 / t2 views - view_names = ["t{}".format(i + 1)] + view_names = [f"t{i + 1}"] if i == 0: view_names.append("t") for view_name in view_names: if not db[view_name].exists(): db.create_view( view_name, - "select * from {}".format(quote_identifier(file_table)), + f"select * from {quote_identifier(file_table)}", ) finally: if should_close_fp and fp: @@ -2373,19 +2366,17 @@ def search( # Check table exists table_obj = db.table(dbtable) if not table_obj.exists(): - raise click.ClickException("Table '{}' does not exist".format(dbtable)) + raise click.ClickException(f"Table '{dbtable}' does not exist") if not table_obj.detect_fts(): raise click.ClickException( - "Table '{}' is not configured for full-text search".format(dbtable) + f"Table '{dbtable}' is not configured for full-text search" ) if column: # Check they all exist table_columns = table_obj.columns_dict for c in column: if c not in table_columns: - raise click.ClickException( - "Table '{}' has no column '{}".format(dbtable, c) - ) + raise click.ClickException(f"Table '{dbtable}' has no column '{c}") sql = table_obj.search_sql(columns=column, order_by=order, limit=limit) if show_sql: click.echo(sql) @@ -2412,7 +2403,7 @@ def search( except click.ClickException as e: if "malformed MATCH expression" in str(e) or "unterminated string" in str(e): raise click.ClickException( - "{}\n\nTry running this again with the --quote option".format(str(e)) + f"{e!s}\n\nTry running this again with the --quote option" ) else: raise @@ -2479,15 +2470,15 @@ def rows( columns = "*" if column: columns = ", ".join(quote_identifier(c) for c in column) - sql = "select {} from {}".format(columns, quote_identifier(dbtable)) + sql = f"select {columns} from {quote_identifier(dbtable)}" if where: sql += " where " + where if order: sql += " order by " + order if limit: - sql += " limit {}".format(limit) + sql += f" limit {limit}" if offset: - sql += " offset {}".format(offset) + sql += f" offset {offset}" ctx.invoke( query, path=path, @@ -2718,6 +2709,11 @@ def schema( multiple=True, help="Drop foreign key constraint for this column", ) +@click.option( + "--strict/--no-strict", + default=None, + help="Enable or disable STRICT mode (default: preserve current mode)", +) @click.option("--sql", is_flag=True, help="Output SQL without executing it") @load_extension_option def transform( @@ -2735,6 +2731,7 @@ def transform( default_none, add_foreign_keys, drop_foreign_keys, + strict, sql, load_extension, ): @@ -2754,7 +2751,7 @@ def transform( for column, ctype in type: if ctype.upper() not in VALID_COLUMN_TYPES: raise click.ClickException( - "column types must be one of {}".format(VALID_COLUMN_TYPES) + f"column types must be one of {VALID_COLUMN_TYPES}" ) types[column] = ctype.upper() @@ -2796,6 +2793,7 @@ def transform( defaults=default_dict, drop_foreign_keys=drop_foreign_keys_value, add_foreign_keys=add_foreign_keys_value, + strict=strict, ): click.echo(line) else: @@ -2809,6 +2807,7 @@ def transform( defaults=default_dict, drop_foreign_keys=drop_foreign_keys_value, add_foreign_keys=add_foreign_keys_value, + strict=strict, ) @@ -2850,12 +2849,12 @@ def extract( db = sqlite_utils.Database(path) _register_db_for_cleanup(db) _load_extensions(db, load_extension) - kwargs: dict[str, Any] = dict( - columns=columns, - table=other_table, - fk_column=fk_column, - rename=dict(rename), - ) + kwargs: dict[str, Any] = { + "columns": columns, + "table": other_table, + "fk_column": fk_column, + "rename": dict(rename), + } try: db.table(table).extract(**kwargs) except (NoTable, InvalidColumns) as e: @@ -2950,7 +2949,7 @@ def yield_paths_and_relative_paths(): with progressbar(paths_and_relative_paths, silent=silent) as bar: def to_insert(): - for path, relative_path in bar: + for file_path, relative_path in bar: row = {} # content_text is special case as it considers 'encoding' @@ -2962,19 +2961,21 @@ def _content_text(p): raise UnicodeDecodeErrorForPath(e, resolved) lookups = dict(FILE_COLUMNS, content_text=_content_text) - if path == "-": + if file_path == "-": stdin_data = sys.stdin.buffer.read() # We only support a subset of columns for this case lookups = { "name": lambda p: name or "-", "path": lambda p: name or "-", - "content": lambda p: stdin_data, - "content_text": lambda p: stdin_data.decode( + "content": lambda p, data=stdin_data: data, + "content_text": lambda p, data=stdin_data: data.decode( encoding or "utf-8" ), - "sha256": lambda p: hashlib.sha256(stdin_data).hexdigest(), - "md5": lambda p: hashlib.md5(stdin_data).hexdigest(), - "size": lambda p: len(stdin_data), + "sha256": lambda p, data=stdin_data: hashlib.sha256( + data + ).hexdigest(), + "md5": lambda p, data=stdin_data: hashlib.md5(data).hexdigest(), + "size": lambda p, data=stdin_data: len(data), } for coldef in column: if ":" in coldef: @@ -2982,7 +2983,7 @@ def _content_text(p): else: colname, coltype = coldef, coldef try: - value = lookups[coltype](path) + value = lookups[coltype](file_path) row[colname] = value except KeyError: raise click.ClickException( @@ -3010,7 +3011,7 @@ def _content_text(p): except UnicodeDecodeErrorForPath as e: raise click.ClickException( UNICODE_ERROR.format( - "Could not read file '{}' as text\n\n{}".format(e.path, e.exception) + f"Could not read file '{e.path}' as text\n\n{e.exception}" ) ) @@ -3188,7 +3189,7 @@ def _generate_convert_help(): for name in recipe_names: fn = getattr(recipes, name) doc = textwrap.dedent(fn.__doc__.rstrip()).replace("\b\n", "") - help += "\n\nr.{}{}\n\n\b{}".format(name, str(inspect.signature(fn)), doc) + help += f"\n\nr.{name}{inspect.signature(fn)!s}\n\n\b{doc}" help += "\n\n" help += textwrap.dedent(""" You can use these recipes like so: @@ -3283,15 +3284,17 @@ def preview(v): return fn(v) if v else v db.conn.create_function("preview_transform", 1, preview) + col_quoted = quote_identifier(columns[0]) + tbl_quoted = quote_identifier(table) sql = """ select - [{column}] as value, - preview_transform([{column}]) as preview - from [{table}]{where} limit 10 + {col} as value, + preview_transform({col}) as preview + from {tbl}{where} limit 10 """.format( - column=columns[0], - table=table, - where=" where {}".format(where) if where is not None else "", + col=col_quoted, + tbl=tbl_quoted, + where=f" where {where}" if where is not None else "", ) for row in db.conn.execute(sql, where_args).fetchall(): click.echo(str(row[0])) @@ -3311,7 +3314,7 @@ def preview(v): def wrapped_fn(value): try: return fn_(value) - except Exception as ex: + except Exception as ex: # noqa: BLE001 print("\nException raised, dropping into pdb...:", ex) pdb.post_mortem(ex.__traceback__) sys.exit(1) @@ -3331,9 +3334,7 @@ def wrapped_fn(value): ) except BadMultiValues as e: raise click.ClickException( - "When using --multi code must return a Python dictionary - returned: {}".format( - repr(e.values) - ) + f"When using --multi code must return a Python dictionary - returned: {e.values!r}" ) @@ -3451,7 +3452,7 @@ def create_spatial_index(db_path, table, column_name, load_extension): def _find_migration_files(migrations): if not migrations: - migrations = [pathlib.Path(".").resolve()] + migrations = [pathlib.Path.cwd()] files = set() for path_str in migrations: path = pathlib.Path(path_str) @@ -3476,7 +3477,7 @@ def _load_migration_sets(files): "__file__": str(filepath), "__name__": "__sqlite_utils_migration__", } - exec(code, namespace) + exec(code, namespace) # noqa: S102 migration_sets.extend( obj for obj in namespace.values() if _compatible_migration_set(obj) ) @@ -3485,17 +3486,17 @@ def _load_migration_sets(files): def _display_migration_list(db, migration_sets): for migration_set in migration_sets: - click.echo("Migrations for: {}".format(migration_set.name)) + click.echo(f"Migrations for: {migration_set.name}") click.echo() click.echo(" Applied:") for migration in migration_set.applied(db): - click.echo(" {} - {}".format(migration.name, migration.applied_at)) + click.echo(f" {migration.name} - {migration.applied_at}") click.echo() click.echo(" Pending:") output = False for migration in migration_set.pending(db): output = True - click.echo(" {}".format(migration.name)) + click.echo(f" {migration.name}") if not output: click.echo(" (none)") click.echo() @@ -3575,7 +3576,7 @@ def migrate(db_path, migrations, stop_before, list_, verbose): prev_schema = db.schema if verbose: - click.echo("Migrating {}".format(db_path)) + click.echo(f"Migrating {db_path}") click.echo("\nSchema before:\n") click.echo(textwrap.indent(prev_schema, " ") or " (empty)") click.echo() @@ -3586,9 +3587,7 @@ def migrate(db_path, migrations, stop_before, list_, verbose): names = {m.name for m in migration_set.pending(db)} names.update(m.name for m in migration_set.applied(db)) known_names.update(names) - known_names.update( - "{}:{}".format(migration_set.name, name) for name in names - ) + known_names.update(f"{migration_set.name}:{name}" for name in names) unknown = [value for value in stop_before if value not in known_names] if unknown: raise click.ClickException( @@ -3644,7 +3643,7 @@ def _render_common(title, values): return "" lines = [title] for value, count in values: - lines.append(" {}: {}".format(count, value)) + lines.append(f" {count}: {value}") return "\n".join(lines) @@ -3714,7 +3713,7 @@ def maybe_json(value): if not isinstance(value, str): return value stripped = value.strip() - if not (stripped.startswith("{") or stripped.startswith("[")): + if not (stripped.startswith(("{", "["))): return value try: return json.loads(stripped) @@ -3732,7 +3731,7 @@ def json_binary(value): def verify_is_dict(doc): if not isinstance(doc, dict): raise click.ClickException( - "Rows must all be dictionaries, got: {}".format(repr(doc)[:1000]) + f"Rows must all be dictionaries, got: {repr(doc)[:1000]}" ) return doc @@ -3760,14 +3759,14 @@ def _register_functions(db, functions): try: functions = pathlib.Path(functions).read_text() except FileNotFoundError: - raise click.ClickException("File not found: {}".format(functions)) + raise click.ClickException(f"File not found: {functions}") sqlite3.enable_callback_tracebacks(True) globals = {} try: - exec(functions, globals) + exec(functions, globals) # noqa: S102 except SyntaxError as ex: - raise click.ClickException("Error in functions definition: {}".format(ex)) + raise click.ClickException(f"Error in functions definition: {ex}") # Register all callables in the locals dict: for name, value in globals.items(): if callable(value) and not name.startswith("_"): @@ -3788,12 +3787,12 @@ def _rows_from_code(code): try: code = pathlib.Path(code).read_text() except FileNotFoundError: - raise click.ClickException("File not found: {}".format(code)) + raise click.ClickException(f"File not found: {code}") namespace = {} try: - exec(code, namespace) + exec(code, namespace) # noqa: S102 except SyntaxError as ex: - raise click.ClickException("Error in --code: {}".format(ex)) + raise click.ClickException(f"Error in --code: {ex}") rows = namespace.get("rows") if callable(rows): rows = rows()