from concurrent.futures import ThreadPoolExecutor
from threading import Event
from time import monotonic, sleep

from django.contrib.auth import get_user_model
from django.core.exceptions import ValidationError
from django.db import connection, connections, IntegrityError, transaction
from django.test import TransactionTestCase

from .assignments import UserOrganizationAssignment as Assignment
from .assignment_services import create_assignment, set_primary_assignment
from . import services
from .test_lifecycle import make_tree, LifecycleConcurrencyTests


class AssignmentConcurrencyTests(TransactionTestCase):
    def setUp(self):
        self.company, self.region, self.center, self.department = make_tree("CCARE")
        self.user = get_user_model().objects.create_user(username="concurrent")
        self.other_user = get_user_model().objects.create_user(username="concurrent-other")

    def run_concurrent(self, first_write, second_write, *, expected, should_block=True):
        first_written, release, attempting = Event(), Event(), Event()
        pid = []

        def owner():
            try:
                with transaction.atomic():
                    first_write()
                    first_written.set()
                    if not release.wait(10):
                        raise AssertionError("Timed out releasing first transaction")
            finally:
                connections.close_all()

        def contender():
            try:
                with connection.cursor() as cursor:
                    cursor.execute("SELECT pg_backend_pid()")
                    pid.append(cursor.fetchone()[0])
                    cursor.execute("SET lock_timeout = '8s'")
                attempting.set()
                try:
                    second_write()
                except ValidationError:
                    return "validation"
                except IntegrityError:
                    return "integrity"
                return "success"
            finally:
                connections.close_all()

        with ThreadPoolExecutor(max_workers=2) as executor:
            first = executor.submit(owner)
            try:
                self.assertTrue(first_written.wait(5))
                second = executor.submit(contender)
                self.assertTrue(attempting.wait(5))
                if should_block:
                    deadline, blocked = monotonic() + 5, False
                    while monotonic() < deadline:
                        with connection.cursor() as cursor:
                            cursor.execute("SELECT cardinality(pg_blocking_pids(%s)) > 0", [pid[0]])
                            blocked = cursor.fetchone()[0]
                        if blocked:
                            break
                        sleep(0.02)
                    self.assertTrue(blocked, "Expected an actual PostgreSQL lock wait")
                else:
                    self.assertEqual(second.result(timeout=5), expected)
            finally:
                release.set()
            first.result(timeout=10)
            self.assertEqual(second.result(timeout=10), expected)

    def test_primary_switches_serialize_on_user(self):
        first = create_assignment(user=self.user, company=self.company)
        second = create_assignment(user=self.user, company=self.company, region=self.region)
        self.run_concurrent(lambda: set_primary_assignment(assignment=first),
                            lambda: set_primary_assignment(assignment=second), expected="success")
        self.assertEqual(list(Assignment.objects.filter(is_primary=True).values_list("pk", flat=True)), [second.pk])

    def test_database_primary_constraint_handles_bypassing_writers(self):
        self.run_concurrent(
            lambda: Assignment.objects.bulk_create([Assignment(user=self.user, company=self.company, is_primary=True)]),
            lambda: Assignment.objects.bulk_create([Assignment(user=self.user, company=self.company, region=self.region, is_primary=True)]),
            expected="integrity",
        )
        self.assertEqual(Assignment.objects.filter(is_primary=True).count(), 1)

    def test_database_nullable_duplicate_constraint_handles_race(self):
        def insert():
            Assignment.objects.bulk_create([Assignment(user=self.user, company=self.company)])
        self.run_concurrent(insert, insert, expected="integrity")
        self.assertEqual(Assignment.objects.count(), 1)

    def test_different_users_assignments_share_company_without_serializing(self):
        self.run_concurrent(
            lambda: create_assignment(user=self.user, company=self.company),
            lambda: create_assignment(user=self.other_user, company=self.company),
            expected="success", should_block=False,
        )
        self.assertEqual(Assignment.objects.count(), 2)

    def test_company_deactivation_blocks_then_rejects_assignment_creation(self):
        LifecycleConcurrencyTests.assert_write_waits_then_rejects(
            self, lambda: services.deactivate_company(company=self.company),
            lambda: create_assignment(user=self.user, company=self.company),
        )
        self.assertEqual(Assignment.objects.count(), 0)

    def test_department_deactivation_blocks_then_rejects_assignment_creation(self):
        LifecycleConcurrencyTests.assert_write_waits_then_rejects(
            self, lambda: services.deactivate_department(department=self.department),
            lambda: create_assignment(user=self.user, company=self.company, department=self.department),
        )
        self.assertEqual(Assignment.objects.count(), 0)

    def test_assignment_commit_then_organization_deactivation_ends_assignment(self):
        self.run_concurrent(
            lambda: create_assignment(user=self.user, company=self.company, is_primary=True),
            lambda: services.deactivate_company(company=self.company), expected="success",
        )
        assignment = Assignment.objects.get()
        self.assertFalse(assignment.is_active)
        self.assertFalse(assignment.is_primary)
