diff --git a/modules/purchase_trade/duplicate.py b/modules/purchase_trade/duplicate.py index 595233e..4184856 100644 --- a/modules/purchase_trade/duplicate.py +++ b/modules/purchase_trade/duplicate.py @@ -290,6 +290,36 @@ class TradeCustomDuplicate(Wizard): 'mirror_clear_pricing_components': False, } + @staticmethod + def _model_has_field(model, field): + fields_ = getattr(model, '_fields', None) + return isinstance(fields_, dict) and field in fields_ + + @staticmethod + def _party_address(party, type_=None): + address_get = getattr(party, 'address_get', None) + if address_get: + address = address_get(type=type_) + if address: + return address + addresses = list(getattr(party, 'addresses', None) or []) + return addresses[0] if addresses else None + + @staticmethod + def _party_address_defaults(model, party): + defaults = {} + if TradeCustomDuplicate._model_has_field(model, 'invoice_party'): + defaults['invoice_party'] = None + if TradeCustomDuplicate._model_has_field(model, 'invoice_address'): + address = TradeCustomDuplicate._party_address(party, type_='invoice') + defaults['invoice_address'] = address.id if address else None + if TradeCustomDuplicate._model_has_field(model, 'shipment_party'): + defaults['shipment_party'] = None + if TradeCustomDuplicate._model_has_field(model, 'shipment_address'): + address = TradeCustomDuplicate._party_address(party, type_='delivery') + defaults['shipment_address'] = address.id if address else None + return defaults + @staticmethod def _copy_contract(record, options): Model = Pool().get(record.__name__) @@ -299,6 +329,8 @@ class TradeCustomDuplicate(Wizard): 'payment_term': options.payment_term.id, 'incoterm': options.incoterm.id, } + default.update(TradeCustomDuplicate._party_address_defaults( + Model, options.party)) new_record, = Model.copy([record], default=default) new_record = Model(new_record.id) TradeCustomDuplicate._apply_party_addresses(new_record, options.party) @@ -341,12 +373,18 @@ class TradeCustomDuplicate(Wizard): @staticmethod def _apply_party_addresses(record, party): - addresses = list(getattr(party, 'addresses', None) or []) - address = addresses[0] if addresses else None + invoice_address = TradeCustomDuplicate._party_address( + party, type_='invoice') + shipment_address = TradeCustomDuplicate._party_address( + party, type_='delivery') + if hasattr(record, 'invoice_party'): + record.invoice_party = None if hasattr(record, 'invoice_address'): - record.invoice_address = address + record.invoice_address = invoice_address + if hasattr(record, 'shipment_party'): + record.shipment_party = None if hasattr(record, 'shipment_address'): - record.shipment_address = address + record.shipment_address = shipment_address @staticmethod def _apply_line_options(record, options): diff --git a/modules/purchase_trade/tests/test_module.py b/modules/purchase_trade/tests/test_module.py index 6e4b843..4ab7784 100644 --- a/modules/purchase_trade/tests/test_module.py +++ b/modules/purchase_trade/tests/test_module.py @@ -8879,6 +8879,54 @@ description self.assertFalse(line.finished) line.save.assert_called_once() + def test_custom_duplicate_copy_defaults_use_new_party_addresses(self): + 'custom duplicate changes party dependent addresses during copy' + invoice_address = Mock(id=20) + shipment_address = Mock(id=30) + party = Mock(id=10) + party.address_get.side_effect = [invoice_address, shipment_address] + options = Mock( + party=party, + currency=Mock(id=40), + payment_term=Mock(id=50), + incoterm=Mock(id=60), + ) + record = Mock(__name__='sale.sale') + created = Mock(id=70) + reloaded = Mock(id=70) + Model = Mock() + Model._fields = { + 'invoice_party': None, + 'invoice_address': None, + 'shipment_party': None, + 'shipment_address': None, + } + Model.copy.return_value = [created] + Model.return_value = reloaded + pool = Mock() + pool.get.return_value = Model + + with patch.object(duplicate_module, 'Pool', return_value=pool): + result = duplicate_module.TradeCustomDuplicate._copy_contract( + record, options) + + Model.copy.assert_called_once_with([record], default={ + 'party': 10, + 'currency': 40, + 'payment_term': 50, + 'incoterm': 60, + 'invoice_party': None, + 'invoice_address': 20, + 'shipment_party': None, + 'shipment_address': 30, + }) + self.assertIs(result, reloaded) + self.assertIsNone(reloaded.invoice_party) + self.assertIs(reloaded.invoice_address, invoice_address) + self.assertIsNone(reloaded.shipment_party) + self.assertIs(reloaded.shipment_address, shipment_address) + reloaded.save.assert_called_once() + def test_custom_duplicate_creates_purchase_open_virtual_lot(self): 'custom duplicate creates the purchase opening lot.qt explicitly' saved_lots = []