.transform(defaults=...)

This commit is contained in:
Simon Willison 2020-09-21 19:57:44 -07:00
commit 785c4a80a2
2 changed files with 45 additions and 3 deletions

View file

@ -336,7 +336,7 @@ class Database:
column_extras.append("PRIMARY KEY") column_extras.append("PRIMARY KEY")
if column_name in not_null: if column_name in not_null:
column_extras.append("NOT NULL") column_extras.append("NOT NULL")
if column_name in defaults: if column_name in defaults and defaults[column_name] is not None:
column_extras.append( column_extras.append(
"DEFAULT {}".format(self.escape(defaults[column_name])) "DEFAULT {}".format(self.escape(defaults[column_name]))
) )
@ -736,9 +736,9 @@ class Table(Queryable):
drop=None, drop=None,
pk=pk, pk=pk,
not_null=not_null, not_null=not_null,
defaults=defaults,
# foreign_keys=foreign_keys, # foreign_keys=foreign_keys,
# column_order=column_order, # column_order=column_order,
# defaults=defaults,
# hash_id=hash_id, # hash_id=hash_id,
# extracts=extracts, # extracts=extracts,
) )
@ -754,6 +754,7 @@ class Table(Queryable):
drop=None, drop=None,
pk=None, pk=None,
not_null=None, not_null=None,
defaults=None,
tmp_suffix=None, tmp_suffix=None,
): ):
columns = columns or {} columns = columns or {}
@ -795,16 +796,27 @@ class Table(Queryable):
elif isinstance(not_null, set): elif isinstance(not_null, set):
create_table_not_null.update(rename.get(k) or k for k in not_null) create_table_not_null.update(rename.get(k) or k for k in not_null)
# defaults=
create_table_defaults = {
(rename.get(c.name) or c.name): c.default_value
for c in self.columns
if c.default_value is not None
}
if defaults is not None:
create_table_defaults.update(
{rename.get(c) or c: v for c, v in defaults.items()}
)
sqls.append( sqls.append(
self.db.create_table_sql( self.db.create_table_sql(
new_table_name, new_table_name,
dict(new_column_pairs), dict(new_column_pairs),
pk=pk, pk=pk,
not_null=create_table_not_null, not_null=create_table_not_null,
defaults=create_table_defaults,
# foreign_keys=foreign_keys, # foreign_keys=foreign_keys,
# column_order=column_order, # column_order=column_order,
# not_null=not_null, # not_null=not_null,
# defaults=defaults,
# hash_id=hash_id, # hash_id=hash_id,
# extracts=extracts, # extracts=extracts,
).strip() ).strip()

View file

@ -129,3 +129,33 @@ def test_transform_add_not_null_with_rename(fresh_db, not_null):
dogs.schema dogs.schema
== 'CREATE TABLE "dogs" (\n [id] INTEGER PRIMARY KEY,\n [name] TEXT,\n [dog_age] TEXT NOT NULL\n)' == 'CREATE TABLE "dogs" (\n [id] INTEGER PRIMARY KEY,\n [name] TEXT,\n [dog_age] TEXT NOT NULL\n)'
) )
def test_transform_defaults(fresh_db):
dogs = fresh_db["dogs"]
dogs.insert({"id": 1, "name": "Cleo", "age": 5}, pk="id")
dogs.transform(defaults={"age": 1})
assert (
dogs.schema
== 'CREATE TABLE "dogs" (\n [id] INTEGER PRIMARY KEY,\n [name] TEXT,\n [age] INTEGER DEFAULT 1\n)'
)
def test_transform_defaults_and_rename_column(fresh_db):
dogs = fresh_db["dogs"]
dogs.insert({"id": 1, "name": "Cleo", "age": 5}, pk="id")
dogs.transform(rename={"age": "dog_age"}, defaults={"age": 1})
assert (
dogs.schema
== 'CREATE TABLE "dogs" (\n [id] INTEGER PRIMARY KEY,\n [name] TEXT,\n [dog_age] INTEGER DEFAULT 1\n)'
)
def test_remove_defaults(fresh_db):
dogs = fresh_db["dogs"]
dogs.insert({"id": 1, "name": "Cleo", "age": 5}, defaults={"age": 1}, pk="id")
dogs.transform(defaults={"age": None})
assert (
dogs.schema
== 'CREATE TABLE "dogs" (\n [id] INTEGER PRIMARY KEY,\n [name] TEXT,\n [age] INTEGER\n)'
)