"""Focused SLA contracts; frozen workflow fixtures are reused without edits."""
from datetime import timedelta, datetime, timezone as tz
from io import StringIO
from unittest.mock import patch
from zoneinfo import ZoneInfo
from concurrent.futures import ThreadPoolExecutor
from threading import Barrier
from django.contrib.auth import get_user_model
from django.contrib.auth.models import Permission
from django.core.exceptions import ValidationError, PermissionDenied
from django.core.management import call_command, CommandError
from django.db import transaction, connections
from django.test import TestCase, TransactionTestCase, Client
from django.urls import reverse
from django.utils import timezone
from apps.service import tests as intake_fixture, test_handover as delivery
from apps.service.services import cancel_service_case
from apps.service.models import ServiceCase
from apps.organization import test_assignment_concurrency as concurrency
from apps.organization.assignment_services import create_assignment
from apps.access.models import Role
from apps.access.services import create_role_assignment, set_role_permissions
from apps.communications.services import save_template
from apps.communications.models import Notification
from apps.customers.services import update_customer
from . import services as s, queries
from .models import SlaPolicy, ServiceSla, SlaEscalation, EscalationCommunication


class Fixture:
    def prepare(self):
        intake_fixture.setup(self)
        self.customer = update_customer(customer=self.customer, primary_mobile="+8801700000097")
        self.admin = get_user_model().objects.create_user(username="sla-admin", is_superuser=True, is_staff=True)
        self.start = timezone.now() - timedelta(hours=2)
        self.case = intake_fixture.intake(self, received_at=self.start)

    def policy(self, **kwargs):
        return s.save_policy(**(dict(actor=self.admin, company=self.company, code="default", name="Acceptance to ready",
            effective_from=timezone.localdate(self.start), target_minutes=60, warning_minutes=15) | kwargs))

    def enroll(self, **kwargs):
        return s.enroll(**(dict(actor=self.admin, service_case=self.case) | kwargs))

    def state(self, row, at):
        return queries.observation(actor=self.admin, sla=row, at=at)["state"]

    def template(self):
        return save_template(actor=self.admin, company=self.company, code="sla-alert", name="Recorded SLA alert",
            channel="SMS", subject="", body="An SLA alert was recorded for case {service_case_reference}, {device_name}, {customer_name}.")

    def grant(self, user, center=None, company=None, permissions=("view_sla", "monitor_sla")):
        path = create_assignment(user=user, company=company or self.company,
            region=center.region if center else None, service_center=center)
        role = Role.objects.create(code="SLA-"+str(user.pk)[:8], name="SLA operator")
        set_role_permissions(role=role, permissions=Permission.objects.filter(content_type__app_label="sla", codename__in=permissions))
        create_role_assignment(user=user, role=role, organization_assignment=path)
        return role


class SlaTests(Fixture, TestCase):
    def setUp(self):
        self.prepare()

    def test_policy_validation(self):
        for kwargs in [dict(target_minutes=0), dict(warning_minutes=61), dict(warning_minutes=-1),
                       dict(effective_until=timezone.localdate(self.start)-timedelta(days=1)), dict(service_center=self.other_center)]:
            with self.subTest(kwargs=kwargs), self.assertRaises(ValidationError):
                self.policy(**kwargs)

    def test_ambiguous_scope_rejected_including_shared_end_date(self):
        day = timezone.localdate(self.start)
        self.policy(effective_until=day)
        with self.assertRaises(ValidationError):
            self.policy(code="overlap", effective_from=day)
        self.policy(code="next", effective_from=day+timedelta(days=1))

    def test_inactive_policy_does_not_conflict_or_match(self):
        self.policy(is_active=False)
        self.assertIsNone(self.enroll())
        winner = self.policy(code="active")
        self.assertEqual(self.enroll().policy, winner)

    def test_precedence_center_over_company_category(self):
        self.policy()
        self.policy(code="category", product_category=self.category)
        winner = self.policy(code="center", service_center=self.center)
        self.assertEqual(self.enroll().policy, winner)

    def test_precedence_center_category_highest(self):
        self.policy()
        self.policy(code="center", service_center=self.center)
        winner = self.policy(code="specific", service_center=self.center, product_category=self.category)
        self.assertEqual(self.enroll().policy, winner)

    def test_company_category_beats_default(self):
        self.policy()
        winner = self.policy(code="category", product_category=self.category)
        self.assertEqual(self.enroll().policy, winner)

    def test_effective_dates_use_acceptance_not_monitor_date(self):
        day = timezone.localdate(self.start)
        self.policy(effective_until=day)
        self.policy(code="future", effective_from=day+timedelta(days=1))
        self.assertEqual(self.enroll().policy.code, "default")

    def test_future_or_expired_policies_do_not_match(self):
        day = timezone.localdate(self.start)
        self.policy(effective_from=day+timedelta(days=1))
        self.policy(code="past", effective_from=day-timedelta(days=3), effective_until=day-timedelta(days=1))
        self.assertIsNone(self.enroll())

    def test_start_due_warning_and_immutable_snapshot(self):
        policy = self.policy()
        row = self.enroll()
        self.assertEqual(row.started_at, self.start)
        self.assertEqual(row.due_at, self.start+timedelta(minutes=60))
        self.assertEqual(row.warning_at, self.start+timedelta(minutes=45))
        self.policy(policy=policy, target_minutes=120, is_active=False)
        self.assertEqual(self.enroll().pk, row.pk)
        row.refresh_from_db()
        self.assertEqual(row.policy_snapshot["target_minutes"], 60)
        for operation in [lambda: row.save(), lambda: row._persist(), lambda: row.delete(),
                          lambda: ServiceSla.objects.filter(pk=row.pk).update(due_at=timezone.now()),
                          lambda: ServiceSla.objects.all().delete()]:
            with self.assertRaises(ValidationError):
                operation()

    def test_exact_warning_due_and_overdue_boundaries(self):
        self.policy()
        row = self.enroll()
        for at, expected in [(row.started_at,"ON_TRACK"), (row.warning_at-timedelta(microseconds=1),"ON_TRACK"),
                             (row.warning_at,"DUE_SOON"), (row.due_at,"DUE_SOON"),
                             (row.due_at+timedelta(microseconds=1),"OVERDUE")]:
            self.assertEqual(self.state(row, at), expected)

    def test_zero_warning_is_deterministic(self):
        self.policy(warning_minutes=0)
        row = self.enroll()
        self.assertEqual(self.state(row, row.due_at-timedelta(microseconds=1)), "ON_TRACK")
        self.assertEqual(self.state(row, row.due_at), "DUE_SOON")

    def test_timezone_equivalent_instants_and_naive_rejected(self):
        self.policy()
        row = self.enroll()
        self.assertEqual(self.state(row, row.due_at.astimezone(ZoneInfo("Asia/Dhaka"))), "DUE_SOON")
        with self.assertRaises(ValidationError):
            self.state(row, datetime(2026, 1, 1))
        with self.assertRaises(ValidationError):
            self.state(row, row.started_at-timedelta(seconds=1))

    def test_local_effective_day_uses_project_timezone(self):
        self.start = datetime(2025, 1, 1, 20, tzinfo=tz.utc)
        self.case = intake_fixture.intake(self, received_at=self.start)
        policy = self.policy(effective_from=datetime(2025, 1, 2).date())
        with timezone.override("America/New_York"):
            self.assertEqual(self.enroll().policy, policy)

    def test_escalation_idempotent_per_level(self):
        self.policy()
        row = self.enroll()
        event, created = s.record_escalation(actor=self.admin, sla=row, at=row.warning_at)
        self.assertTrue(created)
        repeated, created = s.record_escalation(actor=self.admin, sla=row, at=row.due_at)
        self.assertFalse(created)
        self.assertEqual(event.pk, repeated.pk)
        late, created = s.record_escalation(actor=self.admin, sla=row)
        self.assertTrue(created)
        self.assertEqual(late.kind, "OVERDUE")
        self.assertEqual(SlaEscalation.objects.count(), 2)
        self.case.refresh_from_db()
        self.assertEqual(self.case.status, "RECEIVED")

    def test_first_overdue_observation_does_not_invent_due_soon_event(self):
        self.policy()
        event, _ = s.record_escalation(actor=self.admin, sla=self.enroll())
        self.assertEqual(event.kind, "OVERDUE")
        self.assertEqual(SlaEscalation.objects.count(), 1)

    def test_escalation_future_timestamp_rejected(self):
        self.policy()
        with self.assertRaises(ValidationError):
            s.record_escalation(actor=self.admin, sla=self.enroll(), at=timezone.now()+timedelta(hours=1))

    def test_on_track_has_no_escalation(self):
        self.policy(target_minutes=240)
        self.assertEqual(s.record_escalation(actor=self.admin, sla=self.enroll()), (None, False))

    def test_cancellation_stops_clock_and_escalation(self):
        self.policy()
        row = self.enroll()
        case = cancel_service_case(service_case=self.case, reason="Synthetic cancellation", cancelled_by=self.admin)
        observation = queries.observation(actor=self.admin, sla=row, at=timezone.now()+timedelta(days=1))
        self.assertEqual(observation["state"], "CANCELLED")
        self.assertEqual(observation["elapsed"], case.cancelled_at-row.started_at)
        self.assertEqual(s.record_escalation(actor=self.admin, sla=row), (None, False))

    def test_no_retroactive_enrollment_of_cancelled_case(self):
        self.policy()
        cancel_service_case(service_case=self.case, reason="Synthetic cancellation", cancelled_by=self.admin)
        self.assertIsNone(self.enroll())

    def test_staff_and_direct_permissions_do_not_grant_scope(self):
        self.policy()
        self.user.is_staff = True
        self.user.save()
        self.user.user_permissions.set(Permission.objects.filter(content_type__app_label="sla"))
        with self.assertRaises(PermissionDenied):
            self.enroll(actor=self.user)
        self.assertFalse(queries.records(self.user).exists())

    def test_center_scope_and_cross_company_isolation(self):
        self.policy()
        own = self.enroll()
        other = intake_fixture.intake(self, service_center=self.center2)
        self.enroll(service_case=other)
        self.grant(self.user, center=self.center)
        self.assertEqual(list(queries.records(self.user).values_list("pk", flat=True)), [own.pk])
        with self.assertRaises(PermissionDenied):
            self.enroll(actor=self.user, service_case=other)
        foreign = intake_fixture.intake(self, company=self.other_company, service_center=self.other_center, customer=self.outsider)
        with self.assertRaises(PermissionDenied):
            self.enroll(actor=self.user, service_case=foreign)

    def test_center_operator_cannot_manage_company_policy(self):
        self.grant(self.user, center=self.center, permissions=("manage_slapolicy",))
        with self.assertRaises(PermissionDenied):
            self.policy(actor=self.user)

    def test_communication_queued_once_and_not_sent(self):
        self.policy(overdue_template=self.template())
        event, _ = s.record_escalation(actor=self.admin, sla=self.enroll())
        first = s.queue_escalation(actor=self.admin, escalation=event)
        second = s.queue_escalation(actor=self.admin, escalation=event)
        self.assertEqual(first.pk, second.pk)
        self.assertIsNotNone(first.notification_id)
        self.assertEqual(Notification.objects.get().status, "PENDING")
        self.assertEqual(Notification.objects.get().related_id, self.case.pk)
        self.assertEqual(Notification.objects.count(), 1)

    def test_missing_contact_is_a_recorded_queue_failure(self):
        self.customer = update_customer(customer=self.customer, primary_mobile="")
        self.policy(overdue_template=self.template())
        event, _ = s.record_escalation(actor=self.admin, sla=self.enroll())
        result = s.queue_escalation(actor=self.admin, escalation=event)
        self.assertEqual(result.error_code, "NOTIFICATION_REQUEST_FAILED")
        self.assertFalse(Notification.objects.exists())

    def test_policy_post_uses_service_and_preserves_snapshot(self):
        policy = self.policy()
        row = self.enroll()
        self.client.force_login(self.admin)
        response = self.client.post(reverse("sla:policy_edit", args=[policy.pk]), dict(company=self.company.pk,
            code=policy.code, name="Updated", effective_from=policy.effective_from.isoformat(),
            target_minutes=180, warning_minutes=30, is_active="on"))
        self.assertEqual(response.status_code, 302)
        policy.refresh_from_db()
        self.assertEqual(policy.target_minutes, 180)
        row.refresh_from_db()
        self.assertEqual(row.policy_snapshot["target_minutes"], 60)

    def test_escalation_and_communication_evidence_cannot_be_edited(self):
        self.policy(overdue_template=self.template())
        event, _ = s.record_escalation(actor=self.admin, sla=self.enroll())
        result = s.queue_escalation(actor=self.admin, escalation=event)
        for row in (event, result):
            for operation in (row.save, row._persist, row.delete):
                with self.assertRaises(ValidationError):
                    operation()

    def test_communication_failure_isolated_sanitized_and_retryable(self):
        self.policy(overdue_template=self.template())
        row = self.enroll()
        event, _ = s.record_escalation(actor=self.admin, sla=row)
        with patch("apps.sla.services.communications.request_notification", side_effect=RuntimeError("private secret")):
            result = s.queue_escalation(actor=self.admin, escalation=event)
        self.assertEqual(result.error_code, "NOTIFICATION_REQUEST_FAILED")
        self.assertEqual(self.state(row, timezone.now()), "OVERDUE")
        self.assertEqual(SlaEscalation.objects.count(), 1)
        self.assertIsNotNone(s.queue_escalation(actor=self.admin, escalation=event).notification_id)
        self.assertEqual(EscalationCommunication.objects.count(), 2)

    def test_send_permission_is_not_implied_by_monitor_permission(self):
        self.policy(overdue_template=self.template())
        event, _ = s.record_escalation(actor=self.admin, sla=self.enroll())
        self.grant(self.user, center=self.center)
        result = s.queue_escalation(actor=self.user, escalation=event)
        self.assertEqual(result.error_code, "NOTIFICATION_REQUEST_FAILED")
        self.assertFalse(Notification.objects.exists())

    def test_foreign_template_rejected(self):
        template = save_template(actor=self.admin, company=self.other_company, code="foreign", name="Other",
                                 channel="SMS", subject="", body="Synthetic")
        with self.assertRaises(ValidationError):
            self.policy(overdue_template=template)

    def test_dashboard_detail_filter_and_policy_forms(self):
        policy = self.policy()
        row = self.enroll()
        self.client.force_login(self.admin)
        response = self.client.get(reverse("sla:dashboard"))
        self.assertContains(response, self.case.job_number)
        self.assertContains(response, "OVERDUE")
        self.assertContains(self.client.get(reverse("sla:detail", args=[self.case.pk])), "Escalation history")
        self.assertContains(self.client.get(reverse("sla:policy_edit", args=[policy.pk])), "Manage SLA policy")
        self.assertContains(self.client.get(reverse("sla:policy_new")), "Manage SLA policy")
        self.assertEqual(self.client.get(reverse("sla:dashboard"), {"state":"INVALID"}).status_code, 400)
        row.refresh_from_db()
        self.assertEqual(SlaEscalation.objects.count(), 0)

    def test_ui_idor_and_csrf(self):
        self.policy()
        self.enroll()
        self.grant(self.user, center=self.center2)
        self.client.force_login(self.user)
        self.assertEqual(self.client.get(reverse("sla:detail", args=[self.case.pk])).status_code, 404)
        self.assertEqual(self.client.get(reverse("sla:monitor_case", args=[self.case.pk])).status_code, 404)
        self.assertEqual(self.client.get(reverse("sla:dashboard"), {"center":self.center.pk}).status_code, 400)
        client = Client(enforce_csrf_checks=True)
        client.force_login(self.admin)
        self.assertEqual(client.post(reverse("sla:monitor_case", args=[self.case.pk])).status_code, 403)

    def test_case_admin_links_are_scoped(self):
        self.policy()
        self.enroll()
        self.client.force_login(self.admin)
        self.assertContains(self.client.get(reverse("admin:service_servicecase_change", args=[self.case.pk])), "Service SLA")


class CompletionTests(Fixture, TestCase):
    @classmethod
    def setUpTestData(cls):
        delivery.setup_handover(cls, accessories=False)
        cls.admin = get_user_model().objects.create_user(username="sla-admin", is_superuser=True)
        cls.start = cls.case.received_at

    def test_qc_is_not_sla_completion_ready_is_and_pickup_does_not_extend(self):
        self.policy(target_minutes=1440)
        row = self.enroll()
        self.assertEqual(self.state(row, timezone.now()), "ON_TRACK")
        release = delivery.ready(self)
        observation = queries.observation(actor=self.admin, sla=row, at=release.readied_at+timedelta(days=2))
        self.assertEqual(observation["state"], "COMPLETED_ON_TIME")
        self.assertEqual(observation["elapsed"], release.readied_at-row.started_at)
        self.assertEqual(s.record_escalation(actor=self.admin, sla=row), (None, False))

    def test_completion_exactly_due_is_on_time(self):
        self.policy(target_minutes=1, warning_minutes=0)
        row = self.enroll()
        with patch("apps.service.handover_services.timezone.now", return_value=row.due_at):
            release = delivery.ready(self)
        self.assertEqual(release.readied_at, row.due_at)
        self.assertEqual(self.state(row, row.due_at), "COMPLETED_ON_TIME")

    def test_completion_after_due_is_late(self):
        self.policy(target_minutes=1, warning_minutes=0)
        row = self.enroll()
        late = row.due_at+timedelta(seconds=1)
        with patch("apps.service.handover_services.timezone.now", return_value=late):
            delivery.ready(self)
        observation = queries.observation(actor=self.admin, sla=row, at=late+timedelta(days=1))
        self.assertEqual(observation["state"], "COMPLETED_LATE")
        self.assertEqual(observation["overdue"], timedelta(seconds=1))


class MonitoringConcurrencyTests(Fixture, TransactionTestCase):
    run_concurrent = concurrency.AssignmentConcurrencyTests.run_concurrent

    def setUp(self):
        self.prepare()

    def test_concurrent_enrollment_single_snapshot(self):
        self.policy()
        self.run_concurrent(self.enroll, self.enroll, expected="success")
        self.assertEqual(ServiceSla.objects.count(), 1)

    def test_concurrent_policy_overlap_rejected(self):
        self.run_concurrent(self.policy, lambda: self.policy(code="other"), expected="validation")
        self.assertEqual(SlaPolicy.objects.count(), 1)

    def test_policy_edit_and_enrollment_share_company_lock(self):
        policy = self.policy()
        self.run_concurrent(lambda: self.policy(policy=policy, target_minutes=120), self.enroll, expected="success")
        self.assertEqual(ServiceSla.objects.get().policy_snapshot["target_minutes"], 120)

    def test_concurrent_escalation_once(self):
        self.policy()
        row = self.enroll()
        record = lambda: s.record_escalation(actor=self.admin, sla=row)
        self.run_concurrent(record, record, expected="success")
        self.assertEqual(SlaEscalation.objects.count(), 1)

    def test_concurrent_notification_once(self):
        self.policy(overdue_template=self.template())
        event, _ = s.record_escalation(actor=self.admin, sla=self.enroll())
        queue = lambda: s.queue_escalation(actor=self.admin, escalation=event)
        self.run_concurrent(queue, queue, expected="success")
        self.assertEqual(Notification.objects.count(), 1)
        self.assertEqual(EscalationCommunication.objects.count(), 1)

    def test_cancellation_wins_against_escalation(self):
        self.policy()
        row = self.enroll()
        self.run_concurrent(lambda: cancel_service_case(service_case=self.case, reason="Synthetic", cancelled_by=self.admin),
            lambda: s.record_escalation(actor=self.admin, sla=row), expected="success")
        self.assertFalse(SlaEscalation.objects.exists())

    def test_monitor_repeated_command_and_notification_idempotency(self):
        self.policy(overdue_template=self.template())
        out = StringIO()
        call_command("monitor_service_sla", actor=self.admin.pk, stdout=out)
        call_command("monitor_service_sla", actor=self.admin.pk, stdout=out)
        self.assertIn("escalations=1", out.getvalue())
        self.assertIn("escalations=0", out.getvalue())
        self.assertEqual(ServiceSla.objects.count(), 1)
        self.assertEqual(SlaEscalation.objects.count(), 1)
        self.assertEqual(Notification.objects.count(), 1)

    def test_simultaneous_monitor_runs(self):
        self.policy(overdue_template=self.template())
        barrier = Barrier(2)
        original = s.enroll
        def synchronized(**kwargs):
            barrier.wait(timeout=10)
            return original(**kwargs)
        def run():
            try:
                return s.monitor(actor=self.admin)
            finally:
                connections.close_all()
        with patch("apps.sla.services.enroll", side_effect=synchronized), ThreadPoolExecutor(max_workers=2) as executor:
            futures = [executor.submit(run) for _ in range(2)]
            for future in futures:
                future.result(timeout=30)
        self.assertEqual(ServiceSla.objects.count(), 1)
        self.assertEqual(SlaEscalation.objects.count(), 1)
        self.assertEqual(Notification.objects.count(), 1)
        self.assertEqual(EscalationCommunication.objects.count(), 1)

    def test_monitor_failure_does_not_rollback_escalation(self):
        self.policy(overdue_template=self.template())
        with patch("apps.sla.services.communications.request_notification", side_effect=RuntimeError("Synthetic")):
            result = s.monitor(actor=self.admin)
        self.assertEqual(result["communication_failures"], 1)
        self.assertEqual(SlaEscalation.objects.count(), 1)
        self.case.refresh_from_db()
        self.assertEqual(self.case.status, "RECEIVED")

    def test_command_rejects_unauthorized_actor(self):
        with self.assertRaises(CommandError):
            call_command("monitor_service_sla", actor=self.user.pk, stdout=StringIO())

    def test_monitor_requires_commit_boundary(self):
        with transaction.atomic(), self.assertRaises(ValidationError):
            s.monitor(actor=self.admin)

    def test_manual_monitor_endpoint(self):
        self.policy()
        self.client.force_login(self.admin)
        response = self.client.post(reverse("sla:monitor_case", args=[self.case.pk]))
        self.assertEqual(response.status_code, 302)
        self.assertEqual(SlaEscalation.objects.count(), 1)
