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.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 from apps.wpm.models import Handover, WMaterial from apps.wpm.serializers import HandoverSerializer, WMaterialCreateSerializer from apps.wpm.views import MlogbwViewSet from rest_framework.exceptions import ParseError class MlogbwViewSetTests(SimpleTestCase): @patch("apps.wpm.views.MlogViewSet.lock_and_check_can_update") @patch("apps.wpm.views.Mlogbw.cal_count_notok") @patch("apps.wpm.views.Mlogbw.objects.get_or_create") @patch("apps.wpm.views.Mlogb.objects.filter") def test_perform_create_syncs_fix_output_without_route( self, mlogb_filter, mlogbw_get_or_create, cal_count_notok, lock_and_check_can_update, ): wm = object() wpr = SimpleNamespace(wm=wm) material = SimpleNamespace(tracking=Material.MA_TRACKING_SINGLE) mlog = SimpleNamespace( route=None, is_fix=True, cal_mlog_count_from_mlogb=MagicMock(), ) mlogb_in = SimpleNamespace( mlog=mlog, route=None, wm_in=wm, material_in=material, ) mlogb_out = object() mlogb_qs = MagicMock() mlogb_qs.exists.return_value = True mlogb_qs.__iter__.return_value = iter([mlogb_out]) mlogb_filter.return_value = mlogb_qs ins = SimpleNamespace( mlogb=mlogb_in, wpr=wpr, number="2606P1095-3", ) serializer = MagicMock() serializer.save.return_value = ins lock_and_check_can_update.return_value = mlog MlogbwViewSet().perform_create(serializer) mlogbw_get_or_create.assert_called_once_with( mlogb=mlogb_out, wpr=wpr, defaults={ "number": "2606P1095-3", "mlogbw_from": ins, }, ) self.assertEqual(cal_count_notok.call_count, 2) mlog.cal_mlog_count_from_mlogb.assert_called_once_with() 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) self.assertEqual( WmScope.resolve_location(WmScope.MGROUP, mgroup), {"mgroup": mgroup, "belong_dept": dept}, ) self.assertEqual( WmScope.resolve_location(WmScope.DEPT, mgroup), {"mgroup": None, "belong_dept": dept}, ) self.assertEqual( WmScope.resolve_location(WmScope.GLOBAL, mgroup), {"mgroup": None, "belong_dept": None}, ) def test_only_mgroup_scope_requires_handover_mgroup(self): self.assertTrue(WmScope.requires_mgroup(WmScope.MGROUP)) self.assertFalse(WmScope.requires_mgroup(WmScope.DEPT)) self.assertFalse(WmScope.requires_mgroup(WmScope.GLOBAL)) def test_global_scope_key_does_not_access_missing_relations(self): wm = WMaterial() self.assertEqual(wm.belong_dept_or_mgroup_id, ("global", None)) def test_department_and_mgroup_scope_keys_cannot_collide(self): dept_wm = WMaterial(belong_dept_id=10) mgroup_wm = WMaterial(mgroup_id=10) self.assertNotEqual( dept_wm.belong_dept_or_mgroup_id, mgroup_wm.belong_dept_or_mgroup_id, ) def test_manual_create_allows_global_scope(self): serializer = WMaterialCreateSerializer() attrs = { "material": Material(), "count": 1, "batch": "GLOBAL-001", } validated = serializer.validate(attrs) self.assertFalse(serializer.fields["mgroup"].required) self.assertFalse(serializer.fields["belong_dept"].required) self.assertNotIn("mgroup", validated) self.assertNotIn("belong_dept", validated) def test_global_scope_can_be_filtered_explicitly(self): self.assertIn( "isnull", WMaterialFilter.Meta.fields["belong_dept"], ) self.assertIn( "isnull", WMaterialFilter.Meta.fields["mgroup"], ) def test_manual_create_derives_department_from_mgroup(self): dept = SimpleNamespace(id=20) mgroup = SimpleNamespace(id=10, belong_dept=dept) attrs = { "material": Material(), "count": 1, "batch": "MGROUP-001", "mgroup": mgroup, } validated = WMaterialCreateSerializer().validate(attrs) self.assertIs(validated["belong_dept"], dept) def test_global_inventory_cannot_enter_normal_handover(self): wm = WMaterial( material=Material(tracking=Material.MA_TRACKING_BATCH), batch="GLOBAL-001", count=1, ) with self.assertRaisesMessage(ParseError, "全局库存无需正常交接"): HandoverSerializer().validate({ "wm": wm, "count": 1, "type": Handover.H_NORMAL, "mtype": Handover.H_NORMAL, }) def test_global_inventory_split_stays_global(self): wm = WMaterial( material=Material(tracking=Material.MA_TRACKING_BATCH), batch="GLOBAL-001", count=1, ) validated = HandoverSerializer().validate({ "wm": wm, "count": 1, "type": Handover.H_NORMAL, "mtype": Handover.H_DIV, }) self.assertIsNone(validated["send_dept"]) self.assertIsNone(validated["recive_dept"]) self.assertIsNone(validated["recive_mgroup"]) def test_global_inventory_cannot_merge_into_scoped_inventory(self): material = Material(tracking=Material.MA_TRACKING_BATCH) wm = WMaterial(material=material, batch="GLOBAL-001", count=1) target = WMaterial( material=material, batch="GLOBAL-MERGED", count=0, belong_dept_id=20, ) with self.assertRaisesMessage(ParseError, "全局库存合批目标必须是全局库存"): HandoverSerializer().validate({ "wm": wm, "count": 1, "new_wm": target, "type": Handover.H_NORMAL, "mtype": Handover.H_MERGE, }) def test_global_inventory_scrap_requires_receiving_mgroup(self): wm = WMaterial( material=Material(tracking=Material.MA_TRACKING_BATCH), batch="GLOBAL-001", count=1, ) with self.assertRaisesMessage( ParseError, "全局库存返修、报废或改版必须指定接收工段", ): HandoverSerializer().validate({ "wm": wm, "count": 1, "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])