sqlite_utils.utils.flatten() function, closes #500

This commit is contained in:
Simon Willison 2022-10-18 11:00:25 -07:00
commit 34e75ed0dd
5 changed files with 40 additions and 23 deletions

View file

@ -24,6 +24,7 @@ from .utils import (
chunks,
file_progress,
find_spatialite,
flatten as _flatten,
sqlite3,
decode_base64_values,
progressbar,
@ -997,7 +998,7 @@ def insert_upsert_implementation(
"Invalid JSON - use --csv for CSV or --tsv for TSV files"
)
if flatten:
docs = (dict(_flatten(doc)) for doc in docs)
docs = (_flatten(doc) for doc in docs)
if convert:
variable = "row"
@ -1079,15 +1080,6 @@ def insert_upsert_implementation(
db[table].transform(types=tracker.types)
def _flatten(d):
for key, value in d.items():
if isinstance(value, dict):
for key2, value2 in _flatten(value):
yield key + "_" + key2, value2
else:
yield key, value
def _find_variables(tb, vars):
to_find = list(vars)
found = {}
@ -1845,7 +1837,7 @@ def memory(
tracker = TypeTracker()
rows = tracker.wrap(rows)
if flatten:
rows = (dict(_flatten(row)) for row in rows)
rows = (_flatten(row) for row in rows)
db[csv_table].insert_all(rows, alter=True)
if tracker is not None:
db[csv_table].transform(types=tracker.types)

View file

@ -513,3 +513,21 @@ def hash_record(record: Dict, keys: Optional[Iterable[str]] = None):
"utf8"
)
).hexdigest()
def _flatten(d):
for key, value in d.items():
if isinstance(value, dict):
for key2, value2 in _flatten(value):
yield key + "_" + key2, value2
else:
yield key, value
def flatten(row: dict) -> dict:
"""
Turn a nested dict e.g. ``{"a": {"b": 1}}`` into a flat dict: ``{"a_b": 1}``
:param row: A Python dictionary, optionally with nested dictionaries
"""
return dict(_flatten(row))