from __future__ import annotations

import importlib
import sys
import unittest
from datetime import UTC, datetime, timedelta
from types import SimpleNamespace
from typing import Any, cast
from unittest.mock import patch

from app.config import config
from app.core.logging import clear_log_context, set_log_context
from app.llm import LLMOutputKind, LLMResolvedRoute, LLMStructuredResult, LLMTask, LLMTextResult
from app.llm.routing import DefaultModelSelectionPolicy
from app.llm.schemas import HistoriaClinicaStructured
from app.services.clinical_structured_extraction import ClinicalStructuredExtractionService
from app.services.llm_task_cache import (
    MongoLLMTaskCacheRepository,
    build_llm_task_fingerprint,
)


class _FakeCollection:
    def __init__(self) -> None:
        self.docs: list[dict] = []
        self.indexes: list[tuple] = []
        self._id_counter = 0

    def create_index(self, keys, **kwargs):
        self.indexes.append((tuple(keys), dict(kwargs)))

    def find_one(self, query, sort=None):
        now = datetime.now(UTC)
        for doc in reversed(self.docs):
            if doc.get("tipo_documento") != query.get("tipo_documento"):
                continue
            if doc.get("usuario") != query.get("usuario"):
                continue
            if doc.get("task") != query.get("task"):
                continue
            if doc.get("fingerprint") != query.get("fingerprint"):
                continue
            expires_at = doc.get("expires_at")
            if isinstance(expires_at, datetime) and expires_at <= now:
                continue
            return dict(doc)
        return None

    def update_one(self, query, update, upsert=False):
        if "$inc" in update:
            for doc in self.docs:
                if doc.get("_id") == query.get("_id"):
                    for key, value in update["$inc"].items():
                        doc[key] = int(doc.get(key, 0)) + int(value)
                    return SimpleNamespace(matched_count=1, modified_count=1, upserted_id=None)
            return SimpleNamespace(matched_count=0, modified_count=0, upserted_id=None)

        matched = None
        for doc in self.docs:
            if (
                doc.get("tipo_documento") == query.get("tipo_documento")
                and doc.get("usuario") == query.get("usuario")
                and doc.get("task") == query.get("task")
                and doc.get("fingerprint") == query.get("fingerprint")
            ):
                matched = doc
                break
        if matched is None and upsert:
            self._id_counter += 1
            matched = {"_id": self._id_counter}
            self.docs.append(matched)

        if matched is not None:
            for key, value in update.get("$setOnInsert", {}).items():
                matched.setdefault(key, value)
            matched.update(update.get("$set", {}))
        return SimpleNamespace(matched_count=1 if matched else 0, modified_count=1 if matched else 0, upserted_id=None)


class _MemoryCacheRepository:
    def __init__(self) -> None:
        self.cache_version = "v1"
        self.items: dict[tuple[str, str, str], dict] = {}

    def get(self, *, username: str, task: str, fingerprint: str):
        return self.items.get((username, task, fingerprint))

    def upsert(self, *, username: str, task: str, fingerprint: str, provider: str, model: str, payload, metadata=None):
        self.items[(username, task, fingerprint)] = {
            "usuario": username,
            "task": task,
            "fingerprint": fingerprint,
            "provider": provider,
            "model": model,
            "payload": payload,
            "metadata": dict(metadata or {}),
        }


class _StructuredRouter:
    def __init__(self, *, fail: bool = False) -> None:
        self.fail = fail
        self.calls = 0

    def _resolve_route(self, task):
        return LLMResolvedRoute(
            task=task,
            output_kind=LLMOutputKind.STRUCTURED_OBJECT,
            provider="gemini",
            model="gemini-2.5-flash-lite",
        )

    def generate_structured(self, request):
        self.calls += 1
        if self.fail:
            raise AssertionError("LLM should not be called on cache hit")
        payload = HistoriaClinicaStructured(
            patient_name="PACIENTE CACHE",
            resumen_clinico="Resumen cacheable",
            diagnosticos=[],
            procedimientos=[],
            medicamentos=[],
        ).model_dump(by_alias=True, exclude_none=True, exclude_defaults=True)
        return LLMStructuredResult(
            content=payload,
            provider="gemini",
            model="gemini-2.5-flash-lite",
            output_kind=request.output_kind,
        )


class _TextRouter:
    def __init__(self) -> None:
        self.calls: list[tuple[str, str]] = []

    def _resolve_route(self, task):
        model = "gemini-2.5-flash-lite" if task == LLMTask.HISTORIA_CHUNK_SUMMARY else "gemini-2.5-flash"
        return LLMResolvedRoute(task=task, output_kind=LLMOutputKind.CONTROLLED_PLAIN_TEXT, provider="gemini", model=model)

    def generate_text(self, request):
        self.calls.append((request.task.value, request.model or ""))
        text = f"respuesta-{request.task.value}-{len(self.calls)}"
        return LLMTextResult(content=text, provider="gemini", model=request.model or "unset")


class LLMTaskFingerprintTest(unittest.TestCase):
    def test_fingerprint_is_stable_for_same_semantics(self) -> None:
        first = build_llm_task_fingerprint(
            task="clinical_document_extract",
            provider="gemini",
            model="gemini-2.5-flash-lite",
            cache_version="v1",
            prompt_version="p1",
            schema_version="v1",
            payload={"text": "Paciente   uno"},
            metadata={"b": 2, "a": 1},
        )
        second = build_llm_task_fingerprint(
            task="clinical_document_extract",
            provider="gemini",
            model="gemini-2.5-flash-lite",
            cache_version="v1",
            prompt_version="p1",
            schema_version="v1",
            payload={"text": "Paciente uno"},
            metadata={"a": 1, "b": 2},
        )

        self.assertEqual(first, second)

    def test_fingerprint_changes_when_cache_dimensions_change(self) -> None:
        base = build_llm_task_fingerprint(
            task="clinical_document_extract",
            provider="gemini",
            model="gemini-2.5-flash-lite",
            cache_version="v1",
            prompt_version="p1",
            schema_version="v1",
            payload={"text": "Paciente uno"},
        )

        changed_model = build_llm_task_fingerprint(
            task="clinical_document_extract",
            provider="gemini",
            model="gemini-2.5-pro",
            cache_version="v1",
            prompt_version="p1",
            schema_version="v1",
            payload={"text": "Paciente uno"},
        )
        changed_cache_version = build_llm_task_fingerprint(
            task="clinical_document_extract",
            provider="gemini",
            model="gemini-2.5-flash-lite",
            cache_version="v2",
            prompt_version="p1",
            schema_version="v1",
            payload={"text": "Paciente uno"},
        )

        self.assertNotEqual(base, changed_model)
        self.assertNotEqual(base, changed_cache_version)


class ClinicalExtractRoutingPolicyTest(unittest.TestCase):
    def test_factura_route_uses_high_risk_extract_model(self) -> None:
        route = DefaultModelSelectionPolicy().resolve(
            LLMTask.CLINICAL_DOCUMENT_EXTRACT,
            metadata={"document_type": "factura", "risk_level": "high"},
        )

        self.assertEqual(route.provider, "gemini")
        self.assertEqual(route.model, config.GEMINI_MODEL_EXTRACT_HIGH_RISK)
        self.assertEqual(route.metadata["risk_level"], "high")

    def test_low_risk_route_keeps_low_risk_model(self) -> None:
        route = DefaultModelSelectionPolicy().resolve(
            LLMTask.CLINICAL_DOCUMENT_EXTRACT,
            metadata={"document_type": "laboratorio", "risk_level": "low"},
        )

        self.assertEqual(route.model, config.GEMINI_MODEL_EXTRACT_LOW_RISK)


class MongoLLMTaskCacheRepositoryTest(unittest.TestCase):
    def test_upsert_and_get_return_active_hit(self) -> None:
        collection = _FakeCollection()
        repository = MongoLLMTaskCacheRepository(SimpleNamespace(collection=collection))
        repository.ensure_indexes()
        repository.upsert(
            username="auditor",
            task="historia_chunk_summary",
            fingerprint="abc",
            provider="gemini",
            model="gemini-2.5-flash-lite",
            payload={"text": "Resumen"},
            metadata={"prompt_version": "v1"},
        )

        cached = repository.get(username="auditor", task="historia_chunk_summary", fingerprint="abc")

        self.assertIsNotNone(cached)
        cached_doc = cast(dict[str, Any], cached)
        self.assertEqual(cached_doc["payload"]["text"], "Resumen")
        self.assertEqual(collection.docs[0]["hit_count"], 1)
        self.assertEqual(len(collection.indexes), 3)
        self.assertIs(repository.collection, collection)

    def test_get_skips_expired_documents(self) -> None:
        collection = _FakeCollection()
        collection.docs.append(
            {
                "_id": 1,
                "tipo_documento": "llm_task_cache",
                "usuario": "auditor",
                "task": "historia_chunk_summary",
                "fingerprint": "abc",
                "payload": {"text": "viejo"},
                "expires_at": datetime.now(UTC) - timedelta(days=1),
            }
        )
        repository = MongoLLMTaskCacheRepository(SimpleNamespace(collection=collection))

        cached = repository.get(username="auditor", task="historia_chunk_summary", fingerprint="abc")

        self.assertIsNone(cached)

    def test_repository_accepts_direct_collection_backend(self) -> None:
        collection = _FakeCollection()
        repository = MongoLLMTaskCacheRepository(collection)

        repository.upsert(
            username="auditor",
            task="clinical_document_extract",
            fingerprint="fp-001",
            provider="gemini",
            model="gemini-2.5-flash-lite",
            payload={"text": "ok"},
        )

        cached = repository.get(username="auditor", task="clinical_document_extract", fingerprint="fp-001")

        self.assertIsNotNone(cached)
        cached_doc = cast(dict[str, Any], cached)
        self.assertEqual(cached_doc["payload"]["text"], "ok")


class ClinicalStructuredExtractionCacheTest(unittest.TestCase):
    def tearDown(self) -> None:
        clear_log_context()

    def test_extract_reuses_cached_structured_result(self) -> None:
        cache_repository = _MemoryCacheRepository()
        set_log_context({"username": "auditor"})

        first_service = ClinicalStructuredExtractionService(
            llm_router=_StructuredRouter(),
            cache_repository=cache_repository,
        )
        first = first_service.extract(raw_text="Historia clínica ejemplo", document_type="historia_clinica")

        second_service = ClinicalStructuredExtractionService(
            llm_router=_StructuredRouter(fail=True),
            cache_repository=cache_repository,
        )
        second = second_service.extract(raw_text="Historia clínica ejemplo", document_type="historia_clinica")

        self.assertEqual(first.analysis_structured, second.analysis_structured)
        self.assertEqual(second.metrics, {"source": "cache"})
        self.assertEqual(second.analysis_model.patient_name, "PACIENTE CACHE")


class HistoriaSummaryCacheTest(unittest.TestCase):
    def test_prepare_historia_prompt_reuses_chunk_and_final_summary_cache(self) -> None:
        cache_repository = _MemoryCacheRepository()
        router = _TextRouter()
        with patch.object(config, "GROQ_API_KEY", "fake-key"):
            module = sys.modules.get("modules.processing.resumen_google")
            if module is None:
                historia_module = importlib.import_module("modules.processing.resumen_google")
            else:
                historia_module = module

        descripcion = "\n".join(
            f"Evolucion dia {idx}: dolor abdominal persistente, signos vitales estables y plan terapeutico {idx}."
            for idx in range(24)
        )

        with patch.object(config, "HISTORIA_COMPACT_MAX_CHARS", 6000):
            first = historia_module._prepare_historia_prompt_text(
                descripcion,
                summary_threshold=300,
                summary_chunk=250,
                max_chars=500,
                llm_router=router,
                cache_repository=cache_repository,
                username="auditor",
            )
            first_call_count = len(router.calls)
            second = historia_module._prepare_historia_prompt_text(
                descripcion,
                summary_threshold=300,
                summary_chunk=250,
                max_chars=500,
                llm_router=router,
                cache_repository=cache_repository,
                username="auditor",
            )

        self.assertIn("version=v8", first)
        self.assertIn("[RESUMEN PREVIO DEL DOCUMENTO COMPLETO]", first)
        self.assertEqual(first, second)
        self.assertGreater(first_call_count, 1)
        self.assertEqual(len(router.calls), first_call_count)
