from __future__ import annotations

import unittest
from copy import deepcopy
from datetime import UTC, datetime
from types import SimpleNamespace

from fastapi import FastAPI
from httpx import ASGITransport, AsyncClient

from app.auth import get_current_user
from app.routes import rda


class FakeRdaService:
    def __init__(self) -> None:
        self.queue_calls: list[dict[str, object]] = []
        self.latest_calls: list[dict[str, object]] = []
        self.status_calls: list[dict[str, object]] = []
        self.validation_calls: list[dict[str, object]] = []
        self.queue_payload = {
            "case_key": "case-1",
            "artifact_type": "patient",
            "job_id": "job-123",
            "status": "queued",
            "reused": False,
        }
        self.artifact_payload = {
            "case_key": "case-1",
            "artifact_type": "patient",
            "version": 1,
            "schema_version": 1,
            "mapper_version": "mapper-v1",
            "completeness_status": "partial",
            "missing_fields": ["patient.clinical_background"],
            "section_trace": [],
            "payload": {"summary_type": "patient"},
            "source_context_fingerprint": "abc",
            "stored": True,
            "source": "cache",
        }
        self.status_payload = {
            "case_key": "case-1",
            "artifact_type": "patient",
            "status": "completed",
            "job_id": "job-123",
            "error": "",
            "source_context_fingerprint": "abc",
            "resolved_version": 1,
            "reused": True,
            "updated_at": datetime(2026, 5, 9, 10, 30, tzinfo=UTC),
            "source": "status",
        }
        self.validation_payload = {
            "case_key": "case-1",
            "artifact_type": "patient",
            "artifact": deepcopy(self.artifact_payload),
            "status": deepcopy(self.status_payload),
            "validation": {
                "artifact_type": "patient",
                "is_valid": True,
                "error_count": 0,
                "warning_count": 1,
                "findings": [
                    {
                        "severity": "warning",
                        "resource_type": "document_reference",
                        "resource_path": "projection_gaps[0]",
                        "rule_code": "RDA-FHIR-WARN-001",
                        "message": "Projection gap detected: document_reference:not_projected_in_v1.",
                        "observed_value": "document_reference:not_projected_in_v1",
                    }
                ],
            },
        }

    def queue_case_rda(self, username: str, case_key: str, artifact_type: str, *, force: bool = False):
        self.queue_calls.append(
            {
                "username": username,
                "case_key": case_key,
                "artifact_type": str(artifact_type),
                "force": force,
            }
        )
        payload = deepcopy(self.queue_payload)
        payload["case_key"] = case_key
        payload["artifact_type"] = str(artifact_type)
        return SimpleNamespace(model_dump=lambda **kwargs: self._model_dump(payload, **kwargs))

    def get_latest_case_rda(self, username: str, case_key: str, artifact_type: str):
        self.latest_calls.append(
            {
                "username": username,
                "case_key": case_key,
                "artifact_type": str(artifact_type),
            }
        )
        if case_key == "missing":
            return None
        payload = deepcopy(self.artifact_payload)
        payload["case_key"] = case_key
        payload["artifact_type"] = str(artifact_type)
        return SimpleNamespace(model_dump=lambda **kwargs: self._model_dump(payload, **kwargs))

    def get_case_rda_status(self, username: str, case_key: str, artifact_type: str):
        self.status_calls.append(
            {
                "username": username,
                "case_key": case_key,
                "artifact_type": str(artifact_type),
            }
        )
        if case_key == "missing":
            return None
        payload = deepcopy(self.status_payload)
        payload["case_key"] = case_key
        payload["artifact_type"] = str(artifact_type)
        return SimpleNamespace(model_dump=lambda **kwargs: self._model_dump(payload, **kwargs))

    def get_case_rda_validation(self, username: str, case_key: str, artifact_type: str):
        self.validation_calls.append(
            {
                "username": username,
                "case_key": case_key,
                "artifact_type": str(artifact_type),
            }
        )
        if case_key == "missing":
            return None
        payload = deepcopy(self.validation_payload)
        payload["case_key"] = case_key
        payload["artifact_type"] = str(artifact_type)
        if case_key == "queued-only":
            payload["artifact"] = None
            payload["validation"] = None
            payload["status"] = {
                "case_key": case_key,
                "artifact_type": str(artifact_type),
                "status": "queued",
                "job_id": "job-999",
                "error": "",
                "source_context_fingerprint": "",
                "resolved_version": None,
                "reused": False,
                "updated_at": datetime(2026, 5, 9, 10, 35, tzinfo=UTC),
                "source": "status",
            }
        return SimpleNamespace(model_dump=lambda **kwargs: self._model_dump(payload, **kwargs))

    def _model_dump(self, payload: dict[str, object], *, mode: str = "python", **_: object) -> dict[str, object]:
        return deepcopy(payload)


class RdaRoutesTest(unittest.IsolatedAsyncioTestCase):
    def setUp(self) -> None:
        self.fake_services = SimpleNamespace(rda_service=FakeRdaService())
        self.app = FastAPI()
        self.app.include_router(rda.router)
        self.app.state.services = self.fake_services

        async def _override_user():
            return SimpleNamespace(username="tester")

        self.app.dependency_overrides[get_current_user] = _override_user
        self.transport = ASGITransport(app=self.app)

    def tearDown(self) -> None:
        self.app.dependency_overrides.clear()

    async def test_post_queues_patient_rda_job(self) -> None:
        async with AsyncClient(transport=self.transport, base_url="http://testserver") as client:
            response = await client.post("/api/rda/cases/case-1/patient", json={"force": False})

        self.assertEqual(response.status_code, 202)
        self.assertEqual(response.json()["job_id"], "job-123")
        self.assertEqual(response.json()["status"], "queued")
        self.assertEqual(self.fake_services.rda_service.queue_calls[-1]["force"], False)

    async def test_post_requires_authentication(self) -> None:
        self.app.dependency_overrides.clear()

        async with AsyncClient(transport=self.transport, base_url="http://testserver") as client:
            response = await client.post("/api/rda/cases/case-1/patient", json={"force": False})

        self.assertEqual(response.status_code, 401)

    async def test_get_requires_authentication(self) -> None:
        self.app.dependency_overrides.clear()

        async with AsyncClient(transport=self.transport, base_url="http://testserver") as client:
            response = await client.get("/api/rda/cases/case-1/patient")

        self.assertEqual(response.status_code, 401)

    async def test_post_normalizes_case_key_and_propagates_force(self) -> None:
        async with AsyncClient(transport=self.transport, base_url="http://testserver") as client:
            response = await client.post("/api/rda/cases/  case-1  /patient", json={"force": True})

        self.assertEqual(response.status_code, 202)
        self.assertEqual(self.fake_services.rda_service.queue_calls[-1]["case_key"], "case-1")
        self.assertEqual(self.fake_services.rda_service.queue_calls[-1]["force"], True)

    async def test_post_returns_400_for_invalid_json(self) -> None:
        async with AsyncClient(transport=self.transport, base_url="http://testserver") as client:
            response = await client.post(
                "/api/rda/cases/case-1/patient",
                content="{bad",
                headers={"content-type": "application/json"},
            )

        self.assertEqual(response.status_code, 400)
        self.assertEqual(response.json()["detail"], "Payload invalido.")

    async def test_post_returns_400_for_non_object_json_payload(self) -> None:
        async with AsyncClient(transport=self.transport, base_url="http://testserver") as client:
            response = await client.post("/api/rda/cases/case-1/patient", json=[])

        self.assertEqual(response.status_code, 400)
        self.assertEqual(response.json()["detail"], "Payload invalido.")

    async def test_get_returns_404_when_artifact_is_missing(self) -> None:
        async with AsyncClient(transport=self.transport, base_url="http://testserver") as client:
            response = await client.get("/api/rda/cases/missing/patient")

        self.assertEqual(response.status_code, 404)

    async def test_get_status_returns_completed_state(self) -> None:
        async with AsyncClient(transport=self.transport, base_url="http://testserver") as client:
            response = await client.get("/api/rda/cases/case-1/patient/status")

        self.assertEqual(response.status_code, 200)
        self.assertEqual(response.json()["status"], "completed")
        self.assertEqual(response.json()["resolved_version"], 1)
        self.assertEqual(response.json()["updated_at"], "2026-05-09T10:30:00+00:00")

    async def test_get_status_returns_404_when_status_is_missing(self) -> None:
        async with AsyncClient(transport=self.transport, base_url="http://testserver") as client:
            response = await client.get("/api/rda/cases/missing/patient/status")

        self.assertEqual(response.status_code, 404)

    async def test_get_validation_returns_artifact_and_findings(self) -> None:
        async with AsyncClient(transport=self.transport, base_url="http://testserver") as client:
            response = await client.get("/api/rda/cases/case-1/patient/validation")

        self.assertEqual(response.status_code, 200)
        payload = response.json()
        self.assertEqual(payload["artifact"]["artifact_type"], "patient")
        self.assertTrue(payload["validation"]["is_valid"])
        self.assertEqual(payload["validation"]["warning_count"], 1)
        self.assertEqual(payload["validation"]["findings"][0]["rule_code"], "RDA-FHIR-WARN-001")

    async def test_get_validation_returns_status_when_artifact_is_not_ready(self) -> None:
        async with AsyncClient(transport=self.transport, base_url="http://testserver") as client:
            response = await client.get("/api/rda/cases/queued-only/patient/validation")

        self.assertEqual(response.status_code, 200)
        payload = response.json()
        self.assertIsNone(payload["artifact"])
        self.assertIsNone(payload["validation"])
        self.assertEqual(payload["status"]["status"], "queued")

    async def test_get_validation_returns_404_when_neither_artifact_nor_status_exist(self) -> None:
        async with AsyncClient(transport=self.transport, base_url="http://testserver") as client:
            response = await client.get("/api/rda/cases/missing/patient/validation")

        self.assertEqual(response.status_code, 404)
