"""Real PostgreSQL lock waits, fresh dependency validation and rollback."""
from unittest.mock import patch

from django.contrib.auth import get_user_model
from django.contrib.auth.models import Permission
from django.core.exceptions import ValidationError
from django.db import transaction
from django.test import Client, TransactionTestCase
from django.urls import reverse

from apps.access import services as access
from apps.catalog import services as catalog
from apps.devices import services as devices
from apps.organization import services as organization, test_assignment_concurrency as concurrency
from apps.service_catalog import services as taxonomy
from . import repair_services as services, repair_queries as queries
from .models import ServiceCase, ServiceRepairExecution, ServiceRepairAction
from .services import cancel_service_case
from .test_repair import setup_repair, begin, add, perform, complete, abandon, prepared, assert_repair_invariants
from .test_engineer_assignment import reassign


class RepairConcurrencyTests(TransactionTestCase):
    run_concurrent = concurrency.AssignmentConcurrencyTests.run_concurrent

    def setUp(self):
        setup_repair(self)

    def race(self, first, second, expected="validation"):
        self.run_concurrent(first, second, expected=expected)
        assert_repair_invariants(self)

    def cancel(self, **kwargs):
        return cancel_service_case(service_case=self.case, cancelled_by=self.user, **kwargs)

    def remove(self, action):
        return services.remove_repair_action(action=action, actor=self.engineer)

    def update(self, action, **kwargs):
        return services.update_repair_action(action=action, actor=self.engineer, **kwargs)

    def test_two_begin_attempts(self):
        self.race(lambda: begin(self), lambda: begin(self))
        self.assertEqual(ServiceRepairExecution.objects.count(), 1)

    def test_begin_then_cancel(self):
        self.race(lambda: begin(self), self.cancel, "success")
        self.assertEqual(ServiceRepairExecution.objects.get().status, "ABANDONED")

    def test_cancel_then_begin(self):
        self.race(self.cancel, lambda: begin(self))
        self.assertFalse(ServiceRepairExecution.objects.exists())

    def test_revocation_then_begin(self):
        self.race(lambda: access.deactivate_role_assignment(assignment=self.role_assignment), lambda: begin(self))
        self.assertFalse(ServiceRepairExecution.objects.exists())

    def test_begin_then_reassignment_attempt(self):
        self.race(lambda: begin(self), lambda: reassign(self))

    def test_competing_distinct_action_additions(self):
        execution = begin(self)
        self.race(lambda: add(self, execution), lambda: add(self, execution, repair_action=self.action_type2), "success")
        self.assertEqual(queries.active_repair_actions(execution).count(), 2)

    def test_duplicate_action_additions(self):
        execution = begin(self)
        self.race(lambda: add(self, execution), lambda: add(self, execution))
        self.assertEqual(ServiceRepairAction.objects.count(), 1)

    def test_update_then_remove_stale_snapshot(self):
        _, action = prepared(self, performed=False)
        self.race(lambda: self.update(action, note="newer"), lambda: self.remove(action))
        action.refresh_from_db()
        self.assertTrue(action.is_active)
        self.assertEqual(action.note, "newer")

    def test_remove_then_update(self):
        _, action = prepared(self, performed=False)
        self.race(lambda: self.remove(action), lambda: self.update(action, note="stale"))
        action.refresh_from_db()
        self.assertFalse(action.is_active)

    def test_perform_then_remove_stale_snapshot(self):
        _, action = prepared(self, performed=False)
        self.race(lambda: perform(self, action), lambda: self.remove(action))
        action.refresh_from_db()
        self.assertTrue(action.is_active)
        self.assertIsNotNone(action.performed_at)

    def test_remove_then_perform(self):
        _, action = prepared(self, performed=False)
        self.race(lambda: self.remove(action), lambda: perform(self, action))
        action.refresh_from_db()
        self.assertIsNone(action.performed_at)

    def test_perform_twice(self):
        _, action = prepared(self, performed=False)
        self.race(lambda: perform(self, action), lambda: perform(self, action))

    def test_add_then_completion_rejects_changed_plan(self):
        execution, _ = prepared(self)
        self.race(lambda: add(self, execution, repair_action=self.action_type2), lambda: complete(self, execution))
        self.assertEqual(queries.active_repair_actions(execution).count(), 2)

    def test_remove_then_complete(self):
        execution, action = prepared(self)
        self.race(lambda: self.remove(action), lambda: complete(self, execution))
        self.assertFalse(queries.active_repair_actions(execution).exists())

    def test_complete_then_add(self):
        execution, _ = prepared(self)
        self.race(lambda: complete(self, execution), lambda: add(self, execution, repair_action=self.action_type2))

    def test_complete_then_remove(self):
        execution, action = prepared(self)
        self.race(lambda: complete(self, execution), lambda: self.remove(action))

    def test_perform_then_completion_rejects_changed_snapshot(self):
        execution, action = prepared(self, performed=False)
        self.race(lambda: perform(self, action), lambda: complete(self, execution))
        self.assertEqual(queries.current_repair_execution(self.case).pk, execution.pk)

    def test_taxonomy_deactivation_then_completion(self):
        execution, _ = prepared(self)
        self.race(lambda: taxonomy.deactivate_repair_action(repair_action=self.action_type), lambda: complete(self, execution))

    def test_complete_then_taxonomy_deactivation(self):
        execution, _ = prepared(self)
        self.race(lambda: complete(self, execution), lambda: taxonomy.deactivate_repair_action(repair_action=self.action_type), "success")
        self.assertEqual(ServiceRepairExecution.objects.get().outcome, "REPAIRED")

    def test_applicability_removal_then_completion(self):
        execution, _ = prepared(self)
        self.race(lambda: taxonomy.set_repair_action_applicability(repair_action=self.action_type,
            applies_to_all_product_categories=False, product_categories=[]), lambda: complete(self, execution))

    def test_complete_then_applicability_replacement(self):
        execution, _ = prepared(self)
        self.race(lambda: complete(self, execution), lambda: taxonomy.set_repair_action_applicability(repair_action=self.action_type,
            applies_to_all_product_categories=False, product_categories=[self.category]), "success")

    def test_role_assignment_revocation_then_completion(self):
        execution, _ = prepared(self)
        self.race(lambda: access.deactivate_role_assignment(assignment=self.role_assignment), lambda: complete(self, execution))

    def test_permission_revocation_then_completion(self):
        execution, _ = prepared(self)
        self.race(lambda: access.set_role_permissions(role=self.role, permissions=[]), lambda: complete(self, execution))

    def test_user_deactivation_then_completion(self):
        execution, _ = prepared(self)
        def deactivate():
            user = get_user_model().objects.select_for_update().get(pk=self.engineer.pk)
            user.is_active = False
            user.save(update_fields=["is_active"])
        self.race(deactivate, lambda: complete(self, execution))

    def test_company_deactivation_then_completion(self):
        execution, _ = prepared(self)
        self.race(lambda: organization.deactivate_company(company=self.company), lambda: complete(self, execution))

    def test_device_deactivation_then_begin(self):
        self.race(lambda: devices.deactivate_device(device=self.device), lambda: begin(self))

    def test_catalog_deactivation_then_completion(self):
        execution, _ = prepared(self)
        self.race(lambda: catalog.deactivate_category(category=self.category), lambda: complete(self, execution))

    def test_complete_then_cancel_denied(self):
        execution, _ = prepared(self)
        self.race(lambda: complete(self, execution), self.cancel)

    def test_cancel_then_complete(self):
        execution, _ = prepared(self)
        self.race(self.cancel, lambda: complete(self, execution))

    def test_abandon_then_complete(self):
        execution, _ = prepared(self)
        self.race(lambda: abandon(self, execution), lambda: complete(self, execution))

    def test_complete_then_abandon(self):
        execution, _ = prepared(self)
        self.race(lambda: complete(self, execution), lambda: abandon(self, execution))

    def test_abandon_then_cancel(self):
        execution, _ = prepared(self)
        self.race(lambda: abandon(self, execution), self.cancel, "success")

    def test_cancel_then_abandon(self):
        execution, _ = prepared(self)
        self.race(self.cancel, lambda: abandon(self, execution))

    def test_not_repaired_then_explicit_begin(self):
        execution, _ = prepared(self)
        self.race(lambda: complete(self, execution, outcome="NOT_REPAIRED"), lambda: begin(self), "success")
        self.assertEqual(ServiceRepairExecution.objects.count(), 2)
        self.assertEqual(ServiceRepairExecution.objects.filter(status="COMPLETED").get().outcome, "NOT_REPAIRED")

    def test_not_repaired_then_stale_admin_begin_precondition(self):
        execution, _ = prepared(self)
        self.case.refresh_from_db()
        self.race(lambda: complete(self, execution, outcome="NOT_REPAIRED"),
            lambda: begin(self, expected_updated_at=self.case.updated_at))
        self.assertEqual(ServiceRepairExecution.objects.count(), 1)

    def test_aggregate_revision_change_then_action_edit(self):
        execution, action = prepared(self, performed=False)
        execution.refresh_from_db()
        self.race(lambda: services.update_repair_execution(repair_execution=execution, actor=self.engineer, note="newer"),
            lambda: self.update(action, note="stale", expected_updated_at=action.updated_at,
                expected_execution_updated_at=execution.updated_at))
        action.refresh_from_db()
        self.assertEqual(action.note, "")

    def test_successful_completion_vs_stale_admin_submission(self):
        execution, action = prepared(self)
        self.engineer.is_staff = True
        self.engineer.save()
        self.engineer.user_permissions.add(Permission.objects.get(content_type__app_label="service", codename="change_servicecase"))
        client = Client()
        client.force_login(self.engineer)
        url = reverse("admin:service_servicecase_repair", args=[self.case.pk])
        token = client.get(url).context["form"].initial["revision"]
        def submit():
            response = client.post(url, dict(operation="remove", revision=token, action=action.pk))
            self.assertContains(response, "Only the currently eligible assigned engineer")
        self.race(lambda: complete(self, execution), submit, "success")
        action.refresh_from_db()
        self.assertTrue(action.is_active)

    def test_late_completion_failure_then_valid_completion(self):
        execution, _ = prepared(self)
        before = ServiceRepairExecution.objects.values().get(pk=execution.pk)
        def failed_complete():
            with patch.object(ServiceCase, "_persist", side_effect=ValidationError("Synthetic late failure")):
                with self.assertRaises(ValidationError):
                    complete(self, execution)
            self.assertEqual(ServiceRepairExecution.objects.values().get(pk=execution.pk), before)
            ServiceCase.objects.select_for_update().get(pk=self.case.pk)
        self.race(failed_complete, lambda: complete(self, execution), "success")

    def test_cancellation_failure_then_completion(self):
        execution, _ = prepared(self)
        def failed_cancel():
            with self.assertRaises(ValidationError):
                self.cancel(reason="x" * 2001)
            self.assertEqual(queries.current_repair_execution(self.case).pk, execution.pk)
            ServiceCase.objects.select_for_update().get(pk=self.case.pk)
        self.race(failed_cancel, lambda: complete(self, execution), "success")

    def test_begin_rollback_then_begin(self):
        def failed_begin():
            with self.assertRaises(ValidationError):
                with transaction.atomic():
                    begin(self)
                    raise ValidationError("Synthetic rollback")
            self.assertFalse(ServiceRepairExecution.objects.exists())
            ServiceCase.objects.select_for_update().get(pk=self.case.pk)
        self.race(failed_begin, lambda: begin(self), "success")
