"""Reporting HTTP, SQL read-only, scope, export and adversarial-input audit."""
import csv
from io import StringIO
from uuid import uuid4

from django.contrib.auth import get_user_model
from django.db import connection
from django.test import Client, TestCase, override_settings
from django.test.utils import CaptureQueriesContext

from apps.service import tests as intake
from apps.service.services import add_service_case_complaint
from apps.service_catalog.models import ComplaintSymptom
from apps.organization import services as organization
from apps.organization.models import Department
from apps.reporting.tests.test_reports import grant_report
from apps.reporting import permissions as p
from apps.reporting.scope import cases


class ReportingSecurityAudit(TestCase):
    @classmethod
    def setUpTestData(cls):
        intake.setup(cls)
        cls.local_case = intake.intake(cls)
        cls.other_case = intake.intake(cls, service_center=cls.center2)
        cls.foreign_case = intake.intake(cls, company=cls.other_company, service_center=cls.other_center, customer=cls.outsider, device=cls.device2)
        cls.reader = get_user_model().objects.create_user(username="audit-center-reader")
        cls.path, cls.report_role = grant_report(cls.reader, cls.company, cls.center)

    def setUp(self):
        self.client.force_login(self.reader)

    def test_all_case_filters_intersect_scope_and_contradictory_dimensions(self):
        for filters in ({"case": self.foreign_case.pk}, {"case": self.other_case.pk}, {"company": self.other_company.pk},
                        {"service_center": self.center2.pk}, {"model": uuid4()}, {"brand": uuid4()},
                        {"product_category": uuid4()}, {"variant": uuid4()}, {"region": uuid4()}):
            with self.subTest(filters=filters):
                self.assertFalse(cases(self.reader, p.OPERATIONAL, filters).exists())
                response = self.client.get("/reports/", filters | {"export": "csv"})
                self.assertEqual(response.status_code, 200)
                self.assertEqual(len(list(csv.reader(StringIO(b"".join(response.streaming_content).decode("utf-8-sig"))))), 1)

    def test_company_region_center_and_department_only_scopes(self):
        dept = Department.objects.create(company=self.company, code="AUDIT", name="Audit")
        for name, dimensions, expected in (("company", {}, 2), ("region", {"region": self.region}, 2),
                                           ("center", {"center": self.center}, 1), ("department", {"department": dept}, 0)):
            user = get_user_model().objects.create_user(username="audit-" + name)
            grant_report(user, self.company, **dimensions)
            self.assertEqual(cases(user, p.OPERATIONAL, {}).count(), expected, name)

    def test_deactivated_center_is_denied_but_superuser_matches_frozen_bypass(self):
        organization.deactivate_service_center(service_center=self.center)
        self.assertEqual(self.client.get("/reports/").status_code, 403)
        admin = get_user_model().objects.create_superuser(username="audit-active-admin")
        self.assertTrue(cases(admin, p.OPERATIONAL, {"case": self.local_case.pk}).exists())

    def test_deactivated_region_is_denied(self):
        organization.deactivate_region(region=self.region)
        self.assertEqual(self.client.get("/reports/", {"export": "csv"}).status_code, 403)

    def test_deactivated_company_is_denied(self):
        organization.deactivate_company(company=self.company)
        self.assertEqual(self.client.get("/reports/").status_code, 403)

    def test_export_scope_is_per_row_intersection(self):
        user = get_user_model().objects.create_user(username="audit-export-intersection")
        grant_report(user, self.company, permissions=["view_operational_dashboard"])
        grant_report(user, self.company, center=self.center2, permissions=["export_reports"])
        self.client.force_login(user)
        response = self.client.get("/reports/", {"export": "csv"})
        data = list(csv.reader(StringIO(b"".join(response.streaming_content).decode("utf-8-sig"))))
        self.assertEqual([r[0] for r in data[1:]], [str(self.other_case.pk)])

    @override_settings(DEBUG=False)
    def test_bad_inputs_fail_without_traceback_or_sql(self):
        for params in ({"page": "9" * 4000}, {"page": "NaN"}, {"date_from": "2025-02-30"}, {"sort": "(SELECT 1)"},
                       {"company": "' OR 1=1 --"}, {"table": "../../.env"}, {"status": "NOT_A_STATE"},
                       {"root_cause": str(uuid4()), "unknown_root_cause": "yes"}):
            with self.subTest(params=list(params)):
                response = self.client.get("/reports/", params)
                self.assertIn(response.status_code, (400, 404))
                self.assertNotContains(response, "Traceback", status_code=response.status_code)
                self.assertNotContains(response, "SELECT ", status_code=response.status_code)
        self.assertEqual(self.client.get("/reports/?page=1&page=2").status_code, 400)

    def test_xss_and_csv_injection_are_escaped_in_authoritative_labels(self):
        payload = '<script>alert("audit")</script>'
        for index, label in enumerate((payload, '=HYPERLINK("https://invalid.example")', '+1+1', '-1+1', '@SUM(A1)', ' \t=1+1', 'বাংলা')):
            symptom = ComplaintSymptom.objects.create(code="AUDIT-XSS-" + str(index), name=label, applies_to_all_product_categories=True)
            add_service_case_complaint(service_case=self.local_case, complaint_symptom=symptom)
        response = self.client.get("/reports/service/")
        self.assertNotContains(response, payload)
        self.assertContains(response, "&lt;script&gt;")
        exported = self.client.get("/reports/service/", {"export": "csv"})
        content = b"".join(exported.streaming_content)
        self.assertTrue(content.startswith(b"\xef\xbb\xbf"))
        rows = list(csv.reader(StringIO(content.decode("utf-8-sig"))))
        self.assertEqual(rows[0], ["ComplaintSymptom", "Intake complaint records"])
        labels = {row[0] for row in rows[1:]}
        self.assertIn("বাংলা", labels)
        self.assertTrue(any(label.startswith("'=HYPERLINK") for label in labels))
        self.assertFalse(any(label.startswith(("=", "+", "-", "@")) for label in labels))

    def test_login_has_csrf_protection_and_safe_redirect(self):
        client = Client(enforce_csrf_checks=True)
        self.assertEqual(client.post("/reports/login/", {"username": "ignored", "password": "synthetic"}).status_code, 403)
        self.reader.set_password("synthetic-audit-only")
        self.reader.save(update_fields=["password"])
        response = self.client.post("/reports/login/?next=https://invalid.example", {"username": self.reader.username, "password": "synthetic-audit-only", "next": "https://invalid.example"})
        self.assertEqual(response.url, "/reports/")

    def test_get_and_csv_execute_no_database_writes_or_row_locks(self):
        def readonly(execute, sql, params, many, context):
            self.assertTrue(sql.lstrip().upper().startswith("SELECT"), sql[:80])
            self.assertNotIn("FOR UPDATE", sql.upper())
            return execute(sql, params, many, context)
        with connection.execute_wrapper(readonly):
            for section in ("operational", "service", "inventory", "commercial", "management"):
                response = self.client.get(f"/reports/{section}/")
                self.assertEqual(response.status_code, 200)
                response = self.client.get(f"/reports/{section}/", {"export": "csv"})
                list(response.streaming_content)

    def test_pagination_and_csv_growth_are_bounded_and_equivalent(self):
        for _ in range(54):
            intake.intake(self)
        pages = [self.client.get("/reports/", {"page": page}).context["rows"] for page in (1, 2)]
        self.assertEqual([len(rows) for rows in pages], [50, 5])
        ids = [str(row[0]) for page in pages for row in page]
        self.assertEqual(len(set(ids)), 55)
        with CaptureQueriesContext(connection) as captured:
            response = self.client.get("/reports/", {"export": "csv"})
            exported = list(csv.reader(StringIO(b"".join(response.streaming_content).decode("utf-8-sig"))))
        self.assertEqual([r[0] for r in exported[1:]], ids)
        self.assertLessEqual(len(captured), 8)
        self.assertEqual(self.client.get("/reports/", {"page": 3}).status_code, 404)
