from datetime import timedelta

from django.contrib.auth import get_user_model
from django.test import TestCase
from django.utils import timezone

from apps.service import test_diagnosis as diagnosis, test_repair as repair
from apps.service.models import ServiceRepairAction, ServiceDiagnosticFinding
from apps.service_catalog.models import RootCause, RepairAction
from apps.reporting.service_analytics import repair_action_diagnosis_cooccurrence, performed_actions, service_tables


class CooccurrenceTests(TestCase):
    @classmethod
    def setUpTestData(cls):
        fixture = type("Fixture", (), {})()
        diagnosis.setup_diagnosis(fixture)
        fixture.root2 = RootCause.objects.create(code="ROOT-TWO", name="Second root", applies_to_all_product_categories=True)
        assessment = diagnosis.begin(fixture)
        for fault, root in ((fixture.fault, None), (fixture.fault2, None), (fixture.fault, fixture.root), (fixture.fault2, fixture.root2), (fixture.fault, fixture.root2)):
            diagnosis.add(fixture, assessment, fault_diagnosis=fault, root_cause=root)
        fixture.assessment = diagnosis.complete(fixture, assessment)
        fixture.action_type = RepairAction.objects.create(code="ACTION-A", name="Action A", applies_to_all_product_categories=True)
        fixture.action_type2 = RepairAction.objects.create(code="ACTION-B", name="Action B", applies_to_all_product_categories=True)
        fixture.planned_type = RepairAction.objects.create(code="ACTION-C", name="Action C", applies_to_all_product_categories=True)
        fixture.execution, fixture.action = repair.prepared(fixture)
        fixture.action2 = repair.perform(fixture, repair.add(fixture, fixture.execution, repair_action=fixture.action_type2))
        fixture.planned = repair.add(fixture, fixture.execution, repair_action=fixture.planned_type)
        fixture.reader = get_user_model().objects.create_superuser(username="cooccurrence-reader")
        for key, value in vars(fixture).items():
            setattr(cls, key, value)

    def rows(self, **kwargs):
        return list(repair_action_diagnosis_cooccurrence(self.reader, {}, **kwargs))

    def test_multiple_diagnoses_count_each_action_once_per_diagnosis(self):
        rows = self.rows(grouping="diagnosis")
        self.assertEqual(len(rows), 4)
        self.assertEqual({r["performed_actions"] for r in rows}, {1})
        self.assertEqual({r["fault_diagnosis_id"] for r in rows}, {self.fault.pk, self.fault2.pk})

    def test_multiple_causes_duplicate_groupings_are_deduplicated(self):
        rows = self.rows(grouping="root_cause")
        self.assertEqual(len(rows), 6)
        self.assertEqual({r["performed_actions"] for r in rows}, {1})
        self.assertEqual({r["root_cause_id"] for r in rows}, {None, self.root.pk, self.root2.pk})

    def test_combined_grouping_retains_five_distinct_pairs_per_action(self):
        rows = self.rows()
        self.assertEqual(len(rows), 10)
        self.assertEqual({r["performed_actions"] for r in rows}, {1})

    def test_null_cause_explicit_label_without_synthetic_record(self):
        before = RootCause.objects.count()
        unknown = [r for r in self.rows(grouping="root_cause") if r["root_cause_id"] is None]
        self.assertEqual(len(unknown), 2)
        self.assertEqual({r["root_cause_label"] for r in unknown}, {"Unknown / Unconfirmed"})
        self.assertEqual(RootCause.objects.count(), before)

    def test_planned_action_excluded(self):
        self.assertEqual(performed_actions(self.reader, {}).count(), 2)
        self.assertNotIn(self.planned_type.pk, {r["action_id"] for r in self.rows()})

    def test_overlapping_groups_are_not_unique_action_total(self):
        self.assertEqual(performed_actions(self.reader, {}).count(), 2)
        self.assertEqual(sum(r["performed_actions"] for r in self.rows()), 10)
        reports = service_tables(self.reader, {})
        total = next(t for t in reports if t.key == "actions")
        self.assertEqual(sum(r["count"] for r in total.rows), 2)
        overlap = next(t for t in reports if t.key == "cooccurrence_diagnosis_root_cause")
        self.assertIn("MUST NOT be added", overlap.note)
        self.assertIn("Co-occurrence", overlap.title)
        self.client.force_login(self.reader)
        response = self.client.get("/reports/service/", {"table": "cooccurrence_diagnosis_root_cause"})
        self.assertEqual(response.context["cards"]["Distinct performed actions"], 2)

    def test_cooccurrence_links_have_distinct_explicit_labels(self):
        reports = [t for t in service_tables(self.reader, {}) if t.key.startswith("cooccurrence")]
        self.assertEqual(len({t.title for t in reports}), len(reports))
        self.assertTrue(all("Co-occurrence" in t.title for t in reports))

    def test_authoritative_product_and_center_dimensions(self):
        for dimension, expected in (("brand", self.brand.pk), ("product_category", self.category.pk), ("model", self.model.pk), ("variant", None), ("service_center", self.center.pk)):
            self.assertEqual({r["dimension_id"] for r in self.rows(dimension=dimension)}, {expected})

    def test_company_center_and_performed_period_filters(self):
        for filters in ({"company": self.other_company.pk}, {"service_center": self.center2.pk}, {"date_from": timezone.localdate() + timedelta(days=1)}):
            self.assertFalse(repair_action_diagnosis_cooccurrence(self.reader, filters).exists())

    def test_read_only_and_no_manufactured_relationship(self):
        actions = list(ServiceRepairAction.objects.order_by("pk").values())
        findings = list(ServiceDiagnosticFinding.objects.order_by("pk").values())
        self.rows()
        self.assertEqual(list(ServiceRepairAction.objects.order_by("pk").values()), actions)
        self.assertEqual(list(ServiceDiagnosticFinding.objects.order_by("pk").values()), findings)
        self.assertNotIn("finding", {field.name for field in ServiceRepairAction._meta.fields})

    def test_invalid_dimensions_rejected(self):
        with self.assertRaises(ValueError):
            self.rows(grouping="causal")
        with self.assertRaises(ValueError):
            self.rows(dimension="complaint_category")
