diff --git a/apps/utils/models.py b/apps/utils/models.py index 7c08da61..30ae9101 100755 --- a/apps/utils/models.py +++ b/apps/utils/models.py @@ -31,7 +31,27 @@ class SoftDeletableQuerySetMixin(object): return super(SoftDeletableQuerySetMixin, self).delete() -class SoftDeletableQuerySet(SoftDeletableQuerySetMixin, QuerySet): +class SnowflakeQuerySet(QuerySet): + """为 BaseModel 的批量创建统一准备雪花主键。""" + + def bulk_create(self, objs, *args, **kwargs): + objs = list(objs) + for obj in objs: + if not obj.pk: + obj.pk = str(idWorker.get_id()) + else: + obj.pk = str(obj.pk) + return super().bulk_create(objs, *args, **kwargs) + + +class SnowflakeManager(models.Manager.from_queryset(SnowflakeQuerySet)): + pass + + +class SoftDeletableQuerySet( + SoftDeletableQuerySetMixin, + SnowflakeQuerySet, +): pass @@ -111,6 +131,8 @@ class BaseModel(models.Model): is_deleted = models.BooleanField( default=False, verbose_name='删除标记', help_text='删除标记') + objects = SnowflakeManager() + class Meta: abstract = True diff --git a/apps/utils/tests.py b/apps/utils/tests.py new file mode 100644 index 00000000..6301e97c --- /dev/null +++ b/apps/utils/tests.py @@ -0,0 +1,116 @@ +from unittest.mock import patch + +from django.db.models.query import QuerySet +from django.test import SimpleTestCase, TestCase + +from apps.develop.models import Project +from apps.mtm.models import Material +from apps.system.models import User +from apps.wpmw.models import WprDefect + + +class SnowflakeBulkCreateTests(SimpleTestCase): + @patch.object(QuerySet, "bulk_create") + @patch( + "apps.utils.models.idWorker.get_id", + side_effect=[100, 101], + ) + def test_base_model_bulk_create_assigns_missing_ids( + self, + get_id, + bulk_create, + ): + objs = [ + WprDefect(defect_id="10", wpr_id="20"), + WprDefect(defect_id="11", wpr_id="20"), + ] + + WprDefect.objects.bulk_create(objs) + + self.assertEqual([obj.id for obj in objs], ["100", "101"]) + self.assertEqual(get_id.call_count, 2) + bulk_create.assert_called_once() + + @patch.object(QuerySet, "bulk_create") + @patch("apps.utils.models.idWorker.get_id") + def test_base_model_bulk_create_preserves_existing_ids( + self, + get_id, + bulk_create, + ): + objs = [ + WprDefect(id=100, defect_id="10", wpr_id="20"), + WprDefect(id="101", defect_id="11", wpr_id="20"), + ] + + WprDefect.objects.bulk_create(objs) + + self.assertEqual([obj.id for obj in objs], ["100", "101"]) + get_id.assert_not_called() + bulk_create.assert_called_once() + + @patch.object(QuerySet, "bulk_create") + @patch( + "apps.utils.models.idWorker.get_id", + side_effect=[200, 201], + ) + def test_soft_delete_manager_uses_snowflake_bulk_create( + self, + get_id, + bulk_create, + ): + objs = [ + Material(name="物料1"), + Material(name="物料2"), + ] + + Material.objects.bulk_create(objs) + + self.assertEqual([obj.id for obj in objs], ["200", "201"]) + self.assertEqual(get_id.call_count, 2) + bulk_create.assert_called_once() + + @patch.object(QuerySet, "bulk_create") + @patch( + "apps.utils.models.idWorker.get_id", + side_effect=[300, 301], + ) + def test_custom_user_manager_uses_snowflake_bulk_create( + self, + get_id, + bulk_create, + ): + objs = [ + User(username="bulk-user-1"), + User(username="bulk-user-2"), + ] + + User.objects.bulk_create(objs) + + self.assertEqual([obj.id for obj in objs], ["300", "301"]) + self.assertEqual(get_id.call_count, 2) + bulk_create.assert_called_once() + + +class SnowflakeBulkCreateDatabaseTests(TestCase): + @patch( + "apps.utils.models.idWorker.get_id", + side_effect=[900000000000000001, 900000000000000002], + ) + def test_bulk_create_persists_generated_ids(self, get_id): + objs = [ + Project(name="批量项目1", code="bulk-project-1"), + Project(name="批量项目2", code="bulk-project-2"), + ] + + Project.objects.bulk_create(objs) + + self.assertEqual( + [obj.id for obj in objs], + ["900000000000000001", "900000000000000002"], + ) + self.assertEqual( + set(Project.objects.values_list("id", flat=True)), + {"900000000000000001", "900000000000000002"}, + ) + self.assertEqual(get_id.call_count, 2)