"""Real PostgreSQL contenders, including observed pg_blocking_pids waits."""
from django.core.exceptions import ValidationError
from django.db import transaction
from django.test import TransactionTestCase

from apps.catalog import services as catalog_services
from apps.catalog.tests import make_catalog
from apps.organization import test_assignment_concurrency as concurrency_helpers
from . import services as s, queries as q
from .models import PartCategory, SparePart, SparePartCompatibility as Mapping


class PartsConcurrencyTests(TransactionTestCase):
    run_concurrent = concurrency_helpers.AssignmentConcurrencyTests.run_concurrent

    def setUp(self):
        self.brand, self.product_category, self.model, self.variant = make_catalog("PARTS-RACE")
        self.category = s.create_part_category(code="DISPLAY", name="Display")
        self.part = s.create_spare_part(part_code="PART", name="Part", category=self.category,
            serialization_policy="NOT_SERIALIZED")

    def configure(self, *, model=False):
        return s.set_spare_part_compatibility(spare_part=self.part,
            product_models=[self.model] if model else [], product_variants=[] if model else [self.variant])

    def matches(self):
        return q.spare_part_is_compatible_with_variant(spare_part=self.part, product_variant=self.variant)

    def test_same_code_creation_database_unique_race(self):
        def create():
            return s.create_spare_part(part_code="SAME", name="Part", category=self.category,
                serialization_policy="NOT_SERIALIZED")
        self.run_concurrent(create, create, expected="integrity")
        self.assertEqual(SparePart.objects.filter(part_code="SAME").count(), 1)

    def test_same_category_code_creation_database_unique_race(self):
        def create():
            return s.create_part_category(code="SAME", name="Same")
        self.run_concurrent(create, create, expected="integrity")
        self.assertEqual(PartCategory.objects.filter(code="SAME").count(), 1)

    def test_concurrent_replacements_serialize_without_mixed_modes(self):
        self.run_concurrent(lambda: self.configure(model=True), self.configure, expected="success")
        self.assertEqual(Mapping.objects.count(), 2)
        self.assertEqual(list(Mapping.objects.filter(is_active=True).values_list("product_variant_id", flat=True)), [self.variant.pk])

    def test_concurrent_same_mapping_has_one_row(self):
        self.run_concurrent(self.configure, self.configure, expected="success")
        self.assertEqual(Mapping.objects.count(), 1)

    def test_revision_protected_replacement_rejects_stale_writer(self):
        revision = s.revision(self.part)
        self.run_concurrent(self.configure,
            lambda: s.set_spare_part_compatibility(spare_part=self.part, product_models=[self.model],
                product_variants=[], expected_revision=revision), expected="validation")
        self.assertIsNotNone(Mapping.objects.get().product_variant_id)

    def test_part_deactivation_then_configuration_is_retained_not_operational(self):
        self.run_concurrent(lambda: s.deactivate_spare_part(spare_part=self.part), self.configure, expected="success")
        self.assertFalse(self.matches())
        self.assertTrue(Mapping.objects.get().is_active)

    def test_configuration_then_part_deactivation(self):
        self.run_concurrent(self.configure, lambda: s.deactivate_spare_part(spare_part=self.part), expected="success")
        self.assertFalse(self.matches())
        self.assertEqual(Mapping.objects.count(), 1)

    def test_model_deactivation_then_configuration(self):
        self.run_concurrent(lambda: catalog_services.deactivate_product_model(product_model=self.model),
            self.configure, expected="success")
        self.assertFalse(self.matches())
        self.assertEqual(Mapping.objects.count(), 1)

    def test_configuration_then_model_deactivation(self):
        self.run_concurrent(self.configure,
            lambda: catalog_services.deactivate_product_model(product_model=self.model), expected="success")
        self.assertFalse(self.matches())

    def test_variant_deactivation_then_configuration(self):
        self.run_concurrent(lambda: catalog_services.deactivate_variant(variant=self.variant), self.configure, expected="success")
        self.assertFalse(self.matches())

    def test_configuration_then_variant_deactivation(self):
        self.run_concurrent(self.configure, lambda: catalog_services.deactivate_variant(variant=self.variant), expected="success")
        self.assertFalse(self.matches())

    def test_category_deactivation_then_configuration(self):
        self.run_concurrent(lambda: s.deactivate_part_category(part_category=self.category), self.configure, expected="success")
        self.assertFalse(self.matches())

    def test_configuration_then_category_deactivation(self):
        self.run_concurrent(self.configure, lambda: s.deactivate_part_category(part_category=self.category), expected="success")
        self.assertFalse(self.matches())

    def test_stale_admin_revision_after_lifecycle(self):
        # Same signed revision's payload that Admin passes to the locked service.
        revision = s.revision(self.part)
        self.run_concurrent(lambda: s.deactivate_spare_part(spare_part=self.part),
            lambda: s.update_spare_part(spare_part=self.part, name="Stale", expected_revision=revision,
                product_models=[self.model], product_variants=[]), expected="validation")
        self.part.refresh_from_db()
        self.assertFalse(self.part.is_active)
        self.assertEqual(self.part.name, "Part")

    def test_stale_category_admin_revision_after_lifecycle(self):
        revision = s.revision(self.category)
        self.run_concurrent(lambda: s.deactivate_part_category(part_category=self.category),
            lambda: s.update_part_category(part_category=self.category, name="Stale", expected_revision=revision), expected="validation")

    def test_real_admin_submission_waits_then_rejects_lifecycle_race(self):
        from django.contrib.auth import get_user_model
        from django.test import Client
        from django.urls import reverse
        user = get_user_model().objects.create_superuser(username="parts-race-admin", password="test-only")
        client = Client()
        client.force_login(user)
        url = reverse("admin:parts_sparepart_change", args=[self.part.pk])
        page = client.get(url)
        token = page.context["adminform"].form.initial["revision"]
        responses = []
        def submit():
            responses.append(client.post(url, dict(name="Stale", description="", category=str(self.category.pk),
                manufacturer_part_number="", serialization_policy="NOT_SERIALIZED", product_models=[str(self.model.pk)],
                product_variants=[], revision=token, _save="Save"), follow=True))
        self.run_concurrent(lambda: s.deactivate_spare_part(spare_part=self.part), submit, expected="success")
        self.assertContains(responses[0], "Save rejected")
        self.part.refresh_from_db()
        self.assertFalse(self.part.is_active)
        self.assertEqual(self.part.name, "Part")
        self.assertEqual(Mapping.objects.count(), 0)

    def test_failed_replacement_transaction_rolls_back_before_waiting_writer(self):
        self.configure(model=True)
        class Abort(Exception):
            pass
        def rolled_back():
            try:
                with transaction.atomic():
                    self.configure()
                    raise Abort()
            except Abort:
                pass
            # Assert rollback in the owner transaction, then retain its part lock
            # so the second real writer demonstrably waits for this transaction.
            self.assertTrue(Mapping.objects.get(product_model=self.model).is_active)
            SparePart.objects.select_for_update().get(pk=self.part.pk)
        self.run_concurrent(rolled_back, self.configure, expected="success")
        self.assertEqual(Mapping.objects.count(), 2)
        self.assertEqual(Mapping.objects.filter(is_active=True).count(), 1)

    def test_different_parts_share_catalog_without_serializing(self):
        other = s.create_spare_part(part_code="OTHER", name="Other", category=self.category, serialization_policy="NOT_SERIALIZED")
        self.run_concurrent(self.configure,
            lambda: s.set_spare_part_compatibility(spare_part=other, product_models=[], product_variants=[self.variant]),
            expected="success", should_block=False)
        self.assertEqual(Mapping.objects.count(), 2)

    def test_catalog_category_correction_rechecks_ancestor_snapshot(self):
        from apps.catalog.models import ProductCategory
        other = ProductCategory.objects.create(code="CORRECTED", name="Corrected")
        def correct():
            self.model.category = other
            self.model.save(update_fields=["category"])
        self.run_concurrent(correct, self.configure, expected="validation")
        self.assertEqual(Mapping.objects.count(), 0)
