From 362359da7eed7dc6589589122960c0a0d0460d7c Mon Sep 17 00:00:00 2001 From: Simon Willison Date: Wed, 12 Jun 2019 23:10:07 -0700 Subject: [PATCH] not_null= and defaults= arguments to various Python methods, refs #24 --- docs/python-api.rst | 44 ++++++++++++++++- sqlite_utils/db.py | 115 +++++++++++++++++++++++++++++++++---------- tests/test_create.py | 45 +++++++++++++++++ 3 files changed, 178 insertions(+), 26 deletions(-) diff --git a/docs/python-api.rst b/docs/python-api.rst index fa7b691..1243a17 100644 --- a/docs/python-api.rst +++ b/docs/python-api.rst @@ -138,7 +138,7 @@ You can directly create a new table without inserting any data into it using the The first argument here is a dictionary specifying the columns you would like to create. Each column is paired with a Python type indicating the type of column. See :ref:`python_api_add_column` for full details on how these types work. -This method takes optional arguments ``pk=``, ``column_order=`` and ``foreign_keys=``. +This method takes optional arguments ``pk=``, ``column_order=``, ``foreign_keys=``, ``not_null=set()`` and ``defaults=dict()`` - explained below. .. _python_api_foreign_keys: @@ -182,6 +182,48 @@ You can leave off the third item in the tuple to have the referenced column auto ("author_id", "authors") ]) + +.. _python_api_defaults_not_null: + +Setting defaults and not null constraints +========================================= + +Each of the methods that can cause a table to be created take optional arguments ``not_null=set()`` and ``defaults=dict()``. The methods that take these optional arguments are: + +* ``db.create_table(...)`` +* ``table.create(...)`` +* ``table.insert(...)`` +* ``table.insert_all(...)`` +* ``table.upsert(...)`` +* ``table.upsert_all(...)`` + +You can use ``not_null=`` to pass a set of column names that should have a ``NOT NULL`` constraint set on them when they are created. + +You can use ``defaults=`` to pass a dictionary mapping columns to the default value that should be specified in the ``CREATE TABLE`` statement. + +Here's an example that uses these features: + +.. code-block:: python + + db["authors"].insert_all( + [{"id": 1, "name": "Sally", "score": 2}], + pk="id", + not_null={"name", "score"}, + defaults={"score": 1}, + ) + db["authors"].insert({"name": "Dharma"}) + + list(db["authors"].rows) + # Outputs: + # [{'id': 1, 'name': 'Sally', 'score': 2}, + # {'id': 3, 'name': 'Dharma', 'score': 1}] + print(db["authors"].schema) # Outputs: + # CREATE TABLE [authors] ( + # [id] INTEGER PRIMARY KEY, + # [name] TEXT NOT NULL, + # [score] INTEGER NOT NULL DEFAULT 1 + # ) + .. _python_api_bulk_inserts: Bulk inserts diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index e223782..ec9ec95 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -94,6 +94,16 @@ class Database: def __repr__(self): return "".format(self.conn) + def escape(self, value): + # Normally we would use .execute(sql, [params]) for escaping, but + # occasionally that isn't available - most notable when we need + # to include a "... DEFAULT 'value'" in a column definition. + return self.conn.execute( + # Use SQLite itself to correctly escape this string: + "SELECT quote(:value)", + {"value": value}, + ).fetchone()[0] + def table_names(self, fts4=False, fts5=False): where = ["type = 'table'"] if fts4: @@ -154,10 +164,31 @@ class Database: return fks def create_table( - self, name, columns, pk=None, foreign_keys=None, column_order=None, hash_id=None + self, + name, + columns, + pk=None, + foreign_keys=None, + column_order=None, + not_null=None, + defaults=None, + hash_id=None, ): foreign_keys = self.resolve_foreign_keys(name, foreign_keys or []) foreign_keys_by_column = {fk.column: fk for fk in foreign_keys} + # Sanity check not_null, and defaults if provided + not_null = not_null or set() + defaults = defaults or {} + assert all( + n in columns for n in not_null + ), "not_null set {} includes items not in columns {}".format( + repr(not_null), repr(set(columns.keys())) + ) + assert all( + n in columns for n in defaults + ), "defaults set {} includes items not in columns {}".format( + repr(set(defaults)), repr(set(columns.keys())) + ) column_items = list(columns.items()) if column_order is not None: column_items.sort( @@ -175,22 +206,34 @@ class Database: "No such column: {}.{}".format(fk.other_table, fk.other_column) ) extra = "" - columns_sql = ",\n".join( - " [{col_name}] {col_type}{primary_key}{references}".format( - col_name=col_name, - col_type=COLUMN_TYPE_MAPPING[col_type], - primary_key=" PRIMARY KEY" if (pk == col_name) else "", - references=( - " REFERENCES [{other_table}]([{other_column}])".format( - other_table=foreign_keys_by_column[col_name].other_table, - other_column=foreign_keys_by_column[col_name].other_column, + column_defs = [] + for column_name, column_type in column_items: + column_extras = [] + if pk == column_name: + column_extras.append("PRIMARY KEY") + if column_name in not_null: + column_extras.append("NOT NULL") + if column_name in defaults: + column_extras.append( + "DEFAULT {}".format(self.escape(defaults[column_name])) + ) + if column_name in foreign_keys_by_column: + column_extras.append( + "REFERENCES [{other_table}]([{other_column}])".format( + other_table=foreign_keys_by_column[column_name].other_table, + other_column=foreign_keys_by_column[column_name].other_column, ) - if col_name in foreign_keys_by_column - else "" - ), + ) + column_defs.append( + " [{column_name}] {column_type}{column_extras}".format( + column_name=column_name, + column_type=COLUMN_TYPE_MAPPING[column_type], + column_extras=(" " + " ".join(column_extras)) + if column_extras + else "", + ) ) - for col_name, col_type in column_items - ) + columns_sql = ",\n".join(column_defs) sql = """CREATE TABLE [{table}] ( {columns_sql} ){extra}; @@ -308,7 +351,14 @@ class Table: return indexes def create( - self, columns, pk=None, foreign_keys=None, column_order=None, hash_id=None + self, + columns, + pk=None, + foreign_keys=None, + column_order=None, + not_null=None, + defaults=None, + hash_id=None, ): columns = {name: value for (name, value) in columns.items()} with self.db.conn: @@ -318,6 +368,8 @@ class Table: pk=pk, foreign_keys=foreign_keys, column_order=column_order, + not_null=not_null, + defaults=defaults, hash_id=hash_id, ) self.exists = True @@ -367,10 +419,7 @@ class Table: not_null_sql = None if not_null_default is not None: not_null_sql = "NOT NULL DEFAULT {}".format( - # Use SQLite itself to correctly escape this string: - self.db.conn.execute( - "SELECT quote(:default)", {"default": not_null_default} - ).fetchone()[0] + self.db.escape(not_null_default) ) sql = "ALTER TABLE [{table}] ADD COLUMN [{col_name}] {col_type}{not_null_default};".format( table=self.name, @@ -564,8 +613,10 @@ class Table: record, pk=None, foreign_keys=None, - upsert=False, column_order=None, + not_null=None, + defaults=None, + upsert=False, hash_id=None, alter=False, ignore=False, @@ -574,8 +625,10 @@ class Table: [record], pk=pk, foreign_keys=foreign_keys, - upsert=upsert, column_order=column_order, + not_null=not_null, + defaults=defaults, + upsert=upsert, hash_id=hash_id, alter=alter, ignore=ignore, @@ -586,9 +639,11 @@ class Table: records, pk=None, foreign_keys=None, + column_order=None, + not_null=None, + defaults=None, upsert=False, batch_size=100, - column_order=None, hash_id=None, alter=False, ignore=False, @@ -614,6 +669,8 @@ class Table: pk, foreign_keys, column_order=column_order, + not_null=not_null, + defaults=defaults, hash_id=hash_id, ) all_columns = set() @@ -679,6 +736,8 @@ class Table: pk=None, foreign_keys=None, column_order=None, + not_null=None, + defaults=None, hash_id=None, alter=False, ): @@ -686,10 +745,12 @@ class Table: record, pk=pk, foreign_keys=foreign_keys, - upsert=True, column_order=column_order, + not_null=not_null, + defaults=defaults, hash_id=hash_id, alter=alter, + upsert=True, ) def upsert_all( @@ -698,6 +759,8 @@ class Table: pk=None, foreign_keys=None, column_order=None, + not_null=None, + defaults=None, batch_size=100, hash_id=None, alter=False, @@ -707,10 +770,12 @@ class Table: pk=pk, foreign_keys=foreign_keys, column_order=column_order, + not_null=not_null, + defaults=defaults, batch_size=100, - upsert=True, hash_id=hash_id, alter=alter, + upsert=True, ) def add_missing_columns(self, records): diff --git a/tests/test_create.py b/tests/test_create.py index 36a4374..27c1a63 100644 --- a/tests/test_create.py +++ b/tests/test_create.py @@ -55,6 +55,51 @@ def test_create_table(fresh_db): ) == table.schema +def test_create_table_with_bad_defaults(fresh_db): + with pytest.raises(AssertionError): + fresh_db.create_table( + "players", {"name": str, "score": int}, defaults={"mouse": 1} + ) + + +def test_create_table_with_defaults(fresh_db): + table = fresh_db.create_table( + "players", + {"name": str, "score": int}, + defaults={"score": 1, "name": "bob''bob"}, + ) + assert ["players"] == fresh_db.table_names() + assert [{"name": "name", "type": "TEXT"}, {"name": "score", "type": "INTEGER"}] == [ + {"name": col.name, "type": col.type} for col in table.columns + ] + assert ( + "CREATE TABLE [players] (\n [name] TEXT DEFAULT 'bob''''bob',\n [score] INTEGER DEFAULT 1\n)" + ) == table.schema + + +def test_create_table_with_bad_not_null(fresh_db): + with pytest.raises(AssertionError): + fresh_db.create_table( + "players", {"name": str, "score": int}, not_null={"mouse"} + ) + + +def test_create_table_with_not_null(fresh_db): + table = fresh_db.create_table( + "players", + {"name": str, "score": int}, + not_null={"name", "score"}, + defaults={"score": 3}, + ) + assert ["players"] == fresh_db.table_names() + assert [{"name": "name", "type": "TEXT"}, {"name": "score", "type": "INTEGER"}] == [ + {"name": col.name, "type": col.type} for col in table.columns + ] + assert ( + "CREATE TABLE [players] (\n [name] TEXT NOT NULL,\n [score] INTEGER NOT NULL DEFAULT 3\n)" + ) == table.schema + + @pytest.mark.parametrize( "example,expected_columns", (