from django.contrib.auth.models import Permission, Group
from django.core.exceptions import PermissionDenied, ValidationError
from django.test import TestCase, Client
from django.test.utils import CaptureQueriesContext
from django.db import connection
from django.urls import reverse
from apps.access.services import set_role_permissions
from apps.inventory.test_usage import setup_usage
from apps.inventory.tests import grant
from .test_invoice import InvoiceFixture
from . import invoice_services as s, invoice_queries as q
from .models import ServiceInvoice


class InvoiceSecurityTests(InvoiceFixture, TestCase):
    @classmethod
    def setUpTestData(cls): setup_usage(cls)

    def grant_invoice(self, user, company, center=None, names=None):
        role = grant(user, company, center=center)
        permissions = Permission.objects.filter(content_type__app_label="commercial")
        if names is not None: permissions = permissions.filter(codename__in=names)
        set_role_permissions(role=role, permissions=permissions)
        return role

    def test_cross_company_write_and_read_denied(self):
        row = self.billed()
        self.grant_invoice(self.user, self.other_company)
        with self.assertRaises(PermissionDenied): self.reconcile(row, actor=self.user)
        self.assertFalse(q.service_invoices(actor=self.user).exists())

    def test_other_center_does_not_authorize(self):
        row = self.billed()
        self.grant_invoice(self.user, self.company, self.center2)
        with self.assertRaises(PermissionDenied): self.reconcile(row, actor=self.user)
        self.assertFalse(q.invoice_lines(actor=self.user, invoice=row).exists())

    def test_same_path_scope_and_capability_required(self):
        row = self.billed()
        self.grant_invoice(self.user, self.other_company)
        self.grant_invoice(self.user, self.company, names=[])
        with self.assertRaises(PermissionDenied): self.reconcile(row, actor=self.user)

    def test_native_staff_groups_direct_permissions_do_not_grant_scope(self):
        row = self.billed()
        self.user.is_staff = True
        self.user.save()
        permissions = Permission.objects.filter(content_type__app_label="commercial")
        self.user.user_permissions.set(permissions)
        group = Group.objects.create(name="Synthetic billing group")
        group.permissions.set(permissions)
        self.user.groups.add(group)
        with self.assertRaises(PermissionDenied): self.reconcile(row, actor=self.user)
        self.assertFalse(q.service_invoices(actor=self.user).exists())

    def test_engineer_is_not_payer_authority(self):
        row = self.billed(mixed=True)
        with self.assertRaises(PermissionDenied): self.reconcile(row, [self.allocation()], actor=self.engineer)

    def test_manage_does_not_grant_finalize(self):
        row = self.billed()
        self.grant_invoice(self.user, self.company, names=["manage_serviceinvoice"])
        row = self.reconcile(row, [self.allocation()], actor=self.user)
        self.ready()
        with self.assertRaises(PermissionDenied): self.finalize(row, actor=self.user)

    def test_inactive_superuser_denied(self):
        row = self.billed()
        self.actor.is_active = False
        self.actor.save()
        with self.assertRaises(ValidationError): self.reconcile(row)

    def test_returned_source_not_billable(self):
        row = self.billed()
        returned = self.return_unused(self.issue_row)
        with self.assertRaises(ValidationError): self.reconcile(row, [self.allocation(source=returned)])

    def test_old_quotation_line_not_reassignable(self):
        row = self.billed()
        old = self.quote.lines.get()
        self.quote = self.decide(self.submit(self.revise(self.quote)))
        row = self.prepare(invoice=row, expected_revision=s.revision(row))
        with self.assertRaises(ValidationError): self.reconcile(row, [self.allocation(line=old)])

    def test_cross_case_consumption_and_commercial_sources_rejected(self):
        from apps.service.tests import intake
        from apps.service.test_engineer_assignment import assign
        from apps.service import test_diagnosis as diagnosis
        row = self.billed()
        original_case, original_source, original_line = self.case, self.disposition, self.quote.lines.get()
        self.case = intake(self)
        assign(self)
        assessment = diagnosis.begin(self)
        diagnosis.add(self, assessment)
        diagnosis.complete(self, assessment)
        self.quoted()
        foreign_issue = self.prepared_issue()
        foreign_source = self.consume(foreign_issue, self.action)
        foreign_line = self.quote.lines.get()
        self.case = original_case
        for source, line in ((foreign_source, original_line), (original_source, foreign_line)):
            with self.assertRaises(ValidationError):
                self.reconcile(row, [dict(consumption=source, quotation_line=line, quantity=1, reason="Invalid cross-case allocation")])
        row.refresh_from_db()
        self.assertEqual(row.generation, 1)


class InvoiceAdminTests(InvoiceFixture, TestCase):
    @classmethod
    def setUpTestData(cls): setup_usage(cls)

    def setUp(self): self.client.force_login(self.actor)

    def url(self, row): return reverse("admin:commercial_invoice_workflow", args=[row.pk])

    def test_creation_delegates_and_totals_cannot_be_injected(self):
        self.quoted()
        response = self.client.post(reverse("admin:commercial_serviceinvoice_add"), dict(service_case=str(self.case.pk), note="Synthetic invoice", status="FINALIZED", grand_total="999", _save="Save"))
        self.assertEqual(response.status_code, 302)
        row = ServiceInvoice.objects.get()
        self.assertEqual((row.status, row.grand_total), ("DRAFT", 0))

    def test_get_does_not_finalize(self):
        row = self.billed()
        response = self.client.get(self.url(row), dict(operation="finalize"))
        self.assertEqual(response.status_code, 200)
        row.refresh_from_db()
        self.assertEqual(row.status, "DRAFT")

    def test_signed_finalization_and_immutable_display(self):
        row = self.billed()
        self.ready()
        token = self.client.get(self.url(row)).context_data["revision_token"]
        response = self.client.post(self.url(row), dict(operation="finalize", revision_token=token))
        self.assertEqual(response.status_code, 302)
        row.refresh_from_db()
        self.assertEqual(row.status, "FINALIZED")
        response = self.client.get(self.url(row))
        self.assertNotContains(response, 'value="reconcile"')
        self.assertEqual(self.client.post(reverse("admin:commercial_serviceinvoice_delete", args=[row.pk]), {"post":"yes"}).status_code, 403)

    def test_stale_and_tampered_finalization(self):
        row = self.billed()
        token = self.client.get(self.url(row)).context_data["revision_token"]
        row = self.reconcile(row, [self.allocation()])
        self.ready()
        for value in (token, "tampered"):
            self.assertEqual(self.client.post(self.url(row), dict(operation="finalize", revision_token=value)).status_code, 400)
        row.refresh_from_db()
        self.assertEqual(row.status, "DRAFT")

    def test_csrf_required(self):
        row = self.billed()
        client = Client(enforce_csrf_checks=True)
        client.force_login(self.actor)
        token = client.get(self.url(row)).context_data["revision_token"]
        self.assertEqual(client.post(self.url(row), dict(operation="finalize", revision_token=token)).status_code, 403)

    def test_out_of_scope_workflow_is_404(self):
        row = self.billed()
        self.user.is_staff = True
        self.user.save()
        self.user.user_permissions.set(Permission.objects.filter(content_type__app_label="commercial"))
        self.client.force_login(self.user)
        self.assertEqual(self.client.get(self.url(row)).status_code, 404)

    def test_manual_allocation_uses_service_and_fixed_terms(self):
        row = self.billed(mixed=True)
        source = self.quote.lines.get(responsibility="WARRANTY")
        token = self.client.get(self.url(row)).context_data["revision_token"]
        data = dict(operation="reconcile", revision_token=token, note="Synthetic reconciliation")
        for name in ("allocations", "confirmations", "adjustments"):
            data.update({f"{name}-TOTAL_FORMS":"1" if name == "allocations" else "0", f"{name}-INITIAL_FORMS":"0"})
        data.update({"allocations-0-consumption":str(self.disposition.pk), "allocations-0-quotation_line":str(source.pk),
            "allocations-0-quantity":"1", "allocations-0-reason":"Reviewed warranty funding", "allocations-0-responsibility":"CUSTOMER"})
        self.assertEqual(self.client.post(self.url(row), data).status_code, 302)
        row.refresh_from_db()
        self.assertEqual((row.customer_pay_total, row.warranty_covered_total), (0, 100))

    def test_stale_allocation_form_cannot_overwrite_new_payer(self):
        row = self.billed(mixed=True)
        token = self.client.get(self.url(row)).context_data["revision_token"]
        warranty = self.quote.lines.get(responsibility="WARRANTY")
        self.reconcile(row, [self.allocation(line=warranty)])
        response = self.client.post(self.url(row), dict(operation="reconcile", revision_token=token))
        self.assertEqual(response.status_code, 400)
        row.refresh_from_db()
        self.assertEqual((row.customer_pay_total, row.warranty_covered_total), (0, 100))

    def test_reconciliation_display_query_budget_does_not_grow_with_rows(self):
        row = self.billed(mixed=True, consumed=2)
        self.client.get(self.url(row))
        with CaptureQueriesContext(connection) as empty:
            self.assertEqual(self.client.get(self.url(row)).status_code, 200)
        lines = list(self.quote.lines.all())
        row = self.reconcile(row, [self.allocation(line=lines[0]), self.allocation(line=lines[1])])
        with CaptureQueriesContext(connection) as populated:
            self.assertEqual(self.client.get(self.url(row)).status_code, 200)
        self.assertLessEqual(len(populated), len(empty))
        self.assertLessEqual(len(populated), 30)
