mirror of
https://github.com/simonw/sqlite-utils.git
synced 2026-10-08 02:17:08 +02:00
Extracted detect_column_types as suggest_column_types, refs #81
This commit is contained in:
parent
f7289174e6
commit
5ecf3ffdea
3 changed files with 47 additions and 27 deletions
|
|
@ -1,4 +1,4 @@
|
||||||
from .utils import sqlite3, OperationalError
|
from .utils import sqlite3, OperationalError, suggest_column_types
|
||||||
from collections import namedtuple
|
from collections import namedtuple
|
||||||
import datetime
|
import datetime
|
||||||
import hashlib
|
import hashlib
|
||||||
|
|
@ -820,30 +820,6 @@ class Table(Queryable):
|
||||||
)
|
)
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def detect_column_types(self, records):
|
|
||||||
all_column_types = {}
|
|
||||||
for record in records:
|
|
||||||
for key, value in record.items():
|
|
||||||
all_column_types.setdefault(key, set()).add(type(value))
|
|
||||||
column_types = {}
|
|
||||||
for key, types in all_column_types.items():
|
|
||||||
if len(types) == 1:
|
|
||||||
t = list(types)[0]
|
|
||||||
# But if it's list / tuple / dict, use str instead as we
|
|
||||||
# will be storing it as JSON in the table
|
|
||||||
if t in (list, tuple, dict):
|
|
||||||
t = str
|
|
||||||
elif {int, bool}.issuperset(types):
|
|
||||||
t = int
|
|
||||||
elif {int, float, bool}.issuperset(types):
|
|
||||||
t = float
|
|
||||||
elif {bytes, str}.issuperset(types):
|
|
||||||
t = bytes
|
|
||||||
else:
|
|
||||||
t = str
|
|
||||||
column_types[key] = t
|
|
||||||
return column_types
|
|
||||||
|
|
||||||
def search(self, q):
|
def search(self, q):
|
||||||
sql = """
|
sql = """
|
||||||
select * from "{table}" where rowid in (
|
select * from "{table}" where rowid in (
|
||||||
|
|
@ -1006,7 +982,7 @@ class Table(Queryable):
|
||||||
if not self.exists:
|
if not self.exists:
|
||||||
# Use the first batch to derive the table names
|
# Use the first batch to derive the table names
|
||||||
self.create(
|
self.create(
|
||||||
self.detect_column_types(chunk),
|
suggest_column_types(chunk),
|
||||||
pk,
|
pk,
|
||||||
foreign_keys,
|
foreign_keys,
|
||||||
column_order=column_order,
|
column_order=column_order,
|
||||||
|
|
@ -1179,7 +1155,7 @@ class Table(Queryable):
|
||||||
)
|
)
|
||||||
|
|
||||||
def add_missing_columns(self, records):
|
def add_missing_columns(self, records):
|
||||||
needed_columns = self.detect_column_types(records)
|
needed_columns = suggest_column_types(records)
|
||||||
current_columns = self.columns_dict
|
current_columns = self.columns_dict
|
||||||
for col_name, col_type in needed_columns.items():
|
for col_name, col_type in needed_columns.items():
|
||||||
if col_name not in current_columns:
|
if col_name not in current_columns:
|
||||||
|
|
|
||||||
|
|
@ -7,3 +7,28 @@ except ImportError:
|
||||||
import sqlite3
|
import sqlite3
|
||||||
|
|
||||||
OperationalError = sqlite3.OperationalError
|
OperationalError = sqlite3.OperationalError
|
||||||
|
|
||||||
|
|
||||||
|
def suggest_column_types(records):
|
||||||
|
all_column_types = {}
|
||||||
|
for record in records:
|
||||||
|
for key, value in record.items():
|
||||||
|
all_column_types.setdefault(key, set()).add(type(value))
|
||||||
|
column_types = {}
|
||||||
|
for key, types in all_column_types.items():
|
||||||
|
if len(types) == 1:
|
||||||
|
t = list(types)[0]
|
||||||
|
# But if it's list / tuple / dict, use str instead as we
|
||||||
|
# will be storing it as JSON in the table
|
||||||
|
if t in (list, tuple, dict):
|
||||||
|
t = str
|
||||||
|
elif {int, bool}.issuperset(types):
|
||||||
|
t = int
|
||||||
|
elif {int, float, bool}.issuperset(types):
|
||||||
|
t = float
|
||||||
|
elif {bytes, str}.issuperset(types):
|
||||||
|
t = bytes
|
||||||
|
else:
|
||||||
|
t = str
|
||||||
|
column_types[key] = t
|
||||||
|
return column_types
|
||||||
|
|
|
||||||
19
tests/test_suggest_column_types.py
Normal file
19
tests/test_suggest_column_types.py
Normal file
|
|
@ -0,0 +1,19 @@
|
||||||
|
import pytest
|
||||||
|
from sqlite_utils.utils import suggest_column_types
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"records,types",
|
||||||
|
[
|
||||||
|
([{"a": 1}], {"a": int}),
|
||||||
|
([{"a": "baz"}], {"a": str}),
|
||||||
|
([{"a": 1.2}], {"a": float}),
|
||||||
|
([{"a": [1]}], {"a": str}),
|
||||||
|
([{"a": (1,)}], {"a": str}),
|
||||||
|
([{"a": {"b": 1}}], {"a": str}),
|
||||||
|
([{"a": 1}, {"a": 1.1}], {"a": float}),
|
||||||
|
([{"a": b"b"}], {"a": bytes}),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_suggest_column_types(records, types):
|
||||||
|
assert types == suggest_column_types(records)
|
||||||
Loading…
Add table
Add a link
Reference in a new issue