Automated upgrades by Ruff

uvx --with 'ruff>=0.16.0' ruff check . --fix --unsafe-fixes
This commit is contained in:
Simon Willison 2026-07-25 14:20:31 -07:00
commit c0bcc04431
52 changed files with 820 additions and 895 deletions

View file

@ -1,7 +1,6 @@
from .utils import suggest_column_types
from .hookspecs import hookimpl
from .hookspecs import hookspec
from .db import Database
from .hookspecs import hookimpl, hookspec
from .migrations import Migrations
from .utils import suggest_column_types
__all__ = ["Database", "Migrations", "suggest_column_types", "hookimpl", "hookspec"]
__all__ = ["Database", "Migrations", "hookimpl", "hookspec", "suggest_column_types"]

View file

@ -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
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 @@ from sqlite_utils.db import (
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):
@ -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_upsert_implementation(
)
):
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_upsert_implementation(
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_upsert_implementation(
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_upsert_implementation(
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")
@ -1844,7 +1841,7 @@ def rename_table(path, table, new_name, ignore, load_extension):
except sqlite3.OperationalError as ex:
if not ignore:
raise click.ClickException(
'Table "{}" could not be renamed. {}'.format(table, str(ex))
f'Table "{table}" could not be renamed. {ex!s}'
)
@ -1874,10 +1871,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 +1916,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 +1948,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 +2172,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 +2191,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,10 +2368,10 @@ 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
@ -2384,7 +2379,7 @@ def search(
for c in column:
if c not in table_columns:
raise click.ClickException(
"Table '{}' has no column '{}".format(dbtable, c)
f"Table '{dbtable}' has no column '{c}"
)
sql = table_obj.search_sql(columns=column, order_by=order, limit=limit)
if show_sql:
@ -2412,7 +2407,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 +2474,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,
@ -2760,7 +2755,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()
@ -2858,12 +2853,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:
@ -3018,7 +3013,7 @@ def insert_files(
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}"
)
)
@ -3196,7 +3191,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:
@ -3299,7 +3294,7 @@ def convert(
""".format(
column=columns[0],
table=table,
where=" where {}".format(where) if where is not None else "",
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]))
@ -3339,9 +3334,7 @@ def convert(
)
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}"
)
@ -3459,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)
@ -3493,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()
@ -3583,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()
@ -3595,7 +3588,7 @@ def migrate(db_path, migrations, stop_before, list_, verbose):
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
f"{migration_set.name}:{name}" for name in names
)
unknown = [value for value in stop_before if value not in known_names]
if unknown:
@ -3652,7 +3645,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)
@ -3722,7 +3715,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)
@ -3740,7 +3733,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
@ -3768,14 +3761,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)
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("_"):
@ -3796,12 +3789,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)
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()

File diff suppressed because it is too large Load diff

View file

@ -1,8 +1,7 @@
import sqlite3
import click
from pluggy import HookimplMarker
from pluggy import HookspecMarker
from pluggy import HookimplMarker, HookspecMarker
hookspec = HookspecMarker("sqlite_utils")
hookimpl = HookimplMarker("sqlite_utils")

View file

@ -1,7 +1,7 @@
from collections.abc import Iterable
from dataclasses import dataclass
import datetime
from typing import Callable, cast, TYPE_CHECKING
from collections.abc import Callable, Iterable
from dataclasses import dataclass
from typing import TYPE_CHECKING, cast
if TYPE_CHECKING:
from sqlite_utils.db import Database, Table
@ -44,12 +44,10 @@ class Migrations:
"""
def inner(func: Callable) -> Callable:
migration_name = name or getattr(func, "__name__")
migration_name = name or func.__name__
if any(m.name == migration_name for m in self._migrations):
raise ValueError(
"Migration '{}' is already registered in set '{}'".format(
migration_name, self.name
)
f"Migration '{migration_name}' is already registered in set '{self.name}'"
)
self._migrations.append(
self._Migration(migration_name, func, transactional)

View file

@ -1,7 +1,7 @@
from typing import Dict, List, Union
import sys
import pluggy
import sys
from . import hookspecs
pm: pluggy.PluginManager = pluggy.PluginManager("sqlite_utils")
@ -17,13 +17,13 @@ def ensure_plugins_loaded() -> None:
_plugins_loaded = True
def get_plugins() -> List[Dict[str, Union[str, List[str]]]]:
def get_plugins() -> list[dict[str, str | list[str]]]:
ensure_plugins_loaded()
plugins: List[Dict[str, Union[str, List[str]]]] = []
plugins: list[dict[str, str | list[str]]] = []
plugin_to_distinfo = dict(pm.list_plugin_distinfo())
for plugin in pm.get_plugins():
hookcallers = pm.get_hookcallers(plugin) or []
plugin_info: Dict[str, Union[str, List[str]]] = {
plugin_info: dict[str, str | list[str]] = {
"name": plugin.__name__,
"hooks": [h.name for h in hookcallers],
}

View file

@ -1,9 +1,9 @@
from __future__ import annotations
from typing import Callable, Optional
import json
from collections.abc import Callable
from dateutil import parser
import json
IGNORE: object = object()
SET_NULL: object = object()
@ -13,8 +13,8 @@ def parsedate(
value: str,
dayfirst: bool = False,
yearfirst: bool = False,
errors: Optional[object] = None,
) -> Optional[str]:
errors: object | None = None,
) -> str | None:
"""
Parse a date and convert it to ISO date format: yyyy-mm-dd
\b
@ -44,8 +44,8 @@ def parsedatetime(
value: str,
dayfirst: bool = False,
yearfirst: bool = False,
errors: Optional[object] = None,
) -> Optional[str]:
errors: object | None = None,
) -> str | None:
"""
Parse a datetime and convert it to ISO datetime format: yyyy-mm-ddTHH:MM:SS
\b

View file

@ -9,20 +9,11 @@ import itertools
import json
import os
import sys
from collections.abc import Callable, Generator, Iterable, Iterator
from typing import (
TYPE_CHECKING,
Any,
BinaryIO,
Callable,
Dict,
Generator,
Iterable,
Iterator,
List,
Optional,
Set,
Tuple,
Type,
TYPE_CHECKING,
TypeVar,
Union,
cast,
@ -33,8 +24,8 @@ import click
from . import recipes
if TYPE_CHECKING:
import sqlite3 # noqa: F401
from sqlite3 import dbapi2 # noqa: F401
import sqlite3
from sqlite3 import dbapi2
OperationalError = dbapi2.OperationalError
else:
@ -44,7 +35,7 @@ else:
OperationalError = dbapi2.OperationalError
except ImportError:
import sqlite3 # noqa: F401
from sqlite3 import dbapi2 # noqa: F401
from sqlite3 import dbapi2
OperationalError = dbapi2.OperationalError
@ -61,8 +52,8 @@ SPATIALITE_PATHS = (
ORIGINAL_CSV_FIELD_SIZE_LIMIT = csv.field_size_limit()
# Type alias for row dictionaries - values can be various SQLite-compatible types
RowValue = Union[None, int, float, str, bytes, bool, List[str]]
Row = Dict[str, RowValue]
RowValue = Union[None, int, float, str, bytes, bool, list[str]]
Row = dict[str, RowValue]
T = TypeVar("T")
@ -103,7 +94,7 @@ def maximize_csv_field_size_limit() -> None:
field_size_limit = int(field_size_limit / 10)
def find_spatialite() -> Optional[str]:
def find_spatialite() -> str | None:
"""
The ``find_spatialite()`` function searches for the `SpatiaLite <https://www.gaia-gis.it/fossil/libspatialite/index>`__
SQLite extension in some common places. It returns a string path to the location, or ``None`` if SpatiaLite was not found.
@ -132,9 +123,9 @@ def find_spatialite() -> Optional[str]:
def suggest_column_types(
records: Iterable[Dict[str, Any]],
) -> Dict[str, type]:
all_column_types: Dict[str, Set[type]] = {}
records: Iterable[dict[str, Any]],
) -> dict[str, type]:
all_column_types: dict[str, set[type]] = {}
for record in records:
for key, value in record.items():
all_column_types.setdefault(key, set()).add(type(value))
@ -142,9 +133,9 @@ def suggest_column_types(
def types_for_column_types(
all_column_types: Dict[str, Set[type]],
) -> Dict[str, type]:
column_types: Dict[str, type] = {}
all_column_types: dict[str, set[type]],
) -> dict[str, type]:
column_types: dict[str, type] = {}
for key, types in all_column_types.items():
# Ignore null values if at least one other type present:
if len(types) > 1:
@ -153,7 +144,7 @@ def types_for_column_types(
if {None.__class__} == types:
t = str
elif len(types) == 1:
t = list(types)[0]
t = next(iter(types))
# But if it's a subclass of list / tuple / dict, use str
# instead as we will be storing it as JSON in the table
for superclass in (list, tuple, dict):
@ -190,7 +181,7 @@ def column_affinity(column_type: str) -> type:
return float
def decode_base64_values(doc: Dict[str, Any]) -> Dict[str, Any]:
def decode_base64_values(doc: dict[str, Any]) -> dict[str, Any]:
# Looks for '{"$base64": true..., "encoded": ...}' values and decodes them
to_fix = [
k
@ -263,9 +254,9 @@ class RowError(Exception):
def _extra_key_strategy(
reader: Iterable[Dict[Optional[str], object]],
ignore_extras: Optional[bool] = False,
extras_key: Optional[str] = None,
reader: Iterable[dict[str | None, object]],
ignore_extras: bool | None = False,
extras_key: str | None = None,
) -> Iterable[Row]:
# Logic for handling CSV rows with more values than there are headings
for row in reader:
@ -280,7 +271,7 @@ def _extra_key_strategy(
elif not extras_key:
extras = row.pop(None)
raise RowError(
"Row {} contained these extra values: {}".format(row, extras)
f"Row {row} contained these extra values: {extras}"
)
else:
extras_value = row.pop(None)
@ -291,12 +282,12 @@ def _extra_key_strategy(
def rows_from_file(
fp: BinaryIO,
format: Optional[Format] = None,
dialect: Optional[Type[csv.Dialect]] = None,
encoding: Optional[str] = None,
ignore_extras: Optional[bool] = False,
extras_key: Optional[str] = None,
) -> Tuple[Iterable[Row], Format]:
format: Format | None = None,
dialect: type[csv.Dialect] | None = None,
encoding: str | None = None,
ignore_extras: bool | None = False,
extras_key: str | None = None,
) -> tuple[Iterable[Row], Format]:
"""
Load a sequence of dictionaries from a file-like object containing one of four different formats.
@ -363,7 +354,7 @@ def rows_from_file(
)
return (
_extra_key_strategy(
cast(Iterable[Dict[Optional[str], object]], rows),
cast(Iterable[dict[str | None, object]], rows),
ignore_extras,
extras_key,
),
@ -379,7 +370,7 @@ def rows_from_file(
raise TypeError(
"rows_from_file() requires a file-like object that supports peek(), such as io.BytesIO"
)
if first_bytes.startswith(b"[") or first_bytes.startswith(b"{"):
if first_bytes.startswith((b"[", b"{")):
# TODO: Detect newline-JSON
return rows_from_file(buffered, format=Format.JSON)
else:
@ -393,7 +384,7 @@ def rows_from_file(
detected_format = Format.TSV if dialect.delimiter == "\t" else Format.CSV
return (
_extra_key_strategy(
cast(Iterable[Dict[Optional[str], object]], rows),
cast(Iterable[dict[str | None, object]], rows),
ignore_extras,
extras_key,
),
@ -425,9 +416,9 @@ class TypeTracker:
"""
def __init__(self) -> None:
self.trackers: Dict[str, "ValueTracker"] = {}
self.trackers: dict[str, ValueTracker] = {}
def wrap(self, iterator: Iterable[Dict[str, Any]]) -> Iterable[Dict[str, Any]]:
def wrap(self, iterator: Iterable[dict[str, Any]]) -> Iterable[dict[str, Any]]:
"""
Use this to loop through an existing iterator, tracking the column types
as part of the iteration.
@ -441,7 +432,7 @@ class TypeTracker:
yield row
@property
def types(self) -> Dict[str, str]:
def types(self) -> dict[str, str]:
"""
A dictionary mapping column names to their detected types. This can be passed
to the ``db[table_name].transform(types=tracker.types)`` method.
@ -450,16 +441,16 @@ class TypeTracker:
class ValueTracker:
couldbe: Dict[str, Callable[[object], bool]]
couldbe: dict[str, Callable[[object], bool]]
def __init__(self) -> None:
self.couldbe = {key: getattr(self, "test_" + key) for key in self.get_tests()}
@classmethod
def get_tests(cls) -> List[str]:
def get_tests(cls) -> list[str]:
return [
key.split("test_")[-1]
for key in cls.__dict__.keys()
for key in cls.__dict__
if key.startswith("test_")
]
@ -492,7 +483,7 @@ class ValueTracker:
def evaluate(self, value: object) -> None:
if not value or not self.couldbe:
return
not_these: List[str] = []
not_these: list[str] = []
for name, test in self.couldbe.items():
if not test(value):
not_these.append(name)
@ -524,7 +515,7 @@ def progressbar(*args: Iterable[T], **kwargs: Any) -> Generator[Any, None, None]
def _compile_code(
code: str, imports: Iterable[str], variable: str = "value"
) -> Callable[..., Any]:
globals_dict: Dict[str, Any] = {"r": recipes, "recipes": recipes}
globals_dict: dict[str, Any] = {"r": recipes, "recipes": recipes}
# Handle imports first so they're available for all approaches
for import_ in imports:
globals_dict[import_.split(".")[0]] = __import__(import_)
@ -549,13 +540,13 @@ def _compile_code(
body_variants = [code]
# If single line and no 'return', try adding the return
if "\n" not in code and not code.strip().startswith("return "):
body_variants.insert(0, "return {}".format(code))
body_variants.insert(0, f"return {code}")
code_o = None
for variant in body_variants:
new_code = ["def fn({}):".format(variable)]
new_code = [f"def fn({variable}):"]
for line in variant.split("\n"):
new_code.append(" {}".format(line))
new_code.append(f" {line}")
try:
code_o = compile("\n".join(new_code), "<string>", "exec")
break
@ -582,7 +573,7 @@ def chunks(sequence: Iterable[T], size: int) -> Iterable[Iterable[T]]:
yield itertools.chain([item], itertools.islice(iterator, size - 1))
def hash_record(record: Dict[str, Any], keys: Optional[Iterable[str]] = None) -> str:
def hash_record(record: dict[str, Any], keys: Iterable[str] | None = None) -> str:
"""
``record`` should be a Python dictionary. Returns a sha1 hash of the
keys and values in that record.
@ -603,7 +594,7 @@ def hash_record(record: Dict[str, Any], keys: Optional[Iterable[str]] = None) ->
:param record: Record to generate a hash for
:param keys: Subset of keys to use for that hash
"""
to_hash: Dict[str, Any] = record
to_hash: dict[str, Any] = record
if keys is not None:
to_hash = {key: record[key] for key in keys}
return hashlib.sha1(
@ -613,7 +604,7 @@ def hash_record(record: Dict[str, Any], keys: Optional[Iterable[str]] = None) ->
).hexdigest()
def dedupe_keys(keys: Iterable[str]) -> List[str]:
def dedupe_keys(keys: Iterable[str]) -> list[str]:
"""
Rename duplicates in a list of column names so every name is unique,
by appending ``_2``, ``_3``... to later occurrences - skipping any
@ -636,7 +627,7 @@ def dedupe_keys(keys: Iterable[str]) -> List[str]:
new_key = key
suffix = 2
while new_key in seen or new_key in taken:
new_key = "{}_{}".format(key, suffix)
new_key = f"{key}_{suffix}"
suffix += 1
key = new_key
seen.add(key)
@ -644,7 +635,7 @@ def dedupe_keys(keys: Iterable[str]) -> List[str]:
return result
def _flatten(d: Dict[str, Any]) -> Generator[Tuple[str, Any], None, None]:
def _flatten(d: dict[str, Any]) -> Generator[tuple[str, Any], None, None]:
for key, value in d.items():
if isinstance(value, dict):
for key2, value2 in _flatten(value):
@ -653,7 +644,7 @@ def _flatten(d: Dict[str, Any]) -> Generator[Tuple[str, Any], None, None]:
yield key, value
def flatten(row: Dict[str, Any]) -> Dict[str, Any]:
def flatten(row: dict[str, Any]) -> dict[str, Any]:
"""
Turn a nested dict e.g. ``{"a": {"b": 1}}`` into a flat dict: ``{"a_b": 1}``