diff --git a/delivery_carrier_partner_account/models/carrier_account_mixin.py b/delivery_carrier_partner_account/models/carrier_account_mixin.py index 86dffd9..c8f41f6 100644 --- a/delivery_carrier_partner_account/models/carrier_account_mixin.py +++ b/delivery_carrier_partner_account/models/carrier_account_mixin.py @@ -100,32 +100,24 @@ class CarrierAccountMixin(models.AbstractModel): Returns: Tuple containing: - partners: recordset of partners that can have valid accounts - - is_third_party: boolean indicating if this is a third party billing mode - - invalid_partners: for third party mode, partners that cannot own the account """ self.ensure_one() - invalid_partners = self.env["res.partner"].browse() + invalid_partners = self.env["res.partner"] match self.delivery_billing_mode: case "collect": - return ( - self.recipient_id | self.recipient_id.commercial_partner_id, - False, - invalid_partners, - ) + return self.recipient_id | self.recipient_id.commercial_partner_id case "prepaid" | "ppc" | "no charge": - return ( - self.sender_id | self.sender_id.commercial_partner_id, - False, - invalid_partners, - ) + return self.sender_id | self.sender_id.commercial_partner_id case "third party": invalid_partners = ( self.recipient_id | self.recipient_id.commercial_partner_id ) | (self.sender_id | self.sender_id.commercial_partner_id) - return self.env["res.partner"].browse(), True, invalid_partners + return self.env["res.partner"].search( + [("id", "not in", invalid_partners.ids)] + ) case _: - return self.env["res.partner"].browse(), False, invalid_partners + return self.env["res.partner"] @api.depends( "sender_id", @@ -134,40 +126,26 @@ class CarrierAccountMixin(models.AbstractModel): ) def _compute_carrier_id(self): for rec in self: - # Only set carrier if not already set if not rec.carrier_id: rec.carrier_id = rec._get_default_carrier() - continue + elif ( + rec.delivery_billing_mode + and rec.carrier_id not in rec.valid_carrier_ids + and not rec._has_valid_account_for_carrier() + ): + rec.carrier_id = rec._get_default_carrier() - # Don't reset carrier if we're changing billing mode and there's a valid account - if rec.delivery_billing_mode: - partners, is_third_party, invalid_partners = ( - rec._get_valid_carrier_partners() - ) - - if is_third_party: - # For third party, check if there's any account for this carrier - # that doesn't belong to sender or recipient - has_valid_account = bool( - self.env["delivery.carrier.account"].search_count( - [ - ("delivery_carrier_id", "=", rec.carrier_id.id), - ("partner_id", "not in", invalid_partners.ids), - ] - ) - ) - if has_valid_account: - continue - else: - # Check if there's a valid account for this carrier - if any( - account.delivery_carrier_id == rec.carrier_id - for account in partners.mapped("carrier_account_ids") - ): - continue - - # If we get here, there's no valid account for this carrier - rec.carrier_id = rec._get_default_carrier() + def _has_valid_account_for_carrier(self): + """Check if there's a valid account for the current carrier and billing mode.""" + self.ensure_one() + if not self.carrier_id or not self.delivery_billing_mode: + return False + + partners = self._get_valid_carrier_partners() + valid_accounts = partners.mapped("carrier_account_ids").filtered( + lambda account: account.delivery_carrier_id == self.carrier_id + ) + return bool(valid_accounts) def _get_default_carrier(self): self.ensure_one() @@ -253,69 +231,25 @@ class CarrierAccountMixin(models.AbstractModel): ) def _compute_valid_carrier_account_ids(self): for rec in self: - match rec.delivery_billing_mode: - case "collect": - rec.valid_carrier_account_ids = ( - (rec.recipient_id | rec.recipient_id.commercial_partner_id) - .mapped("carrier_account_ids") - .filtered( - lambda account: account.delivery_carrier_id - == rec.carrier_id - ) - ) - case "third party": - rec.valid_carrier_account_ids = self.env[ - "delivery.carrier.account" - ].search( - [ - ("delivery_carrier_id", "=", rec.carrier_id.id), - ( - "partner_id", - "not in", - [ - rec.sender_id.id, - rec.sender_id.commercial_partner_id.id, - rec.recipient_id.id, - rec.recipient_id.commercial_partner_id.id, - ], - ), - ] - ) - case "prepaid" | "ppc" | "no charge": - rec.valid_carrier_account_ids = ( - (rec.sender_id | rec.sender_id.commercial_partner_id) - .mapped("carrier_account_ids") - .filtered( - lambda account: account.delivery_carrier_id - == rec.carrier_id - ) - ) - case _: - rec.valid_carrier_account_ids = self.env[ - "delivery.carrier.account" - ].search([]) + partners = rec._get_valid_carrier_partners() + if rec.carrier_id and partners: + rec.valid_carrier_account_ids = partners.mapped( + "carrier_account_ids" + ).filtered( + lambda account: account.delivery_carrier_id == rec.carrier_id + ) + else: + rec.valid_carrier_account_ids = partners.mapped("carrier_account_ids") @api.depends("delivery_billing_mode", "sender_id", "recipient_id") def _compute_valid_carrier_ids(self): """Compute all valid carriers for the current billing mode and partners.""" for rec in self: - partners, is_third_party, invalid_partners = ( - rec._get_valid_carrier_partners() + partners = rec._get_valid_carrier_partners() + rec.valid_carrier_ids = partners.mapped( + "carrier_account_ids.delivery_carrier_id" ) - if is_third_party: - # For third party, all carriers with accounts not belonging to sender/recipient are valid - accounts = self.env["delivery.carrier.account"].search( - [ - ("partner_id", "not in", invalid_partners.ids), - ] - ) - rec.valid_carrier_ids = accounts.mapped("delivery_carrier_id") - else: - rec.valid_carrier_ids = partners.mapped( - "carrier_account_ids.delivery_carrier_id" - ) - def _on_carrier_fields_changed(self): pass @@ -325,8 +259,20 @@ class CarrierAccountMixin(models.AbstractModel): if ( rec.carrier_account_id and rec.delivery_billing_mode + and rec.valid_carrier_account_ids # Use the computed field directly and rec.carrier_account_id not in rec.valid_carrier_account_ids ): + _logger.warning( + "Billing mode: %s, sender: %s (commercial: %s), recipient: %s (commercial: %s), account: %s, id: %s, valid: %s", + rec.delivery_billing_mode, + rec.sender_id.name, + rec.sender_id.commercial_partner_id.name, + rec.recipient_id.name, + rec.recipient_id.commercial_partner_id.name, + rec.carrier_account_id.partner_id.name, + rec.carrier_account_id.id, + rec.valid_carrier_account_ids.ids, + ) if rec.delivery_billing_mode == "collect": raise UserError( diff --git a/delivery_carrier_partner_account/tests/test_sale_order.py b/delivery_carrier_partner_account/tests/test_sale_order.py index 3002678..bcb6337 100644 --- a/delivery_carrier_partner_account/tests/test_sale_order.py +++ b/delivery_carrier_partner_account/tests/test_sale_order.py @@ -1,4 +1,5 @@ from .test_carrier_account_common import TestCarrierAccountCommon +from odoo.tests import Form class TestSalesOrder(TestCarrierAccountCommon): @@ -21,12 +22,13 @@ class TestSalesOrder(TestCarrierAccountCommon): def test_prepaid_sale_order_line_gets_proper_name(self): order = self.env["sale.order"].create({"partner_id": self.client_partner.id}) wiz = self._get_shipping_wizard(order) - wiz.carrier_id = self.delivery_carrier_1 - wiz.delivery_billing_mode = "prepaid" + with Form(wiz) as form: + form.carrier_id = self.delivery_carrier_2 + form.delivery_billing_mode = "prepaid" wiz.button_confirm() self.assertEqual( order.order_line[0].name, - f"{self.delivery_carrier_1.name} [PREPAID]", + f"{self.delivery_carrier_2.name} [PREPAID]", ) def test_third_party_sale_order_line_gets_proper_name(self): diff --git a/delivery_carrier_partner_account/wizard/choose_delivery_carrier.py b/delivery_carrier_partner_account/wizard/choose_delivery_carrier.py index 29926b1..52bcd0e 100644 --- a/delivery_carrier_partner_account/wizard/choose_delivery_carrier.py +++ b/delivery_carrier_partner_account/wizard/choose_delivery_carrier.py @@ -15,9 +15,25 @@ class ChooseDeliveryCarrier(models.TransientModel): if self.delivery_billing_mode: vals.update(delivery_billing_mode=self.delivery_billing_mode) if self.carrier_account_id: - vals.update(carrier_account_id=self.carrier_account_id) + vals.update(carrier_account_id=self.carrier_account_id.id) if self.carrier_id: vals.update(carrier_id=self.carrier_id.id) - self.order_id.with_context(no_carrier_update=True).write(vals) + + # Ensure we have a valid carrier account + if ( + self.carrier_id + and self.delivery_billing_mode + and not self.carrier_account_id + ): + # Force recompute of valid carrier accounts + self._compute_valid_carrier_account_ids() + default_account = self._get_default_carrier_account() + if default_account: + vals.update(carrier_account_id=default_account.id) + + # Write values to the order before calling super + if vals: + self.order_id.with_context(no_carrier_update=True).write(vals) + res = super().button_confirm() return res