from __future__ import annotations

import unittest

from app.scripts.purge_data import (
    CONFIRMATION_TEXT,
    PROTECTED_COLLECTIONS,
    confirm_purge,
    get_target_collection_names,
    purge_database,
)


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


class _FakeCollection:
    def __init__(self, documents: list[dict[str, object]]) -> None:
        self.documents = list(documents)
        self.delete_calls: list[dict[str, object]] = []

    def delete_many(self, query: dict[str, object]) -> _FakeDeleteResult:
        self.delete_calls.append(dict(query))
        deleted_count = len(self.documents)
        self.documents.clear()
        return _FakeDeleteResult(deleted_count)


class _FakeDatabase(dict):
    def __init__(self, collections: dict[str, _FakeCollection]) -> None:
        super().__init__(collections)
        self.name = "EpicrisisTest"

    def list_collection_names(self) -> list[str]:
        return list(self.keys())


class PurgeDataScriptTest(unittest.TestCase):
    def test_get_target_collection_names_excludes_protected_collections(self) -> None:
        targets = get_target_collection_names(
            ["users", "schema_migrations", "audit_events", "processing_batches"],
        )

        self.assertEqual(targets, ["audit_events", "processing_batches"])

    def test_purge_database_deletes_documents_from_non_protected_collections_only(self) -> None:
        database = _FakeDatabase(
            {
                "users": _FakeCollection([{"_id": 1}]),
                "schema_migrations": _FakeCollection([{"version": "20260622_01"}]),
                "audit_events": _FakeCollection([{"_id": 2}, {"_id": 3}]),
                "processing_batches": _FakeCollection([{"_id": 4}]),
            }
        )

        summary = purge_database(database, database_name=database.name)

        self.assertEqual(summary.database_name, "EpicrisisTest")
        self.assertEqual(summary.protected_collections, tuple(sorted(PROTECTED_COLLECTIONS)))
        self.assertEqual(
            [item.name for item in summary.purged_collections],
            ["audit_events", "processing_batches"],
        )
        self.assertEqual(summary.total_deleted, 3)
        self.assertEqual(database["users"].documents, [{"_id": 1}])
        self.assertEqual(database["schema_migrations"].documents, [{"version": "20260622_01"}])
        self.assertEqual(database["audit_events"].documents, [])
        self.assertEqual(database["processing_batches"].documents, [])
        self.assertEqual(database["audit_events"].delete_calls, [{}])
        self.assertEqual(database["processing_batches"].delete_calls, [{}])

    def test_confirm_purge_returns_false_when_confirmation_text_does_not_match(self) -> None:
        outputs: list[str] = []

        confirmed = confirm_purge(
            database_name="EpicrisisTest",
            input_fn=lambda _prompt: "cancelar",
            output_fn=outputs.append,
        )

        self.assertFalse(confirmed)
        self.assertEqual(len(outputs), 3)
        self.assertIn("EpicrisisTest", outputs[0])
        self.assertIn("schema_migrations", outputs[1])
        self.assertIn("users", outputs[1])

    def test_confirm_purge_returns_true_when_confirmation_text_matches(self) -> None:
        confirmed = confirm_purge(
            database_name="EpicrisisTest",
            input_fn=lambda _prompt: CONFIRMATION_TEXT,
            output_fn=lambda _message: None,
        )

        self.assertTrue(confirmed)


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