from unittest.mock import patch

from django.test import TransactionTestCase

from apps.catalog.models import ProductCategory
from apps.organization import test_assignment_concurrency as concurrency
from .models import ComplaintSymptom, ComplaintSymptomProductCategory as Mapping
from .services import set_complaint_applicability


class ApplicabilityConcurrencyTests(TransactionTestCase):
    run_concurrent = concurrency.AssignmentConcurrencyTests.run_concurrent

    def setUp(self):
        self.complaint = ComplaintSymptom.objects.create(code="EXAMPLE", name="Example")
        self.a = ProductCategory.objects.create(code="A", name="A")
        self.b = ProductCategory.objects.create(code="B", name="B")
        self.c = ProductCategory.objects.create(code="C", name="C")

    def configure(self, categories=(), global_mode=False):
        return set_complaint_applicability(complaint=self.complaint,
            applies_to_all_product_categories=global_mode, product_categories=categories)

    def test_competing_replacements_serialize_without_mixed_result(self):
        self.run_concurrent(lambda: self.configure([self.a, self.b]),
            lambda: self.configure([self.b, self.c]), expected="success")
        self.assertEqual(set(Mapping.objects.values_list("product_category_id", flat=True)), {self.b.pk, self.c.pk})
        self.complaint.refresh_from_db()
        self.assertFalse(self.complaint.applies_to_all_product_categories)

    def test_global_then_restricted_service_update(self):
        self.run_concurrent(lambda: self.configure(global_mode=True), lambda: self.configure([self.a]), expected="success")
        self.complaint.refresh_from_db()
        self.assertFalse(self.complaint.applies_to_all_product_categories)
        self.assertEqual(Mapping.objects.get().product_category, self.a)

    def test_restricted_then_global_service_update(self):
        self.run_concurrent(lambda: self.configure([self.a, self.b]), lambda: self.configure(global_mode=True), expected="success")
        self.complaint.refresh_from_db()
        self.assertTrue(self.complaint.applies_to_all_product_categories)
        self.assertFalse(Mapping.objects.exists())

    def test_global_switch_rejects_competing_direct_mapping_creation(self):
        self.run_concurrent(lambda: self.configure(global_mode=True),
            lambda: Mapping.objects.create(complaint_symptom=self.complaint, product_category=self.a), expected="validation")
        self.assertFalse(Mapping.objects.exists())

    def test_direct_mapping_then_global_service_clears_mapping(self):
        self.run_concurrent(lambda: Mapping.objects.create(complaint_symptom=self.complaint, product_category=self.a),
            lambda: self.configure(global_mode=True), expected="success")
        self.complaint.refresh_from_db()
        self.assertTrue(self.complaint.applies_to_all_product_categories)
        self.assertFalse(Mapping.objects.exists())

    def test_failed_waiting_replacement_restores_committed_configuration(self):
        original = Mapping.save
        def fail_selected(mapping, *args, **kwargs):
            if mapping.product_category_id == self.c.pk:
                from django.core.exceptions import ValidationError
                raise ValidationError("simulated mapping failure")
            return original(mapping, *args, **kwargs)
        with patch.object(Mapping, "save", fail_selected):
            self.run_concurrent(lambda: self.configure([self.a, self.b]), lambda: self.configure([self.c]), expected="validation")
        self.assertEqual(set(Mapping.objects.values_list("product_category_id", flat=True)), {self.a.pk, self.b.pk})

    def test_different_complaints_do_not_serialize(self):
        other = ComplaintSymptom.objects.create(code="OTHER", name="Other")
        self.run_concurrent(lambda: self.configure([self.a]),
            lambda: set_complaint_applicability(complaint=other, applies_to_all_product_categories=False, product_categories=[self.a]),
            expected="success", should_block=False)
        self.assertEqual(Mapping.objects.count(), 2)
