Rename transform_table to just transform, refs #114

This commit is contained in:
Simon Willison 2020-09-21 16:14:05 -07:00
commit e3f3bea60a
2 changed files with 9 additions and 11 deletions

View file

@ -387,7 +387,7 @@ class Database:
hash_id=hash_id, hash_id=hash_id,
extracts=extracts, extracts=extracts,
) )
self.conn.execute(sql) self.execute(sql)
return self.table( return self.table(
name, name,
pk=pk, pk=pk,
@ -710,7 +710,7 @@ class Table(Queryable):
) )
return self return self
def transform_table( def transform(
self, self,
columns=None, columns=None,
rename=None, rename=None,
@ -724,7 +724,7 @@ class Table(Queryable):
extracts=None, extracts=None,
): ):
assert self.exists(), "Cannot transform a table that doesn't exist yet" assert self.exists(), "Cannot transform a table that doesn't exist yet"
sqls = self.transform_table_sql( sqls = self.transform_sql(
columns=columns, columns=columns,
rename=rename, rename=rename,
change_type=change_type, change_type=change_type,
@ -741,7 +741,7 @@ class Table(Queryable):
self.db.conn.execute(sql) self.db.conn.execute(sql)
return self return self
def transform_table_sql( def transform_sql(
self, self,
columns=None, columns=None,
rename=None, rename=None,
@ -783,15 +783,13 @@ class Table(Queryable):
) )
) )
# Copy across data, respecting any renamed columns # Copy across data, respecting any renamed columns
new_columns = set(columns.keys()) new_columns = columns.keys()
columns_to_copy = new_columns.intersection(previous_columns) columns_to_copy = set(new_columns).intersection(previous_columns)
copy_sql = "INSERT INTO [{new_table}] ({new_cols}) SELECT {old_cols} FROM [{old_table}]".format( copy_sql = "INSERT INTO [{new_table}] ({new_cols}) SELECT {old_cols} FROM [{old_table}]".format(
new_table=new_table_name, new_table=new_table_name,
old_table=self.name, old_table=self.name,
old_cols=", ".join(sorted("[{}]".format(col) for col in columns_to_copy)), old_cols=", ".join(sorted("[{}]".format(col) for col in columns_to_copy)),
new_cols=", ".join( new_cols=", ".join(sorted("[{}]".format(col) for col in new_columns)),
sorted("[{}]".format(col) for col in new_columns)
),
) )
sqls.append(copy_sql) sqls.append(copy_sql)
# Drop the old table # Drop the old table

View file

@ -24,9 +24,9 @@ import pytest
), ),
], ],
) )
def test_transform_table_sql(fresh_db, params, expected_sql): def test_transform_sql(fresh_db, params, expected_sql):
dogs = fresh_db["dogs"] dogs = fresh_db["dogs"]
dogs.insert({"id": 1, "name": "Cleo", "age": "5"}, pk="id") dogs.insert({"id": 1, "name": "Cleo", "age": "5"}, pk="id")
params["tmp_suffix"] = "suffix" params["tmp_suffix"] = "suffix"
sql = dogs.transform_table_sql(**params) sql = dogs.transform_sql(**params)
assert sql == expected_sql assert sql == expected_sql