"""Real PostgreSQL QC serialization and fresh-state validation."""
from unittest.mock import patch

from django.contrib.auth import get_user_model
from django.core.exceptions import ValidationError
from django.db import transaction
from django.test import TransactionTestCase

from apps.access import services as access
from apps.organization import services as organization, assignment_services as postings, test_assignment_concurrency as concurrency
from apps.catalog import services as catalog
from apps.devices import services as devices
from . import quality_control_services as services, quality_control_queries as queries
from .models import ServiceCase, ServiceQualityControl
from .services import cancel_service_case
from .test_quality_control import setup_qc, pending, begin, prepared, fill, check, pass_qc, fail_qc, abandon, invariants
from . import test_repair as repair


class QualityControlConcurrencyTests(TransactionTestCase):
    run_concurrent = concurrency.AssignmentConcurrencyTests.run_concurrent

    def setUp(self):
        setup_qc(self)

    def race(self, first, second, expected="validation"):
        self.run_concurrent(first, second, expected=expected)
        invariants(self)

    def cancel(self):
        return cancel_service_case(service_case=self.case, cancelled_by=self.user)

    def rejected_then_hold_case(self, operation):
        with self.assertRaises(ValidationError):
            operation()
        ServiceCase.objects.select_for_update().get(pk=self.case.pk)

    def test_two_inspectors_begin(self):
        pending(self)
        self.race(lambda: begin(self), lambda: begin(self, inspector=self.inspector2))
        self.assertEqual(ServiceQualityControl.objects.count(), 1)
        self.assertEqual(ServiceQualityControl.objects.get().inspector, self.inspector)

    def test_same_inspector_duplicate_begin(self):
        pending(self)
        self.race(lambda: begin(self), lambda: begin(self))
        self.assertEqual(ServiceQualityControl.objects.count(), 1)

    def test_pass_then_fail(self):
        row = prepared(self, filled=True)
        self.race(lambda: pass_qc(self, row), lambda: fail_qc(self, row))
        self.assertEqual(ServiceQualityControl.objects.get().outcome, "PASSED")

    def test_fail_then_pass(self):
        row = prepared(self, filled=True)
        def failure():
            check(self, row, result="FAIL")
            fail_qc(self, row)
        self.race(failure, lambda: pass_qc(self, row))
        self.assertEqual(ServiceQualityControl.objects.get().outcome, "FAILED")

    def test_pass_then_cancellation_denied(self):
        row = prepared(self, filled=True)
        self.race(lambda: pass_qc(self, row), self.cancel)

    def test_cancellation_denied_then_pass(self):
        row = prepared(self, filled=True)
        self.race(lambda: self.rejected_then_hold_case(self.cancel), lambda: pass_qc(self, row), "success")

    def test_fail_then_rework_cancellation_allowed(self):
        row = prepared(self)
        check(self, row, result="FAIL")
        self.race(lambda: fail_qc(self, row), self.cancel, "success")
        self.case.refresh_from_db()
        self.assertEqual(self.case.status, "CANCELLED")
        self.assertEqual(ServiceQualityControl.objects.get().outcome, "FAILED")

    def test_cancellation_denied_then_fail(self):
        row = prepared(self)
        check(self, row, result="FAIL")
        self.race(lambda: self.rejected_then_hold_case(self.cancel), lambda: fail_qc(self, row), "success")

    def test_begin_then_repair_start_rejected(self):
        pending(self)
        self.race(lambda: begin(self), lambda: repair.begin(self))

    def test_repair_start_denied_then_qc_begin(self):
        pending(self)
        self.race(lambda: self.rejected_then_hold_case(lambda: repair.begin(self)), lambda: begin(self), "success")

    def test_revocation_then_begin(self):
        pending(self)
        self.race(lambda: access.deactivate_role_assignment(assignment=self.qc_role_assignment), lambda: begin(self))
        self.assertFalse(ServiceQualityControl.objects.exists())

    def test_begin_then_revocation_preserves_open_attempt(self):
        pending(self)
        self.race(lambda: begin(self), lambda: access.deactivate_role_assignment(assignment=self.qc_role_assignment), "success")
        self.assertIsNotNone(queries.current_quality_control(self.case))
        self.assertFalse(queries.qc_cases_for_inspector(self.inspector).exists())

    def test_permission_revocation_then_pass(self):
        row = prepared(self, filled=True)
        self.race(lambda: access.set_role_permissions(role=self.qc_role, permissions=[]), lambda: pass_qc(self, row))

    def test_pass_then_permission_revocation_preserves_history(self):
        row = prepared(self, filled=True)
        self.race(lambda: pass_qc(self, row), lambda: access.set_role_permissions(role=self.qc_role, permissions=[]), "success")
        self.assertEqual(ServiceQualityControl.objects.get().outcome, "PASSED")

    def test_posting_deactivation_then_fail(self):
        row = prepared(self)
        check(self, row, result="FAIL")
        self.race(lambda: postings.deactivate_assignment(assignment=self.qc_path), lambda: fail_qc(self, row))

    def test_user_deactivation_then_pass(self):
        row = prepared(self, filled=True)
        def deactivate():
            user = get_user_model().objects.select_for_update().get(pk=self.inspector.pk)
            user.is_active = False
            user.save(update_fields=["is_active"])
        self.race(deactivate, lambda: pass_qc(self, row))

    def test_check_update_then_stale_completion(self):
        row = prepared(self, filled=True)
        expected = row.updated_at
        self.race(lambda: check(self, row, result="FAIL"), lambda: pass_qc(self, row, expected_updated_at=expected))
        self.assertEqual(ServiceQualityControl.objects.get().status, "IN_PROGRESS")

    def test_completion_then_stale_check_edit(self):
        row = prepared(self, filled=True)
        expected = row.updated_at
        self.race(lambda: pass_qc(self, row), lambda: check(self, row, result="FAIL", expected_updated_at=expected))
        self.assertEqual(queries.quality_control_checks(row).get(check_code="BOOT").result, "PASS")

    def test_competing_check_writes_require_aggregate_revision(self):
        row = prepared(self)
        expected = row.updated_at
        self.race(lambda: check(self, row, result="FAIL"), lambda: check(self, row, result="PASS", expected_updated_at=expected))
        self.assertEqual(queries.quality_control_checks(row).get(check_code="BOOT").result, "FAIL")

    def test_abandon_then_pass(self):
        row = prepared(self, filled=True)
        self.race(lambda: abandon(self, row), lambda: pass_qc(self, row))
        self.assertEqual(ServiceQualityControl.objects.get().status, "ABANDONED")

    def test_pass_then_abandon(self):
        row = prepared(self, filled=True)
        self.race(lambda: pass_qc(self, row), lambda: abandon(self, row))

    def test_abandon_then_new_inspector_attempt(self):
        row = prepared(self)
        def new_attempt():
            return begin(self, inspector=self.inspector2, expected_updated_at=self.case.updated_at)
        # Stale case revision cannot cross abandonment even if the old form was pending.
        self.race(lambda: abandon(self, row), new_attempt)
        self.case.refresh_from_db()
        new = begin(self, inspector=self.inspector2)
        self.assertNotEqual(row.pk, new.pk)

    def test_fail_then_new_repair_then_new_qc_cycle(self):
        row = prepared(self)
        check(self, row, result="FAIL")
        self.race(lambda: fail_qc(self, row), lambda: repair.begin(self), "success")
        execution = repair.queries.current_repair_execution(self.case)
        action = repair.add(self, execution)
        repair.perform(self, action)
        repair.complete(self, execution)
        next_qc = prepared(self, filled=True)
        pass_qc(self, next_qc)
        self.assertNotEqual(row.repair_execution_id, next_qc.repair_execution_id)
        self.assertEqual(ServiceQualityControl.objects.count(), 2)
        invariants(self)

    def test_self_qc_competes_with_valid_inspector(self):
        access.create_role_assignment(user=self.engineer, role=self.qc_role, organization_assignment=self.path)
        pending(self)
        self.race(lambda: begin(self), lambda: begin(self, inspector=self.engineer))
        self.assertEqual(ServiceQualityControl.objects.get().inspector, self.inspector)

    def invalid_scope_candidate(self, **scope):
        user = get_user_model().objects.create_user(username="out-of-scope-qc")
        path = postings.create_assignment(user=user, **scope)
        access.create_role_assignment(user=user, role=self.qc_role, organization_assignment=path)
        return user

    def test_cross_company_candidate_competes_with_valid_inspector(self):
        outsider = self.invalid_scope_candidate(company=self.other_company)
        pending(self)
        self.race(lambda: begin(self), lambda: begin(self, inspector=outsider))

    def test_wrong_center_candidate_competes_with_valid_inspector(self):
        outsider = self.invalid_scope_candidate(company=self.company, region=self.region, service_center=self.center2)
        pending(self)
        self.race(lambda: begin(self), lambda: begin(self, inspector=outsider))

    def test_late_pass_failure_then_pass(self):
        row = prepared(self, filled=True)
        before = ServiceQualityControl.objects.values().get(pk=row.pk)
        def failure():
            with patch.object(ServiceCase, "_persist", side_effect=ValidationError("Synthetic late failure")):
                with self.assertRaises(ValidationError):
                    pass_qc(self, row)
            self.assertEqual(ServiceQualityControl.objects.values().get(pk=row.pk), before)
            ServiceCase.objects.select_for_update().get(pk=self.case.pk)
        self.race(failure, lambda: pass_qc(self, row), "success")

    def test_late_fail_failure_then_fail(self):
        row = prepared(self)
        check(self, row, result="FAIL")
        before = ServiceQualityControl.objects.values().get(pk=row.pk)
        def failure():
            with patch.object(ServiceCase, "_persist", side_effect=ValidationError("Synthetic late failure")):
                with self.assertRaises(ValidationError):
                    fail_qc(self, row)
            self.assertEqual(ServiceQualityControl.objects.values().get(pk=row.pk), before)
            ServiceCase.objects.select_for_update().get(pk=self.case.pk)
        self.race(failure, lambda: fail_qc(self, row), "success")

    def test_begin_rollback_then_begin(self):
        pending(self)
        def failure():
            with self.assertRaises(ValidationError):
                with transaction.atomic():
                    begin(self)
                    raise ValidationError("Synthetic rollback")
            self.assertFalse(ServiceQualityControl.objects.exists())
            ServiceCase.objects.select_for_update().get(pk=self.case.pk)
        self.race(failure, lambda: begin(self), "success")

    def test_company_deactivation_then_begin(self):
        pending(self)
        self.race(lambda: organization.deactivate_company(company=self.company), lambda: begin(self))

    def test_center_deactivation_then_pass(self):
        row = prepared(self, filled=True)
        self.race(lambda: organization.deactivate_service_center(service_center=self.center), lambda: pass_qc(self, row))

    def test_device_deactivation_then_pass(self):
        row = prepared(self, filled=True)
        self.race(lambda: devices.deactivate_device(device=self.device), lambda: pass_qc(self, row))

    def test_catalog_deactivation_then_begin(self):
        pending(self)
        self.race(lambda: catalog.deactivate_category(category=self.category), lambda: begin(self))
