Skip to content

Commit 637a7d5

Browse files
committed
fix and test impl
1 parent 1f540e4 commit 637a7d5

5 files changed

Lines changed: 260 additions & 48 deletions

File tree

src/polymorphic/managers.py

Lines changed: 20 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -2,8 +2,6 @@
22
The manager class for use in the models.
33
"""
44

5-
import inspect
6-
75
from django.contrib.contenttypes.models import ContentType
86
from django.db import DEFAULT_DB_ALIAS, models
97

@@ -63,33 +61,33 @@ def create_from_super(self, obj, **kwargs):
6361
6462
returns obj as an instance of cls.
6563
"""
66-
cls = self.model
64+
from .models import PolymorphicModel
65+
66+
# ensure we have the most derived real instance
67+
if isinstance(obj, PolymorphicModel):
68+
obj = obj.get_real_instance()
69+
70+
parent_ptr = self.model._meta.parents.get(type(obj), None)
6771

68-
scls = inspect.getmro(cls)[1]
69-
if scls is not type(obj):
72+
if not parent_ptr:
7073
raise TypeError(
7174
"create_from_super can only be used if obj is one level of inheritance up from cls"
7275
)
73-
74-
parent_link_field = None
75-
for parent, field in cls._meta.parents.items():
76-
if parent is scls:
77-
parent_link_field = field
78-
break
79-
if parent_link_field is None:
80-
raise TypeError(f"Could not find parent link field for {scls.__name__}")
81-
kwargs[parent_link_field.get_attname()] = obj.id
76+
kwargs[parent_ptr.get_attname()] = obj.pk
8277

8378
# create the new base class with only fields that apply to it.
84-
nobj = cls(**kwargs)
85-
nobj.save_base(raw=True)
79+
ctype = ContentType.objects.db_manager(
80+
using=(obj._state.db or DEFAULT_DB_ALIAS)
81+
).get_for_model(self.model)
82+
nobj = self.model(**kwargs, polymorphic_ctype=ctype)
83+
nobj.save_base(raw=True, using=obj._state.db or DEFAULT_DB_ALIAS, force_insert=True)
8684
# force update the content type, but first we need to
8785
# retrieve a clean copy from the db to fill in the null
8886
# fields otherwise they would be overwritten.
89-
nobj = obj.__class__.objects.using(obj._state.db or DEFAULT_DB_ALIAS).get(pk=obj.pk)
90-
nobj.polymorphic_ctype = ContentType.objects.db_manager(
91-
using=(obj._state.db or DEFAULT_DB_ALIAS)
92-
).get_for_model(cls)
93-
nobj.save()
87+
if isinstance(obj, PolymorphicModel):
88+
parent = obj.__class__.objects.using(obj._state.db or DEFAULT_DB_ALIAS).get(pk=obj.pk)
89+
parent.polymorphic_ctype = ctype
90+
parent.save()
9491

95-
return nobj.get_real_instance() # cast to cls
92+
nobj.refresh_from_db() # cast to cls
93+
return nobj

src/polymorphic/tests/migrations/0001_initial.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
# Generated by Django 4.2 on 2025-12-12 18:24
1+
# Generated by Django 4.2 on 2025-12-13 21:30
22

33
from django.conf import settings
44
from django.db import migrations, models

src/polymorphic/tests/test_multidb.py

Lines changed: 139 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -137,3 +137,142 @@ def run():
137137

138138
# Ensure no queries are made using the default database.
139139
self.assertNumQueries(0, run)
140+
141+
def test_create_from_super(self):
142+
# run create test 3 times because initial implementation
143+
# would fail after first success.
144+
from polymorphic.tests.models import (
145+
NormalBase,
146+
NormalExtension,
147+
PolyExtension,
148+
PolyExtChild,
149+
)
150+
151+
nb = NormalBase.objects.db_manager("secondary").create(nb_field=1)
152+
ne = NormalExtension.objects.db_manager("secondary").create(nb_field=2, ne_field="ne2")
153+
154+
with self.assertRaises(TypeError):
155+
PolyExtension.objects.db_manager("secondary").create_from_super(nb, poly_ext_field=3)
156+
157+
pe = PolyExtension.objects.db_manager("secondary").create_from_super(ne, poly_ext_field=3)
158+
159+
ne.refresh_from_db()
160+
self.assertEqual(type(ne), NormalExtension)
161+
self.assertEqual(type(pe), PolyExtension)
162+
self.assertEqual(pe.pk, ne.pk)
163+
164+
self.assertEqual(pe.nb_field, 2)
165+
self.assertEqual(pe.ne_field, "ne2")
166+
self.assertEqual(pe.poly_ext_field, 3)
167+
pe.refresh_from_db()
168+
self.assertEqual(pe.nb_field, 2)
169+
self.assertEqual(pe.ne_field, "ne2")
170+
self.assertEqual(pe.poly_ext_field, 3)
171+
172+
pc = PolyExtChild.objects.db_manager("secondary").create_from_super(
173+
pe, poly_child_field="pcf6"
174+
)
175+
176+
pe.refresh_from_db()
177+
ne.refresh_from_db()
178+
self.assertEqual(type(ne), NormalExtension)
179+
self.assertEqual(type(pe), PolyExtension)
180+
self.assertEqual(pe.pk, ne.pk)
181+
self.assertEqual(pe.pk, pc.pk)
182+
183+
self.assertEqual(pc.nb_field, 2)
184+
self.assertEqual(pc.ne_field, "ne2")
185+
self.assertEqual(pc.poly_ext_field, 3)
186+
pc.refresh_from_db()
187+
self.assertEqual(pc.nb_field, 2)
188+
self.assertEqual(pc.ne_field, "ne2")
189+
self.assertEqual(pc.poly_ext_field, 3)
190+
self.assertEqual(pc.poly_child_field, "pcf6")
191+
192+
self.assertEqual(
193+
pe.polymorphic_ctype,
194+
ContentType.objects.db_manager("secondary").get_for_model(PolyExtChild),
195+
)
196+
self.assertEqual(
197+
pc.polymorphic_ctype,
198+
ContentType.objects.db_manager("secondary").get_for_model(PolyExtChild),
199+
)
200+
201+
self.assertEqual(set(PolyExtension.objects.db_manager("secondary").all()), {pc})
202+
203+
a1 = Model2A.objects.db_manager("secondary").create(field1="A1a")
204+
a2 = Model2A.objects.db_manager("secondary").create(field1="A1b")
205+
206+
b1 = Model2B.objects.db_manager("secondary").create(field1="B1a", field2="B2a")
207+
b2 = Model2B.objects.db_manager("secondary").create(field1="B1b", field2="B2b")
208+
209+
c1 = Model2C.objects.db_manager("secondary").create(
210+
field1="C1a", field2="C2a", field3="C3a"
211+
)
212+
c2 = Model2C.objects.db_manager("secondary").create(
213+
field1="C1b", field2="C2b", field3="C3b"
214+
)
215+
216+
d1 = Model2D.objects.db_manager("secondary").create(
217+
field1="D1a", field2="D2a", field3="D3a", field4="D4a"
218+
)
219+
d2 = Model2D.objects.db_manager("secondary").create(
220+
field1="D1b", field2="D2b", field3="D3b", field4="D4b"
221+
)
222+
223+
with self.assertRaises(TypeError):
224+
Model2D.objects.db_manager("secondary").create_from_super(
225+
b1, field3="D3x", field4="D4x"
226+
)
227+
228+
b1_of_c = Model2B.objects.db_manager("secondary").non_polymorphic().get(pk=c1.pk)
229+
with self.assertRaises(TypeError):
230+
Model2C.objects.db_manager("secondary").create_from_super(b1_of_c, field3="C3x")
231+
232+
self.assertEqual(
233+
c1.polymorphic_ctype,
234+
ContentType.objects.db_manager("secondary").get_for_model(Model2C),
235+
)
236+
dfs1 = Model2D.objects.db_manager("secondary").create_from_super(b1_of_c, field4="D4x")
237+
self.assertEqual(type(dfs1), Model2D)
238+
self.assertEqual(dfs1.pk, c1.pk)
239+
self.assertEqual(dfs1.field1, "C1a")
240+
self.assertEqual(dfs1.field2, "C2a")
241+
self.assertEqual(dfs1.field3, "C3a")
242+
self.assertEqual(dfs1.field4, "D4x")
243+
self.assertEqual(
244+
dfs1.polymorphic_ctype,
245+
ContentType.objects.db_manager("secondary").get_for_model(Model2D),
246+
)
247+
c1.refresh_from_db()
248+
self.assertEqual(
249+
c1.polymorphic_ctype,
250+
ContentType.objects.db_manager("secondary").get_for_model(Model2D),
251+
)
252+
253+
self.assertEqual(
254+
b2.polymorphic_ctype,
255+
ContentType.objects.db_manager("secondary").get_for_model(Model2B),
256+
)
257+
cfs1 = Model2C.objects.db_manager("secondary").create_from_super(b2, field3="C3y")
258+
self.assertEqual(type(cfs1), Model2C)
259+
self.assertEqual(cfs1.pk, b2.pk)
260+
self.assertEqual(cfs1.field1, "B1b")
261+
self.assertEqual(cfs1.field2, "B2b")
262+
self.assertEqual(cfs1.field3, "C3y")
263+
b2.refresh_from_db()
264+
self.assertEqual(
265+
b2.polymorphic_ctype,
266+
ContentType.objects.db_manager("secondary").get_for_model(Model2C),
267+
)
268+
self.assertEqual(
269+
cfs1.polymorphic_ctype,
270+
ContentType.objects.db_manager("secondary").get_for_model(Model2C),
271+
)
272+
273+
self.assertEqual(
274+
set(Model2A.objects.db_manager("secondary").all()),
275+
{a1, a2, b1, dfs1, cfs1, c2, d1, d2},
276+
)
277+
278+
self.assertEqual(Model2A.objects.count(), 0)

src/polymorphic/tests/test_orm.py

Lines changed: 100 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@
1111

1212
from polymorphic import query_translate
1313
from polymorphic.managers import PolymorphicManager
14-
from polymorphic.models import PolymorphicTypeInvalid, PolymorphicTypeUndefined
14+
from polymorphic.models import PolymorphicModel, PolymorphicTypeInvalid, PolymorphicTypeUndefined
1515
from polymorphic.tests.models import (
1616
ArtProject,
1717
Base,
@@ -24,6 +24,7 @@
2424
CustomPkBase,
2525
CustomPkInherit,
2626
Enhance_Base,
27+
Enhance_Plain,
2728
Enhance_Inherit,
2829
InlineParent,
2930
InlineModelA,
@@ -1745,3 +1746,101 @@ def test_polymorphic_extension(self):
17451746
}
17461747
assert set(PolyExtension.objects.all()) == {poly_ext, child_ext}
17471748
assert set(PolyExtChild.objects.all()) == {child_ext}
1749+
1750+
def test_create_from_super(self):
1751+
# run create test 3 times because initial implementation
1752+
# would fail after first success.
1753+
from polymorphic.tests.models import (
1754+
NormalBase,
1755+
NormalExtension,
1756+
PolyExtension,
1757+
PolyExtChild,
1758+
)
1759+
1760+
nb = NormalBase.objects.create(nb_field=1)
1761+
ne = NormalExtension.objects.create(nb_field=2, ne_field="ne2")
1762+
1763+
with self.assertRaises(TypeError):
1764+
PolyExtension.objects.create_from_super(nb, poly_ext_field=3)
1765+
1766+
pe = PolyExtension.objects.create_from_super(ne, poly_ext_field=3)
1767+
1768+
ne.refresh_from_db()
1769+
self.assertEqual(type(ne), NormalExtension)
1770+
self.assertEqual(type(pe), PolyExtension)
1771+
self.assertEqual(pe.pk, ne.pk)
1772+
1773+
self.assertEqual(pe.nb_field, 2)
1774+
self.assertEqual(pe.ne_field, "ne2")
1775+
self.assertEqual(pe.poly_ext_field, 3)
1776+
pe.refresh_from_db()
1777+
self.assertEqual(pe.nb_field, 2)
1778+
self.assertEqual(pe.ne_field, "ne2")
1779+
self.assertEqual(pe.poly_ext_field, 3)
1780+
1781+
pc = PolyExtChild.objects.create_from_super(pe, poly_child_field="pcf6")
1782+
1783+
pe.refresh_from_db()
1784+
ne.refresh_from_db()
1785+
self.assertEqual(type(ne), NormalExtension)
1786+
self.assertEqual(type(pe), PolyExtension)
1787+
self.assertEqual(pe.pk, ne.pk)
1788+
self.assertEqual(pe.pk, pc.pk)
1789+
1790+
self.assertEqual(pc.nb_field, 2)
1791+
self.assertEqual(pc.ne_field, "ne2")
1792+
self.assertEqual(pc.poly_ext_field, 3)
1793+
pc.refresh_from_db()
1794+
self.assertEqual(pc.nb_field, 2)
1795+
self.assertEqual(pc.ne_field, "ne2")
1796+
self.assertEqual(pc.poly_ext_field, 3)
1797+
self.assertEqual(pc.poly_child_field, "pcf6")
1798+
1799+
self.assertEqual(pe.polymorphic_ctype, ContentType.objects.get_for_model(PolyExtChild))
1800+
self.assertEqual(pc.polymorphic_ctype, ContentType.objects.get_for_model(PolyExtChild))
1801+
1802+
self.assertEqual(set(PolyExtension.objects.all()), {pc})
1803+
1804+
a1 = Model2A.objects.create(field1="A1a")
1805+
a2 = Model2A.objects.create(field1="A1b")
1806+
1807+
b1 = Model2B.objects.create(field1="B1a", field2="B2a")
1808+
b2 = Model2B.objects.create(field1="B1b", field2="B2b")
1809+
1810+
c1 = Model2C.objects.create(field1="C1a", field2="C2a", field3="C3a")
1811+
c2 = Model2C.objects.create(field1="C1b", field2="C2b", field3="C3b")
1812+
1813+
d1 = Model2D.objects.create(field1="D1a", field2="D2a", field3="D3a", field4="D4a")
1814+
d2 = Model2D.objects.create(field1="D1b", field2="D2b", field3="D3b", field4="D4b")
1815+
1816+
with self.assertRaises(TypeError):
1817+
Model2D.objects.create_from_super(b1, field3="D3x", field4="D4x")
1818+
1819+
b1_of_c = Model2B.objects.non_polymorphic().get(pk=c1.pk)
1820+
with self.assertRaises(TypeError):
1821+
Model2C.objects.create_from_super(b1_of_c, field3="C3x")
1822+
1823+
self.assertEqual(c1.polymorphic_ctype, ContentType.objects.get_for_model(Model2C))
1824+
dfs1 = Model2D.objects.create_from_super(b1_of_c, field4="D4x")
1825+
self.assertEqual(type(dfs1), Model2D)
1826+
self.assertEqual(dfs1.pk, c1.pk)
1827+
self.assertEqual(dfs1.field1, "C1a")
1828+
self.assertEqual(dfs1.field2, "C2a")
1829+
self.assertEqual(dfs1.field3, "C3a")
1830+
self.assertEqual(dfs1.field4, "D4x")
1831+
self.assertEqual(dfs1.polymorphic_ctype, ContentType.objects.get_for_model(Model2D))
1832+
c1.refresh_from_db()
1833+
self.assertEqual(c1.polymorphic_ctype, ContentType.objects.get_for_model(Model2D))
1834+
1835+
self.assertEqual(b2.polymorphic_ctype, ContentType.objects.get_for_model(Model2B))
1836+
cfs1 = Model2C.objects.create_from_super(b2, field3="C3y")
1837+
self.assertEqual(type(cfs1), Model2C)
1838+
self.assertEqual(cfs1.pk, b2.pk)
1839+
self.assertEqual(cfs1.field1, "B1b")
1840+
self.assertEqual(cfs1.field2, "B2b")
1841+
self.assertEqual(cfs1.field3, "C3y")
1842+
b2.refresh_from_db()
1843+
self.assertEqual(b2.polymorphic_ctype, ContentType.objects.get_for_model(Model2C))
1844+
self.assertEqual(cfs1.polymorphic_ctype, ContentType.objects.get_for_model(Model2C))
1845+
1846+
self.assertEqual(set(Model2A.objects.all()), {a1, a2, b1, dfs1, cfs1, c2, d1, d2})

src/polymorphic/tests/test_recasting.py

Lines changed: 0 additions & 24 deletions
This file was deleted.

0 commit comments

Comments
 (0)