"""Independent Phase 3D reconciliation using synthetic, service-created evidence."""
from collections import Counter, defaultdict
import csv
from io import StringIO
from datetime import date, datetime, timedelta, timezone as utc
from decimal import Decimal
from uuid import uuid4
from zoneinfo import ZoneInfo

from django.contrib.auth import get_user_model
from django.db import connection
from django.test import TestCase
from django.test.utils import CaptureQueriesContext
from django.utils import timezone

from apps.commercial.test_payment import PaymentFixture, shared_payment_data
from apps.commercial import models as money
from apps.inventory import models as stock_models, services as stock
from apps.reporting import service_analytics as service, service_dashboard as dashboard
from apps.reporting import commercial_analytics as commercial, inventory_analytics as inventory
from apps.reporting.filters import period
from apps.service import models as technical, tests as intake, test_diagnosis as diagnosis
from apps.service import test_repair as repair, test_quality_control as qc, test_engineer_assignment as engineer
from apps.service.services import add_service_case_complaint, cancel_service_case
from apps.service_catalog.models import ComplaintSymptom


def selected(builder, actor, key, filters=None):
    return next(t for t in builder(actor, filters or {}) if t.key == key)


class CommercialReconciliationAudit(PaymentFixture, TestCase):
    @classmethod
    def setUpTestData(cls):
        shared_payment_data(cls)

    def test_invoice_and_quote_values_reconcile_independent_lines(self):
        invoice = selected(commercial.commercial_tables, self.actor, "invoice_values").rows.get()
        lines = list(money.InvoiceLine.objects.filter(invoice=self.invoice, is_active=True))
        self.assertEqual(invoice["value"], sum(line.total for line in lines))
        for key, payer in (("customer", "CUSTOMER"), ("warranty", "WARRANTY"), ("company", "COMPANY")):
            self.assertEqual(invoice[key], sum(line.total for line in lines if line.responsibility == payer))
        quote = selected(commercial.commercial_tables, self.actor, "quotation_values").rows.get()
        expected = sum(money.QuotationLine.objects.filter(quotation=self.quote, is_active=True).values_list("total", flat=True))
        self.assertEqual(quote["value"], expected)
        for event in ("created", "submitted", "approved"):
            self.assertEqual(selected(commercial.commercial_tables, self.actor, "quotations_" + event).rows.get()["count"], 1)

    def test_payment_reversal_due_release_and_closed_history_reconcile(self):
        first = self.pay(amount="30")
        self.pay(amount="20")
        self.reverse_payment(first)
        self.release_due()
        self.return_unused(self.issue_row)
        self.deliver()
        from apps.service.test_handover import close
        self.case.refresh_from_db()
        close(self)
        payments = list(money.ServicePayment.objects.all())
        reversals = sum(money.PaymentReversal.objects.values_list("amount", flat=True))
        valid = sum(p.amount for p in payments if p.status == "POSTED")
        report = commercial.collection_totals(self.actor, {}).get()
        self.assertEqual((report["received"], report["reversed"], report["net_collected"]), (sum(p.amount for p in payments), reversals, valid))
        settlement = commercial.settlements(self.actor, {}).get()
        self.assertEqual((settlement.balance_due, settlement.clearance), (self.invoice.customer_pay_total - valid, "DUE_RELEASE"))
        self.assertEqual(money.ServicePaymentReceipt.objects.count(), 2)
        counts = dashboard.dashboard(self.actor, {})
        self.assertEqual((counts["current_open"], counts["closed_in_period"], counts["delivered_in_period"]), (0, 1, 1))

    def test_usage_reservations_and_ledger_reconcile_without_netting_unused(self):
        self.return_unused(self.issue_row)
        expected = Counter()
        for item in stock_models.PartsDisposition.objects.all():
            expected[item.kind] += item.quantity
        row = selected(inventory.inventory_tables, self.actor, "usage_parts").rows.get()
        self.assertEqual((row["consumed"], row["returned_unused"], row["net_consumption"]), (expected["CONSUMED"], expected["RETURNED"], expected["CONSUMED"]))
        ledger = Counter()
        for item in stock_models.StockLedgerEntry.objects.all():
            ledger[item.location_id, item.spare_part_id] += item.quantity_delta
        for item in inventory.positions(self.actor, {}):
            self.assertEqual(item.on_hand, ledger[item.location_id, item.spare_part_id])
        expected_reservations = Counter()
        for item in stock_models.StockReservation.objects.all():
            expected_reservations[item.spare_part_id, item.status] += item.quantity
        actual = {(r["spare_part_id"], r["status"]): r["quantity"] for r in selected(inventory.inventory_tables, self.actor, "reservations").rows}
        self.assertEqual(actual, expected_reservations)

    def test_all_tables_are_read_only_including_stream_iteration(self):
        from apps.reporting.management_reports import management_tables
        from apps.reporting.exports import csv_response
        def read_only(execute, sql, params, many, context):
            self.assertTrue(sql.lstrip().upper().startswith("SELECT"), sql[:120])
            self.assertNotIn("FOR UPDATE", sql.upper())
            self.assertNotIn("FOR SHARE", sql.upper())
            return execute(sql, params, many, context)
        with connection.execute_wrapper(read_only):
            for builder in (dashboard.dashboard_tables, service.service_tables, inventory.inventory_tables, commercial.commercial_tables, management_tables):
                for report in builder(self.actor, {}):
                    list(report.rows[:50])
                    list(csv_response(report).streaming_content)

    def test_performance_all_representative_pages_and_streaming(self):
        self.pay(amount="1")
        self.client.force_login(self.actor)
        reports = [("operational", "cases"), ("operational", "engineer_workload"), ("service", "complaints_complaint"),
                   ("service", "root_causes"), ("service", "cooccurrence_diagnosis_root_cause"), ("service", "diagnosis_records"),
                   ("inventory", "positions"), ("commercial", "quotations"), ("commercial", "invoices"),
                   ("commercial", "payments"), ("commercial", "outstanding")]
        for section, key in reports:
            with self.subTest(table=key), CaptureQueriesContext(connection) as captured:
                response = self.client.get(f"/reports/{section}/", {"table": key})
                self.assertEqual(response.status_code, 200)
            self.assertLessEqual(len(captured), 24 if section == "operational" else 10)
            print(f"Phase 3D audit page {key}: {len(captured)} queries")
        with CaptureQueriesContext(connection) as captured:
            response = self.client.get("/reports/commercial/", {"table": "payments", "export": "csv"})
            self.assertTrue(response.streaming)
            list(response.streaming_content)
        self.assertLessEqual(len(captured), 8)
        print(f"Phase 3D audit payment CSV including iteration: {len(captured)} queries")

    def test_payment_population_growth_has_constant_queries_and_csv_is_lazy(self):
        self.pay(amount="1")
        self.client.force_login(self.actor)
        def measured():
            with CaptureQueriesContext(connection) as captured:
                response = self.client.get("/reports/commercial/", {"table": "payments"})
            return len(captured), len(response.context["rows"])
        initial, _ = measured()
        for _ in range(19):
            self.pay(amount="1")
        final, count = measured()
        self.assertEqual((initial, count), (final, 20))
        from apps.reporting.exports import csv_response
        report = selected(commercial.commercial_tables, self.actor, "payments")
        with CaptureQueriesContext(connection) as captured:
            response = csv_response(report)
            stream = iter(response.streaming_content)
            next(stream)  # header does not fetch transaction rows
        self.assertEqual(len(captured), 0)
        with CaptureQueriesContext(connection) as captured:
            self.assertEqual(len(list(stream)), 20)
        self.assertEqual(len(captured), 1)
        self.assertIsNone(report.rows._result_cache)


class ServiceHistoryAudit(TestCase):
    @classmethod
    def setUpTestData(cls):
        diagnosis.setup_diagnosis(cls)
        cls.reader = get_user_model().objects.create_superuser(username="audit-service-reader")

    def test_complaint_filter_restricts_diagnosis_complaint_groups(self):
        engineer.unassign(self, reason="Prepare synthetic intake complaints")
        second = ComplaintSymptom.objects.create(code="AUDIT-SECOND", name="Second complaint", applies_to_all_product_categories=True)
        for symptom in (self.symptom, second):
            add_service_case_complaint(service_case=self.case, complaint_symptom=symptom)
        self.case.refresh_from_db()
        self.assignment = engineer.assign(self)
        assessment = diagnosis.begin(self)
        diagnosis.add(self, assessment)
        diagnosis.complete(self, assessment)
        rows = selected(service.service_tables, self.reader, "diagnosis_complaint", {"complaint": self.symptom.pk}).rows
        self.assertEqual({r["assessment__service_case__complaints__complaint_symptom_id"] for r in rows}, {self.symptom.pk})

    def test_abandoned_diagnosis_and_reassignment_do_not_inflate_workload(self):
        abandoned = diagnosis.begin(self)
        diagnosis.add(self, abandoned)
        diagnosis.abandon(self, abandoned)
        self.case.refresh_from_db()
        engineer.reassign(self, reason="Synthetic audit reassignment")
        self.case.refresh_from_db()
        self.assertEqual(list(service.diagnosis_rows(self.reader, {})), [])
        queue = selected(dashboard.dashboard_tables, self.reader, "engineer_queue").rows.get()
        self.assertEqual(queue["engineer_id"], self.engineer2.pk)
        engineer.unassign(self, reason="Synthetic audit unassignment")
        self.assertFalse(selected(dashboard.dashboard_tables, self.reader, "engineer_queue").rows.exists())
        self.case.refresh_from_db()
        cancel_service_case(service_case=self.case, cancelled_by=self.user, reason="Synthetic audit cancellation")
        self.assertEqual(dashboard.dashboard(self.reader, {})["cancelled_in_period"], 1)

    def test_findings_reconcile_independent_counter_and_null_bucket(self):
        assessment = diagnosis.begin(self)
        for fault, root in ((self.fault, None), (self.fault2, None), (self.fault, self.root)):
            diagnosis.add(self, assessment, fault_diagnosis=fault, root_cause=root)
        diagnosis.complete(self, assessment)
        findings = list(technical.ServiceDiagnosticFinding.objects.filter(assessment=assessment))
        for key, field in (("diagnoses", "fault_diagnosis_id"), ("root_causes", "root_cause_id")):
            expected = Counter(getattr(f, field) for f in findings)
            actual = {r[field]: r["count"] for r in selected(service.service_tables, self.reader, key).rows}
            self.assertEqual(actual, expected)
        row = selected(service.service_tables, self.reader, "root_causes", {"unknown_root_cause": "yes"}).rows.get()
        self.assertEqual((row["root_cause_label"], row["count"]), ("Unknown / Unconfirmed", 2))

    def test_failed_repair_then_success_counts_attempts_not_current_cases(self):
        from apps.service_catalog.models import RepairAction
        self.assessment = diagnosis.begin(self)
        diagnosis.add(self, self.assessment)
        self.assessment = diagnosis.complete(self, self.assessment)
        self.action_type = RepairAction.objects.create(code="AUDIT-REPAIR", name="Audit repair", applies_to_all_product_categories=True)
        first, _ = repair.prepared(self)
        repair.complete(self, first, outcome="NOT_REPAIRED")
        self.case.refresh_from_db()
        second, _ = repair.prepared(self)
        repair.complete(self, second)
        outcomes = {r["outcome"]: r["count"] for r in selected(service.service_tables, self.reader, "repair_completed").rows}
        self.assertEqual(outcomes, {"NOT_REPAIRED": 1, "REPAIRED": 1})
        self.assertEqual(service.service_rates(self.reader, {})[0]["percent"], 50)
        self.assertEqual(selected(dashboard.dashboard_tables, self.reader, "workflow").rows.get(), {"status": "REPAIRED", "count": 1})


class CalendarAudit(TestCase):
    @classmethod
    def setUpTestData(cls):
        intake.setup(cls)

    def test_local_day_boundaries_across_year_month_and_dst(self):
        for zone, day, hours in (("Asia/Dhaka", date(2025, 12, 31), 24), ("Asia/Dhaka", date(2024, 2, 29), 24),
                                 ("America/New_York", date(2025, 3, 9), 23), ("America/New_York", date(2025, 11, 2), 25)):
            with self.subTest(zone=zone, day=day), timezone.override(ZoneInfo(zone)):
                start = timezone.make_aware(datetime.combine(day, datetime.min.time())).astimezone(utc.utc)
                end = timezone.make_aware(datetime.combine(day + timedelta(days=1), datetime.min.time())).astimezone(utc.utc)
                self.assertEqual((end - start).total_seconds(), hours * 3600)
                records = [intake.intake(self, received_at=t) for t in (start - timedelta(microseconds=1), start, end - timedelta(microseconds=1), end)]
                rows = technical.ServiceCase.objects.filter(pk__in=[r.pk for r in records])
                ids = set(period(rows, "received_at", {"date_from": day, "date_to": day}).values_list("pk", flat=True))
                self.assertEqual(ids, {records[1].pk, records[2].pk})


from apps.inventory.test_documents import DocumentFixture
from apps.inventory.test_control import ControlFixture


class InventoryLedgerAudit(ControlFixture, DocumentFixture):
    def test_drafts_posting_transfers_and_adjustments_reconcile_signed_entries(self):
        receipt = self.receipt(quantity=9)
        self.assertFalse(selected(inventory.inventory_tables, self.actor, "movements").rows.exists())
        self.post(receipt)
        transfer = self.transfer(quantity=3)
        self.assertFalse(selected(inventory.inventory_tables, self.actor, "transfers").rows.exists())
        self.dispatch(transfer)
        self.receive_transfer(transfer)
        self.adjust(-1)
        self.assertEqual(selected(inventory.inventory_tables, self.actor, "transfers").rows.get()["quantity"], 3)
        expected = Counter()
        for entry in stock_models.StockLedgerEntry.objects.select_related("movement"):
            expected[entry.location_id, entry.movement.kind, entry.spare_part_id] += entry.quantity_delta
        actual = {(row["location_id"], row["movement__kind"], row["spare_part_id"]): row["delta"]
                  for row in selected(inventory.inventory_tables, self.actor, "movements").rows}
        self.assertEqual(actual, expected)
        self.assertEqual(sum(actual.values()), 8)
        self.assertEqual(sum(row.on_hand for row in inventory.positions(self.actor, {})), 8)

    def test_reconciled_count_variance_and_serialized_stock(self):
        self.receive(5)
        self.draft_count()  # drafts do not appear as completed counts
        count = self.reconcile(self.record_count(self.start_count(), 3))
        row = selected(inventory.inventory_tables, self.actor, "count_variances").rows.get()
        self.assertEqual((row["pk"], row["expected_quantity"], row["counted_quantity"], row["variance"]), (count.pk, 5, 3, -2))
        unit = stock.register_serialized_unit(actor=self.actor, company=self.company, spare_part=self.serial_part, identifier="AUDIT-STOCK-SERIAL")
        self.receive(1, spare_part=self.serial_part, units=[unit])
        rows = {r.spare_part_id: r for r in inventory.positions(self.actor, {})}
        self.assertEqual((rows[self.part.pk].on_hand, rows[self.part.pk].nonserialized_quantity), (3, 3))
        self.assertEqual((rows[self.serial_part.pk].on_hand, rows[self.serial_part.pk].serialized_units), (1, 1))


class MultiAttemptAudit(TestCase):
    @classmethod
    def setUpTestData(cls):
        qc.setup_qc(cls)
        cls.reader = get_user_model().objects.create_superuser(username="audit-attempt-reader")

    def test_abandoned_failed_and_passed_qc_reconcile_independent_attempts(self):
        abandoned = qc.prepared(self)
        qc.abandon(self, abandoned)
        self.case.refresh_from_db()
        failed = qc.begin(self)
        qc.check(self, failed, result="FAIL")
        qc.fail_qc(self, failed)
        self.case.refresh_from_db()
        abandoned_repair, _ = repair.prepared(self, performed=False)
        repair.abandon(self, abandoned_repair)
        self.case.refresh_from_db()
        execution, _ = repair.prepared(self)
        repair.complete(self, execution)
        self.case.refresh_from_db()
        passed = qc.prepared(self, filled=True)
        passed = qc.pass_qc(self, passed)
        attempts = list(technical.ServiceQualityControl.objects.all())
        expected = Counter((a.outcome, a.status) for a in attempts)
        actual = {(r["outcome"], r["status"]): r["count"] for r in selected(service.service_tables, self.reader, "qc_attempts").rows}
        self.assertEqual(actual, expected)
        completed = [a for a in attempts if a.status == "COMPLETED"]
        self.assertEqual(sum(r["count"] for r in selected(service.service_tables, self.reader, "qc_completed").rows), len(completed))
        self.assertEqual(service.service_rates(self.reader, {})[1]["numerator"], 0)
        self.assertEqual(dashboard.dashboard(self.reader, {})["current_open"], 1)
        self.assertEqual(selected(dashboard.dashboard_tables, self.reader, "workflow").rows.get(), {"status": "QC_PASSED", "count": 1})
        durations = dashboard.turnaround(self.reader, {})
        self.assertEqual([r["count"] for r in durations[:3]], [1, 2, 1])
        self.assertEqual(durations[2]["average"], passed.completed_at - passed.repair_execution.completed_at)


class MixedAllocationAudit(PaymentFixture, TestCase):
    @classmethod
    def setUpTestData(cls):
        from apps.inventory.test_usage import setup_usage
        setup_usage(cls)

    def test_manual_mixed_allocation_is_authoritative_and_drafts_excluded(self):
        invoice = self.billed(mixed=True, consumed=2)
        self.assertFalse(selected(commercial.commercial_tables, self.actor, "invoice_values").rows.exists())
        lines = list(self.quote.lines.order_by("responsibility"))
        invoice = self.reconcile(invoice, [self.allocation(line=line) for line in lines])
        self.ready()
        self.invoice = self.finalize(invoice)
        self.pay(amount="40")
        rows = list(money.InvoiceAllocation.objects.select_related("line"))
        self.assertEqual({row.mode for row in rows}, {"MANUAL"})
        result = selected(commercial.commercial_tables, self.actor, "invoice_values").rows.get()
        self.assertEqual((result["customer"], result["warranty"], result["value"]), (100, 100, 200))
        self.assertEqual(commercial.settlements(self.actor, {}).get().balance_due, 60)

    def test_company_only_invoice_and_recovered_component_remain_separate(self):
        from apps.inventory.usage_services import recover_defective_component
        invoice = self.billed(responsibility="COMPANY")
        quarantine = stock.create_location(actor=self.actor, company=self.company, service_center=self.center, code="AUDIT-RECOVERY", name="Audit recovery", location_type="QUARANTINE")
        recover_defective_component(actor=self.engineer, repair_action=self.action, location=quarantine,
                                   component_description="Synthetic removed component", command_key=uuid4(), replacement=self.disposition)
        self.ready()
        self.invoice = self.finalize(invoice)
        row = commercial.settlements(self.actor, {}).get()
        self.assertEqual((row.company_covered_total, row.customer_pay_total, row.clearance), (100, 0, "NO_CUSTOMER_DUE"))
        recovery = selected(inventory.inventory_tables, self.actor, "recovery").rows.get()
        self.assertEqual(recovery["quantity"], sum(stock_models.DefectiveRecovery.objects.values_list("quantity", flat=True)))
        self.assertFalse(stock_models.StockLedgerEntry.objects.filter(location=quarantine).exists())


class CooccurrenceReconciliationAudit(TestCase):
    @classmethod
    def setUpTestData(cls):
        from apps.reporting.tests.test_cooccurrence import CooccurrenceTests
        CooccurrenceTests.setUpTestData.__func__(cls)

    def test_raw_action_assessment_sets_reconcile_every_group(self):
        findings = list(technical.ServiceDiagnosticFinding.objects.filter(removed_at=None))
        actions = list(technical.ServiceRepairAction.objects.filter(performed_at__isnull=False).select_related("repair_execution"))
        for grouping in ("diagnosis", "root_cause", "diagnosis_root_cause"):
            expected = defaultdict(set)
            for action in actions:
                for finding in findings:
                    if finding.assessment_id != action.repair_execution.diagnostic_assessment_id:
                        continue
                    key = (action.repair_action_id,)
                    if grouping != "root_cause":
                        key += (finding.fault_diagnosis_id,)
                    if grouping != "diagnosis":
                        key += (finding.root_cause_id,)
                    expected[key].add(action.pk)
            actual = {}
            for row in service.repair_action_diagnosis_cooccurrence(self.reader, {}, grouping=grouping):
                key = (row["action_id"],)
                if grouping != "root_cause":
                    key += (row["fault_diagnosis_id"],)
                if grouping != "diagnosis":
                    key += (row["root_cause_id"],)
                actual[key] = row["performed_actions"]
            self.assertEqual(actual, {key: len(ids) for key, ids in expected.items()})
        self.assertEqual(service.performed_actions(self.reader, {}).count(), len(actions))

    def test_unknown_root_drilldown_and_csv_preserve_explicit_bucket(self):
        self.client.force_login(self.reader)
        response = self.client.get("/reports/service/", {"table": "root_causes", "unknown_root_cause": "yes"})
        self.assertEqual(response.context["rows"], [["Unknown / Unconfirmed", 2]])
        link = response.context["linked_rows"][0][1]
        drilled = self.client.get("/reports/service/" + link)
        self.assertEqual(len(drilled.context["rows"]), 2)
        self.assertEqual({row[3] for row in drilled.context["rows"]}, {"Unknown / Unconfirmed"})
        exported = self.client.get("/reports/service/" + link + "&export=csv")
        rows = list(csv.reader(StringIO(b"".join(exported.streaming_content).decode("utf-8-sig"))))
        self.assertEqual({row[3] for row in rows[1:]}, {"Unknown / Unconfirmed"})
        self.assertEqual({row[0] for row in rows[1:]}, {str(row[0]) for row in drilled.context["rows"]})
