diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index d7fe44a..4f2b906 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -78,6 +78,7 @@ class NoObviousTable(Exception): class BadPrimaryKey(Exception): pass + class NotFoundError(Exception): pass @@ -208,11 +209,15 @@ class Database: raise AlterError( "No such column: {}.{}".format(fk.other_table, fk.other_column) ) - extra = "" + column_defs = [] + # ensure pk is a tuple + single_pk = None + if isinstance(pk, str): + single_pk = pk for column_name, column_type in column_items: column_extras = [] - if pk == column_name: + if column_name == single_pk: column_extras.append("PRIMARY KEY") if column_name in not_null: column_extras.append("NOT NULL") @@ -236,12 +241,17 @@ class Database: else "", ) ) + extra_pk = "" + if single_pk is None and pk and len(pk) > 1: + extra_pk = ",\n PRIMARY KEY ({pks})".format( + pks=", ".join(["[{}]".format(p) for p in pk]) + ) columns_sql = ",\n".join(column_defs) sql = """CREATE TABLE [{table}] ( -{columns_sql} -){extra}; +{columns_sql}{extra_pk} +); """.format( - table=name, columns_sql=columns_sql, extra=extra + table=name, columns_sql=columns_sql, extra_pk=extra_pk ) self.conn.execute(sql) return self[name] @@ -401,7 +411,7 @@ class Table: pk_names = pks last_pk = pk_values wheres = ["[{}] = ?".format(pk_name) for pk_name in pk_names] - rows = self.rows_where(' and '.join(wheres), pk_values) + rows = self.rows_where(" and ".join(wheres), pk_values) try: row = list(rows)[0] self.last_pk = last_pk @@ -813,12 +823,13 @@ class Table: self.last_pk = None # self.last_rowid will be 0 if a "INSERT OR IGNORE" happened if (hash_id or pk) and self.last_rowid: - self.last_pk = self.db.conn.execute( - "select [{}] from [{}] where rowid = ?".format( - hash_id or pk, self.name - ), - (self.last_rowid,), - ).fetchone()[0] + row = list(self.rows_where("rowid = ?", [self.last_rowid]))[0] + if hash_id: + self.last_pk = row[hash_id] + elif isinstance(pk, str): + self.last_pk = row[pk] + else: + self.last_pk = tuple(row[p] for p in pk) return self def upsert( diff --git a/tests/test_create.py b/tests/test_create.py index 5e76d0c..91ad169 100644 --- a/tests/test_create.py +++ b/tests/test_create.py @@ -55,6 +55,21 @@ def test_create_table(fresh_db): ) == table.schema +def test_create_table_compound_primary_key(fresh_db): + table = fresh_db.create_table( + "test_table", {"id1": str, "id2": str, "value": int}, pk=("id1", "id2") + ) + assert ( + "CREATE TABLE [test_table] (\n" + " [id1] TEXT,\n" + " [id2] TEXT,\n" + " [value] INTEGER,\n" + " PRIMARY KEY ([id1], [id2])\n" + ")" + ) == table.schema + assert ["id1", "id2"] == table.pks + + def test_create_table_with_bad_defaults(fresh_db): with pytest.raises(AssertionError): fresh_db.create_table( @@ -122,6 +137,13 @@ def test_create_table_from_example(fresh_db, example, expected_columns): ] +def test_create_table_from_example_with_compound_primary_keys(fresh_db): + record = {"name": "Zhang", "group": "staff", "employee_id": 2} + table = fresh_db["people"].insert(record, pk=("group", "employee_id")) + assert ["group", "employee_id"] == table.pks + assert record == table.get(("staff", 2)) + + def test_create_table_column_order(fresh_db): fresh_db["table"].insert( collections.OrderedDict(