from __future__ import annotations

import unittest
from types import SimpleNamespace

from bson import ObjectId

from app.individual_ingestion.domain.models import UPLOAD_STATUS_ELIMINADO
from app.services.case_deletion_service import CaseDeletionError, CaseDeletionService


class FakeDeleteResult:
    def __init__(self, deleted_count: int) -> None:
        self.deleted_count = deleted_count


class FakeCollection:
    def __init__(self, items: list[dict]) -> None:
        self.items = list(items)

    def find_one(self, query: dict) -> dict | None:
        for item in self.items:
            if _matches(item, query):
                return dict(item)
        return None

    def find(self, query: dict) -> list[dict]:
        return [dict(item) for item in self.items if _matches(item, query)]

    def delete_one(self, query: dict) -> FakeDeleteResult:
        for index, item in enumerate(self.items):
            if _matches(item, query):
                self.items.pop(index)
                return FakeDeleteResult(1)
        return FakeDeleteResult(0)

    def delete_many(self, query: dict) -> FakeDeleteResult:
        kept: list[dict] = []
        deleted = 0
        for item in self.items:
            if _matches(item, query):
                deleted += 1
                continue
            kept.append(item)
        self.items = kept
        return FakeDeleteResult(deleted)

    def count_documents(self, query: dict) -> int:
        return sum(1 for item in self.items if _matches(item, query))


def _matches(item: dict, query: dict) -> bool:
    for key, value in query.items():
        current = item.get(key)
        if isinstance(value, dict):
            if "$in" in value and current not in value["$in"]:
                return False
            if "$nin" in value and current in value["$nin"]:
                return False
            continue
        if current != value:
            return False
    return True


class FakeIndividualUploadInvalidation:
    def __init__(self, records: list[dict]) -> None:
        self.records = records

    def _invalidate(self, predicate, *, reason: str, document_id: str = "", case_key: str = "") -> int:
        changed = 0
        for record in self.records:
            if record.get("status") == UPLOAD_STATUS_ELIMINADO or not predicate(record):
                continue
            record.update(
                {
                    "status": UPLOAD_STATUS_ELIMINADO,
                    "deletion_reason": reason,
                    "deletion_document_id": document_id,
                    "deletion_case_key": case_key,
                    "stored_path": "",
                    "extracted_text": "",
                }
            )
            changed += 1
        return changed

    def invalidate_by_clinical_document_id(self, **kwargs) -> int:
        return self._invalidate(
            lambda record: record.get("clinical_document_id") == kwargs["clinical_document_id"],
            reason=kwargs["reason"],
            document_id=kwargs.get("document_id", ""),
            case_key=kwargs.get("case_key", ""),
        )

    def invalidate_by_case_key(self, **kwargs) -> int:
        case_key = kwargs["case_key"]
        return self._invalidate(
            lambda record: case_key
            in {
                record.get("case_key", ""),
                record.get("provided_case_key", ""),
                record.get("session_case_key", ""),
            },
            reason=kwargs["reason"],
            document_id=kwargs.get("document_id", ""),
            case_key=case_key,
        )

    def invalidate_by_batch_file_id(self, **kwargs) -> int:
        return self._invalidate(
            lambda record: record.get("batch_file_id") == kwargs["batch_file_id"],
            reason=kwargs["reason"],
        )


class CaseDeletionServiceTest(unittest.TestCase):
    def setUp(self) -> None:
        self.base_id_a = ObjectId()
        self.base_id_b = ObjectId()
        self.derived_id = ObjectId()
        self.template_id = ObjectId()

    def _build_service(
        self,
        analysis_docs: list[dict],
        runtime_docs: list[dict],
        uploads: list[dict] | None = None,
    ) -> CaseDeletionService:
        mongo_analyses = SimpleNamespace(collection=FakeCollection(analysis_docs))
        mongo_database = {"processing_batch_cases": FakeCollection(runtime_docs)}
        return CaseDeletionService(
            mongo_analyses=mongo_analyses,
            mongo_database=mongo_database,
            individual_upload_invalidation=FakeIndividualUploadInvalidation(uploads or []),
        )

    def test_delete_document_removes_only_requested_base_document(self) -> None:
        uploads = [
            {
                "_id": "upload-a",
                "clinical_document_id": str(self.base_id_a),
                "case_key": "CASE-1",
                "status": "completado",
                "stored_path": "/tmp/a.pdf",
            },
            {
                "_id": "upload-b",
                "clinical_document_id": str(self.base_id_b),
                "case_key": "CASE-1",
                "status": "completado",
            },
        ]
        service = self._build_service(
            [
                {"_id": self.base_id_a, "usuario": "auditor", "case_key": "CASE-1", "tipo_documento": "factura"},
                {"_id": self.base_id_b, "usuario": "auditor", "case_key": "CASE-1", "tipo_documento": "laboratorio"},
                {
                    "_id": self.derived_id,
                    "usuario": "auditor",
                    "case_key": "CASE-1",
                    "tipo_documento": "rips_case_payload",
                },
            ],
            [{"_id": ObjectId(), "usuario": "auditor", "case_key": "CASE-1"}],
            uploads,
        )

        payload = service.delete_document(username="auditor", document_id=str(self.base_id_a))

        self.assertEqual(payload["deleted_document_id"], str(self.base_id_a))
        self.assertEqual(payload["remaining_documents"], 1)
        self.assertFalse(payload["case_deleted"])
        self.assertEqual(service.mongo_analyses.collection.count_documents({"usuario": "auditor"}), 2)
        self.assertEqual(service.case_runtime_collection.count_documents({"usuario": "auditor"}), 1)
        self.assertEqual(uploads[0]["status"], UPLOAD_STATUS_ELIMINADO)
        self.assertEqual(uploads[1]["status"], "completado")

    def test_delete_last_document_invalidates_all_case_uploads_including_pending(self) -> None:
        uploads = [
            {
                "clinical_document_id": str(self.base_id_a),
                "case_key": "CASE-1",
                "status": "completado",
            },
            {
                "provided_case_key": "CASE-1",
                "session_case_key": "CASE-1",
                "status": "precheck_en_cola",
                "stored_path": "/tmp/pending.pdf",
                "extracted_text": "pendiente",
            },
        ]
        service = self._build_service(
            [
                {"_id": self.base_id_a, "usuario": "auditor", "case_key": "CASE-1", "tipo_documento": "factura"},
            ],
            [],
            uploads,
        )

        payload = service.delete_document(username="auditor", document_id=str(self.base_id_a))

        self.assertTrue(payload["case_deleted"])
        self.assertTrue(all(item["status"] == UPLOAD_STATUS_ELIMINADO for item in uploads))
        self.assertTrue(all(item.get("stored_path", "") == "" for item in uploads))

    def test_delete_document_rejects_derived_artifact(self) -> None:
        service = self._build_service(
            [
                {
                    "_id": self.derived_id,
                    "usuario": "auditor",
                    "case_key": "CASE-1",
                    "tipo_documento": "rda_case_artifact",
                }
            ],
            [],
        )

        with self.assertRaises(CaseDeletionError) as ctx:
            service.delete_document(username="auditor", document_id=str(self.derived_id))

        self.assertEqual(ctx.exception.status_code, 400)

    def test_delete_document_last_base_cleans_derivatives_and_runtime(self) -> None:
        service = self._build_service(
            [
                {"_id": self.base_id_a, "usuario": "auditor", "case_key": "CASE-1", "tipo_documento": "factura"},
                {
                    "_id": self.derived_id,
                    "usuario": "auditor",
                    "case_key": "CASE-1",
                    "tipo_documento": "epicrisis_case_cache",
                },
                {
                    "_id": self.template_id,
                    "usuario": "auditor",
                    "case_key": "CASE-1",
                    "tipo_documento": "rips_operational_template",
                },
            ],
            [{"_id": ObjectId(), "usuario": "auditor", "case_key": "CASE-1"}],
        )

        payload = service.delete_document(username="auditor", document_id=str(self.base_id_a))

        self.assertTrue(payload["case_deleted"])
        self.assertEqual(payload["remaining_documents"], 0)
        remaining_types = {item["tipo_documento"] for item in service.mongo_analyses.collection.items}
        self.assertEqual(remaining_types, {"rips_operational_template"})
        self.assertEqual(service.case_runtime_collection.count_documents({"usuario": "auditor"}), 0)

    def test_delete_case_removes_base_documents_and_derivatives_but_keeps_template(self) -> None:
        uploads = [
            {"provided_case_key": "CASE-1", "status": "esperando_confirmacion"},
            {"session_case_key": "CASE-1", "status": "procesando"},
        ]
        service = self._build_service(
            [
                {"_id": self.base_id_a, "usuario": "auditor", "case_key": "CASE-1", "tipo_documento": "factura"},
                {"_id": self.base_id_b, "usuario": "auditor", "case_key": "CASE-1", "tipo_documento": "laboratorio"},
                {
                    "_id": self.derived_id,
                    "usuario": "auditor",
                    "case_key": "CASE-1",
                    "tipo_documento": "rda_case_job_status",
                },
                {
                    "_id": self.template_id,
                    "usuario": "auditor",
                    "case_key": "CASE-1",
                    "tipo_documento": "rips_operational_template",
                },
            ],
            [{"_id": ObjectId(), "usuario": "auditor", "case_key": "CASE-1"}],
            uploads,
        )

        payload = service.delete_case(username="auditor", case_key="CASE-1")

        self.assertEqual(payload["deleted_documents"], 2)
        self.assertEqual(payload["deleted_derived_artifacts"], 2)
        remaining_types = {item["tipo_documento"] for item in service.mongo_analyses.collection.items}
        self.assertEqual(remaining_types, {"rips_operational_template"})
        self.assertEqual(service.case_runtime_collection.count_documents({"usuario": "auditor"}), 0)
        self.assertTrue(all(item["status"] == UPLOAD_STATUS_ELIMINADO for item in uploads))

    def test_delete_case_rejects_missing_case(self) -> None:
        service = self._build_service([], [])

        with self.assertRaises(CaseDeletionError) as ctx:
            service.delete_case(username="auditor", case_key="CASE-404")

        self.assertEqual(ctx.exception.status_code, 404)


if __name__ == "__main__":
    unittest.main()
