from django.test import TransactionTestCase
from apps.organization import test_assignment_concurrency as concurrency_helpers
from apps.parts import services as parts
from .test_requests import RequestFixture, setup_requests
from .test_control import ControlFixture
from . import services as stock, control_services as s
from .models import StockAdjustment


class ControlConcurrencyTests(ControlFixture, RequestFixture, TransactionTestCase):
    run_concurrent = concurrency_helpers.AssignmentConcurrencyTests.run_concurrent

    def setUp(self):
        setup_requests(self)

    def test_reservation_first_prevents_loss_adjustment(self):
        self.receive(quantity=1)
        row = self.approve(self.request())
        self.run_concurrent(lambda: self.reserve(row), lambda: self.adjust(-1), expected="validation")

    def test_loss_adjustment_first_prevents_reservation(self):
        self.receive(quantity=1)
        row = self.approve(self.request())
        self.run_concurrent(lambda: self.adjust(-1), lambda: self.reserve(row), expected="validation")

    def test_adjustment_first_prevents_transfer(self):
        self.receive(quantity=1)
        self.run_concurrent(lambda: self.adjust(-1), self.move, expected="validation")

    def test_transfer_first_prevents_adjustment(self):
        self.receive(quantity=1)
        self.run_concurrent(self.move, lambda: self.adjust(-1), expected="validation")

    def test_count_start_first_freezes_receipt(self):
        count = self.draft_count()
        self.run_concurrent(lambda: self.start_count(count), self.receive, expected="validation")

    def test_receipt_first_is_included_in_count_snapshot(self):
        count = self.draft_count()
        self.run_concurrent(self.receive, lambda: self.start_count(count), expected="success")
        count.refresh_from_db()
        self.assertEqual(count.expected_quantity, 2)

    def test_count_start_first_blocks_reservation(self):
        self.receive()
        request = self.approve(self.request())
        count = self.draft_count()
        self.run_concurrent(lambda: self.start_count(count), lambda: self.reserve(request), expected="validation")

    def test_simultaneous_reconciliation_posts_once(self):
        self.receive()
        count = self.record_count(self.start_count(), 1)
        self.run_concurrent(lambda: self.reconcile(count), lambda: self.reconcile(count), expected="validation")
        self.assertEqual(StockAdjustment.objects.count(), 1)

    def test_two_counts_cannot_start_same_position(self):
        first, second = self.draft_count(), self.draft_count()
        self.run_concurrent(lambda: self.start_count(first), lambda: self.start_count(second), expected="validation")

    def test_serial_unit_cannot_be_lost_twice(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="LOSS-RACE")
        self.receive(quantity=1, units=[unit])
        self.run_concurrent(lambda: self.adjust(-1, units=[unit]), lambda: self.adjust(-1, units=[unit]), expected="validation")

    def test_serial_unit_cannot_be_found_twice(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-RACE")
        self.run_concurrent(lambda: self.adjust(1, units=[unit]), lambda: self.adjust(1, units=[unit]), expected="validation")

    def test_release_first_allows_count_shortage_reconciliation(self):
        self.receive(quantity=1)
        reservation = self.reserve(self.approve(self.request()))
        count = self.record_count(self.start_count(), 0)
        self.run_concurrent(lambda: self.release(reservation), lambda: self.reconcile(count), expected="success")

    def test_stale_adjustment_revision_rejected_after_competing_receipt(self):
        revision = s.position_revision(self.location, self.part)
        self.run_concurrent(self.receive, lambda: self.adjust(1, expected_revision=revision), expected="validation")

    def test_count_first_prevents_empty_location_deactivation(self):
        count = self.draft_count()
        self.run_concurrent(lambda: self.start_count(count), lambda: stock.deactivate_location(actor=self.actor, location=self.location), expected="validation")

    def test_found_count_unit_received_elsewhere_before_reconciliation(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="COUNT-FOUND-RACE")
        count = self.record_count(self.start_count(), 1, units=[unit])
        self.run_concurrent(lambda: self.receive(destination=self.destination, quantity=1, units=[unit]), lambda: self.reconcile(count), expected="validation")

    def test_reconciliation_first_prevents_receiving_same_found_unit_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="COUNT-FOUND-FIRST")
        count = self.record_count(self.start_count(), 1, units=[unit])
        self.run_concurrent(lambda: self.reconcile(count), lambda: self.receive(destination=self.destination, quantity=1, units=[unit]), expected="validation")
