From 5f81ef949428dd0a85c0b23a621bd157e617e2e7 Mon Sep 17 00:00:00 2001 From: caoqianming Date: Tue, 28 Jul 2026 15:27:39 +0800 Subject: [PATCH] Prevent duplicate workshop inventory creation --- apps/inm/services.py | 16 +++--- apps/inm/views.py | 7 ++- apps/qm/services.py | 29 +++++++--- apps/utils/models.py | 8 ++- apps/wpm/models.py | 119 +++++++++++++++++++++++++++++++++++++++- apps/wpm/serializers.py | 15 +++++ apps/wpm/services.py | 22 +++++--- apps/wpm/tests.py | 83 +++++++++++++++++++++++++++- 8 files changed, 269 insertions(+), 30 deletions(-) diff --git a/apps/inm/services.py b/apps/inm/services.py index 6fb4203c..4504b620 100644 --- a/apps/inm/services.py +++ b/apps/inm/services.py @@ -7,8 +7,10 @@ from apps.wpm.models import WMaterial, BatchSt, BatchLog from apps.wpm.services_2 import ana_batch_thread from apps.wpmw.models import Wpr from apps.qm.models import Ftest, Defect +from django.db import transaction from django.db.models import Count, Q +@transaction.atomic def do_out(item: MIOItem, is_reverse: bool = False): """ 生产领料到车间 @@ -45,7 +47,7 @@ def do_out(item: MIOItem, is_reverse: bool = False): if is_zhj: try: - mb = MaterialBatch.objects.get( + mb = MaterialBatch.objects.select_for_update().get( material=item.material, warehouse=item.warehouse, batch=item.batch, @@ -82,7 +84,7 @@ def do_out(item: MIOItem, is_reverse: bool = False): mb = None if not is_zhj: try: - mb = MaterialBatch.objects.get( + mb = MaterialBatch.objects.select_for_update().get( material=xmaterial, warehouse=item.warehouse, batch=xbatch, @@ -99,7 +101,7 @@ def do_out(item: MIOItem, is_reverse: bool = False): if xmaterial.into_wm: # 领到车间库存(或工段) - wm, new_create = WMaterial.objects.get_or_create( + wm, new_create = WMaterial.locked_get_or_create_inventory( batch=xbatch, material=xmaterial, belong_dept=belong_dept, mgroup=mgroup, state=state, defect=defect) @@ -107,7 +109,7 @@ def do_out(item: MIOItem, is_reverse: bool = False): wm.create_by = do_user wm.batch_ofrom = mb.batch if mb else None wm.material_ofrom = mb.material if mb else None - wm.count = wm.count + item.count + wm.count = wm.count + xcount wm.update_by = do_user wm.save() @@ -130,6 +132,7 @@ def do_out(item: MIOItem, is_reverse: bool = False): ana_batch_thread(xbatches) +@transaction.atomic def do_in(item: MIOItem): """ 生产入库后更新车间物料 @@ -184,9 +187,9 @@ def do_in(item: MIOItem): xbatchs.append(xbatch) if xmaterial.into_wm: if xwm: - wm = xwm + wm = WMaterial.objects.select_for_update().get(pk=xwm.pk) else: - wm_qs = WMaterial.objects.filter( + wm_qs = WMaterial.objects.select_for_update().filter( batch=xbatch, material=xmaterial, belong_dept=belong_dept, @@ -488,4 +491,3 @@ class InmService: # 若该出入库记录已无明细,自动删除 if not MIOItem.objects.filter(mio=mio).exists(): mio.delete() - \ No newline at end of file diff --git a/apps/inm/views.py b/apps/inm/views.py index 8e96d177..c6984703 100644 --- a/apps/inm/views.py +++ b/apps/inm/views.py @@ -255,7 +255,8 @@ class MIOViewSet(CustomModelViewSet): 提交 """ - ins:MIO = self.get_object() + current = self.get_object() + ins = MIO.objects.select_for_update().get(pk=current.pk) if ins.inout_date is None: raise ParseError('出入库日期未填写') if ins.state != MIO.MIO_CREATE: @@ -276,7 +277,8 @@ class MIOViewSet(CustomModelViewSet): 撤回 """ - ins = self.get_object() + current = self.get_object() + ins = MIO.objects.select_for_update().get(pk=current.pk) user = self.request.user if ins.state != MIO.MIO_SUBMITED: raise ParseError('记录状态异常') @@ -586,4 +588,3 @@ class MIOItemwViewSet(CustomModelViewSet): if ftest: ftest.delete() self.cal_mioitem_count(mioitem) - \ No newline at end of file diff --git a/apps/qm/services.py b/apps/qm/services.py index 8de5259e..fdae864f 100644 --- a/apps/qm/services.py +++ b/apps/qm/services.py @@ -7,6 +7,7 @@ from apps.wf.models import Ticket from apps.qm.models import NotOkOption, Defect from apps.wpm.services_2 import ana_batch_thread from apps.inm.models import MaterialBatch +from django.db import transaction def ftestwork_submit_validate(ins: FtestWork): wm:WMaterial = ins.wm @@ -21,8 +22,15 @@ def ftestwork_submit_validate(ins: FtestWork): raise ParseError("不合格数不可大于批次数量") +@transaction.atomic def ftestwork_submit(ins:FtestWork, user: User): - wm:WMaterial = ins.wm + ins = FtestWork.objects.select_for_update().get(pk=ins.pk) + if ins.submit_time is not None: + raise ParseError('该检验工作已提交') + wm = ( + WMaterial.objects.select_for_update().get(pk=ins.wm_id) + if ins.wm_id else None + ) fwd_qs = FtestworkDefect.objects.filter(ftestwork=ins) if wm and ins.need_update_wm: if ins.qct is None and not fwd_qs.exists(): @@ -46,7 +54,7 @@ def ftestwork_submit(ins:FtestWork, user: User): need_move_count = need_move_count + v count_ok = ins.count_ok - need_move_count if count_ok > 0: - wm, new_create = WMaterial.objects.get_or_create( + wm, new_create = WMaterial.locked_get_or_create_inventory( material=wm.material, batch=wm.batch, mgroup=wm.mgroup, @@ -77,7 +85,7 @@ def ftestwork_submit(ins:FtestWork, user: User): astate = WMaterial.WM_NOTOK if NotOkOption.get_extra_info(notok_sign)['cate'] == 'ok_b': astate = WMaterial.WM_OK - wm_n, new_create = WMaterial.objects.get_or_create( + wm_n, new_create = WMaterial.locked_get_or_create_inventory( material=wm.material, batch=wm.batch, mgroup=wm.mgroup, @@ -110,7 +118,7 @@ def ftestwork_submit(ins:FtestWork, user: User): wmstate = WMaterial.WM_OK if item.defect.okcate == Defect.DEFECT_NOTOK: wmstate = WMaterial.WM_NOTOK - wmx, new_create = WMaterial.objects.get_or_create( + wmx, new_create = WMaterial.locked_get_or_create_inventory( material=wm.material, batch=wm.batch, mgroup=wm.mgroup, @@ -127,7 +135,7 @@ def ftestwork_submit(ins:FtestWork, user: User): wmx.save() if ins.mb: - mb:MaterialBatch = ins.mb + mb = MaterialBatch.objects.select_for_update().get(pk=ins.mb_id) for item in fwd_qs: item:FtestworkDefect = item if item.count > 0: @@ -158,8 +166,15 @@ def ftestwork_submit(ins:FtestWork, user: User): ana_batch_thread(xbatchs=[ins.batch]) +@transaction.atomic def ftestwork_revert(ins: FtestWork): - wm:WMaterial = ins.wm + ins = FtestWork.objects.select_for_update().get(pk=ins.pk) + if ins.submit_time is None: + raise ParseError('该检验工作未提交') + wm = ( + WMaterial.objects.select_for_update().get(pk=ins.wm_id) + if ins.wm_id else None + ) if wm and ins.need_update_wm: fwd_qs = FtestworkDefect.objects.filter(ftestwork=ins) for item in fwd_qs: @@ -213,4 +228,4 @@ def bind_ftestwork(ticket: Ticket, transition, new_ticket_data: dict): def ftestwork_audit_end(ticket: Ticket): ins = FtestWork.objects.get(id=ticket.ticket_data['t_id']) - ftestwork_submit(ins, ticket.create_by) \ No newline at end of file + ftestwork_submit(ins, ticket.create_by) diff --git a/apps/utils/models.py b/apps/utils/models.py index 6fc66c5b..bfae6a5e 100755 --- a/apps/utils/models.py +++ b/apps/utils/models.py @@ -154,8 +154,12 @@ class BaseModel(models.Model): @classmethod def locked_get_or_create(cls, defaults: dict, **kwargs): """ - 仅用于事务内 - 并发安全的 get_or_create + 仅用于事务内锁定已存在的记录。 + + PostgreSQL 无法通过 select_for_update 锁定不存在的记录,因此该方法 + 不保证首次创建并发安全。需要防止首次重复创建的业务应提供稳定的业务 + 键并使用专用 advisory lock;车间库存使用 + WMaterial.locked_get_or_create_inventory。 """ if not connection.in_atomic_block: raise RuntimeError("locked_get_or_create 必须在事务中调用") diff --git a/apps/wpm/models.py b/apps/wpm/models.py index 7407713c..2275a67a 100644 --- a/apps/wpm/models.py +++ b/apps/wpm/models.py @@ -11,8 +11,9 @@ from django.db.models import Sum, Subquery from django.utils.translation import gettext_lazy as _ from rest_framework.exceptions import ParseError from django.db.models import Count -from django.db import transaction +from django.db import connection, transaction from django.db.models import Max +import json import re from django.db.models import Q, F import django.utils.timezone as timezone @@ -131,6 +132,122 @@ class WMaterial(CommonBDModel): number_from = models.TextField("来源于个号", null=True, blank=True) is_manual = models.BooleanField('手动创建', default=False) + INVENTORY_KEY_FIELDS = ( + 'material', + 'batch', + 'mgroup', + 'belong_dept', + 'state', + 'defect', + 'notok_sign', + 'material_origin', + 'state_origin', + ) + + @classmethod + def _normalize_inventory_lookup(cls, **kwargs): + """生成唯一、完整的库存业务键,避免省略 NULL 字段产生不同锁键。""" + unknown_fields = set(kwargs) - set(cls.INVENTORY_KEY_FIELDS) + if unknown_fields: + fields = ', '.join(sorted(unknown_fields)) + raise TypeError(f'不支持的车间库存定位字段: {fields}') + + if kwargs.get('material') is None or kwargs.get('batch') is None: + raise ValueError('车间库存业务键必须包含 material 和 batch') + + lookup = { + field: kwargs.get(field) + for field in cls.INVENTORY_KEY_FIELDS + } + if lookup['state'] is None: + lookup['state'] = cls._meta.get_field('state').get_default() + + mgroup = lookup['mgroup'] + belong_dept = lookup['belong_dept'] + if mgroup is not None: + mgroup_dept_id = getattr(mgroup, 'belong_dept_id', None) + belong_dept_id = getattr(belong_dept, 'pk', belong_dept) + if belong_dept is None: + lookup['belong_dept'] = mgroup.belong_dept + elif mgroup_dept_id != belong_dept_id: + raise ValueError('车间库存的工段与所属部门不匹配') + + return lookup + + @classmethod + def _inventory_advisory_lock_payload(cls, lookup): + lock_values = {} + for name in cls.INVENTORY_KEY_FIELDS: + field = cls._meta.get_field(name) + value = lookup[name] + if field.is_relation and value is not None: + value = getattr(value, 'pk', value) + lock_values[field.attname] = value + return json.dumps( + {'model': cls._meta.label_lower, 'lookup': lock_values}, + sort_keys=True, + ensure_ascii=False, + default=str, + separators=(',', ':'), + ) + + @classmethod + def locked_get_or_create_inventory(cls, defaults=None, **kwargs): + """ + 在事务中按完整库存业务键获取或创建记录。 + + 已存在记录使用行锁;首次创建使用 PostgreSQL 事务级 advisory lock, + 并在取得锁后重新查询,避免两个事务同时创建相同库存。 + """ + if not connection.in_atomic_block: + raise RuntimeError( + 'locked_get_or_create_inventory 必须在事务中调用' + ) + if connection.vendor != 'postgresql': + raise RuntimeError( + 'locked_get_or_create_inventory 仅支持 PostgreSQL' + ) + + defaults = defaults or {} + lookup = cls._normalize_inventory_lookup(**kwargs) + create_defaults = { + key: value + for key, value in defaults.items() + if key not in cls.INVENTORY_KEY_FIELDS + } + + rows = list( + cls.objects.select_for_update().filter(**lookup)[:2] + ) + if len(rows) > 1: + raise RuntimeError( + f'{cls.__name__} 数据异常:库存业务键 {lookup} 命中多条' + ) + if rows: + return rows[0], False + + lock_payload = cls._inventory_advisory_lock_payload(lookup) + with connection.cursor() as cursor: + cursor.execute( + 'SELECT pg_advisory_xact_lock(hashtextextended(%s, 0))', + [lock_payload], + ) + + rows = list( + cls.objects.select_for_update().filter(**lookup)[:2] + ) + if len(rows) > 1: + raise RuntimeError( + f'{cls.__name__} 数据异常:库存业务键 {lookup} 命中多条' + ) + if rows: + return rows[0], False + + return cls.objects.create( + **lookup, + **create_defaults, + ), True + def delete(self, *args, **kwargs): if not self.is_manual: raise ParseError('只能删除手动创建的车间库存') diff --git a/apps/wpm/serializers.py b/apps/wpm/serializers.py index 42a5a66c..d6d94b54 100644 --- a/apps/wpm/serializers.py +++ b/apps/wpm/serializers.py @@ -247,6 +247,21 @@ class WMaterialCreateSerializer(CustomModelSerializer): attrs['belong_dept'] = mgroup.belong_dept return attrs + @transaction.atomic + def create(self, validated_data): + lookup = { + field: validated_data.pop(field) + for field in WMaterial.INVENTORY_KEY_FIELDS + if field in validated_data + } + instance, created = WMaterial.locked_get_or_create_inventory( + **lookup, + defaults=validated_data, + ) + if not created: + raise serializers.ValidationError('相同业务键的车间库存已存在') + return instance + class MlogbDefectSerializer(CustomModelSerializer): defect_name = serializers.CharField(source="defect.name", read_only=True) diff --git a/apps/wpm/services.py b/apps/wpm/services.py index 35010d0f..2fd9e219 100644 --- a/apps/wpm/services.py +++ b/apps/wpm/services.py @@ -298,7 +298,8 @@ def mlog_submit(mlog: Mlog, user: User, now: Union[datetime.datetime, None]): 'state': c_state, **stored_location, } - wm, is_create = WMaterial.locked_get_or_create(**lookup, defaults={}) + wm, is_create = WMaterial.locked_get_or_create_inventory( + **lookup, defaults={}) wm.count = wm.count + count if is_create: wm.create_by = user @@ -388,7 +389,8 @@ def mlog_submit(mlog: Mlog, user: User, now: Union[datetime.datetime, None]): lookup['defect'] = notok_sign_or_defect elif notok_sign_or_defect is not None: lookup['notok_sign'] = notok_sign_or_defect - wm, is_create2 = WMaterial.locked_get_or_create(**lookup, defaults={}) + wm, is_create2 = WMaterial.locked_get_or_create_inventory( + **lookup, defaults={}) wm.count = wm.count + mo_count wm.count_eweight = mo_count_eweight wm.update_by = user @@ -617,7 +619,8 @@ def mlog_revert(mlog: Mlog, user: User, now: Union[datetime.datetime, None]): 'state': WMaterial.WM_OK, **stored_location, } - wm, _ = WMaterial.locked_get_or_create(**lookup, defaults={}) + wm, _ = WMaterial.locked_get_or_create_inventory( + **lookup, defaults={}) wm.count = wm.count + mi_count wm.update_by = user wm.save() @@ -644,7 +647,8 @@ def mlog_revert(mlog: Mlog, user: User, now: Union[datetime.datetime, None]): 'state': c_state, **stored_location, } - wm, is_create = WMaterial.locked_get_or_create(**lookup, defaults={}) + wm, is_create = WMaterial.locked_get_or_create_inventory( + **lookup, defaults={}) wm.count = wm.count - count if wm.count < 0: raise ParseError('加工前不良数量大于库存量') @@ -879,7 +883,7 @@ def handover_submit(handover:Handover, user: User, now: Union[datetime.datetime, if wm_to.state != wm_from.state or wm_to.material != wm_from.material or not defect_ok: raise ParseError("正常合并到的车间库存状态或物料异常") else: - wm_to, _ = WMaterial.locked_get_or_create( + wm_to, _ = WMaterial.locked_get_or_create_inventory( batch=batch, material=material, mgroup=recive_mgroup, @@ -903,7 +907,7 @@ def handover_submit(handover:Handover, user: User, now: Union[datetime.datetime, if wm_to.state != WMaterial.WM_REPAIR or wm_to.material != wm_from.material or wm_to.defect != wm_from.defect: raise ParseError("返修合并到的车间库存状态或物料异常") elif recive_mgroup: - wm_to, _ = WMaterial.locked_get_or_create( + wm_to, _ = WMaterial.locked_get_or_create_inventory( batch=batch, material=material, mgroup=recive_mgroup, @@ -927,7 +931,7 @@ def handover_submit(handover:Handover, user: User, now: Union[datetime.datetime, if wm_to.state != WMaterial.WM_SCRAP or wm_to.material != wm_from.material or wm_to.defect != wm_from.defect: raise ParseError("报废合并到的车间库存状态或物料异常") elif recive_mgroup: - wm_to, _ = WMaterial.locked_get_or_create( + wm_to, _ = WMaterial.locked_get_or_create_inventory( batch=batch, material=material, mgroup=recive_mgroup, @@ -950,7 +954,7 @@ def handover_submit(handover:Handover, user: User, now: Union[datetime.datetime, if wm_to.material != handover.material_changed or wm_to.state != handover.state_changed: raise ParseError("改版合并到的车间库存状态或物料异常") elif handover.recive_mgroup: - wm_to, _ = WMaterial.locked_get_or_create( + wm_to, _ = WMaterial.locked_get_or_create_inventory( batch=batch, material=handover.material_changed, state=handover.state_changed, @@ -975,7 +979,7 @@ def handover_submit(handover:Handover, user: User, now: Union[datetime.datetime, if mtype == Handover.H_MERGE and handover.new_wm: wm_to = WMaterial.objects.select_for_update().get(id=handover.new_wm.id) else: - wm_to, _ = WMaterial.locked_get_or_create( + wm_to, _ = WMaterial.locked_get_or_create_inventory( batch=batch, material=material, mgroup=recive_mgroup, diff --git a/apps/wpm/tests.py b/apps/wpm/tests.py index d935f054..ac8e412b 100644 --- a/apps/wpm/tests.py +++ b/apps/wpm/tests.py @@ -1,7 +1,12 @@ +from concurrent.futures import ThreadPoolExecutor +from decimal import Decimal +from threading import Barrier from types import SimpleNamespace from unittest.mock import MagicMock, patch -from django.test import SimpleTestCase +from django.db import connection, connections, transaction +from django.test import SimpleTestCase, TransactionTestCase +from unittest import skipUnless from apps.mtm.models import Material, WmScope from apps.wpm.filters import WMaterialFilter @@ -66,6 +71,32 @@ class MlogbwViewSetTests(SimpleTestCase): class WMaterialScopeTests(SimpleTestCase): + def test_inventory_key_normalizes_omitted_nullable_fields(self): + material = Material(id="100", name="测试物料") + + omitted = WMaterial._normalize_inventory_lookup( + material=material, + batch="BATCH-001", + state=WMaterial.WM_OK, + ) + explicit = WMaterial._normalize_inventory_lookup( + material=material, + batch="BATCH-001", + mgroup=None, + belong_dept=None, + state=WMaterial.WM_OK, + defect=None, + notok_sign=None, + material_origin=None, + state_origin=None, + ) + + self.assertEqual(omitted, explicit) + self.assertEqual( + WMaterial._inventory_advisory_lock_payload(omitted), + WMaterial._inventory_advisory_lock_payload(explicit), + ) + def test_scope_resolver_returns_consistent_location_fields(self): dept = object() mgroup = SimpleNamespace(belong_dept=dept) @@ -210,3 +241,53 @@ class WMaterialScopeTests(SimpleTestCase): "type": Handover.H_SCRAP, "mtype": Handover.H_NORMAL, }) + + +@skipUnless( + connection.vendor == "postgresql", + "advisory lock concurrency test requires PostgreSQL", +) +class WMaterialConcurrencyTests(TransactionTestCase): + def test_concurrent_first_create_uses_one_inventory_record(self): + material = Material.objects.create(name="并发库存测试物料") + barrier = Barrier(2) + + def create_inventory(): + connections.close_all() + try: + barrier.wait() + with transaction.atomic(): + wm, created = ( + WMaterial.locked_get_or_create_inventory( + material=material, + batch="CONCURRENT-001", + state=WMaterial.WM_OK, + defaults={"count": Decimal("0")}, + ) + ) + wm.count += Decimal("1") + wm.save(update_fields=["count"]) + return created + finally: + connections.close_all() + + with ThreadPoolExecutor(max_workers=2) as executor: + created_results = list(executor.map( + lambda _: create_inventory(), + range(2), + )) + + queryset = WMaterial.objects.filter( + material=material, + batch="CONCURRENT-001", + mgroup=None, + belong_dept=None, + state=WMaterial.WM_OK, + defect=None, + notok_sign=None, + material_origin=None, + state_origin=None, + ) + self.assertEqual(queryset.count(), 1) + self.assertEqual(queryset.get().count, Decimal("2")) + self.assertCountEqual(created_results, [True, False])