New conversions= feature, closes #77

Pull request: #78
This commit is contained in:
Simon Willison 2020-01-30 16:24:30 -08:00 committed by GitHub
commit e8b2b7383b
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
3 changed files with 126 additions and 4 deletions

View file

@ -846,6 +846,68 @@ For example:
""").fetchall()
# Returns [('Felton, CA',)]
.. _python_api_conversions:
Converting column values using SQL functions
============================================
Sometimes it can be useful to run values through a SQL function prior to inserting them. A simple example might be converting a value to upper case while it is being inserted.
The ``conversions={...}`` parameter can be used to specify custom SQL to be used as part of a ``INSERT`` or ``UPDATE`` SQL statement.
You can specify an upper case conversion for a specific column like so:
.. code-block:: python
db["example"].insert({
"name": "The Bigfoot Discovery Museum"
}, conversions={"name": "upper(?)"})
# list(db["example"].rows) now returns:
# [{'name': 'THE BIGFOOT DISCOVERY MUSEUM'}]
The dictionary key is the column name to be converted. The value is the SQL fragment to use, with a ``?`` placeholder for the original value.
A more useful example: if you are working with `SpatiaLite <https://www.gaia-gis.it/fossil/libspatialite/index>`__ you may find yourself wanting to create geometry values from a WKT value. Code to do that could look like this:
.. code-block:: python
import sqlite3
import sqlite_utils
from shapely.geometry import shape
import requests
# Open a database and load the SpatiaLite extension:
import sqlite3
conn = sqlite3.connect("places.db")
conn.enable_load_extension(True)
conn.load_extension("/usr/local/lib/mod_spatialite.dylib")
# Use sqlite-utils to create a places table:
db = sqlite_utils.Database(conn)
places = db["places"].create({"id": int, "name": str,})
# Add a SpatiaLite 'geometry' column:
db.conn.execute("select InitSpatialMetadata(1)")
db.conn.execute(
"SELECT AddGeometryColumn('places', 'geometry', 4326, 'MULTIPOLYGON', 2);"
)
# Fetch some GeoJSON from Who's On First:
geojson = requests.get(
"https://data.whosonfirst.org/404/227/475/404227475.geojson"
).json()
# Convert to "Well Known Text" format using shapely
wkt = shape(geojson["geometry"]).wkt
# Insert the record, converting the WKT to a SpatiaLite geometry:
db["places"].insert(
{"name": "Wales", "geometry": wkt},
conversions={"geometry": "GeomFromText(?, 4326)"},
)
.. _python_api_introspection:
Introspection

View file

@ -483,6 +483,7 @@ class Table(Queryable):
ignore=False,
replace=False,
extracts=None,
conversions=None,
):
super().__init__(db, name)
self.exists = self.name in self.db.table_names()
@ -498,6 +499,7 @@ class Table(Queryable):
ignore=ignore,
replace=replace,
extracts=extracts,
conversions=conversions or {},
)
def __repr__(self):
@ -876,8 +878,9 @@ class Table(Queryable):
sql += " where " + where
self.db.conn.execute(sql, where_args or [])
def update(self, pk_values, updates=None, alter=False):
def update(self, pk_values, updates=None, alter=False, conversions=None):
updates = updates or {}
conversions = conversions or {}
if not isinstance(pk_values, (list, tuple)):
pk_values = [pk_values]
# Sanity check that the record exists (raises error if not):
@ -888,7 +891,7 @@ class Table(Queryable):
sets = []
wheres = []
for key, value in updates.items():
sets.append("[{}] = ?".format(key))
sets.append("[{}] = {}".format(key, conversions.get(key, "?")))
args.append(value)
wheres = ["[{}] = ?".format(pk_name) for pk_name in self.pks]
args.extend(pk_values)
@ -924,6 +927,7 @@ class Table(Queryable):
ignore=DEFAULT,
replace=DEFAULT,
extracts=DEFAULT,
conversions=DEFAULT,
):
return self.insert_all(
[record],
@ -937,6 +941,7 @@ class Table(Queryable):
ignore=ignore,
replace=replace,
extracts=extracts,
conversions=conversions,
)
def insert_all(
@ -953,6 +958,7 @@ class Table(Queryable):
ignore=DEFAULT,
replace=DEFAULT,
extracts=DEFAULT,
conversions=DEFAULT,
upsert=False,
):
"""
@ -973,6 +979,7 @@ class Table(Queryable):
ignore = self.value_or_default("ignore", ignore)
replace = self.value_or_default("replace", replace)
extracts = self.value_or_default("extracts", extracts)
conversions = self.value_or_default("conversions", conversions)
assert not (hash_id and pk), "Use either pk= or hash_id="
assert not (
@ -1053,7 +1060,10 @@ class Table(Queryable):
set_cols = [col for col in all_columns if col not in pks]
sql2 = "UPDATE [{table}] SET {pairs} WHERE {wheres}".format(
table=self.name,
pairs=", ".join("[{}] = ?".format(col) for col in set_cols),
pairs=", ".join(
"[{}] = {}".format(col, conversions.get(col, "?"))
for col in set_cols
),
wheres=" AND ".join("[{}] = ?".format(pk) for pk in pks),
)
queries_and_params.append(
@ -1079,7 +1089,9 @@ class Table(Queryable):
"""
({placeholders})
""".format(
placeholders=", ".join(["?"] * len(all_columns))
placeholders=", ".join(
[conversions.get(col, "?") for col in all_columns]
)
)
for record in chunk
),
@ -1122,6 +1134,7 @@ class Table(Queryable):
hash_id=DEFAULT,
alter=DEFAULT,
extracts=DEFAULT,
conversions=DEFAULT,
):
return self.upsert_all(
[record],
@ -1133,6 +1146,7 @@ class Table(Queryable):
hash_id=hash_id,
alter=alter,
extracts=extracts,
conversions=conversions,
)
def upsert_all(
@ -1147,6 +1161,7 @@ class Table(Queryable):
hash_id=DEFAULT,
alter=DEFAULT,
extracts=DEFAULT,
conversions=DEFAULT,
):
return self.insert_all(
records,
@ -1159,6 +1174,7 @@ class Table(Queryable):
hash_id=hash_id,
alter=alter,
extracts=extracts,
conversions=conversions,
upsert=True,
)

44
tests/test_conversions.py Normal file
View file

@ -0,0 +1,44 @@
import pytest
def test_insert_conversion(fresh_db):
table = fresh_db["table"]
table.insert({"foo": "bar"}, conversions={"foo": "upper(?)"})
assert [{"foo": "BAR"}] == list(table.rows)
def test_insert_all_conversion(fresh_db):
table = fresh_db["table"]
table.insert_all([{"foo": "bar"}], conversions={"foo": "upper(?)"})
assert [{"foo": "BAR"}] == list(table.rows)
def test_upsert_conversion(fresh_db):
table = fresh_db["table"]
table.upsert({"id": 1, "foo": "bar"}, pk="id", conversions={"foo": "upper(?)"})
assert [{"id": 1, "foo": "BAR"}] == list(table.rows)
table.upsert(
{"id": 1, "bar": "baz"}, pk="id", conversions={"bar": "upper(?)"}, alter=True
)
assert [{"id": 1, "foo": "BAR", "bar": "BAZ"}] == list(table.rows)
def test_upsert_all_conversion(fresh_db):
table = fresh_db["table"]
table.upsert_all(
[{"id": 1, "foo": "bar"}], pk="id", conversions={"foo": "upper(?)"}
)
assert [{"id": 1, "foo": "BAR"}] == list(table.rows)
def test_update_conversion(fresh_db):
table = fresh_db["table"]
table.insert({"id": 5, "foo": "bar"}, pk="id")
table.update(5, {"foo": "baz"}, conversions={"foo": "upper(?)"})
assert [{"id": 5, "foo": "BAZ"}] == list(table.rows)
def test_table_constructor_conversion(fresh_db):
table = fresh_db.table("table", conversions={"bar": "upper(?)"})
table.insert({"bar": "baz"})
assert [{"bar": "BAZ"}] == list(table.rows)