"""Deterministic committed-read schedules using independent PostgreSQL connections."""
from concurrent.futures import ThreadPoolExecutor
from threading import Event
from uuid import uuid4

from django.contrib.auth import get_user_model
from django.db import connection, connections, transaction
from django.test import TransactionTestCase

from apps.commercial.test_payment import PaymentFixture, setup_payment
from apps.reporting import commercial_analytics as commercial, inventory_analytics as inventory
from apps.reporting import service_dashboard as dashboard, permissions as p
from apps.reporting.scope import cases
from apps.reporting.tests.test_reports import grant_report
from apps.inventory import services as stock, tests as stock_fixture
from apps.organization.assignment_services import deactivate_assignment
from apps.service import test_diagnosis as diagnosis
from apps.service.models import ServiceCase


class CommittedReadSchedule:
    def concurrent_write(self, write, read_before_commit, read_after_commit):
        self.assertEqual(connection.vendor, "postgresql")
        written, release = Event(), Event()
        def writer():
            try:
                with transaction.atomic():
                    with connection.cursor() as cursor:
                        cursor.execute("SET LOCAL lock_timeout = '4s'")
                        cursor.execute("SET LOCAL statement_timeout = '12s'")
                    write()
                    written.set()
                    if not release.wait(20):
                        raise AssertionError("Audit reader failed to release writer")
            finally:
                connections.close_all()
        with ThreadPoolExecutor(max_workers=1) as pool:
            future = pool.submit(writer)
            try:
                if not written.wait(15):
                    future.result(timeout=1)
                    self.fail("Writer did not reach uncommitted state")
                with connection.cursor() as cursor:
                    cursor.execute("SET statement_timeout = '5s'")
                read_before_commit()
            finally:
                release.set()
                with connection.cursor() as cursor:
                    cursor.execute("SET statement_timeout = 0")
            future.result(timeout=20)
        read_after_commit()


class ReportingPaymentConcurrency(CommittedReadSchedule, PaymentFixture, TransactionTestCase):
    def setUp(self):
        setup_payment(self)

    def test_payment_post_is_atomic_to_settlement_report_without_blocking(self):
        def check(paid, balance):
            row = commercial.settlements(self.actor, {}).get()
            self.assertEqual((row.paid_amount, row.balance_due), (paid, balance))
        self.concurrent_write(lambda: self.pay(amount="40"), lambda: check(0, 100), lambda: check(40, 60))

    def test_reversal_never_exposes_void_payment_with_valid_settlement(self):
        payment = self.pay()
        def check(net, reversed_amount):
            row = commercial.collection_totals(self.actor, {}).get()
            self.assertEqual((row["net_collected"], row["valid_allocations"], row["reversed"], row["receipts"]), (net, net, reversed_amount, 1))
        self.concurrent_write(lambda: self.reverse_payment(payment), lambda: check(100, 0), lambda: check(0, 100))


class ReportingInventoryConcurrency(CommittedReadSchedule, TransactionTestCase):
    def setUp(self):
        stock_fixture.make_inventory_fixture(self)

    def test_serialized_transfer_is_consistent_with_ledger_per_statement(self):
        unit = stock.register_serialized_unit(actor=self.actor, company=self.company, spare_part=self.serial_part, identifier="AUDIT-RACE-SERIAL")
        stock.receive_stock(actor=self.actor, destination=self.location, spare_part=self.serial_part, quantity=1, units=[unit], reference="AUDIT", idempotency_key=uuid4())
        def move():
            stock.move_stock(actor=self.actor, source=self.location, destination=self.destination, spare_part=self.serial_part, quantity=1, units=[unit], reference="AUDIT", idempotency_key=uuid4())
        def check(at):
            rows = list(inventory.positions(self.actor, {}))
            self.assertEqual(sum(row.on_hand for row in rows), 1)
            self.assertEqual(sum(row.serialized_units for row in rows), 1)
            for row in rows:
                self.assertEqual(row.on_hand, row.serialized_units)
                self.assertEqual(row.location.company_id, self.company.pk)
                if row.on_hand:
                    self.assertEqual(row.location_id, at)
        self.concurrent_write(move, lambda: check(self.location.pk), lambda: check(self.destination.pk))


class ReportingLifecycleConcurrency(CommittedReadSchedule, TransactionTestCase):
    def setUp(self):
        diagnosis.setup_diagnosis(self)
        self.reader = get_user_model().objects.create_user(username="audit-concurrent-reader")
        self.path, _ = grant_report(self.reader, self.company, self.center)

    def test_diagnosis_transition_does_not_block_or_duplicate_current_case(self):
        def check(status):
            report = next(t for t in dashboard.dashboard_tables(self.reader, {}) if t.key == "workflow")
            self.assertEqual(list(report.rows), [{"status": status, "count": 1}])
            row = cases(self.reader, p.OPERATIONAL, {}).get()
            self.assertEqual((row.company_id, row.service_center.company_id), (self.company.pk, self.company.pk))
        self.concurrent_write(lambda: diagnosis.begin(self), lambda: check("ASSIGNED"), lambda: check("DIAGNOSING"))

    def test_lazy_report_rechecks_scope_after_committed_revocation(self):
        lazy = cases(self.reader, p.OPERATIONAL, {})
        self.concurrent_write(lambda: deactivate_assignment(assignment=self.path),
                              lambda: self.assertTrue(lazy.exists()), lambda: self.assertFalse(lazy.exists()))
        self.assertEqual(ServiceCase.objects.count(), 1)
