diff --git a/docs/python-api.rst b/docs/python-api.rst index db988fe..a1f0461 100644 --- a/docs/python-api.rst +++ b/docs/python-api.rst @@ -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 `__ 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 diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index 0a81a60..b061e98 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -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, ) diff --git a/tests/test_conversions.py b/tests/test_conversions.py new file mode 100644 index 0000000..ebe2a50 --- /dev/null +++ b/tests/test_conversions.py @@ -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)