Allow iterables other than lists in m2m records (#189)

* Allow iterables other than Lists in m2m records
* Add test for iterable m2m records

Thanks, @adamwolf!
This commit is contained in:
Adam Wolf 2020-10-27 11:24:21 -05:00 • committed by GitHub
commit f045d8559a
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
2 changed files with 64 additions and 26 deletions

View file

@ -1,5 +1,6 @@
from .utils import sqlite3, OperationalError, suggest_column_types, column_affinity from .utils import sqlite3, OperationalError, suggest_column_types, column_affinity
from collections import namedtuple, OrderedDict from collections import namedtuple, OrderedDict
from collections.abc import Mapping
import contextlib import contextlib
import datetime import datetime
import decimal import decimal
@ -1772,15 +1773,15 @@ class Table(Queryable):
return pk return pk
def m2m( def m2m(
self, other_table, record_or_list=None, pk=DEFAULT, lookup=None, m2m_table=None self, other_table, record_or_iterable=None, pk=DEFAULT, lookup=None, m2m_table=None
): ):
if isinstance(other_table, str): if isinstance(other_table, str):
other_table = self.db.table(other_table, pk=pk) other_table = self.db.table(other_table, pk=pk)
our_id = self.last_pk our_id = self.last_pk
if lookup is not None: if lookup is not None:
assert record_or_list is None, "Provide lookup= or record, not both" assert record_or_iterable is None, "Provide lookup= or record, not both"
else: else:
assert record_or_list is not None, "Provide lookup= or record, not both" assert record_or_iterable is not None, "Provide lookup= or record, not both"
tables = list(sorted([self.name, other_table.name])) tables = list(sorted([self.name, other_table.name]))
columns = ["{}_id".format(t) for t in tables] columns = ["{}_id".format(t) for t in tables]
if m2m_table is not None: if m2m_table is not None:
@ -1801,10 +1802,11 @@ class Table(Queryable):
m2m_table_name = m2m_table or "{}_{}".format(*tables) m2m_table_name = m2m_table or "{}_{}".format(*tables)
m2m_table = self.db.table(m2m_table_name, pk=columns, foreign_keys=columns) m2m_table = self.db.table(m2m_table_name, pk=columns, foreign_keys=columns)
if lookup is None: if lookup is None:
# if records is only one record, put the record in a list
records = ( records = (
[record_or_list] [record_or_iterable]
if not isinstance(record_or_list, (list, tuple)) if isinstance(record_or_iterable, Mapping)
else record_or_list else record_or_iterable
) )
# Ensure each record exists in other table # Ensure each record exists in other table
for record in records: for record in records:

View file

@ -43,6 +43,42 @@ def test_insert_m2m_list(fresh_db):
] == dogs_humans.foreign_keys ] == dogs_humans.foreign_keys
def test_insert_m2m_iterable(fresh_db):
iterable_records = ({"id": 1, "name": "Phineas"}, {"id": 2, "name": "Ferb"})
def iterable():
for record in iterable_records:
yield record
platypuses = fresh_db["platypuses"]
platypuses.insert({"id": 1, "name": "Perry"}, pk="id").m2m(
"humans",
iterable(),
pk="id",
)
assert {"platypuses", "humans", "humans_platypuses"} == set(fresh_db.table_names())
humans = fresh_db["humans"]
humans_platypuses = fresh_db["humans_platypuses"]
assert [{"humans_id": 1, "platypuses_id": 1}, {"humans_id": 2, "platypuses_id": 1}] == list(
humans_platypuses.rows
)
assert [{"id": 1, "name": "Phineas"}, {"id": 2, "name": "Ferb"}] == list(
humans.rows
)
assert [
ForeignKey(
table="humans_platypuses", column="platypuses_id", other_table="platypuses", other_column="id"
),
ForeignKey(
table="humans_platypuses",
column="humans_id",
other_table="humans",
other_column="id",
),
] == humans_platypuses.foreign_keys
def test_m2m_with_table_objects(fresh_db): def test_m2m_with_table_objects(fresh_db):
dogs = fresh_db.table("dogs", pk="id") dogs = fresh_db.table("dogs", pk="id")
humans = fresh_db.table("humans", pk="id") humans = fresh_db.table("humans", pk="id")