mirror of
https://github.com/simonw/sqlite-utils.git
synced 2026-09-28 21:04:13 +02:00
sqlite_utils.utils.flatten() function, closes #500
This commit is contained in:
parent
d792dad1cf
commit
34e75ed0dd
5 changed files with 40 additions and 23 deletions
|
|
@ -101,3 +101,10 @@ sqlite_utils.utils.chunks
|
||||||
-------------------------
|
-------------------------
|
||||||
|
|
||||||
.. autofunction:: sqlite_utils.utils.chunks
|
.. autofunction:: sqlite_utils.utils.chunks
|
||||||
|
|
||||||
|
.. _reference_utils_flatten:
|
||||||
|
|
||||||
|
sqlite_utils.utils.flatten
|
||||||
|
--------------------------
|
||||||
|
|
||||||
|
.. autofunction:: sqlite_utils.utils.flatten
|
||||||
|
|
|
||||||
|
|
@ -24,6 +24,7 @@ from .utils import (
|
||||||
chunks,
|
chunks,
|
||||||
file_progress,
|
file_progress,
|
||||||
find_spatialite,
|
find_spatialite,
|
||||||
|
flatten as _flatten,
|
||||||
sqlite3,
|
sqlite3,
|
||||||
decode_base64_values,
|
decode_base64_values,
|
||||||
progressbar,
|
progressbar,
|
||||||
|
|
@ -997,7 +998,7 @@ def insert_upsert_implementation(
|
||||||
"Invalid JSON - use --csv for CSV or --tsv for TSV files"
|
"Invalid JSON - use --csv for CSV or --tsv for TSV files"
|
||||||
)
|
)
|
||||||
if flatten:
|
if flatten:
|
||||||
docs = (dict(_flatten(doc)) for doc in docs)
|
docs = (_flatten(doc) for doc in docs)
|
||||||
|
|
||||||
if convert:
|
if convert:
|
||||||
variable = "row"
|
variable = "row"
|
||||||
|
|
@ -1079,15 +1080,6 @@ def insert_upsert_implementation(
|
||||||
db[table].transform(types=tracker.types)
|
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):
|
def _find_variables(tb, vars):
|
||||||
to_find = list(vars)
|
to_find = list(vars)
|
||||||
found = {}
|
found = {}
|
||||||
|
|
@ -1845,7 +1837,7 @@ def memory(
|
||||||
tracker = TypeTracker()
|
tracker = TypeTracker()
|
||||||
rows = tracker.wrap(rows)
|
rows = tracker.wrap(rows)
|
||||||
if flatten:
|
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)
|
db[csv_table].insert_all(rows, alter=True)
|
||||||
if tracker is not None:
|
if tracker is not None:
|
||||||
db[csv_table].transform(types=tracker.types)
|
db[csv_table].transform(types=tracker.types)
|
||||||
|
|
|
||||||
|
|
@ -513,3 +513,21 @@ def hash_record(record: Dict, keys: Optional[Iterable[str]] = None):
|
||||||
"utf8"
|
"utf8"
|
||||||
)
|
)
|
||||||
).hexdigest()
|
).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))
|
||||||
|
|
|
||||||
|
|
@ -2176,18 +2176,6 @@ def test_upsert_detect_types(tmpdir, option):
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
"input,expected",
|
|
||||||
(
|
|
||||||
({"foo": {"bar": 1}}, {"foo_bar": 1}),
|
|
||||||
({"foo": {"bar": [1, 2, {"baz": 3}]}}, {"foo_bar": [1, 2, {"baz": 3}]}),
|
|
||||||
({"foo": {"bar": 1, "baz": {"three": 3}}}, {"foo_bar": 1, "foo_baz_three": 3}),
|
|
||||||
),
|
|
||||||
)
|
|
||||||
def test_flatten_helper(input, expected):
|
|
||||||
assert dict(cli._flatten(input)) == expected
|
|
||||||
|
|
||||||
|
|
||||||
def test_integer_overflow_error(tmpdir):
|
def test_integer_overflow_error(tmpdir):
|
||||||
db_path = str(tmpdir / "test.db")
|
db_path = str(tmpdir / "test.db")
|
||||||
result = CliRunner().invoke(
|
result = CliRunner().invoke(
|
||||||
|
|
|
||||||
|
|
@ -71,3 +71,15 @@ def test_maximize_csv_field_size_limit():
|
||||||
assert len(rows_list2) == 1
|
assert len(rows_list2) == 1
|
||||||
assert rows_list2[0]["id"] == "1"
|
assert rows_list2[0]["id"] == "1"
|
||||||
assert rows_list2[0]["text"] == long_value
|
assert rows_list2[0]["text"] == long_value
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"input,expected",
|
||||||
|
(
|
||||||
|
({"foo": {"bar": 1}}, {"foo_bar": 1}),
|
||||||
|
({"foo": {"bar": [1, 2, {"baz": 3}]}}, {"foo_bar": [1, 2, {"baz": 3}]}),
|
||||||
|
({"foo": {"bar": 1, "baz": {"three": 3}}}, {"foo_bar": 1, "foo_baz_three": 3}),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
def test_flatten(input, expected):
|
||||||
|
assert utils.flatten(input) == expected
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue