import uuid
from unittest.mock import patch

from django.core.exceptions import ValidationError
from django.db import IntegrityError, transaction
from django.db.models.deletion import ProtectedError
from django.test import TestCase

from apps.catalog import services as catalog_services
from apps.catalog.models import ProductModel, ProductVariant
from apps.catalog.tests import make_catalog
from . import queries as q, services as s
from .models import PartCategory, SparePart, SparePartCompatibility as Mapping


class PartsFixture(TestCase):
    @classmethod
    def setUpTestData(cls):
        cls.brand, cls.product_category, cls.model, cls.variant = make_catalog("PARTS-A")
        cls.other_brand, cls.other_product_category, cls.other_model, cls.other_variant = make_catalog("PARTS-B")
        cls.variant2 = ProductVariant.objects.create(product_model=cls.model, code="V2", name="Second")
        cls.category = s.create_part_category(code="DISPLAY", name="Display")
        cls.part = s.create_spare_part(part_code="DISPLAY-01", name="Display assembly", category=cls.category,
            serialization_policy="REQUIRED_SERIAL")

    def configure(self, models=(), variants=()):
        return s.set_spare_part_compatibility(spare_part=self.part, product_models=models, product_variants=variants)

    def matches_model(self, model=None):
        return q.spare_part_is_compatible_with_model(spare_part=self.part, product_model=model or self.model)

    def matches_variant(self, variant=None):
        return q.spare_part_is_compatible_with_variant(spare_part=self.part, product_variant=variant or self.variant)


class CategoryTests(PartsFixture):
    def test_creation_normalizes_and_timestamps(self):
        obj = s.create_part_category(code=" battery-1 ", name=" Battery ", description=" Details ")
        self.assertEqual((obj.code, obj.name, obj.description), ("BATTERY-1", "Battery", "Details"))
        self.assertIsInstance(obj.pk, uuid.UUID)
        self.assertTrue(obj.is_active)
        self.assertIsNotNone(obj.created_at)
        self.assertIsNotNone(obj.updated_at)

    def test_duplicate_code_case_normalized(self):
        with self.assertRaises(ValidationError):
            s.create_part_category(code=" display ", name="Other")

    def test_invalid_codes(self):
        for value in ("", " ", "bad code", "_START", "x" * 65, "café", None, 123):
            with self.subTest(value=value), self.assertRaises(ValidationError):
                s.create_part_category(code=value, name="Valid")

    def test_invalid_name_description(self):
        for field, value in (("name", " "), ("name", "x" * 201), ("name", None), ("description", "x" * 4001)):
            with self.subTest(field=field, value=value), self.assertRaises(ValidationError):
                s.create_part_category(**{"code": "NEW", "name": "Valid", field: value})

    def test_update_trim_and_immutable_code(self):
        current = s.update_part_category(part_category=self.category, name=" New name ", description=" Notes ")
        self.assertEqual((current.name, current.description, current.code), ("New name", "Notes", "DISPLAY"))
        self.assertGreater(current.updated_at, self.category.updated_at)

    def test_lifecycle_does_not_overwrite_metadata(self):
        self.category.name = "unsaved"
        obj = s.deactivate_part_category(part_category=self.category)
        self.assertFalse(obj.is_active)
        self.assertEqual(obj.name, "Display")
        self.assertTrue(s.reactivate_part_category(part_category=self.category).is_active)

    def test_stale_update_rejected_after_lifecycle(self):
        rev = s.revision(self.category)
        s.deactivate_part_category(part_category=self.category)
        with self.assertRaises(ValidationError):
            s.update_part_category(part_category=self.category, name="Stale", expected_revision=rev)

    def test_deactivated_code_remains_reserved(self):
        s.deactivate_part_category(part_category=self.category)
        with self.assertRaises(ValidationError):
            s.create_part_category(code="DISPLAY", name="Duplicate")

    def test_protected_category_foreign_key(self):
        # Exercise Django's FK protection beneath the public deletion guard.
        from django.db.models import Model
        with self.assertRaises(ProtectedError):
            Model.delete(self.category)

    def test_database_code_and_name_checks(self):
        for values in ({"code": "lower"}, {"name": ""}, {"name": " leading"}):
            with self.subTest(values=values), self.assertRaises(IntegrityError), transaction.atomic():
                PartCategory.objects.filter(pk=self.category.pk).update(**values)


class SparePartTests(PartsFixture):
    def test_creation_normalization_and_no_inventory(self):
        obj = s.create_spare_part(part_code=" batt-001 ", name=" Battery ", category=self.category,
            serialization_policy="OPTIONAL_SERIAL", manufacturer_part_number=" abc/123 ", description=" Text ")
        self.assertEqual((obj.part_code, obj.name, obj.manufacturer_part_number, obj.description),
            ("BATT-001", "Battery", "abc/123", "Text"))
        self.assertIsInstance(obj.pk, uuid.UUID)
        self.assertTrue(obj.is_active)
        for field in ("quantity", "stock", "warehouse", "company", "repair_action", "selling_price", "supplier"):
            self.assertNotIn(field, {f.name for f in obj._meta.get_fields()})

    def test_all_serialization_policies(self):
        for policy in SparePart.SerializationPolicy.values:
            obj = s.update_spare_part(spare_part=self.part, serialization_policy=policy)
            self.assertEqual(obj.serialization_policy, policy)

    def test_invalid_policy(self):
        for policy in ("", "INVALID", "required_serial", None):
            with self.subTest(policy=policy), self.assertRaises(ValidationError):
                s.update_spare_part(spare_part=self.part, serialization_policy=policy)

    def test_category_required_saved_and_correct_type(self):
        for category in (None, PartCategory(code="X", name="X"), self.model):
            with self.subTest(category=category), self.assertRaises(ValidationError):
                s.update_spare_part(spare_part=self.part, category=category)

    def test_missing_category_rejected(self):
        self.category.pk = uuid.uuid4()
        with self.assertRaises(ValidationError):
            s.update_spare_part(spare_part=self.part, category=self.category)

    def test_inactive_master_configuration_is_non_operational(self):
        s.deactivate_part_category(part_category=self.category)
        obj = s.create_spare_part(part_code="NEW", name="New", category=self.category,
            serialization_policy="NOT_SERIALIZED", product_models=[self.model])
        self.assertTrue(obj.is_active)
        self.assertFalse(q.spare_part_is_compatible_with_model(spare_part=obj, product_model=self.model))

    def test_duplicate_part_code(self):
        with self.assertRaises(ValidationError):
            s.create_spare_part(part_code=" display-01 ", name="Duplicate", category=self.category,
                serialization_policy="NOT_SERIALIZED")

    def test_invalid_part_fields(self):
        for field, value in (("part_code", "bad code"), ("part_code", "x" * 65), ("name", " "),
                ("name", "x" * 201), ("description", "x" * 4001), ("manufacturer_part_number", "x" * 129)):
            data = dict(part_code="NEW", name="New", category=self.category, serialization_policy="NOT_SERIALIZED")
            data[field] = value
            with self.subTest(field=field), self.assertRaises(ValidationError):
                s.create_spare_part(**data)

    def test_update_uses_current_not_unsaved_fields(self):
        self.part.part_code = "UNSAVED"
        self.part.is_active = False
        obj = s.update_spare_part(spare_part=self.part, name=" Updated ")
        self.assertEqual((obj.part_code, obj.name, obj.is_active), ("DISPLAY-01", "Updated", True))

    def test_category_correction(self):
        other = s.create_part_category(code="OTHER", name="Other")
        self.assertEqual(s.update_spare_part(spare_part=self.part, category=other).category, other)

    def test_lifecycle_and_aba_revision(self):
        old = s.revision(self.part)
        self.assertFalse(s.deactivate_spare_part(spare_part=self.part).is_active)
        self.assertTrue(s.reactivate_spare_part(spare_part=self.part).is_active)
        with self.assertRaises(ValidationError):
            s.update_spare_part(spare_part=self.part, name="Stale", expected_revision=old)

    def test_lifecycle_preconditions(self):
        old = s.revision(self.part)
        s.update_spare_part(spare_part=self.part, name="New")
        with self.assertRaises(ValidationError):
            s.deactivate_spare_part(spare_part=self.part, expected_revision=old)

    def test_direct_save_delete_and_queryset_delete_rejected(self):
        self.configure(models=[self.model])
        for obj in (self.part, self.category, Mapping.objects.get()):
            with self.subTest(model=type(obj)), self.assertRaises(ValidationError):
                obj.save()
            with self.assertRaises(ValidationError):
                obj.delete()
            with self.assertRaises(ValidationError):
                type(obj).objects.filter(pk=obj.pk).delete()

    def test_database_rejects_invalid_policy_code_name_and_category(self):
        for values in ({"serialization_policy": "BAD"}, {"part_code": "lower"}, {"name": ""},
                {"name": " spaced "}, {"category_id": None}):
            with self.subTest(values=values), self.assertRaises(IntegrityError), transaction.atomic():
                SparePart.objects.filter(pk=self.part.pk).update(**values)

    def test_database_unique_code(self):
        with self.assertRaises(IntegrityError), transaction.atomic():
            SparePart.objects.bulk_create([SparePart(part_code=self.part.part_code, name="Duplicate",
                category=self.category, serialization_policy="NOT_SERIALIZED")])


class CompatibilityTests(PartsFixture):
    def test_empty_configuration_denies(self):
        self.assertFalse(self.matches_model())
        self.assertFalse(self.matches_variant())

    def test_model_wide_covers_model_and_variants(self):
        self.configure(models=[self.model])
        self.assertTrue(self.matches_model())
        self.assertTrue(self.matches_variant())
        self.assertTrue(self.matches_variant(self.variant2))
        self.assertFalse(self.matches_model(self.other_model))
        self.assertFalse(self.matches_variant(self.other_variant))

    def test_future_variant_inherits_explicit_model_mapping(self):
        self.configure(models=[self.model])
        new = ProductVariant.objects.create(product_model=self.model, code="NEW", name="Future")
        self.assertTrue(self.matches_variant(new))

    def test_exact_variant_does_not_imply_model_or_sibling(self):
        self.configure(variants=[self.variant])
        self.assertFalse(self.matches_model())
        self.assertTrue(self.matches_variant())
        self.assertFalse(self.matches_variant(self.variant2))
        self.assertEqual(list(q.compatible_spare_parts_for_model(self.model)), [])

    def test_all_current_variants_never_infer_model_or_future_variant(self):
        self.configure(variants=[self.variant, self.variant2])
        self.assertFalse(self.matches_model())
        self.assertTrue(self.matches_variant())
        self.assertTrue(self.matches_variant(self.variant2))
        new = ProductVariant.objects.create(product_model=self.model, code="NEW", name="Future")
        self.assertFalse(self.matches_variant(new))

    def test_mixed_modes_for_different_models(self):
        self.configure(models=[self.model], variants=[self.other_variant])
        self.assertTrue(self.matches_model())
        self.assertFalse(self.matches_model(self.other_model))
        self.assertTrue(self.matches_variant(self.other_variant))

    def test_contradictory_configuration_rejected(self):
        with self.assertRaises(ValidationError):
            self.configure(models=[self.model], variants=[self.variant])
        self.assertEqual(Mapping.objects.count(), 0)

    def test_duplicate_inputs_are_deduplicated(self):
        self.configure(models=[self.model, self.model], variants=[self.other_variant, self.other_variant])
        self.assertEqual(Mapping.objects.count(), 2)

    def test_replacement_retains_mapping_identity(self):
        self.configure(models=[self.model])
        original = Mapping.objects.get()
        self.configure(variants=[self.variant])
        original.refresh_from_db()
        self.assertFalse(original.is_active)
        self.assertFalse(self.matches_model())
        self.assertTrue(self.matches_variant())
        self.configure(models=[self.model])
        original.refresh_from_db()
        self.assertTrue(original.is_active)
        self.assertEqual(Mapping.objects.count(), 2)

    def test_clear_retains_rows(self):
        self.configure(models=[self.model])
        self.configure()
        self.assertEqual(Mapping.objects.count(), 1)
        self.assertFalse(Mapping.objects.get().is_active)
        self.assertFalse(self.matches_variant())

    def test_part_deactivation_retains_mapping_and_reactivation_restores(self):
        self.configure(models=[self.model])
        s.deactivate_spare_part(spare_part=self.part)
        self.assertFalse(self.matches_model())
        self.assertTrue(Mapping.objects.get().is_active)
        s.reactivate_spare_part(spare_part=self.part)
        self.assertTrue(self.matches_model())

    def test_category_masks_parts_without_cascading_individual_flags(self):
        self.configure(models=[self.model])
        s.deactivate_part_category(part_category=self.category)
        self.assertFalse(self.matches_model())
        self.part.refresh_from_db()
        self.assertTrue(self.part.is_active)
        self.assertTrue(Mapping.objects.get().is_active)
        s.reactivate_part_category(part_category=self.category)
        self.assertTrue(self.matches_model())

    def test_category_reactivation_does_not_reactivate_individually_inactive_part(self):
        self.configure(models=[self.model])
        s.deactivate_spare_part(spare_part=self.part)
        s.deactivate_part_category(part_category=self.category)
        s.reactivate_part_category(part_category=self.category)
        self.assertFalse(self.matches_model())

    def test_product_model_lifecycle_retains_mapping(self):
        self.configure(models=[self.model])
        catalog_services.deactivate_product_model(product_model=self.model)
        self.assertFalse(self.matches_model())
        self.assertFalse(self.matches_variant())
        self.assertTrue(Mapping.objects.get().is_active)
        catalog_services.reactivate_product_model(product_model=self.model)
        self.assertTrue(self.matches_model())
        self.assertFalse(self.matches_variant())  # Frozen catalog does not cascade reactivation.
        catalog_services.reactivate_variant(variant=self.variant)
        self.assertTrue(self.matches_variant())

    def test_variant_lifecycle_exact_and_model_wide(self):
        for models, variants in (([], [self.variant]), ([self.model], [])):
            self.configure(models=models, variants=variants)
            catalog_services.deactivate_variant(variant=self.variant)
            self.assertFalse(self.matches_variant())
            catalog_services.reactivate_variant(variant=self.variant)
            self.assertTrue(self.matches_variant())

    def test_inactive_brand_and_product_category_deny_freshly(self):
        self.configure(models=[self.model])
        for obj in (self.brand, self.product_category):
            # Also defend against inconsistent privileged writes to ancestors.
            type(obj).objects.filter(pk=obj.pk).update(is_active=False)
            self.assertFalse(self.matches_model())
            self.assertFalse(self.matches_variant())
            type(obj).objects.filter(pk=obj.pk).update(is_active=True)
            self.assertTrue(self.matches_variant())

    def test_can_configure_inactive_catalog_but_never_select_it(self):
        catalog_services.deactivate_product_model(product_model=self.model)
        self.configure(variants=[self.variant])
        self.assertFalse(self.matches_variant())
        self.assertEqual(Mapping.objects.count(), 1)

    def test_invalid_reference_rejected_without_partial_change(self):
        self.configure(models=[self.model])
        for models, variants in (([ProductModel()], []), ([self.variant], []), ([], [self.model]), (None, [])):
            with self.subTest(models=models), self.assertRaises(ValidationError):
                self.configure(models=models, variants=variants)
            self.assertTrue(self.matches_model())

    def test_missing_reference_rejected_freshly(self):
        self.model.pk = uuid.uuid4()
        with self.assertRaises(ValidationError):
            self.configure(models=[self.model])

    def test_requires_complete_replacement_inputs(self):
        with self.assertRaises(ValidationError):
            s.update_spare_part(spare_part=self.part, product_models=[self.model])

    def test_revision_changes_on_compatibility_replacement(self):
        old = s.revision(self.part)
        self.configure(models=[self.model])
        with self.assertRaises(ValidationError):
            s.update_spare_part(spare_part=self.part, name="Stale", expected_revision=old)

    def test_rollback_after_mapping_write_failure(self):
        self.configure(models=[self.model])
        original = Mapping._persist
        def fail_new(mapping):
            if mapping._state.adding:
                raise ValidationError("Injected failure after old mapping retirement")
            return original(mapping)
        with patch.object(Mapping, "_persist", fail_new), self.assertRaises(ValidationError):
            s.update_spare_part(spare_part=self.part, name="Rolled back", product_models=[], product_variants=[self.variant])
        self.part.refresh_from_db()
        self.assertEqual(self.part.name, "Display assembly")
        self.assertTrue(self.matches_model())
        self.assertEqual(Mapping.objects.count(), 1)

    def test_database_mapping_requires_exactly_one_target(self):
        for fields in ({}, {"product_model": self.model, "product_variant": self.variant}):
            with self.subTest(fields=fields), self.assertRaises(IntegrityError), transaction.atomic():
                Mapping.objects.bulk_create([Mapping(spare_part=self.part, **fields)])

    def test_database_mapping_unique_including_retired(self):
        for field, value in (("product_model", self.model), ("product_variant", self.variant)):
            with transaction.atomic():
                Mapping.objects.bulk_create([Mapping(spare_part=self.part, is_active=False, **{field: value})])
            with self.subTest(field=field), self.assertRaises(IntegrityError), transaction.atomic():
                Mapping.objects.bulk_create([Mapping(spare_part=self.part, **{field: value})])

    def test_mapping_endpoints_protected(self):
        self.configure(variants=[self.variant])
        self.configure()
        from django.db.models import Model
        for obj in (self.part, self.variant):
            with self.subTest(obj=obj), self.assertRaises(ProtectedError):
                Model.delete(obj)


class QueryTests(PartsFixture):
    def test_variant_union_order_and_no_duplicates(self):
        self.configure(models=[self.model])
        other = s.create_spare_part(part_code="AAA", name="Another", category=self.category,
            serialization_policy="NOT_SERIALIZED", product_variants=[self.variant])
        self.assertEqual(list(q.compatible_spare_parts_for_variant(self.variant)), [other, self.part])
        self.assertEqual(list(q.compatible_spare_parts_for_model(self.model)), [self.part])

    def test_lazy_queries_and_one_query_evaluation_with_related_category(self):
        self.configure(models=[self.model])
        with self.assertNumQueries(0):
            models = q.compatible_spare_parts_for_model(self.model)
            variants = q.compatible_spare_parts_for_variant(self.variant)
        with self.assertNumQueries(1):
            self.assertEqual([(p.part_code, p.category.name) for p in models], [(self.part.part_code, "Display")])
        with self.assertNumQueries(1):
            self.assertEqual([(p.part_code, p.category.name) for p in variants], [(self.part.part_code, "Display")])

    def test_boolean_query_budget(self):
        self.configure(variants=[self.variant])
        with self.assertNumQueries(1):
            self.assertFalse(self.matches_model())
        with self.assertNumQueries(1):
            self.assertTrue(self.matches_variant())

    def test_query_budget_does_not_grow_with_parts(self):
        for index in range(8):
            s.create_spare_part(part_code=f"PART-{index}", name="Part", category=self.category,
                serialization_policy="NOT_SERIALIZED", product_models=[self.model])
        with self.assertNumQueries(1):
            self.assertEqual(len([(p.name, p.category.name) for p in q.compatible_spare_parts_for_variant(self.variant)]), 8)

    def test_queries_read_lifecycle_at_evaluation_not_construction(self):
        self.configure(models=[self.model])
        pending = q.compatible_spare_parts_for_model(self.model)
        s.deactivate_spare_part(spare_part=self.part)
        self.assertEqual(list(pending), [])

    def test_search_case_insensitive_and_active_only(self):
        self.assertEqual(list(q.search_spare_parts(" display ")), [self.part])
        self.assertEqual(list(q.search_spare_parts("' OR 1=1 --")), [])
        s.update_spare_part(spare_part=self.part, manufacturer_part_number="ACME-88")
        self.assertEqual(list(q.search_spare_parts("acme-88")), [self.part])
        s.deactivate_part_category(part_category=self.category)
        self.assertEqual(list(q.search_spare_parts("display")), [])
        self.assertEqual(list(q.active_part_categories()), [])

    def test_missing_saved_target_fails_closed(self):
        self.configure(models=[self.model])
        self.model.pk = uuid.uuid4()
        self.assertFalse(self.matches_model())

    def test_unsaved_wrong_type_or_wrong_database_rejected(self):
        for obj in (ProductModel(), self.variant):
            with self.assertRaises(ValidationError):
                q.compatible_spare_parts_for_model(obj)
        self.model._state.db = "other"
        with self.assertRaises(ValidationError):
            q.compatible_spare_parts_for_model(self.model)

    def test_global_query_has_no_organization_or_auth_joins(self):
        sql = str(q.compatible_spare_parts_for_variant(self.variant).query)
        for token in ("organization_", "access_", "accounts_", "service_servicecase"):
            self.assertNotIn(token, sql)
