Type signatures for .create_table() and .create_table_sql() and .create() and Table.__init__

Closes #314
This commit is contained in:
Simon Willison 2021-08-18 15:25:18 -07:00
commit c79737bb4f
2 changed files with 62 additions and 43 deletions

View file

@ -147,6 +147,15 @@ XIndexColumn = namedtuple(
Trigger = namedtuple("Trigger", ("name", "table", "sql")) Trigger = namedtuple("Trigger", ("name", "table", "sql"))
ForeignKeysType = Union[
Iterable[str],
Iterable[ForeignKey],
Iterable[Tuple[str, str]],
Iterable[Tuple[str, str, str]],
Iterable[Tuple[str, str, str, str]],
]
class Default: class Default:
pass pass
@ -572,18 +581,22 @@ class Database:
) -> List[dict]: ) -> List[dict]:
return list(self.query(sql, params)) return list(self.query(sql, params))
def resolve_foreign_keys(self, name, foreign_keys): def resolve_foreign_keys(
# foreign_keys may be a list of strcolumn names, a list of ForeignKey tuples, self, name: str, foreign_keys: ForeignKeysType
) -> List[ForeignKey]:
# foreign_keys may be a list of column names, a list of ForeignKey tuples,
# a list of tuple-pairs or a list of tuple-triples. We want to turn # a list of tuple-pairs or a list of tuple-triples. We want to turn
# it into a list of ForeignKey tuples # it into a list of ForeignKey tuples
table = cast(Table, self[name])
if all(isinstance(fk, ForeignKey) for fk in foreign_keys): if all(isinstance(fk, ForeignKey) for fk in foreign_keys):
return foreign_keys return cast(List[ForeignKey], foreign_keys)
if all(isinstance(fk, str) for fk in foreign_keys): if all(isinstance(fk, str) for fk in foreign_keys):
# It's a list of columns # It's a list of columns
fks = [] fks = []
for column in foreign_keys: for column in foreign_keys:
other_table = self[name].guess_foreign_table(column) column = cast(str, column)
other_column = self[name].guess_foreign_column(other_table) other_table = table.guess_foreign_table(column)
other_column = table.guess_foreign_column(other_table)
fks.append(ForeignKey(name, column, other_table, other_column)) fks.append(ForeignKey(name, column, other_table, other_column))
return fks return fks
assert all( assert all(
@ -596,6 +609,7 @@ class Database:
3, 3,
), "foreign_keys= should be a list of tuple pairs or triples" ), "foreign_keys= should be a list of tuple pairs or triples"
if len(tuple_or_list) == 3: if len(tuple_or_list) == 3:
tuple_or_list = cast(Tuple[str, str, str], tuple_or_list)
fks.append( fks.append(
ForeignKey( ForeignKey(
name, tuple_or_list[0], tuple_or_list[1], tuple_or_list[2] name, tuple_or_list[0], tuple_or_list[1], tuple_or_list[2]
@ -608,7 +622,7 @@ class Database:
name, name,
tuple_or_list[0], tuple_or_list[0],
tuple_or_list[1], tuple_or_list[1],
self[name].guess_foreign_column(tuple_or_list[1]), table.guess_foreign_column(tuple_or_list[1]),
) )
) )
return fks return fks
@ -618,12 +632,12 @@ class Database:
name: str, name: str,
columns: Dict[str, Any], columns: Dict[str, Any],
pk: Optional[Any] = None, pk: Optional[Any] = None,
foreign_keys=None, foreign_keys: Optional[ForeignKeysType] = None,
column_order=None, column_order: Optional[List[str]] = None,
not_null=None, not_null: Iterable[str] = None,
defaults=None, defaults: Optional[Dict[str, Any]] = None,
hash_id=None, hash_id: Optional[Any] = None,
extracts=None, extracts: Optional[Union[Dict[str, str], List[str]]] = None,
) -> str: ) -> str:
"Returns the SQL ``CREATE TABLE`` statement for creating the specified table." "Returns the SQL ``CREATE TABLE`` statement for creating the specified table."
foreign_keys = self.resolve_foreign_keys(name, foreign_keys or []) foreign_keys = self.resolve_foreign_keys(name, foreign_keys or [])
@ -656,9 +670,11 @@ class Database:
validate_column_names(columns.keys()) validate_column_names(columns.keys())
column_items = list(columns.items()) column_items = list(columns.items())
if column_order is not None: if column_order is not None:
column_items.sort(
key=lambda p: column_order.index(p[0]) if p[0] in column_order else 999 def sort_key(p):
) return column_order.index(p[0]) if p[0] in column_order else 999
column_items.sort(key=sort_key)
if hash_id: if hash_id:
column_items.insert(0, (hash_id, str)) column_items.insert(0, (hash_id, str))
pk = hash_id pk = hash_id
@ -725,12 +741,12 @@ class Database:
name: str, name: str,
columns: Dict[str, Any], columns: Dict[str, Any],
pk: Optional[Any] = None, pk: Optional[Any] = None,
foreign_keys=None, foreign_keys: Optional[ForeignKeysType] = None,
column_order=None, column_order: Optional[List[str]] = None,
not_null=None, not_null: Iterable[str] = None,
defaults=None, defaults: Optional[Dict[str, Any]] = None,
hash_id=None, hash_id: Optional[Any] = None,
extracts=None, extracts: Optional[Union[Dict[str, str], List[str]]] = None,
) -> "Table": ) -> "Table":
""" """
Create a table with the specified name and the specified ``{column_name: type}`` columns. Create a table with the specified name and the specified ``{column_name: type}`` columns.
@ -1021,19 +1037,19 @@ class Table(Queryable):
self, self,
db: Database, db: Database,
name: str, name: str,
pk=None, pk: Optional[Any] = None,
foreign_keys=None, foreign_keys: Optional[ForeignKeysType] = None,
column_order=None, column_order: Optional[List[str]] = None,
not_null=None, not_null: Iterable[str] = None,
defaults=None, defaults: Optional[Dict[str, Any]] = None,
batch_size=100, batch_size: int = 100,
hash_id=None, hash_id: Optional[Any] = None,
alter=False, alter: bool = False,
ignore=False, ignore: bool = False,
replace=False, replace: bool = False,
extracts=None, extracts: Optional[Union[Dict[str, str], List[str]]] = None,
conversions=None, conversions: Optional[dict] = None,
columns=None, columns: Optional[Union[Dict[str, Any]]] = None,
): ):
super().__init__(db, name) super().__init__(db, name)
self._defaults = dict( self._defaults = dict(
@ -1202,14 +1218,14 @@ class Table(Queryable):
def create( def create(
self, self,
columns, columns: Dict[str, Any],
pk=None, pk: Optional[Any] = None,
foreign_keys=None, foreign_keys: Optional[ForeignKeysType] = None,
column_order=None, column_order: Optional[List[str]] = None,
not_null=None, not_null: Iterable[str] = None,
defaults=None, defaults: Optional[Dict[str, Any]] = None,
hash_id=None, hash_id: Optional[Any] = None,
extracts=None, extracts: Optional[Union[Dict[str, str], List[str]]] = None,
) -> "Table": ) -> "Table":
""" """
Create a table with the specified columns. Create a table with the specified columns.
@ -2914,7 +2930,9 @@ def _hash(record):
).hexdigest() ).hexdigest()
def resolve_extracts(extracts): def resolve_extracts(
extracts: Optional[Union[Dict[str, str], List[str], Tuple[str]]]
) -> dict:
if extracts is None: if extracts is None:
extracts = {} extracts = {}
if isinstance(extracts, (list, tuple)): if isinstance(extracts, (list, tuple)):

View file

@ -13,6 +13,7 @@ def test_tracer():
("PRAGMA recursive_triggers=on;", None), ("PRAGMA recursive_triggers=on;", None),
("select name from sqlite_master where type = 'view'", None), ("select name from sqlite_master where type = 'view'", None),
("select name from sqlite_master where type = 'table'", None), ("select name from sqlite_master where type = 'table'", None),
("select name from sqlite_master where type = 'view'", None),
("CREATE TABLE [dogs] (\n [name] TEXT\n);\n ", None), ("CREATE TABLE [dogs] (\n [name] TEXT\n);\n ", None),
("select name from sqlite_master where type = 'view'", None), ("select name from sqlite_master where type = 'view'", None),
("INSERT INTO [dogs] ([name]) VALUES (?);", ["Cleopaws"]), ("INSERT INTO [dogs] ([name]) VALUES (?);", ["Cleopaws"]),