import uuid
from unittest.mock import patch
from django.core.exceptions import PermissionDenied, ValidationError
from django.db import IntegrityError, transaction
from django.test import TestCase
from apps.parts import services as parts
from . import control_services as s, services as stock, queries as q
from .models import StockCount, StockAdjustment, StockMovement
from .tests import InventoryFixture
from .test_requests import RequestFixture, setup_requests


class ControlFixture:
    def adjust(self, delta=-1, **kwargs):
        return s.adjust_stock(**dict(actor=self.actor, location=self.location, spare_part=self.part, quantity_delta=delta,
            reason="CORRECTION", command_key=uuid.uuid4(), reference="CONTROL-TEST", note="Documented physical verification") | kwargs)

    def draft_count(self, **kwargs):
        return s.create_stock_count(**dict(actor=self.actor, location=self.location, spare_part=self.part) | kwargs)

    def start_count(self, count=None):
        count = count or self.draft_count()
        return s.start_stock_count(actor=self.actor, count=count, expected_revision=count.updated_at.isoformat())

    def record_count(self, count, quantity, units=()):
        return s.record_stock_count(actor=self.actor, count=count, counted_quantity=quantity, units=units, note="Physical count",
            expected_revision=count.updated_at.isoformat())

    def reconcile(self, count):
        return s.reconcile_stock_count(actor=self.actor, count=count, expected_revision=count.updated_at.isoformat())


class ControlTests(ControlFixture, InventoryFixture):
    def test_center_adjustment_scope_cannot_claim_unlocated_company_units(self):
        from .tests import grant
        grant(self.staff, self.company, center=self.center, permission_names=["adjust_stock"])
        unit = stock.register_serialized_unit(actor=self.actor, company=self.company, spare_part=self.serial_part, identifier="COMPANY-POOL")
        with self.assertRaises(PermissionDenied):
            self.adjust(1, actor=self.staff, location=self.destination, spare_part=self.serial_part, units=[unit])
        self.adjust(1, actor=self.staff, location=self.destination)

    def test_found_unit_requires_scope_at_its_loss_origin(self):
        from .tests import grant
        grant(self.staff, self.company, center=self.center, permission_names=["adjust_stock"])
        unit = stock.register_serialized_unit(actor=self.actor, company=self.company, spare_part=self.serial_part, identifier="LOST-AT-WAREHOUSE")
        self.receive(1, spare_part=self.serial_part, units=[unit])
        self.adjust(-1, spare_part=self.serial_part, units=[unit])
        with self.assertRaises(PermissionDenied):
            self.adjust(1, actor=self.staff, location=self.destination, spare_part=self.serial_part, units=[unit])

    def test_positive_adjustment_is_immutable_ledger_evidence(self):
        row = self.adjust(3, reason="FOUND")
        self.assertEqual(self.balance(), 3)
        self.assertEqual(row.movement.kind, "ADJUST_IN")
        self.assertEqual(row.movement.entries.get().quantity_delta, 3)
        with self.assertRaises(IntegrityError), transaction.atomic():
            StockAdjustment.objects.filter(pk=row.pk).update(reason="CORRECTION")

    def test_loss_posts_negative_movement(self):
        self.receive(3)
        row = self.adjust(-2, reason="LOSS")
        self.assertEqual(self.balance(), 1)
        self.assertEqual(row.movement.kind, "ADJUST_OUT")

    def test_cannot_adjust_below_zero(self):
        with self.assertRaises(ValidationError):
            self.adjust(-1)
        self.assertFalse(StockAdjustment.objects.exists())

    def test_adjustment_reasons_and_bounds(self):
        for delta, reason in ((1, "LOSS"), (-1, "FOUND"), (1, "COUNT_VARIANCE"), (0, "CORRECTION"), (True, "FOUND"), (1000000001, "FOUND"), (1, "ARBITRARY")):
            with self.subTest(delta=delta, reason=reason), self.assertRaises(ValidationError):
                self.adjust(delta, reason=reason)

    def test_adjustment_requires_explanation(self):
        with self.assertRaises(ValidationError):
            self.adjust(1, note="")

    def test_stale_position_revision_rejected(self):
        revision = s.position_revision(self.location, self.part)
        self.receive()
        with self.assertRaises(ValidationError):
            self.adjust(1, expected_revision=revision)

    def test_duplicate_command_key_rejected(self):
        key = uuid.uuid4()
        self.adjust(1, command_key=key)
        with self.assertRaises((ValidationError, IntegrityError)), transaction.atomic():
            self.adjust(1, command_key=key)
        self.assertEqual(self.balance(), 1)

    def test_adjustment_scope_required(self):
        with self.assertRaises(PermissionDenied):
            self.adjust(1, actor=self.staff)

    def test_failed_adjustment_rolls_back_movement(self):
        with patch.object(StockAdjustment, "_persist", side_effect=ValidationError("abort")), self.assertRaises(ValidationError):
            self.adjust(1)
        self.assertFalse(StockMovement.objects.exists())

    def test_draft_count_does_not_change_stock_or_lock_policy(self):
        self.draft_count()
        parts.update_spare_part(spare_part=self.part, serialization_policy="OPTIONAL_SERIAL")
        self.assertFalse(StockMovement.objects.exists())

    def test_count_snapshots_stock_and_freezes_movement(self):
        self.receive(5)
        count = self.start_count()
        self.assertEqual(count.expected_quantity, 5)
        for operation in (self.receive, self.move, lambda: self.adjust(1)):
            with self.assertRaises(ValidationError):
                operation()
        self.assertEqual(self.balance(), 5)

    def test_only_one_active_count_per_position(self):
        self.start_count()
        with self.assertRaises(ValidationError):
            self.start_count()

    def test_count_positive_variance_posts_adjustment(self):
        self.receive(5)
        count = self.reconcile(self.record_count(self.start_count(), 7))
        self.assertEqual(count.status, "RECONCILED")
        self.assertEqual(count.adjustments.get().quantity_delta, 2)
        self.assertEqual(self.balance(), 7)
        self.receive(1)
        self.assertEqual(self.balance(), 8)

    def test_count_negative_variance_posts_adjustment(self):
        self.receive(5)
        count = self.reconcile(self.record_count(self.start_count(), 3))
        self.assertEqual(count.adjustments.get().quantity_delta, -2)
        self.assertEqual(self.balance(), 3)

    def test_zero_variance_posts_no_synthetic_movement(self):
        self.receive(5)
        count = self.reconcile(self.record_count(self.start_count(), 5))
        self.assertFalse(count.adjustments.exists())
        self.assertEqual(StockMovement.objects.count(), 1)

    def test_cancellation_releases_count_freeze(self):
        count = self.start_count()
        s.cancel_stock_count(actor=self.actor, count=count, reason="Recount later", expected_revision=count.updated_at.isoformat())
        self.receive(1)
        self.assertEqual(self.balance(), 1)

    def test_count_can_cancel_after_part_deactivation(self):
        count = self.start_count()
        parts.deactivate_spare_part(spare_part=self.part)
        s.cancel_stock_count(actor=self.actor, count=count, reason="Inactive part", expected_revision=count.updated_at.isoformat())

    def test_stale_count_record_and_reconcile_rejected(self):
        count = self.start_count()
        self.record_count(count, 0)
        with self.assertRaises(ValidationError):
            self.record_count(count, 1)
        with self.assertRaises(ValidationError):
            self.reconcile(count)

    def test_count_requires_observation_before_reconcile(self):
        with self.assertRaises(ValidationError):
            self.reconcile(self.start_count())

    def test_completed_count_snapshot_cannot_be_rewritten(self):
        count = self.reconcile(self.record_count(self.start_count(), 0))
        with self.assertRaises(IntegrityError), transaction.atomic():
            StockCount.objects.filter(pk=count.pk).update(expected_quantity=1)
        with self.assertRaises(ValidationError):
            self.reconcile(count)

    def test_count_prevents_location_deactivation(self):
        self.start_count()
        with self.assertRaises(ValidationError):
            stock.deactivate_location(actor=self.actor, location=self.location)

    def test_failed_reconciliation_keeps_count_active_and_stock_unchanged(self):
        self.receive(2)
        count = self.record_count(self.start_count(), 1)
        with patch.object(StockAdjustment, "_persist", side_effect=ValidationError("abort")), self.assertRaises(ValidationError):
            self.reconcile(count)
        self.assertEqual(self.balance(), 2)
        count.refresh_from_db()
        self.assertEqual(count.status, "COUNTING")

    def test_serial_loss_and_found_preserve_identity(self):
        parts.update_spare_part(spare_part=self.part, serialization_policy="REQUIRED_SERIAL")
        unit = stock.register_serialized_unit(actor=self.actor, company=self.company, spare_part=self.part, identifier="FOUND-UNIT")
        self.receive(1, units=[unit])
        self.adjust(-1, reason="LOSS", units=[unit])
        unit.refresh_from_db()
        self.assertEqual(unit.state, "REMOVED")
        self.adjust(1, reason="FOUND", units=[unit])
        unit.refresh_from_db()
        self.assertEqual((unit.state, unit.current_location_id), ("IN_STOCK", self.location.pk))

    def test_serial_count_missing_and_found_reconciles_equal_quantity(self):
        parts.update_spare_part(spare_part=self.part, serialization_policy="REQUIRED_SERIAL")
        first = stock.register_serialized_unit(actor=self.actor, company=self.company, spare_part=self.part, identifier="MISSING")
        found = stock.register_serialized_unit(actor=self.actor, company=self.company, spare_part=self.part, identifier="FOUND")
        self.receive(1, units=[first])
        count = self.reconcile(self.record_count(self.start_count(), 1, units=[found]))
        self.assertEqual(count.adjustments.count(), 2)
        self.assertEqual(self.balance(), 1)
        first.refresh_from_db()
        found.refresh_from_db()
        self.assertEqual((first.state, found.state), ("REMOVED", "IN_STOCK"))

    def test_required_serial_count_cannot_invent_anonymous_stock(self):
        parts.update_spare_part(spare_part=self.part, serialization_policy="REQUIRED_SERIAL")
        count = self.start_count()
        with self.assertRaises(ValidationError):
            self.record_count(count, 1)

    def test_count_cannot_claim_unit_held_elsewhere(self):
        parts.update_spare_part(spare_part=self.part, serialization_policy="REQUIRED_SERIAL")
        unit = stock.register_serialized_unit(actor=self.actor, company=self.company, spare_part=self.part, identifier="OTHER-LOCATION")
        self.receive(1, destination=self.destination, units=[unit])
        with self.assertRaises(ValidationError):
            self.record_count(self.start_count(), 1, units=[unit])


class ControlReservationTests(ControlFixture, RequestFixture, TestCase):
    @classmethod
    def setUpTestData(cls):
        setup_requests(cls)

    def test_adjustment_cannot_spend_reserved_stock(self):
        self.receive(quantity=1)
        self.reserve(self.approve(self.request()))
        with self.assertRaises(ValidationError):
            self.adjust(-1)

    def test_count_prevents_new_reservations_but_allows_release(self):
        self.receive()
        request = self.approve(self.request())
        reservation = self.reserve(request)
        self.start_count()
        with self.assertRaises(ValidationError):
            self.reserve(request)
        self.release(reservation)

    def test_count_shortage_requires_reservation_release(self):
        self.receive(quantity=1)
        reservation = self.reserve(self.approve(self.request()))
        count = self.record_count(self.start_count(), 0)
        with self.assertRaises(ValidationError):
            self.reconcile(count)
        self.release(reservation)
        self.reconcile(count)
        self.assertEqual(q.stock_on_hand(actor=self.actor, location=self.location, spare_part=self.part), 0)

    def test_superuser_cannot_reserve_for_inactive_case_center_from_central_stock(self):
        from apps.organization.services import deactivate_service_center
        request = self.approve(self.request())
        self.receive(destination=self.destination)
        deactivate_service_center(service_center=self.center)
        with self.assertRaises(ValidationError):
            self.reserve(request, location=self.destination)
